#mjx Make pytree registration optional.
This is a noop for MJX itself, but is useful for custom types extending MJX types. This CL also ensures unused kwargs in `__init_subclass__` are passed to the parent class. PiperOrigin-RevId: 728215373 Change-Id: I362d3fb8d4b7d679a3a00d96866d46a52d711369
This commit is contained in:
committed by
Copybara-Service
parent
36142cddee
commit
982330d15a
@@ -34,7 +34,7 @@ def _jax_in_args(typ) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def dataclass(clz: _T) -> _T:
|
||||
def dataclass(clz: _T, register_as_pytree: bool) -> _T:
|
||||
"""Wraps a dataclass with metadata for which fields are pytrees.
|
||||
|
||||
This is based off flax.struct.dataclass, but instead of using field
|
||||
@@ -48,54 +48,51 @@ def dataclass(clz: _T) -> _T:
|
||||
the resulting dataclass, registered with Jax
|
||||
"""
|
||||
data_clz = dataclasses.dataclass(frozen=True)(clz)
|
||||
meta_fields, data_fields = [], []
|
||||
for field in dataclasses.fields(data_clz):
|
||||
if _jax_in_args(field.type):
|
||||
data_fields.append(field)
|
||||
else:
|
||||
meta_fields.append(field)
|
||||
data_clz.replace = dataclasses.replace
|
||||
|
||||
def replace(self, **updates):
|
||||
"""Returns a new object replacing the specified fields with new values."""
|
||||
return dataclasses.replace(self, **updates)
|
||||
|
||||
data_clz.replace = replace
|
||||
|
||||
def iterate_clz_with_keys(x):
|
||||
def to_meta(field, obj):
|
||||
val = getattr(obj, field.name)
|
||||
# numpy arrays are not hashable so return raw bytes instead
|
||||
if isinstance(val, np.ndarray):
|
||||
return (val.tobytes(), val.dtype, val.shape)
|
||||
if register_as_pytree:
|
||||
meta_fields, data_fields = [], []
|
||||
for field in dataclasses.fields(data_clz):
|
||||
if _jax_in_args(field.type):
|
||||
data_fields.append(field)
|
||||
else:
|
||||
return val
|
||||
meta_fields.append(field)
|
||||
|
||||
def to_data(field, obj):
|
||||
return (jax.tree_util.GetAttrKey(field.name), getattr(obj, field.name))
|
||||
def iterate_clz_with_keys(x):
|
||||
def to_meta(field, obj):
|
||||
val = getattr(obj, field.name)
|
||||
# numpy arrays are not hashable so return raw bytes instead
|
||||
if isinstance(val, np.ndarray):
|
||||
return (val.tobytes(), val.dtype, val.shape)
|
||||
else:
|
||||
return val
|
||||
|
||||
data = tuple(to_data(f, x) for f in data_fields)
|
||||
meta = tuple(to_meta(f, x) for f in meta_fields)
|
||||
return data, meta
|
||||
def to_data(field, obj):
|
||||
return (jax.tree_util.GetAttrKey(field.name), getattr(obj, field.name))
|
||||
|
||||
def clz_from_iterable(meta, data):
|
||||
data = tuple(to_data(f, x) for f in data_fields)
|
||||
meta = tuple(to_meta(f, x) for f in meta_fields)
|
||||
return data, meta
|
||||
|
||||
def from_meta(field, meta):
|
||||
if field.type is np.ndarray:
|
||||
arr = np.frombuffer(meta[0], dtype=meta[1]).reshape(meta[2])
|
||||
return (field.name, arr)
|
||||
else:
|
||||
return (field.name, meta)
|
||||
def clz_from_iterable(meta, data):
|
||||
|
||||
from_data = lambda field, meta: (field.name, meta)
|
||||
def from_meta(field, meta):
|
||||
if field.type is np.ndarray:
|
||||
arr = np.frombuffer(meta[0], dtype=meta[1]).reshape(meta[2])
|
||||
return (field.name, arr)
|
||||
else:
|
||||
return (field.name, meta)
|
||||
|
||||
meta_args = tuple(from_meta(f, m) for f, m in zip(meta_fields, meta))
|
||||
data_args = tuple(from_data(f, m) for f, m in zip(data_fields, data))
|
||||
from_data = lambda field, meta: (field.name, meta)
|
||||
|
||||
return data_clz(**dict(meta_args + data_args))
|
||||
meta_args = tuple(from_meta(f, m) for f, m in zip(meta_fields, meta))
|
||||
data_args = tuple(from_data(f, m) for f, m in zip(data_fields, data))
|
||||
|
||||
jax.tree_util.register_pytree_with_keys(
|
||||
data_clz, iterate_clz_with_keys, clz_from_iterable
|
||||
)
|
||||
return data_clz(**dict(meta_args + data_args))
|
||||
|
||||
jax.tree_util.register_pytree_with_keys(
|
||||
data_clz, iterate_clz_with_keys, clz_from_iterable
|
||||
)
|
||||
|
||||
return data_clz
|
||||
|
||||
@@ -109,8 +106,9 @@ class PyTreeNode:
|
||||
This base class additionally avoids type checking errors when using PyType.
|
||||
"""
|
||||
|
||||
def __init_subclass__(cls):
|
||||
dataclass(cls)
|
||||
def __init_subclass__(cls, register_as_pytree: bool = True, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
dataclass(cls, register_as_pytree=register_as_pytree)
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
# stub for pytype
|
||||
|
||||
Reference in New Issue
Block a user