#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:
Google DeepMind
2025-02-18 08:12:02 -08:00
committed by Copybara-Service
parent 36142cddee
commit 982330d15a
+39 -41
View File
@@ -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