diff --git a/mjx/mujoco/mjx/codegen/generate_warp_types.py b/mjx/mujoco/mjx/codegen/generate_warp_types.py index 466ee281..98b57215 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_types.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_types.py @@ -297,6 +297,21 @@ _FLATTEN_UNFLATTEN = """ return cls(*children) """ +_NESTED_DATACLASS_MANUAL_METHODS = { + 'TileSet': """ + def __eq__(self, other) -> bool: + if self.__class__ is not other.__class__: + return NotImplemented + return self.size == other.size and np.array_equal( + np.asarray(self.adr), np.asarray(other.adr) + ) + + def __hash__(self) -> int: + adr = np.asarray(self.adr) + return hash((self.size, adr.dtype.str, adr.shape, adr.tobytes())) +""", +} + def write_nested_dataclass(target_fpath: epath.Path, cls: Any): new_class_body = _build_new_class_body_ast( @@ -307,6 +322,7 @@ def write_nested_dataclass(target_fpath: epath.Path, cls: Any): ) cls_str = '\n'.join([' ' + ast.unparse(node) for node in new_class_body]) cls_str = cls_str.replace('jax.Array', 'np.ndarray') + manual_methods = _NESTED_DATACLASS_MANUAL_METHODS.get(cls.__name__, '') with target_fpath.open('a') as f: f.write(f''' @dataclasses.dataclass(frozen=True) @@ -314,7 +330,7 @@ def write_nested_dataclass(target_fpath: epath.Path, cls: Any): class {cls.__name__}: """{cls.__doc__}""" {cls_str} -{_FLATTEN_UNFLATTEN} +{manual_methods}{_FLATTEN_UNFLATTEN} ''') diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 7130b49b..3186b1a3 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -55,6 +55,19 @@ class TileSet: adr: np.ndarray size: int + # Manually kept in this generated shim until TileSet method generation is + # needed more broadly. Keep this in sync with mujoco_warp._src.types.TileSet. + def __eq__(self, other) -> bool: + if self.__class__ is not other.__class__: + return NotImplemented + return self.size == other.size and np.array_equal( + np.asarray(self.adr), np.asarray(other.adr) + ) + + def __hash__(self) -> int: + adr = np.asarray(self.adr) + return hash((self.size, adr.dtype.str, adr.shape, adr.tobytes())) + def tree_flatten(self): children = list((getattr(self, k) for k in self.__dataclass_fields__)) return (children, None) diff --git a/mjx/mujoco/mjx/warp/types_test.py b/mjx/mujoco/mjx/warp/types_test.py new file mode 100644 index 00000000..6d1bb9af --- /dev/null +++ b/mjx/mujoco/mjx/warp/types_test.py @@ -0,0 +1,38 @@ +# Copyright 2026 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Tests for generated MJX Warp types.""" + +from absl.testing import absltest +from absl.testing import parameterized +import numpy as np + +from mujoco.mjx.warp import types + + +class TileSetTest(parameterized.TestCase): + + @parameterized.parameters(1, 3) + def test_structural_equality_and_hash(self, count): + tile_a = types.TileSet(np.arange(count) * 6, 16) + tile_b = types.TileSet(np.arange(count) * 6, 16) + tile_c = types.TileSet(np.arange(count) * 6 + 1, 16) + + self.assertEqual(tile_a, tile_b) + self.assertEqual(hash(tile_a), hash(tile_b)) + self.assertNotEqual(tile_a, tile_c) + + +if __name__ == '__main__': + absltest.main()