Merge pull request #3299 from tkelestemur:tarik/mjx-warp-nested-dataclass-methods

PiperOrigin-RevId: 928341876
Change-Id: I0c89bcbbd95a14fadbf33d7785bf212c2c074c99
This commit is contained in:
Copybara-Service
2026-06-07 22:50:49 -07:00
3 changed files with 68 additions and 1 deletions
+17 -1
View File
@@ -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}
''')
+13
View File
@@ -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)
+38
View File
@@ -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()