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 app
|
||||||
from absl import flags
|
from absl import flags
|
||||||
from etils import epath
|
from etils import epath
|
||||||
|
import numpy as np
|
||||||
|
import warp as wp
|
||||||
|
|
||||||
import mujoco
|
import mujoco
|
||||||
from mujoco.mjx.codegen import file
|
from mujoco.mjx.codegen import file
|
||||||
import mujoco.mjx.third_party.mujoco_warp as mjwarp
|
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
|
from mujoco.mjx.third_party.warp._src.jax_experimental import ffi
|
||||||
|
|
||||||
|
|
||||||
_MJX_WARP_TYPES_OUT_FPATH = flags.DEFINE_string(
|
_MJX_WARP_TYPES_OUT_FPATH = flags.DEFINE_string(
|
||||||
'mjx_warp_types_out_path',
|
'mjx_warp_types_out_path',
|
||||||
'third_party/py/mujoco/mjx/warp/types.py',
|
'third_party/py/mujoco/mjx/warp/types.py',
|
||||||
@@ -47,6 +47,7 @@ _MJX_TYPES_PATH = flags.DEFINE_string(
|
|||||||
)
|
)
|
||||||
|
|
||||||
_DATA_SHAPE_PROPERTY_FIELD = 'cacc'
|
_DATA_SHAPE_PROPERTY_FIELD = 'cacc'
|
||||||
|
_DOCSTRING_LINE_LENGTH = 80
|
||||||
_DUMMY_XML = """
|
_DUMMY_XML = """
|
||||||
<mujoco>
|
<mujoco>
|
||||||
<worldbody>
|
<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):
|
def _to_py_string(value, indent=0):
|
||||||
"""Converts a dictionary/set/tuple/type to a Python code string."""
|
"""Converts a dictionary/set/tuple/type to a Python code string."""
|
||||||
indent_str = ' ' * indent
|
indent_str = ' ' * indent
|
||||||
@@ -111,8 +135,10 @@ def _get_target_annotation_node(
|
|||||||
if annotation == np.ndarray:
|
if annotation == np.ndarray:
|
||||||
return _ast_parse_type('np.ndarray')
|
return _ast_parse_type('np.ndarray')
|
||||||
|
|
||||||
if (isinstance(annotation, wp.array) or
|
if (
|
||||||
type(annotation).__name__ == '_ArrayAnnotation'):
|
isinstance(annotation, wp.array)
|
||||||
|
or type(annotation).__name__ == '_ArrayAnnotation'
|
||||||
|
):
|
||||||
return _ast_parse_type('jax.Array')
|
return _ast_parse_type('jax.Array')
|
||||||
|
|
||||||
if annotation in (int, float, bool):
|
if annotation in (int, float, bool):
|
||||||
@@ -168,9 +194,10 @@ def _get_annotations_recursive(
|
|||||||
flattened = {}
|
flattened = {}
|
||||||
for key, annotation in annotations.items():
|
for key, annotation in annotations.items():
|
||||||
full_key = f'{prefix}{key}'
|
full_key = f'{prefix}{key}'
|
||||||
if hasattr(
|
if (
|
||||||
annotation, '__annotations__'
|
hasattr(annotation, '__annotations__')
|
||||||
) and 'mujoco_warp' in annotation.__module__:
|
and 'mujoco_warp' in annotation.__module__
|
||||||
|
):
|
||||||
nested = _get_annotations_recursive(
|
nested = _get_annotations_recursive(
|
||||||
dict(annotation.__annotations__), prefix=f'{full_key}__'
|
dict(annotation.__annotations__), prefix=f'{full_key}__'
|
||||||
)
|
)
|
||||||
@@ -287,7 +314,7 @@ PyTreeNode = mjx_dataclasses.PyTreeNode
|
|||||||
|
|
||||||
def _as_numpy_array(value):
|
def _as_numpy_array(value):
|
||||||
if hasattr(value, 'numpy'):
|
if hasattr(value, 'numpy'):
|
||||||
value = value.numpy()
|
return value.numpy()
|
||||||
return np.asarray(value)
|
return np.asarray(value)
|
||||||
'''
|
'''
|
||||||
target_fpath.write_text(header)
|
target_fpath.write_text(header)
|
||||||
@@ -309,6 +336,23 @@ _FLATTEN_UNFLATTEN = """
|
|||||||
class _NumpyMethodAdapter(ast.NodeTransformer):
|
class _NumpyMethodAdapter(ast.NodeTransformer):
|
||||||
"""Adapts copied Warp array methods to MJX numpy-backed fields."""
|
"""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
|
def visit_Call(self, node: ast.Call) -> ast.AST: # pylint: disable=invalid-name
|
||||||
self.generic_visit(node)
|
self.generic_visit(node)
|
||||||
if (
|
if (
|
||||||
@@ -325,6 +369,10 @@ class _NumpyMethodAdapter(ast.NodeTransformer):
|
|||||||
),
|
),
|
||||||
node,
|
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
|
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
|
textwrap.indent(ast.unparse(node), ' ') for node in new_class_body
|
||||||
)
|
)
|
||||||
cls_str = cls_str.replace('jax.Array', 'np.ndarray')
|
cls_str = cls_str.replace('jax.Array', 'np.ndarray')
|
||||||
|
docstring = _format_docstring(cls.__doc__ or '')
|
||||||
with target_fpath.open('a') as f:
|
with target_fpath.open('a') as f:
|
||||||
f.write(f'''
|
f.write(f'''
|
||||||
@dataclasses.dataclass(frozen=True)
|
@dataclasses.dataclass(frozen=True)
|
||||||
@tree_util.register_pytree_node_class
|
@tree_util.register_pytree_node_class
|
||||||
class {cls.__name__}:
|
class {cls.__name__}:
|
||||||
"""{cls.__doc__}"""
|
"""{docstring}"""
|
||||||
{cls_str}
|
{cls_str}
|
||||||
{_FLATTEN_UNFLATTEN}
|
{_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]:
|
def _get_meta_fields(cls_name: str) -> Set[str]:
|
||||||
"""Returns the set of fields that should be meta-fields in the pytree."""
|
"""Returns the set of fields that should be meta-fields in the pytree."""
|
||||||
m = mujoco.MjModel.from_xml_string(_DUMMY_XML)
|
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:
|
if f.type in (int, float, bool) and add_static:
|
||||||
s.add(prefix + f.name)
|
s.add(prefix + f.name)
|
||||||
continue
|
continue
|
||||||
if not (isinstance(f.type, wp.array) or
|
if not (
|
||||||
type(f.type).__name__ == '_ArrayAnnotation'):
|
isinstance(f.type, wp.array)
|
||||||
|
or type(f.type).__name__ == '_ArrayAnnotation'
|
||||||
|
):
|
||||||
continue
|
continue
|
||||||
if cond_fn(attr):
|
if cond_fn(attr):
|
||||||
s.add(prefix + f.name)
|
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:
|
def _is_ffi_compatible(wp_type: Any) -> bool:
|
||||||
"""Returns True if the type is an array, scalar, or variadic tuple."""
|
"""Returns True if the type is an array, scalar, or variadic tuple."""
|
||||||
if (isinstance(wp_type, wp.array) or
|
if (
|
||||||
type(wp_type).__name__ == '_ArrayAnnotation'):
|
isinstance(wp_type, wp.array)
|
||||||
|
or type(wp_type).__name__ == '_ArrayAnnotation'
|
||||||
|
):
|
||||||
return True
|
return True
|
||||||
if wp_type in wp._src.types.value_types:
|
if wp_type in wp._src.types.value_types:
|
||||||
return True
|
return True
|
||||||
@@ -600,6 +719,9 @@ def main(argv):
|
|||||||
base_path = file.get_base_path()
|
base_path = file.get_base_path()
|
||||||
target_fpath = base_path / _MJX_WARP_TYPES_OUT_FPATH.value
|
target_fpath = base_path / _MJX_WARP_TYPES_OUT_FPATH.value
|
||||||
mjx_types_fpath = base_path / _MJX_TYPES_PATH.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)
|
write_header(target_fpath)
|
||||||
# TODO(btaba): consider automated grabbing of nested dataclasses from mjwarp.
|
# 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('Statistic', target_fpath, mjx_types_fpath, set_diff=False)
|
||||||
write_core_cls(
|
write_core_cls(
|
||||||
'Option', target_fpath, mjx_types_fpath,
|
'Option',
|
||||||
|
target_fpath,
|
||||||
|
mjx_types_fpath,
|
||||||
extra_annotations={'graph_mode': ffi.GraphMode},
|
extra_annotations={'graph_mode': ffi.GraphMode},
|
||||||
)
|
)
|
||||||
write_core_cls('Model', target_fpath, mjx_types_fpath)
|
write_core_cls('Model', target_fpath, mjx_types_fpath)
|
||||||
@@ -619,6 +743,9 @@ def main(argv):
|
|||||||
|
|
||||||
file.write_license(target_fpath)
|
file.write_license(target_fpath)
|
||||||
file.format_file(target_fpath)
|
file.format_file(target_fpath)
|
||||||
|
_restore_compact_derived_docstring_spacing(
|
||||||
|
target_fpath, compact_derived_docstring_classes
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
|
|||||||
@@ -49,7 +49,7 @@ PyTreeNode = mjx_dataclasses.PyTreeNode
|
|||||||
|
|
||||||
def _as_numpy_array(value):
|
def _as_numpy_array(value):
|
||||||
if hasattr(value, 'numpy'):
|
if hasattr(value, 'numpy'):
|
||||||
value = value.numpy()
|
return value.numpy()
|
||||||
return np.asarray(value)
|
return np.asarray(value)
|
||||||
|
|
||||||
|
|
||||||
@@ -72,12 +72,11 @@ class TileSet:
|
|||||||
if self.__class__ is not other.__class__:
|
if self.__class__ is not other.__class__:
|
||||||
return NotImplemented
|
return NotImplemented
|
||||||
return self.size == other.size and np.array_equal(
|
return self.size == other.size and np.array_equal(
|
||||||
np.asarray(_as_numpy_array(self.adr)),
|
_as_numpy_array(self.adr), _as_numpy_array(other.adr)
|
||||||
np.asarray(_as_numpy_array(other.adr)),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def __hash__(self) -> int:
|
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()))
|
return hash((self.size, adr.dtype.str, adr.shape, adr.tobytes()))
|
||||||
|
|
||||||
def tree_flatten(self):
|
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