Stabilize MJX Warp types codegen

This commit is contained in:
Tarik Kelestemur
2026-06-01 10:45:22 -04:00
parent f97b9f6e3b
commit c80d8fb35d
3 changed files with 183 additions and 19 deletions
+142 -15
View File
@@ -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 = """
<mujoco>
<worldbody>
@@ -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__':
+3 -4
View File
@@ -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):
+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()