Open-source MJX-Warp codegen to make external contribution easier.

PiperOrigin-RevId: 910200635
Change-Id: I0218f8e281413803500c98fc29d0589c4e90d6b1
This commit is contained in:
Baruch Tabanpour
2026-05-04 13:12:54 -07:00
committed by Copybara-Service
parent 0aeea2e4f6
commit a51a7bf062
10 changed files with 1682 additions and 2 deletions
+23
View File
@@ -0,0 +1,23 @@
# MJX Warp Codegen
Generates the MJX-Warp shim layer in `mujoco/mjx/warp/` by reading the vendored
`mujoco_warp` source in `mujoco/mjx/third_party/mujoco_warp/`.
## Setup
From the root `mjx/` directory, once you install [`uv`](https://docs.astral.sh/uv/getting-started/installation/), install the latest MuJoCo and local MJX:
```bash
uv venv .venv --default-index https://pypi.org/simple
source .venv/bin/activate
uv pip install --upgrade --force-reinstall mujoco --default-index https://pypi.org/simple --extra-index-url https://py.mujoco.org/
uv pip install -e ".[warp,dev]" --default-index https://pypi.org/simple
```
## Run codegen
From the root `mjx/` directory:
```bash
bash mujoco/mjx/codegen/update_for_mujoco_warp.sh
```
+14
View File
@@ -0,0 +1,14 @@
# 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.
# ==============================================================================
+98
View File
@@ -0,0 +1,98 @@
# 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.
# ==============================================================================
"""File tools."""
import ast
import os
import subprocess
from typing import Dict
from absl import logging
from etils import epath
LICENSE_TEXT = """
# 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.
# ==============================================================================
"""
def get_base_path() -> epath.Path:
"""Resolves the base workspace path."""
base_path = os.environ.get('BUILD_WORKSPACE_DIRECTORY')
if base_path:
return epath.Path(base_path)
# Assume this file is at <root>/mujoco/mjx/codegen/file.py.
root = epath.Path(os.path.abspath(__file__)).parents[3]
if not (root / 'mujoco' / 'mjx' / 'codegen').is_dir():
raise RuntimeError(
f'Unexpected codegen layout, resolved root: {root}'
)
return root
def format_file(target_fpath: epath.Path):
"""Formats a Python file."""
logging.info('Running pyink on: %s', target_fpath)
subprocess.run(
['pyink', str(target_fpath)],
check=True,
text=True,
capture_output=True,
)
logging.info('Running isort on: %s', target_fpath)
subprocess.run(
['isort', str(target_fpath)],
check=True,
text=True,
capture_output=True,
)
def write_license(target_fpath: epath.Path):
"""Writes license to the target file."""
src = target_fpath.read_text()
target_fpath.write_text(LICENSE_TEXT + src)
def get_cls_type_annotations(src: str) -> Dict[str, Dict[str, str]]:
"""Return classes with their field annotation strings from source code."""
ret = {}
tree = ast.parse(src)
class Visitor(ast.NodeVisitor):
def visit_ClassDef(self, node: ast.ClassDef): # pylint: disable=invalid-name
class_name = node.name
ret[class_name] = {}
for item in node.body:
if not isinstance(item, ast.AnnAssign):
continue
field_name = item.target.id # pytype: disable=attribute-error
annotation_str = ast.unparse(item.annotation).strip()
ret[class_name][field_name] = annotation_str
Visitor().visit(tree)
return ret
@@ -0,0 +1,536 @@
# 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.
# ==============================================================================
"""Creates a shim between JAX and Warp for a given function."""
import enum
import inspect
import re
from typing import Dict, List, Sequence
from absl import app
from absl import flags
from absl import logging
from etils import epath
from mujoco.mjx._src import types as mjx_types
from mujoco.mjx.codegen import file
from mujoco.mjx.codegen import trace
import jax
from mujoco.mjx.third_party import mujoco_warp # pylint: disable=unused-import
_MJWARP_FUNCTION = flags.DEFINE_string(
'mjwarp_function',
'third_party/py/mujoco_warp/_src/smooth.py:kinematics',
'Function to create the shim for.',
)
_MJWARP_TYPES = flags.DEFINE_string(
'mjwarp_types',
'third_party/py/mujoco_warp/_src/types.py',
'Path to the mjwarp types file.',
)
_MJX_WARP_OUTPUT_PATH = flags.DEFINE_string(
'mjx_warp_output_path',
'third_party/py/mujoco/mjx/warp',
'Path to the output file.',
)
_ONLY_PUBLIC_OUTPUT_FIELDS = flags.DEFINE_bool(
'only_public_output_fields',
False,
'Whether to keep only public fields in the output.',
)
_APPEND_TO_OUTPUT_FILE = flags.DEFINE_bool(
'append_to_output_file',
False,
'Whether to append to the output file.',
)
_RENDER_CONTEXT_BUFFER_NAME = '_MJX_RENDER_CONTEXT_BUFFERS'
def _clean_type(type_: str):
# check for enums
if type_ in [
name
for name, obj in inspect.getmembers(mujoco_warp, inspect.isclass)
if issubclass(obj, (enum.IntEnum, enum.IntFlag))
]:
return 'int'
types_to_prefix = (
'vec5',
'vec8',
'vec8i',
'vec10',
'vec10f',
'vec11',
'TileSet',
'BlockDim',
'vec_pluginattr',
)
m = re.match(r'array\((.*)\)', type_)
if m: # match custom mujoco_warp array annotation types
args_str = m.group(1)
args = [a.strip() for a in args_str.split(',')]
ndim, dtype = len(args) - 1, args[-1]
dims = {1: '', 2: '2d', 3: '3d', 4: '4d'}
if ndim not in dims:
raise ValueError(f'Unsupported array ndim: {ndim} for type: {dtype}')
type_ = f'wp.array{dims[ndim]}[{dtype}]'
for t in types_to_prefix:
type_ = re.sub(rf'\b{t}\b', f'mjwp_types.{t}', type_)
return type_
def _get_stage_fields(
field_usage: trace.FieldUsage,
) -> tuple[list[str], list[str]]:
"""Returns stage_in and stage_out fields after tracing.
stage_in:
* Model/ModelWarp jax.Array input fields
* Data jax.Array input fields
* Option/OptionWarp jax.Array input fields
stage_out:
* Data jax.Array output fields
Args:
field_usage: FieldUsage object
Returns:
A tuple of (stage_in, stage_out) field name lists.
"""
stage_in = []
stage_out = []
def is_jax_array(cls, field):
if cls is None:
return False
return cls.__annotations__.get(field) is jax.Array
ModelWarp = getattr(mjx_types, 'ModelWarp', None)
OptionWarp = getattr(mjx_types, 'OptionWarp', None)
# stage_in: Model/ModelWarp jax.Array input fields
for field in field_usage.model_fields:
if is_jax_array(mjx_types.Model, field):
stage_in.append(field)
elif is_jax_array(ModelWarp, field):
stage_in.append(field)
# stage_in: Option/OptionWarp jax.Array input fields
elif field.startswith('opt__'):
sub_field = field.split('opt__')[-1]
if is_jax_array(mjx_types.Option, sub_field):
stage_in.append(field)
elif is_jax_array(OptionWarp, sub_field):
stage_in.append(field)
# stage_in: Data jax.Array input fields
for field in field_usage.data_fields:
if is_jax_array(mjx_types.Data, field):
stage_in.append(field)
# stage_out: Data jax.Array output fields
for field in field_usage.data_out_fields:
if is_jax_array(mjx_types.Data, field):
stage_out.append(field)
return sorted(stage_in), sorted(stage_out)
def _top_level_imports(field_usage: trace.FieldUsage):
"""Returns top-level imports."""
imports = '''
"""DO NOT EDIT. This file is auto-generated."""
import dataclasses
import functools
from mujoco.mjx._src import types
from mujoco.mjx.warp import ffi
import mujoco.mjx.third_party.mujoco_warp as mjwarp
import warp as wp
import jax
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
'''
if field_usage.render_context_in_caller:
imports += (
"""
from mujoco.mjx.warp.render_context import """
+ _RENDER_CONTEXT_BUFFER_NAME
+ """
from mujoco.mjx.warp.render_context import RenderContextPytree
"""
)
return imports
def _global_assignments():
"""Returns global assignments."""
assignments = ''
for attr, cls in (
('_m', 'mjwarp.Model'),
('_d', 'mjwarp.Data'),
('_o', 'mjwarp.Option'),
('_s', 'mjwarp.Statistic'),
('_c', 'mjwarp.Contact'),
('_e', 'mjwarp.Constraint'),
('_cb', 'mjwp_types.Callback'),
):
assignments += (
f'{attr} = {cls}(**{{f.name: None for f in dataclasses.fields({cls}) if'
' f.init})\n'
)
return assignments
def _warp_function(
fn_name: str,
field_usage: trace.FieldUsage,
mjwarp_field_info: Dict[str, trace.FieldInfo],
mjx_warp_field_info: Dict[str, trace.FieldInfo],
):
"""Returns warp function arguments, assignments, and call."""
# create warp function.
fn_args_model, fn_assignments = [('nworld: int,', (-1, ''))], []
if field_usage.model_fields:
for f in field_usage.model_fields:
if f not in mjwarp_field_info:
raise AssertionError(f'Field {f} not found in mjwarp_field_info.')
info = mjwarp_field_info[f]
expected_type = _clean_type(info.expected_type)
fn_args_model.append((f'{f}: {expected_type},', info.param_order))
fn_assignments.append(f' _m.{f.replace('__', '.')} = {f}')
fn_args_model = sorted(fn_args_model, key=lambda x: x[1])
fn_args_model = ['# Model'] + [f[0] for f in fn_args_model]
fn_args_data = []
if field_usage.data_fields:
for f in field_usage.data_fields:
if f not in mjwarp_field_info:
raise AssertionError(f'Field {f} not found in mjwarp_field_info.')
if f == 'nworld':
continue # this gets set manually below
j = trace.FieldInfo(f, 'jax.Array', (0, ''))
is_jax_arr = mjx_warp_field_info.get(f, j).expected_type == 'jax.Array'
is_out = is_jax_arr
info = mjwarp_field_info[f]
param_order = info.param_order
expected_type = _clean_type(info.expected_type)
fn_args_data.append((f'{f}: {expected_type},', (is_out, param_order)))
fn_assignments.append(f' _d.{f.replace('__', '.')} = {f}')
fn_args_data = sorted(fn_args_data, key=lambda x: x[1])
fn_args_data = ['# Data'] + [f[0] for f in fn_args_data]
fn_assignments.append(' _d.nworld = nworld')
render_context_args = []
render_context_call_arg = ''
if field_usage.render_context_in_caller:
render_context_args = ['# Registry', 'rc_id: int,']
render_context_call_arg = ', render_context'
fn_assignments.append(
f' render_context = {_RENDER_CONTEXT_BUFFER_NAME}[(rc_id, wp.get_device().ordinal)]'
)
if fn_name == 'render':
render_context_args.append('rgb: wp.array2d[wp.uint32],')
render_context_args.append('depth: wp.array2d[wp.float32],')
fn_assignments.append(' render_context.rgb_data = rgb')
fn_assignments.append(' render_context.depth_data = depth')
else:
fn_assignments.append(' dummy.zero_()')
fn_call = f'mjwarp.{fn_name}(_m, _d{render_context_call_arg})'
fn_args_raw = fn_args_model + fn_args_data + render_context_args
# create a dummy output if there are no output fields
needs_dummy_output = not field_usage.data_out_fields
if needs_dummy_output and fn_name != 'render':
fn_args_raw.append('# Dummy output')
fn_args_raw.append('dummy: wp.array[int],')
return fn_args_raw, fn_assignments, fn_call
def _jax_shim_fn(
fn_name: str,
field_usage: trace.FieldUsage,
warp_fn_args: List[str],
mjwarp_field_info: Dict[str, trace.FieldInfo],
):
"""Generates a JAX shim for the Warp function."""
num_outputs = 0
output_dims = []
jax_args = []
tree_replace = []
in_out_argnames = []
has_side_effect = False
stage_in_fields, stage_out_fields = _get_stage_fields(field_usage)
stage_in_argnames = [f"'{f}'" for f in stage_in_fields]
stage_out_argnames = [f"'{f}'" for f in stage_out_fields]
for arg in warp_fn_args:
if 'nworld' in arg:
jax_args.append('d.qpos.shape[0]')
continue
if arg in ('rc_id', 'dummy'):
continue
if arg in ('rgb', 'depth') and fn_name == 'render':
num_outputs += 1
continue
arg_jax = arg
if mjwarp_field_info[arg].param_source == 'Data':
if arg.split('__')[0] not in mjx_types.Data.__annotations__:
arg_jax = f'_impl.{arg}'
arg_jax = 'd.' + arg_jax
elif mjwarp_field_info[arg].param_source == 'Model':
arg_jax = (
arg.replace('__', '.')
if arg.startswith('opt') or arg.startswith('stat')
else arg
)
public_field = arg.split('__')[0] in mjx_types.Model.__annotations__
if not public_field:
arg_jax = f'_impl.{arg_jax}'
if (
arg.startswith('opt')
and arg.split('__')[-1] not in mjx_types.Option.__annotations__
):
arg_jax = arg_jax.replace('opt', 'opt._impl')
arg_jax = 'm.' + arg_jax
else:
raise ValueError(
f'Unknown param source: {mjwarp_field_info[arg].param_source}'
)
if arg in field_usage.data_out_fields:
# all out fields are in_out, since JAX already allocated them
in_out_argnames.append(f"'{arg}'")
num_outputs += 1
output_dims.append(f"'{arg}': {arg_jax}.shape")
if '_impl' not in arg_jax or not _ONLY_PUBLIC_OUTPUT_FIELDS.value:
tree_replace.append(f'"{arg_jax[2:]}": out[{num_outputs - 1}]')
if arg == 'geom_dataid':
jax_args.append(f'jax.numpy.expand_dims({arg_jax}, 0)')
else:
jax_args.append(arg_jax)
if field_usage.render_context_in_caller:
jax_args.append('ctx.key')
needs_dummy_output = not field_usage.data_out_fields
if needs_dummy_output and fn_name != 'render':
num_outputs = 1
output_dims = ["'dummy': (d.qpos.shape[0],)"]
has_side_effect = True
if fn_name == 'render':
output_dims = [
"'rgb': render_ctx.rgb_data_shape",
"'depth': render_ctx.depth_data_shape",
]
tree_replace = []
render_ctx_param = (
'ctx: RenderContextPytree' if field_usage.render_context_in_caller else ''
)
fn_args = ['m: types.Model', 'd: types.Data']
if render_ctx_param:
fn_args.append(render_ctx_param)
return (
fn_args,
jax_args,
output_dims,
num_outputs,
tree_replace,
in_out_argnames,
stage_in_argnames,
stage_out_argnames,
has_side_effect,
)
def create_jax_warp_shim(
fn_name: str,
field_usage: trace.FieldUsage,
mjwarp_field_info: Dict[str, trace.FieldInfo],
mjx_warp_field_info: Dict[str, trace.FieldInfo],
out_fpath: epath.Path,
):
"""Creates a JAX-wrapped MJWarp function."""
src = ''
old_src = (
out_fpath.read_text()
if out_fpath.exists()
else ''
if out_fpath.exists()
else ''
)
# create top-level imports.
if not _APPEND_TO_OUTPUT_FILE.value:
src += _top_level_imports(field_usage) + '\n\n'
# create global assignments.
assignments = _global_assignments()
already_in_src = re.sub(r'\s+', '', assignments) in re.sub(
r'\s+', '', old_src
)
if not already_in_src or not _APPEND_TO_OUTPUT_FILE.value:
src += assignments
# create warp function.
fn_args_raw, fn_assignments, fn_call = _warp_function(
fn_name, field_usage, mjwarp_field_info, mjx_warp_field_info
)
fn_args_raw_str = '\n'.join([' ' + arg for arg in fn_args_raw])
warp_fn_args = [arg.split(':')[0] for arg in fn_args_raw if '#' not in arg] # pytype: disable=attribute-error
src += f"""
@ffi.format_args_for_warp
def _{fn_name}_shim(
{fn_args_raw_str}
):
_m.stat = _s
_m.opt = _o
_m.callback = _cb
_d.efc = _e
_d.contact = _c
{'\n'.join(fn_assignments)}
{fn_call}
"""
src += '\n\n'
# create private jax function.
(
fn_args,
jax_args,
output_dims,
num_outputs,
tree_replace,
in_out_argnames,
stage_in_argnames,
stage_out_argnames,
has_side_effect,
) = _jax_shim_fn(fn_name, field_usage, warp_fn_args, mjwarp_field_info)
render_ctx_line = ''
return_stmt = 'return d'
if fn_name == 'render':
render_ctx_line = f' render_ctx = _MJX_RENDER_CONTEXT_BUFFERS[(ctx.key, None)]\n'
return_stmt = 'return out'
output_dims_str = '{' + ','.join(output_dims) + '}'
data_tree_replace = f"d = d.tree_replace({{ {','.join(tree_replace)} }})"
src += f"""
def _{fn_name}_jax_impl({','.join(fn_args)}):
{render_ctx_line} output_dims = {output_dims_str}
jf = ffi.jax_callable_variadic_tuple(
_{fn_name}_shim, num_outputs={num_outputs},
output_dims=output_dims,
vmap_method=None,
in_out_argnames=set([{','.join(in_out_argnames)}]),
stage_in_argnames=set([{','.join(stage_in_argnames)}]),
stage_out_argnames=set([{','.join(stage_out_argnames)}]),
graph_mode=m.opt._impl.graph_mode,
has_side_effect={has_side_effect},
)
out = jf({','.join(jax_args)})
{data_tree_replace}
{return_stmt}
"""
src += '\n'
# create public jax functions.
fn_args_no_annotation = [arg.split(':')[0] for arg in fn_args]
fn_call_str = ','.join(fn_args_no_annotation)
marshal_decorator = '@ffi.marshal_jax_warp_callable'
marshal_vmap_decorator = '@ffi.marshal_custom_vmap'
vmap_return_stmt = f'd = {fn_name}({fn_call_str})\n return d, is_batched[1]'
if fn_name == 'render':
marshal_decorator = (
'@functools.partial('
'ffi.marshal_jax_warp_callable, tree_map_output=True)'
)
marshal_vmap_decorator = (
'@functools.partial(ffi.marshal_custom_vmap, tree_map_output=True)'
)
vmap_return_stmt = (
f'out = {fn_name}({fn_call_str})\n return out, [True, True]'
)
src += f"""
@jax.custom_batching.custom_vmap
{marshal_decorator}
def {fn_name}({','.join(fn_args)}):
return _{fn_name}_jax_impl({','.join(fn_args_no_annotation)})
@{fn_name}.def_vmap
{marshal_vmap_decorator}
def {fn_name}_vmap(unused_axis_size, is_batched, {','.join(fn_args)}):
{vmap_return_stmt}
"""
src += '\n'
if _APPEND_TO_OUTPUT_FILE.value:
src = old_src + '\n\n' + src
out_fpath.write_text(src)
def main(argv: Sequence[str]) -> None:
del argv
logging.set_verbosity(logging.DEBUG)
# Get mjwarp field annotations.
fpath = epath.Path(_MJWARP_TYPES.value)
mjwarp_field_info = trace.get_mjwarp_field_info(
fpath.read_text(), file.get_cls_type_annotations
)
# Trace function to get field usage.
fpath, fn_name = _MJWARP_FUNCTION.value.split(':')
field_usage = trace.trace_function(fpath, fn_name, mjwarp_field_info)
base_path = file.get_base_path()
types_fpath = base_path / _MJX_WARP_OUTPUT_PATH.value / 'types.py'
mjx_warp_field_info = trace.get_mjx_warp_field_info(
types_fpath.read_text(), file.get_cls_type_annotations
)
base_path = file.get_base_path()
target_fpath = (
base_path / _MJX_WARP_OUTPUT_PATH.value / epath.Path(fpath).name
)
create_jax_warp_shim(
fn_name, field_usage, mjwarp_field_info, mjx_warp_field_info, target_fpath
)
if not _APPEND_TO_OUTPUT_FILE.value:
file.write_license(target_fpath)
file.format_file(target_fpath)
if __name__ == '__main__':
app.run(main)
@@ -0,0 +1,570 @@
# 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.
# ==============================================================================
"""Generate types for MJX warp integration."""
import ast
import dataclasses
import enum
import logging
import typing
from typing import Any, Callable, Dict, List, Optional, Set
from absl import app
from absl import flags
from etils import epath
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',
'Path to write the mjWarp types into.',
)
_MJX_TYPES_PATH = flags.DEFINE_string(
'mjx_types_path',
'third_party/py/mujoco/mjx/_src/types.py',
'Path to read the MJX types from.',
)
_DATA_SHAPE_PROPERTY_FIELD = 'cacc'
_DUMMY_XML = """
<mujoco>
<worldbody>
<body>
<joint type="free"/>
<geom pos="0 0 0" size="0.2" type="sphere"/>
</body>
<body >
<joint type="free"/>
<geom pos="0 0.3 0" size="0.11" type="sphere"/>
</body>
</worldbody>
</mujoco>
"""
def _to_py_string(value, indent=0):
"""Converts a dictionary/set/tuple/type to a Python code string."""
indent_str = ' ' * indent
next_indent_str = ' ' * (indent + 1)
if isinstance(value, tuple):
items = [_to_py_string(item, indent) for item in value]
return f'({', '.join(items)})'
if isinstance(value, type):
if value.__module__ == 'builtins':
return value.__name__
return f'{value.__module__}.{value.__name__}'
if isinstance(value, dict):
items = [
f'\n{next_indent_str}{repr(k)}: {_to_py_string(v, indent + 1)}'
for k, v in sorted(value.items(), key=lambda x: x[0])
]
return f'{{{','.join(items)}\n{indent_str}}}'
if isinstance(value, set):
items = sorted([_to_py_string(item, indent) for item in value])
items = [f'\n{next_indent_str}{item}' for item in items]
return f'{{{",".join(items)}\n{indent_str}}}'
return repr(value)
def _ast_parse_type(type_repr: str) -> ast.expr:
"""Parses a string representation of a type into an AST node."""
try:
return ast.parse(type_repr, mode='eval').body
except SyntaxError as e:
raise ValueError(f'Failed to parse type repr "{type_repr}": {e}') from e
def _get_target_annotation_node(
key: str,
target_annotations: Dict[str, Any],
) -> ast.expr:
"""Determines the AST node for the target type annotation for MJX."""
annotation = target_annotations.get(key)
if annotation == np.ndarray:
return _ast_parse_type('np.ndarray')
if (isinstance(annotation, wp.array) or
type(annotation).__name__ == '_ArrayAnnotation'):
return _ast_parse_type('jax.Array')
if annotation in (int, float, bool):
return _ast_parse_type(annotation.__name__)
if annotation is ffi.GraphMode:
return _ast_parse_type('GraphMode')
if isinstance(annotation, type) and issubclass(annotation, enum.Enum):
return _ast_parse_type('int')
if dataclasses.is_dataclass(annotation):
return _ast_parse_type(annotation.__name__)
is_tuple = typing.get_origin(annotation) == tuple
if is_tuple and typing.get_args(annotation)[1] != ...:
raise NotImplementedError(
'Only variadic tuples are supported. Got annotation type'
f' {annotation} for key {key}.'
)
if is_tuple and typing.get_args(annotation)[0] in (int, float, bool):
type_ = typing.get_args(annotation)[0].__name__
return _ast_parse_type(f'Tuple[{type_}, ...]')
if is_tuple and (
isinstance(typing.get_args(annotation)[0], wp.array)
or type(typing.get_args(annotation)[0]).__name__ == '_ArrayAnnotation'
):
return _ast_parse_type('Tuple[np.ndarray, ...]')
if is_tuple and dataclasses.is_dataclass(typing.get_args(annotation)[0]):
cls_type = typing.get_args(annotation)[0].__name__
return _ast_parse_type(f'Tuple[{cls_type}, ...]')
raise NotImplementedError(
f'Unhandled annotation type {annotation} for key {key}.'
)
def _get_annotations_recursive(
annotations: Dict[str, Any], prefix: str = ''
) -> Dict[str, Any]:
"""Recursively flattens type annotations, handling nested classes.
Args:
annotations: A dictionary of type annotations (field_name: type).
prefix: The prefix to add to each key, for nested classes.
Returns:
A dictionary of flattened annotations (e.g., 'contact__dist').
"""
flattened = {}
for key, annotation in annotations.items():
full_key = f'{prefix}{key}'
if hasattr(
annotation, '__annotations__'
) and 'mujoco_warp' in annotation.__module__:
nested = _get_annotations_recursive(
dict(annotation.__annotations__), prefix=f'{full_key}__'
)
flattened.update(nested)
else:
flattened[full_key] = annotation # Leaf node.
return flattened
def _build_new_class_body_ast(
keys: Set[str],
cls_name: str,
target_annotations: Dict[str, Any],
shape_property: Optional[str] = None,
add_docstring: bool = True,
) -> List[ast.AST]:
"""Builds the list of AST nodes for the new class body."""
new_body_nodes: List[ast.AST] = []
if add_docstring:
docstring = f'Derived fields from {cls_name}.'
new_body_nodes.append(ast.Expr(value=ast.Constant(value=docstring)))
# Sort keys alphabetically before creating AST nodes
sorted_keys = sorted(list(keys))
for key in sorted_keys:
annotation_node = _get_target_annotation_node(key, target_annotations)
new_body_nodes.append(
ast.AnnAssign(
target=ast.Name(id=key, ctx=ast.Store()),
annotation=annotation_node,
simple=1, # No value assignment
)
)
if shape_property is not None:
property_string = (
f'shape = property(lambda self: self.{shape_property}.shape)'
)
shape_property_node = ast.parse(property_string).body[0]
new_body_nodes.append(shape_property_node)
return new_body_nodes
def _write_class_in_file(
target_fpath: epath.Path,
target_cls_name: str,
target_base_name: str,
new_body_ast: List[ast.AST],
) -> None:
"""Reads target file, writes the specified class, and saves."""
target_code = target_fpath.read_text()
target_tree = ast.parse(target_code)
if target_cls_name in target_code:
raise ValueError(
f'Class {target_cls_name} already exists in file: {target_fpath}'
)
new_class_def = ast.ClassDef(
name=target_cls_name,
bases=[ast.Name(id=target_base_name, ctx=ast.Load())],
body=new_body_ast, # pytype: disable=wrong-arg-types
decorator_list=[],
keywords=[],
type_params=[],
)
target_tree.body.append(new_class_def)
ast.fix_missing_locations(target_tree)
modified_code = ast.unparse(target_tree)
logging.info('Writing modified code back to: %s', target_fpath)
with target_fpath.open('w') as f:
f.write(modified_code)
logging.info('File successfully rewritten.')
def write_header(target_fpath: epath.Path):
"""Writes imports and pre-defined class definitions to types.py."""
header = '''
"""MJX Warp types.
DO NOT EDIT. This file is auto-generated.
"""
import dataclasses
import typing
from typing import Tuple
import jax
from jax import tree_util
from jax.interpreters import batching
from mujoco.mjx._src import dataclasses as mjx_dataclasses
import numpy as np
if typing.TYPE_CHECKING:
GraphMode = int # Type alias for pytype.
@dataclasses.dataclass
class Callback:
pass
else:
try:
from warp._src.jax_experimental.ffi import GraphMode
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
Callback = mjwp_types.Callback
except ImportError:
GraphMode = int # Fallback when warp not installed.
Callback = None
PyTreeNode = mjx_dataclasses.PyTreeNode
'''
target_fpath.write_text(header)
_FLATTEN_UNFLATTEN = """
# flatten/unflatten all fields for custom jax_callable, but prevent the parent
# PyTreeNode from putting these fields on device.
def tree_flatten(self):
children = list(getattr(self, k) for k in self.__dataclass_fields__)
return (children, None)
@classmethod
def tree_unflatten(cls, aux_data, children):
del aux_data
return cls(*children)
"""
def write_nested_dataclass(target_fpath: epath.Path, cls: Any):
new_class_body = _build_new_class_body_ast(
set(cls.__annotations__.keys()),
cls.__name__,
dict(cls.__annotations__),
add_docstring=False,
)
cls_str = '\n'.join([' ' + 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'''
@dataclasses.dataclass(frozen=True)
@tree_util.register_pytree_node_class
class {cls.__name__}:
"""{cls.__doc__}"""
{cls_str}
{_FLATTEN_UNFLATTEN}
''')
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)
d = mujoco.MjData(m)
mujoco.mj_step(m, d)
if cls_name == 'Data':
with wp.ScopedDevice('cpu'):
dw = mjwarp.put_data(m, d, nworld=113, naconmax=1113, njmax=1113)
cond_fn = lambda x: x.shape[0] not in {113, 1113}
dw_meta = _get_fields_with_cond(dw, cond_fn)
data_non_vmap = _get_non_vmap_data_fields()
return {k for k in dw_meta if k not in data_non_vmap}
with wp.ScopedDevice('cpu'):
mw = mjwarp.put_model(m)
cond_fn = lambda x: not hasattr(x, '_is_batched')
mw_meta = _get_fields_with_cond(mw, cond_fn)
if cls_name == 'Model':
return mw_meta
if cls_name == 'Option':
return {k[len('opt__') :] for k in mw_meta if k.startswith('opt')}
if cls_name == 'Statistic':
return {k[len('stat__') :] for k in mw_meta if k.startswith('stat')}
raise NotImplementedError(f'Unhandled class name {cls_name}.')
def write_core_cls(
cls_name: str,
target_fpath: epath.Path,
mjx_types_fpath: epath.Path,
flatten_fields: bool = False,
set_diff: bool = True,
extra_annotations: dict[str, type] | None = None,
):
"""Writes a core API class (e.g. Model/Data/Option/Statistic)."""
cls = {
'Model': mjwarp.Model,
'Data': mjwarp.Data,
'Option': mjwarp.Option,
'Statistic': mjwarp.Statistic,
}[cls_name]
annotations = dict(cls.__annotations__) # pytype: disable=attribute-error
if flatten_fields:
annotations = _get_annotations_recursive(annotations)
meta_fields = _get_meta_fields(cls_name)
for k, v in annotations.items():
if k not in meta_fields:
continue
if isinstance(v, wp.array) or type(v).__name__ == '_ArrayAnnotation':
annotations[k] = np.ndarray
warp_keys = annotations.keys()
mjx_annotations = file.get_cls_type_annotations(mjx_types_fpath.read_text())[
cls_name
]
keys = warp_keys
if set_diff:
# Take the set difference between warp and mjx annotation keys.
keys = warp_keys - mjx_annotations.keys()
if extra_annotations:
annotations.update(extra_annotations)
keys = set(keys) | extra_annotations.keys()
if not keys:
raise ValueError('No derived keys found')
shape_property = None
if cls_name == 'Data':
shape_property = _DATA_SHAPE_PROPERTY_FIELD
new_class_body = _build_new_class_body_ast(
keys,
cls_name,
annotations,
shape_property=shape_property,
)
_write_class_in_file(
target_fpath=target_fpath,
target_cls_name=f'{cls_name}Warp',
target_base_name='PyTreeNode',
new_body_ast=new_class_body,
)
def _get_fields_with_cond(
d: Any,
cond_fn: Callable[[Any], bool],
s: Optional[Set[str]] = None,
prefix: str = '',
not_in: bool = False,
add_static: bool = False,
) -> Set[str]:
"""Recursively finds fields given a condition on the leading dimensions."""
if s is None:
s = set()
for f in dataclasses.fields(d):
attr = getattr(d, f.name)
if dataclasses.is_dataclass(f.type):
s = _get_fields_with_cond(attr, cond_fn, s, f.name + '__', not_in)
continue
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'):
continue
if cond_fn(attr):
s.add(prefix + f.name)
return s
def _get_non_vmap_data_fields() -> Set[str]:
"""Returns the fields that are not be vmapped but are still jax.Array."""
m = mujoco.MjModel.from_xml_string(_DUMMY_XML)
d = mujoco.MjData(m)
mujoco.mj_step(m, d)
dw = mjwarp.put_data(m, d, nworld=113, naconmax=1113, njmax=1113)
cond_fn = lambda x: x.shape[0] not in {113}
non_vmap = _get_fields_with_cond(dw, cond_fn, add_static=True)
return non_vmap
def write_register_vmappable(target_fpath: epath.Path):
"""Writes register_vmappable to types.py."""
data_non_vmap = _get_non_vmap_data_fields()
with target_fpath.open('a') as f:
f.write('\nDATA_NON_VMAP =' + _to_py_string(data_non_vmap))
f.write("""\n
def _to_elt(cont, _, d, axis):
return DataWarp(**{f.name: cont(getattr(d, f.name), axis)
if f.name not in DATA_NON_VMAP
else getattr(d, f.name) for f in DataWarp.fields()})
def _from_elt(cont, axis_size, d, axis_dest):
return DataWarp(**{f.name: cont(axis_size, getattr(d, f.name), axis_dest)
if f.name not in DATA_NON_VMAP
else getattr(d, f.name) for f in DataWarp.fields()})
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'):
return True
if wp_type in wp._src.types.value_types:
return True
if typing.get_origin(wp_type) is tuple:
return True
return False
def _to_jax_ndim(name: str, wp_type: Any) -> int:
if typing.get_origin(wp_type) is tuple:
if typing.get_args(wp_type)[1] != ...:
raise NotImplementedError('Only variadic tuples are supported.')
return -1 # signals that dim should be untouched in downstream code.
ffi_arg = ffi.FfiArg(name, wp_type)
return ffi_arg.jax_ndim
def write_ndim_annotations(target_fpath: epath.Path):
"""Writes a dictionary of ndim for every warp field."""
ndim_annotations = {}
for cls in [mjwarp.Model, mjwarp.Data, mjwarp.Option, mjwarp.Statistic]:
ndim_annotations[cls.__name__] = {}
annotations = _get_annotations_recursive(cls.__annotations__)
for name, type_ in annotations.items():
if not _is_ffi_compatible(type_):
continue
ndim = _to_jax_ndim(name, type_)
ndim_annotations[cls.__name__][name] = ndim
with target_fpath.open('a') as f:
f.write('\n_NDIM = ' + _to_py_string(ndim_annotations))
def write_nworld_leading_dim(target_fpath: epath.Path):
"""Writes a dictionary of which MJW fields are batched."""
# TODO(btaba): check that batch fields have MJX jax.Array annotations, and
# that non-batch fields have np.ndarray annotations. Fail early.
m = mujoco.MjModel.from_xml_string(_DUMMY_XML)
d = mujoco.MjData(m)
batched = {}
for cls in [mjwarp.Model, mjwarp.Data, mjwarp.Option, mjwarp.Statistic]:
batched[cls.__name__] = {}
if cls.__name__ == 'Data':
with wp.ScopedDevice('cpu'):
dw = mjwarp.put_data(m, d, nworld=113, naconmax=1113, njmax=1113)
cond_fn = lambda x: x.shape[0] == 113
batched_fields = _get_fields_with_cond(dw, cond_fn)
else:
with wp.ScopedDevice('cpu'):
obj = mjwarp.put_model(m)
if cls.__name__ == 'Option':
obj = obj.opt
elif cls.__name__ == 'Statistic':
obj = obj.stat
cond_fn = lambda x: hasattr(x, '_is_batched')
batched_fields = _get_fields_with_cond(obj, cond_fn)
all_annotations = _get_annotations_recursive(cls.__annotations__)
for name, type_ in all_annotations.items():
if not _is_ffi_compatible(type_):
continue
batched[cls.__name__][name] = name in batched_fields
with target_fpath.open('a') as f:
f.write('\n_BATCH_DIM = ' + _to_py_string(batched))
def main(argv):
del 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
write_header(target_fpath)
# TODO(btaba): consider automated grabbing of nested dataclasses from mjwarp.
write_nested_dataclass(target_fpath, mjwarp._src.types.TileSet)
write_nested_dataclass(target_fpath, mjwarp._src.types.BlockDim)
write_core_cls('Statistic', target_fpath, mjx_types_fpath, set_diff=False)
write_core_cls(
'Option', target_fpath, mjx_types_fpath,
extra_annotations={'graph_mode': ffi.GraphMode},
)
write_core_cls('Model', target_fpath, mjx_types_fpath)
write_core_cls('Data', target_fpath, mjx_types_fpath, flatten_fields=True)
write_register_vmappable(target_fpath)
write_ndim_annotations(target_fpath)
write_nworld_leading_dim(target_fpath)
file.write_license(target_fpath)
file.format_file(target_fpath)
if __name__ == '__main__':
app.run(main)
+318
View File
@@ -0,0 +1,318 @@
# 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.
# ==============================================================================
"""Static AST tracing to find MuJoCo Model and Data field usages."""
import ast
import dataclasses
import functools
import importlib
import importlib.util
import os
from typing import Dict, Optional, Sequence, Set, Tuple
from absl import logging
from etils import epath
from mujoco.mjx.codegen import file
def _get_imported_module_names(fpath: epath.Path) -> Sequence[Tuple[str, str]]:
"""Returns set of (fully qualified module_name, alias) tuples."""
module_names = set()
tree = ast.parse(fpath.read_text(), filename=fpath)
for node in ast.walk(tree):
if isinstance(node, ast.Import):
for alias in node.names:
module_names.add((alias.name, alias.name))
elif isinstance(node, ast.ImportFrom):
if node.module:
for name in node.names:
name_ = name.name if not name.asname else name.asname
module_names.add((f'{node.module}.{name.name}', f'{name_}'))
return list(module_names)
def _resolve_module_name_to_fpath(fully_qualified_name: str) -> Optional[str]:
"""Resolves a fully qualified name to its file path using importlib."""
name_parts = fully_qualified_name.split('.')
for i in range(len(name_parts), 0, -1):
module_name_to_try = '.'.join(name_parts[:i])
try:
spec = importlib.util.find_spec(module_name_to_try)
if spec and spec.origin and spec.origin != 'built-in':
return os.path.abspath(spec.origin)
except (ModuleNotFoundError, ImportError):
continue
def _get_imported_module_fpaths(fpath: epath.Path) -> Dict[str, str]:
"""Returns the file paths of all imported modules."""
all_imported_names = _get_imported_module_names(fpath)
all_resolved_fpaths = {}
for fully_qualified_name, alias in all_imported_names:
fpath = _resolve_module_name_to_fpath(fully_qualified_name)
if fpath:
all_resolved_fpaths[alias] = fpath
return all_resolved_fpaths
@dataclasses.dataclass
class FieldInfo:
param_source: str
expected_type: str
param_order: tuple[int, str]
class _FunctionFieldUsageVisitor(ast.NodeVisitor):
"""AST visitor to find attribute usages on 'm' and 'd' variables."""
def __init__(
self,
current_fpath: epath.Path,
visited_fns: Set[Tuple[str, str]],
mjwarp_field_info: Dict[str, FieldInfo],
):
self.model_fields = set()
self.data_fields = set()
self.data_out_fields = set()
self._current_fpath = current_fpath.as_posix()
self._visited_fns = visited_fns
self._mjwarp_field_info = mjwarp_field_info
self._module_fpaths = _get_imported_module_fpaths(current_fpath)
self._in_outputs_context = False
def visit_FunctionDef(self, node: ast.FunctionDef):
"""Visits nested function definitions."""
self.generic_visit(node)
def add_field_usage(self, node: ast.Attribute, is_output: bool):
"""Adds field to the appropriate set."""
attr_parts = []
curr_node = node
while isinstance(curr_node, ast.Attribute):
attr_parts.append(curr_node.attr)
curr_node = curr_node.value
if isinstance(curr_node, ast.Name):
attr_parts.reverse()
full_attribute_str = '__'.join(attr_parts)
in_field_info = full_attribute_str in self._mjwarp_field_info
if in_field_info:
if curr_node.id == 'm':
self.model_fields.add(full_attribute_str)
if curr_node.id == 'd':
self.data_fields.add(full_attribute_str)
if curr_node.id == 'd' and is_output:
self.data_out_fields.add(full_attribute_str)
def visit_Attribute(self, node: ast.Attribute):
self.add_field_usage(node, self._in_outputs_context)
self.generic_visit(node)
def visit_keyword(self, node: ast.keyword):
"""Visit a keyword argument node (e.g., outputs=[...])."""
previous_in_outputs_context = self._in_outputs_context
if node.arg == 'outputs':
self._in_outputs_context = True
try:
self.generic_visit(node)
finally:
self._in_outputs_context = previous_in_outputs_context
def recurse_trace(self, next_fpath: str, called_fn_name: str):
"""Recursively trace into a function."""
try:
field_usage = trace_function(
next_fpath,
called_fn_name,
self._mjwarp_field_info,
self._visited_fns,
)
self.model_fields.update(field_usage.model_fields)
self.data_fields.update(field_usage.data_fields)
self.data_out_fields.update(field_usage.data_out_fields)
except ValueError as e:
logging.warning(
'Could not trace function %s in %s: %s',
called_fn_name,
next_fpath,
e,
)
def visit_Call(self, node: ast.Call):
"""Visit a function call node and recursively find all attribute usages."""
if isinstance(node.func, ast.Name):
called_fn_name = node.func.id
key = (hash(self._current_fpath), called_fn_name)
if key not in self._visited_fns:
self._visited_fns.add(key)
next_fpath = self._module_fpaths.get(
called_fn_name, self._current_fpath
)
self.recurse_trace(next_fpath, called_fn_name)
elif isinstance(node.func, ast.Attribute):
parts = []
current = node.func
while isinstance(current, ast.Attribute):
parts.append(current.attr)
current = current.value
if isinstance(current, ast.Name):
parts.append(current.id)
parts = parts[::-1]
if len(parts) == 2 and parts[0] == 'wp' and parts[1] == 'copy':
if len(node.args) != 2:
raise ValueError(f'wp.copy() must have 2 arguments, got {node.args}.')
out_node, in_node = node.args
self.add_field_usage(out_node, True)
self.add_field_usage(in_node, False)
return
for arg in node.args:
if isinstance(arg, ast.Attribute):
self.add_field_usage(arg, is_output=True)
called_fn_name = '.'.join(parts[1:])
key = (hash(self._current_fpath), called_fn_name)
next_fpath = self._module_fpaths.get(parts[0])
if next_fpath and key not in self._visited_fns:
self._visited_fns.add(key)
self.recurse_trace(next_fpath, called_fn_name)
self.generic_visit(node)
@dataclasses.dataclass
class FieldUsage:
model_fields: Sequence[str] = dataclasses.field(default_factory=list)
data_fields: Sequence[str] = dataclasses.field(default_factory=list)
data_out_fields: Sequence[str] = dataclasses.field(default_factory=list)
render_context_in_caller: bool = False
def trace_function(
fpath: str,
fn: str,
mjwarp_field_info: Dict[str, FieldInfo],
visited_fns: Set[Tuple[str, str]] | None = None,
) -> FieldUsage:
"""Traces the function statically to find usages of model and data fields."""
base_path = file.get_base_path()
fpath = base_path / fpath
logging.info('Tracing function: "%s" in "%s"', fn, fpath)
src = fpath.read_text()
parsed_ast = ast.parse(src, filename=str(fpath))
target_fn_nodes = (
node
for node in parsed_ast.body
if isinstance(node, ast.FunctionDef) and node.name == fn
)
target_fn_node = next(target_fn_nodes, None)
if not target_fn_node:
raise ValueError(f'Function "{fn}" not found in "{fpath}".')
args_node = target_fn_node.args
args_tuple = tuple(
map(functools.partial(ast.get_source_segment, src), args_node.args)
)
check_args_tuple = (
('m: Model', 'd: Data'),
('m: types.Model', 'd: types.Data'),
('m', 'd'),
('m: Model',),
('m: types.Model',),
('m',),
)
if (
args_tuple[:2] not in check_args_tuple
and args_tuple[:1] not in check_args_tuple
):
raise ValueError(
f'Function "{fn}" in "{fpath}" must have arguments in'
f' {check_args_tuple} got {args_tuple}.'
)
if visited_fns is None:
visited_fns = set()
visitor = _FunctionFieldUsageVisitor(fpath, visited_fns, mjwarp_field_info)
for body in target_fn_node.body:
visitor.visit(body)
render_context_in_caller = False
if len(target_fn_node.args.args) > 2:
third_param = target_fn_node.args.args[2]
if third_param.annotation:
annotation_str = ast.unparse(third_param.annotation)
render_context_in_caller = 'RenderContext' in annotation_str
logging.info(
'End trace function "%s". Output fields: %s, RenderContext: %s',
fn, visitor.data_out_fields, render_context_in_caller
)
return FieldUsage(
model_fields=sorted(list(visitor.model_fields)),
data_fields=sorted(list(visitor.data_fields)),
data_out_fields=sorted(list(visitor.data_out_fields)),
render_context_in_caller=render_context_in_caller,
)
def get_mjwarp_field_info(src: str, get_cls_type_annotations) -> Dict[str, FieldInfo]:
"""Return field info for mujoco_warp/_src/types.py."""
dataclass_map = {
'opt': 'Option',
'stat': 'Statistic',
'efc': 'Constraint',
'contact': 'Contact',
}
field_info = {}
type_classes = get_cls_type_annotations(src)
for field, typ in type_classes['Model'].items():
if field == 'callback':
continue
if field in field_info:
raise AssertionError(f'Field {field} is duplicated in Model.')
if field in dataclass_map:
for sfield, styp in type_classes[dataclass_map[field]].items():
field_name = field + '__' + sfield
field_info[field_name] = FieldInfo('Model', styp, (1, field_name))
else:
field_info[field] = FieldInfo('Model', typ, (0, field))
for field, typ in type_classes['Data'].items():
if field in field_info:
raise AssertionError(f'Field {field} is duplicated.')
if field in dataclass_map:
for sfield, styp in type_classes[dataclass_map[field]].items():
field_name = field + '__' + sfield
field_info[field_name] = FieldInfo('Data', styp, (3, field_name))
else:
field_info[field] = FieldInfo('Data', typ, (2, field))
return field_info
def get_mjx_warp_field_info(src: str, get_cls_type_annotations) -> Dict[str, FieldInfo]:
"""Return field info for mjx/warp/types.py."""
field_info = {}
type_classes = get_cls_type_annotations(src)
for field, typ in type_classes['DataWarp'].items():
field_info[field] = FieldInfo('Data', typ, (0, field))
for field, typ in type_classes['ModelWarp'].items():
field_info[field] = FieldInfo('Model', typ, (1, field))
return field_info
+119
View File
@@ -0,0 +1,119 @@
#!/bin/bash
# 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.
# ==============================================================================
set -euo pipefail
# --- Logging helpers ---------------------------------------------------------
log_stage() { echo -e "\n\033[1;34m==== $1 ====\033[0m"; }
log_ok() { echo -e "\033[1;32m ✓ $1\033[0m"; }
log_fail() { echo -e "\033[1;31m ✗ $1\033[0m"; }
run_shim() {
local label="$1"; shift
log_stage "Generating shim: ${label}"
echo "$*"
local output
if output=$("$@" --logtostderr 2>&1); then
# Show only Python-level log lines (filters noisy C++ infra logs).
echo "$output" | grep -E '\.py' || true
log_ok "${label}"
else
echo "$output"
log_fail "${label}"
exit 1
fi
}
# --- Path setup --------------------------------------------------------------
mjwarp_base="mujoco/mjx/third_party/mujoco_warp/_src"
mjx_base="mujoco/mjx"
# Derived paths (shared).
mjwarp="${mjwarp_base}"
mjx_warp_out="${mjx_base}/warp"
mjx_types="${mjx_base}/_src/types.py"
mjx_warp_types="${mjx_base}/warp/types.py"
log_stage "Path configuration"
echo " mjwarp_base = ${mjwarp_base}"
echo " mjx_base = ${mjx_base}"
echo " mjx_types = ${mjx_types}"
echo " output dir = ${mjx_warp_out}"
# --- Stage 1: Generate warp types -------------------------------------------
log_stage "Stage 1/3: Generating warp types"
python mujoco/mjx/codegen/generate_warp_types.py \
--mjx_warp_types_out_path=${mjx_warp_types} \
--mjx_types_path=${mjx_types}
log_ok "Warp types written to ${mjx_warp_types}"
# --- Stage 2: Build shim generator ------------------------------------------
log_stage "Stage 2/3: Building shim generator"
generate_warp_shim="python mujoco/mjx/codegen/generate_warp_shim.py"
log_ok "Shim generator ready"
# --- Stage 3: Generate shim code for each function --------------------------
log_stage "Stage 3/3: Generating shim code"
# Smooth.
run_shim "smooth:kinematics" ${generate_warp_shim} \
--mjwarp_function=${mjwarp}/smooth.py:kinematics \
--mjx_warp_output_path=${mjx_warp_out}/ \
--mjwarp_types=${mjwarp}/types.py
run_shim "smooth:tendon" ${generate_warp_shim} \
--mjwarp_function=${mjwarp}/smooth.py:tendon \
--mjx_warp_output_path=${mjx_warp_out}/ \
--mjwarp_types=${mjwarp}/types.py \
--append_to_output_file=True
run_shim "smooth:com_pos" ${generate_warp_shim} \
--mjwarp_function=${mjwarp}/smooth.py:com_pos \
--mjx_warp_output_path=${mjx_warp_out}/ \
--mjwarp_types=${mjwarp}/types.py \
--append_to_output_file=True
# Collision.
run_shim "collision_driver:collision" ${generate_warp_shim} \
--mjwarp_function=${mjwarp}/collision_driver.py:collision \
--mjx_warp_output_path=${mjx_warp_out}/ \
--mjwarp_types=${mjwarp}/types.py
# Forward.
run_shim "forward:forward" ${generate_warp_shim} \
--mjwarp_function=${mjwarp}/forward.py:forward \
--mjx_warp_output_path=${mjx_warp_out}/ \
--mjwarp_types=${mjwarp}/types.py
run_shim "forward:step" ${generate_warp_shim} \
--mjwarp_function=${mjwarp}/forward.py:step \
--mjx_warp_output_path=${mjx_warp_out}/ \
--mjwarp_types=${mjwarp}/types.py \
--append_to_output_file=True
# Render and bvh.
run_shim "bvh:refit_bvh" ${generate_warp_shim} \
--mjwarp_function=${mjwarp}/bvh.py:refit_bvh \
--mjx_warp_output_path=${mjx_warp_out}/ \
--mjwarp_types=${mjwarp}/types.py
run_shim "render:render" ${generate_warp_shim} \
--mjwarp_function=${mjwarp}/render.py:render \
--mjx_warp_output_path=${mjx_warp_out}/ \
--mjwarp_types=${mjwarp}/types.py
log_stage "Done"
log_ok "All shims generated successfully"
-1
View File
@@ -46,7 +46,6 @@ _cb = mjwp_types.Callback(
**{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init}
)
@ffi.format_args_for_warp
def _collision_shim(
# Model
-1
View File
@@ -46,7 +46,6 @@ _cb = mjwp_types.Callback(
**{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init}
)
@ffi.format_args_for_warp
def _forward_shim(
# Model
+4
View File
@@ -39,6 +39,10 @@ dependencies = [
warp = [
"warp-lang==1.12.1",
]
dev = [
"isort",
"pyink",
]
[project.scripts]
mjx-testspeed = "mujoco.mjx.testspeed:main"