Merge pull request #3299 from tkelestemur:tarik/mjx-warp-nested-dataclass-methods
PiperOrigin-RevId: 928341876 Change-Id: I0c89bcbbd95a14fadbf33d7785bf212c2c074c99
This commit is contained in:
@@ -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}
|
||||
''')
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user