diff --git a/mjx/mujoco/mjx/_src/dataclasses.py b/mjx/mujoco/mjx/_src/dataclasses.py index e82d6047..fad9ce7c 100644 --- a/mjx/mujoco/mjx/_src/dataclasses.py +++ b/mjx/mujoco/mjx/_src/dataclasses.py @@ -16,15 +16,42 @@ import copy import dataclasses +import hashlib import typing from typing import Dict, Optional, Sequence, Tuple, TypeVar, Union -import warnings import jax import numpy as np _T = TypeVar('_T') +class _NumPyArrayHashWrapper: + """A wrapper for NumPy arrays to make them hashable based on content. + + This class is used to allow NumPy arrays to be part of the metadata in a Jax + PyTree registration, as metadata must be hashable. The hash is based on the + array's content, dtype, and shape. + """ + __slots__ = ('_hash_key', 'array') + + def __init__(self, arr: np.ndarray): + if arr.size == 0: + h = hashlib.sha256(b'').hexdigest() + else: + contiguous = np.ascontiguousarray(arr) + h = hashlib.sha256(contiguous.data.cast('B')).hexdigest() + self._hash_key = (h, arr.dtype, arr.shape) + self.array = arr + + def __hash__(self): + return hash(self._hash_key) + + def __eq__(self, other): + if not isinstance(other, _NumPyArrayHashWrapper): + return NotImplemented + return self._hash_key == other._hash_key + + def _jax_in_args(typ) -> bool: if typ is jax.Array: return True @@ -63,17 +90,15 @@ def dataclass(clz: _T, register_as_pytree: bool) -> _T: def to_meta(field, obj): val = getattr(obj, field.name) if isinstance(val, np.ndarray): - # numpy arrays are not hashable so return raw bytes instead - return (val.tobytes(), val.dtype, val.shape) + return _NumPyArrayHashWrapper(val) if typing.get_origin(field.type) == tuple: - # variadic tuples of numpy arrays type_args = typing.get_args(field.type) if ( len(type_args) == 2 and type_args[0] == np.ndarray and type_args[1] == ... ): - return tuple((v.tobytes(), v.dtype, v.shape) for v in val) + return tuple(_NumPyArrayHashWrapper(v) for v in val) return val def to_data(field, obj): @@ -87,8 +112,7 @@ def dataclass(clz: _T, register_as_pytree: bool) -> _T: 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) + return (field.name, meta.array) if typing.get_origin(field.type) == tuple: type_args = typing.get_args(field.type) if ( @@ -98,9 +122,7 @@ def dataclass(clz: _T, register_as_pytree: bool) -> _T: ): return ( field.name, - tuple( - np.frombuffer(m[0], dtype=m[1]).reshape(m[2]) for m in meta - ), + tuple(m.array for m in meta), ) return (field.name, meta) diff --git a/mjx/mujoco/mjx/_src/dataclasses_test.py b/mjx/mujoco/mjx/_src/dataclasses_test.py index b5727941..a3cd4350 100644 --- a/mjx/mujoco/mjx/_src/dataclasses_test.py +++ b/mjx/mujoco/mjx/_src/dataclasses_test.py @@ -30,6 +30,12 @@ class Obj(dataclasses.PyTreeNode): f: tuple[jax.Array, ...] +class LargeArrayNode(dataclasses.PyTreeNode): + array_a: np.ndarray + array_b: np.ndarray + array_c: np.ndarray + + class DataclassesTest(absltest.TestCase): def test_pytree_structure(self): @@ -64,6 +70,48 @@ class DataclassesTest(absltest.TestCase): # ensure hashable meta hash(meta) + def test_metadata_equality(self): + array_size = 1_000 + data = np.random.rand(array_size, 3).astype(np.float32) + tex = np.random.randint(0, 255, (1024, 1024, 3), dtype=np.uint8) + obj1 = LargeArrayNode( + array_a=data.copy(), + array_b=data.copy(), + array_c=tex.copy(), + ) + obj2 = LargeArrayNode( + array_a=data.copy(), + array_b=data.copy(), + array_c=tex.copy(), + ) + + _, meta1 = jax.tree_util.tree_flatten(obj1) + _, meta2 = jax.tree_util.tree_flatten(obj2) + + self.assertEqual( + meta1, + meta2, + 'Two objects with identical numpy data should produce equal pytree' + ' metadata (same trace cache key).', + ) + + def test_metadata_does_not_copy_arrays(self): + obj = LargeArrayNode( + array_a=np.random.rand(100, 3).astype(np.float32), + array_b=np.random.rand(100, 3).astype(np.float32), + array_c=np.random.randint(0, 255, (16, 16, 3), dtype=np.uint8), + ) + + leaves, meta = jax.tree_util.tree_flatten(obj) + reconstructed = meta.unflatten(leaves) + + self.assertIs( + reconstructed.array_a, + obj.array_a, + 'Metadata should reference the original array.', + ) + self.assertIs(reconstructed.array_c, obj.array_c) + if __name__ == '__main__': absltest.main()