From f97b9f6e3b35cbd2b6bc66a0861d4ef842ae09de Mon Sep 17 00:00:00 2001 From: Tarik Kelestemur Date: Fri, 29 May 2026 21:39:14 -0400 Subject: [PATCH 1/4] TileSet MJX codegen fix --- mjx/mujoco/mjx/codegen/generate_warp_types.py | 54 ++++++++++++++++++- mjx/mujoco/mjx/warp/types.py | 18 +++++++ 2 files changed, 71 insertions(+), 1 deletion(-) diff --git a/mjx/mujoco/mjx/codegen/generate_warp_types.py b/mjx/mujoco/mjx/codegen/generate_warp_types.py index 466ee281..a95db234 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_types.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_types.py @@ -17,7 +17,9 @@ import ast import dataclasses import enum +import inspect import logging +import textwrap import typing from typing import Any, Callable, Dict, List, Optional, Set @@ -281,6 +283,12 @@ else: Callback = None PyTreeNode = mjx_dataclasses.PyTreeNode + + +def _as_numpy_array(value): + if hasattr(value, 'numpy'): + value = value.numpy() + return np.asarray(value) ''' target_fpath.write_text(header) @@ -298,6 +306,47 @@ _FLATTEN_UNFLATTEN = """ """ +class _NumpyMethodAdapter(ast.NodeTransformer): + """Adapts copied Warp array methods to MJX numpy-backed fields.""" + + def visit_Call(self, node: ast.Call) -> ast.AST: # pylint: disable=invalid-name + self.generic_visit(node) + if ( + isinstance(node.func, ast.Attribute) + and node.func.attr == 'numpy' + and not node.args + and not node.keywords + ): + return ast.copy_location( + ast.Call( + func=ast.Name(id='_as_numpy_array', ctx=ast.Load()), + args=[node.func.value], + keywords=[], + ), + node, + ) + return node + + +def _get_explicit_method_nodes(cls: Any) -> List[ast.FunctionDef]: + """Returns explicit methods from a source class, adapted for MJX fields.""" + source = textwrap.dedent(inspect.getsource(cls)) + tree = ast.parse(source) + class_def = next(node for node in tree.body if isinstance(node, ast.ClassDef)) + + methods: List[ast.FunctionDef] = [] + for node in class_def.body: + if not isinstance(node, ast.FunctionDef): + continue + if node.name in {'tree_flatten', 'tree_unflatten'}: + continue + methods.append( + typing.cast(ast.FunctionDef, _NumpyMethodAdapter().visit(node)) + ) + + return methods + + def write_nested_dataclass(target_fpath: epath.Path, cls: Any): new_class_body = _build_new_class_body_ast( set(cls.__annotations__.keys()), @@ -305,7 +354,10 @@ def write_nested_dataclass(target_fpath: epath.Path, cls: Any): dict(cls.__annotations__), add_docstring=False, ) - cls_str = '\n'.join([' ' + ast.unparse(node) for node in new_class_body]) + new_class_body.extend(_get_explicit_method_nodes(cls)) + cls_str = '\n'.join( + textwrap.indent(ast.unparse(node), ' ') for node in new_class_body + ) cls_str = cls_str.replace('jax.Array', 'np.ndarray') with target_fpath.open('a') as f: f.write(f''' diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 8f952897..fd4c6da1 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -47,6 +47,12 @@ else: PyTreeNode = mjx_dataclasses.PyTreeNode +def _as_numpy_array(value): + if hasattr(value, 'numpy'): + value = value.numpy() + return np.asarray(value) + + @dataclasses.dataclass(frozen=True) @tree_util.register_pytree_node_class class TileSet: @@ -62,6 +68,18 @@ class TileSet: adr: np.ndarray size: int + 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(_as_numpy_array(self.adr)), + np.asarray(_as_numpy_array(other.adr)), + ) + + def __hash__(self) -> int: + adr = np.asarray(_as_numpy_array(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) From c80d8fb35dbef8d21aab71327aef2731d00a8c07 Mon Sep 17 00:00:00 2001 From: Tarik Kelestemur Date: Mon, 1 Jun 2026 10:45:22 -0400 Subject: [PATCH 2/4] Stabilize MJX Warp types codegen --- mjx/mujoco/mjx/codegen/generate_warp_types.py | 157 ++++++++++++++++-- mjx/mujoco/mjx/warp/types.py | 7 +- mjx/mujoco/mjx/warp/types_test.py | 38 +++++ 3 files changed, 183 insertions(+), 19 deletions(-) create mode 100644 mjx/mujoco/mjx/warp/types_test.py diff --git a/mjx/mujoco/mjx/codegen/generate_warp_types.py b/mjx/mujoco/mjx/codegen/generate_warp_types.py index a95db234..7719291c 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_types.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_types.py @@ -26,14 +26,14 @@ from typing import Any, Callable, Dict, List, Optional, Set from absl import app from absl import flags from etils import epath +import numpy as np +import warp as wp + import mujoco from mujoco.mjx.codegen import file import mujoco.mjx.third_party.mujoco_warp as mjwarp -import numpy as np -import warp as wp from mujoco.mjx.third_party.warp._src.jax_experimental import ffi - _MJX_WARP_TYPES_OUT_FPATH = flags.DEFINE_string( 'mjx_warp_types_out_path', 'third_party/py/mujoco/mjx/warp/types.py', @@ -47,6 +47,7 @@ _MJX_TYPES_PATH = flags.DEFINE_string( ) _DATA_SHAPE_PROPERTY_FIELD = 'cacc' +_DOCSTRING_LINE_LENGTH = 80 _DUMMY_XML = """ @@ -63,6 +64,29 @@ _DUMMY_XML = """ """ +def _format_docstring(docstring: str) -> str: + """Wraps docstring body lines to match checked-in generated files.""" + formatted_lines = [] + for line in docstring.splitlines(): + if not line.strip(): + formatted_lines.append(line) + continue + + indent = line[: len(line) - len(line.lstrip())] + continuation_indent = indent + (' ' if indent else '') + formatted_lines.append( + textwrap.fill( + line, + width=_DOCSTRING_LINE_LENGTH, + subsequent_indent=continuation_indent, + break_long_words=False, + break_on_hyphens=False, + ) + ) + + return '\n'.join(formatted_lines) + + def _to_py_string(value, indent=0): """Converts a dictionary/set/tuple/type to a Python code string.""" indent_str = ' ' * indent @@ -111,8 +135,10 @@ def _get_target_annotation_node( if annotation == np.ndarray: return _ast_parse_type('np.ndarray') - if (isinstance(annotation, wp.array) or - type(annotation).__name__ == '_ArrayAnnotation'): + if ( + isinstance(annotation, wp.array) + or type(annotation).__name__ == '_ArrayAnnotation' + ): return _ast_parse_type('jax.Array') if annotation in (int, float, bool): @@ -168,9 +194,10 @@ def _get_annotations_recursive( flattened = {} for key, annotation in annotations.items(): full_key = f'{prefix}{key}' - if hasattr( - annotation, '__annotations__' - ) and 'mujoco_warp' in annotation.__module__: + if ( + hasattr(annotation, '__annotations__') + and 'mujoco_warp' in annotation.__module__ + ): nested = _get_annotations_recursive( dict(annotation.__annotations__), prefix=f'{full_key}__' ) @@ -287,7 +314,7 @@ PyTreeNode = mjx_dataclasses.PyTreeNode def _as_numpy_array(value): if hasattr(value, 'numpy'): - value = value.numpy() + return value.numpy() return np.asarray(value) ''' target_fpath.write_text(header) @@ -309,6 +336,23 @@ _FLATTEN_UNFLATTEN = """ class _NumpyMethodAdapter(ast.NodeTransformer): """Adapts copied Warp array methods to MJX numpy-backed fields.""" + def _is_asarray_call(self, node: ast.Call) -> bool: + return ( + isinstance(node.func, ast.Attribute) + and isinstance(node.func.value, ast.Name) + and node.func.value.id == 'np' + and node.func.attr == 'asarray' + and len(node.args) == 1 + and not node.keywords + ) + + def _is_as_numpy_array_call(self, node: ast.AST) -> bool: + return ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Name) + and node.func.id == '_as_numpy_array' + ) + def visit_Call(self, node: ast.Call) -> ast.AST: # pylint: disable=invalid-name self.generic_visit(node) if ( @@ -325,6 +369,10 @@ class _NumpyMethodAdapter(ast.NodeTransformer): ), node, ) + if self._is_asarray_call(node) and self._is_as_numpy_array_call( + node.args[0] + ): + return ast.copy_location(node.args[0], node) return node @@ -359,17 +407,84 @@ def write_nested_dataclass(target_fpath: epath.Path, cls: Any): textwrap.indent(ast.unparse(node), ' ') for node in new_class_body ) cls_str = cls_str.replace('jax.Array', 'np.ndarray') + docstring = _format_docstring(cls.__doc__ or '') with target_fpath.open('a') as f: f.write(f''' @dataclasses.dataclass(frozen=True) @tree_util.register_pytree_node_class class {cls.__name__}: - """{cls.__doc__}""" + """{docstring}""" {cls_str} {_FLATTEN_UNFLATTEN} ''') +def _get_class_name(line: str) -> str | None: + """Returns the class name from a generated class definition line.""" + if not line.startswith('class '): + return None + return line.removeprefix('class ').split('(', maxsplit=1)[0] + + +def _get_compact_derived_docstring_classes( + target_fpath: epath.Path, +) -> Set[str]: + """Finds generated classes with no blank line after their docstring.""" + compact_classes = set() + if not target_fpath.exists(): + return compact_classes + + lines = target_fpath.read_text().splitlines() + for i, line in enumerate(lines[:-2]): + cls_name = _get_class_name(line) + if cls_name is None: + continue + if ( + '(PyTreeNode)' in line + and lines[i + 1].startswith(' """Derived fields from ') + and lines[i + 1].rstrip().endswith('."""') + and lines[i + 2].startswith(' ') + and ':' in lines[i + 2] + ): + compact_classes.add(cls_name) + + return compact_classes + + +def _restore_compact_derived_docstring_spacing( + target_fpath: epath.Path, + compact_classes: Set[str], +): + """Restores compact generated PyTreeNode field class spacing.""" + if not compact_classes: + return + + lines = target_fpath.read_text().splitlines(keepends=True) + result = [] + i = 0 + current_class = None + while i < len(lines): + cls_name = _get_class_name(lines[i]) + if cls_name is not None: + current_class = cls_name + + result.append(lines[i]) + if ( + current_class in compact_classes + and lines[i].startswith(' """Derived fields from ') + and lines[i].rstrip().endswith('."""') + and i + 2 < len(lines) + and lines[i + 1].strip() == '' + and lines[i + 2].startswith(' ') + and ':' in lines[i + 2] + ): + i += 2 + continue + i += 1 + + target_fpath.write_text(''.join(result)) + + def _get_meta_fields(cls_name: str) -> Set[str]: """Returns the set of fields that should be meta-fields in the pytree.""" m = mujoco.MjModel.from_xml_string(_DUMMY_XML) @@ -482,8 +597,10 @@ def _get_fields_with_cond( if f.type in (int, float, bool) and add_static: s.add(prefix + f.name) continue - if not (isinstance(f.type, wp.array) or - type(f.type).__name__ == '_ArrayAnnotation'): + if not ( + isinstance(f.type, wp.array) + or type(f.type).__name__ == '_ArrayAnnotation' + ): continue if cond_fn(attr): s.add(prefix + f.name) @@ -522,8 +639,10 @@ batching.register_vmappable(DataWarp, int, int, _to_elt, _from_elt, None) def _is_ffi_compatible(wp_type: Any) -> bool: """Returns True if the type is an array, scalar, or variadic tuple.""" - if (isinstance(wp_type, wp.array) or - type(wp_type).__name__ == '_ArrayAnnotation'): + if ( + isinstance(wp_type, wp.array) + or type(wp_type).__name__ == '_ArrayAnnotation' + ): return True if wp_type in wp._src.types.value_types: return True @@ -600,6 +719,9 @@ def main(argv): base_path = file.get_base_path() target_fpath = base_path / _MJX_WARP_TYPES_OUT_FPATH.value mjx_types_fpath = base_path / _MJX_TYPES_PATH.value + compact_derived_docstring_classes = _get_compact_derived_docstring_classes( + target_fpath + ) write_header(target_fpath) # TODO(btaba): consider automated grabbing of nested dataclasses from mjwarp. @@ -608,7 +730,9 @@ def main(argv): write_core_cls('Statistic', target_fpath, mjx_types_fpath, set_diff=False) write_core_cls( - 'Option', target_fpath, mjx_types_fpath, + 'Option', + target_fpath, + mjx_types_fpath, extra_annotations={'graph_mode': ffi.GraphMode}, ) write_core_cls('Model', target_fpath, mjx_types_fpath) @@ -619,6 +743,9 @@ def main(argv): file.write_license(target_fpath) file.format_file(target_fpath) + _restore_compact_derived_docstring_spacing( + target_fpath, compact_derived_docstring_classes + ) if __name__ == '__main__': diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index fd4c6da1..8a575f20 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -49,7 +49,7 @@ PyTreeNode = mjx_dataclasses.PyTreeNode def _as_numpy_array(value): if hasattr(value, 'numpy'): - value = value.numpy() + return value.numpy() return np.asarray(value) @@ -72,12 +72,11 @@ class TileSet: if self.__class__ is not other.__class__: return NotImplemented return self.size == other.size and np.array_equal( - np.asarray(_as_numpy_array(self.adr)), - np.asarray(_as_numpy_array(other.adr)), + _as_numpy_array(self.adr), _as_numpy_array(other.adr) ) def __hash__(self) -> int: - adr = np.asarray(_as_numpy_array(self.adr)) + adr = _as_numpy_array(self.adr) return hash((self.size, adr.dtype.str, adr.shape, adr.tobytes())) def tree_flatten(self): 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() From b4dd8e2d9cf21bca0c4f1513ca90ff1be95fbae9 Mon Sep 17 00:00:00 2001 From: Tarik Kelestemur Date: Wed, 3 Jun 2026 12:31:23 -0400 Subject: [PATCH 3/4] Emit TileSet methods in MJX Warp codegen --- mjx/mujoco/mjx/codegen/generate_warp_types.py | 41 ++++++++++++++++++- 1 file changed, 40 insertions(+), 1 deletion(-) diff --git a/mjx/mujoco/mjx/codegen/generate_warp_types.py b/mjx/mujoco/mjx/codegen/generate_warp_types.py index 466ee281..4e13054b 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_types.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_types.py @@ -297,6 +297,30 @@ _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())) +""", +} + +_NESTED_DATACLASS_MANUAL_METHOD_NOTES = { + 'TileSet': ( + ' # Manually kept in this generated shim until TileSet method ' + 'generation is\n' + ' # needed more broadly. Keep this in sync with ' + 'mujoco_warp._src.types.TileSet.\n' + ), +} + def write_nested_dataclass(target_fpath: epath.Path, cls: Any): new_class_body = _build_new_class_body_ast( @@ -307,6 +331,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,10 +339,23 @@ def write_nested_dataclass(target_fpath: epath.Path, cls: Any): class {cls.__name__}: """{cls.__doc__}""" {cls_str} -{_FLATTEN_UNFLATTEN} +{manual_methods}{_FLATTEN_UNFLATTEN} ''') +def _write_manual_method_notes(target_fpath: epath.Path): + """Restores method comments stripped by AST-based file rewrites.""" + src = target_fpath.read_text() + for cls_name, note in _NESTED_DATACLASS_MANUAL_METHOD_NOTES.items(): + class_start = src.index(f'class {cls_name}:') + method_start = src.index( + ' def __eq__(self, other) -> bool:\n', class_start + ) + if note not in src[class_start:method_start]: + src = src[:method_start] + note + src[method_start:] + target_fpath.write_text(src) + + def _get_meta_fields(cls_name: str) -> Set[str]: """Returns the set of fields that should be meta-fields in the pytree.""" m = mujoco.MjModel.from_xml_string(_DUMMY_XML) @@ -565,6 +603,7 @@ def main(argv): write_ndim_annotations(target_fpath) write_nworld_leading_dim(target_fpath) + _write_manual_method_notes(target_fpath) file.write_license(target_fpath) file.format_file(target_fpath) From 7078f38163b9828699e3af4fc345be4e5c795525 Mon Sep 17 00:00:00 2001 From: btaba Date: Thu, 4 Jun 2026 12:28:38 -0700 Subject: [PATCH 4/4] Update generate_warp_types.py --- mjx/mujoco/mjx/codegen/generate_warp_types.py | 23 ------------------- 1 file changed, 23 deletions(-) diff --git a/mjx/mujoco/mjx/codegen/generate_warp_types.py b/mjx/mujoco/mjx/codegen/generate_warp_types.py index 4e13054b..98b57215 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_types.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_types.py @@ -312,15 +312,6 @@ _NESTED_DATACLASS_MANUAL_METHODS = { """, } -_NESTED_DATACLASS_MANUAL_METHOD_NOTES = { - 'TileSet': ( - ' # Manually kept in this generated shim until TileSet method ' - 'generation is\n' - ' # needed more broadly. Keep this in sync with ' - 'mujoco_warp._src.types.TileSet.\n' - ), -} - def write_nested_dataclass(target_fpath: epath.Path, cls: Any): new_class_body = _build_new_class_body_ast( @@ -343,19 +334,6 @@ class {cls.__name__}: ''') -def _write_manual_method_notes(target_fpath: epath.Path): - """Restores method comments stripped by AST-based file rewrites.""" - src = target_fpath.read_text() - for cls_name, note in _NESTED_DATACLASS_MANUAL_METHOD_NOTES.items(): - class_start = src.index(f'class {cls_name}:') - method_start = src.index( - ' def __eq__(self, other) -> bool:\n', class_start - ) - if note not in src[class_start:method_start]: - src = src[:method_start] + note + src[method_start:] - target_fpath.write_text(src) - - def _get_meta_fields(cls_name: str) -> Set[str]: """Returns the set of fields that should be meta-fields in the pytree.""" m = mujoco.MjModel.from_xml_string(_DUMMY_XML) @@ -603,7 +581,6 @@ def main(argv): write_ndim_annotations(target_fpath) write_nworld_leading_dim(target_fpath) - _write_manual_method_notes(target_fpath) file.write_license(target_fpath) file.format_file(target_fpath)