From 982330d15a84bf17af7248b5a6c830a73b43eb24 Mon Sep 17 00:00:00 2001 From: Google DeepMind Date: Tue, 18 Feb 2025 08:12:02 -0800 Subject: [PATCH] #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 --- mjx/mujoco/mjx/_src/dataclasses.py | 80 +++++++++++++++--------------- 1 file changed, 39 insertions(+), 41 deletions(-) diff --git a/mjx/mujoco/mjx/_src/dataclasses.py b/mjx/mujoco/mjx/_src/dataclasses.py index 0db4a141..85a18e60 100644 --- a/mjx/mujoco/mjx/_src/dataclasses.py +++ b/mjx/mujoco/mjx/_src/dataclasses.py @@ -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