Stabilize MJX Warp types codegen
This commit is contained in:
@@ -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__':
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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