From 9e5fdd698390a900c61ee4624f3c50471a6adf19 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Thu, 7 Aug 2025 10:09:22 -0700 Subject: [PATCH] Add warp graph capture and allow graph conditional. PiperOrigin-RevId: 792208987 Change-Id: I139ff1db734b66734d588c3e7fe73426659cb484 --- mjx/mujoco/mjx/_src/io.py | 5 ++--- mjx/mujoco/mjx/_src/io_test.py | 12 ++++++------ mjx/mujoco/mjx/warp/collision_driver.py | 1 - mjx/mujoco/mjx/warp/ffi.py | 20 ++++++++++---------- mjx/mujoco/mjx/warp/forward.py | 2 -- mjx/mujoco/mjx/warp/smooth.py | 1 - mjx/mujoco/mjx/warp/types.py | 4 ++-- 7 files changed, 20 insertions(+), 25 deletions(-) diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index fe2dabd6..c3600c2c 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -455,7 +455,6 @@ def _put_model_warp( with wp.ScopedDevice('cpu'): # pylint: disable=undefined-variable mw = mjwp.put_model(m) # pylint: disable=undefined-variable - mw.opt.graph_conditional = False fields = {f.name for f in types.Model.fields() if f.name != '_impl'} fields = {f: getattr(m, f) for f in fields} @@ -865,7 +864,7 @@ def _make_data_warp( if not hasattr(dw, k): raise ValueError(f'Public data field {k} not found in Warp data.') field = _wp_to_np_type(getattr(dw, k)) - if mjxw.types.BATCH_DIM['Data'][k]: + if mjxw.types._BATCH_DIM['Data'][k]: # pylint: disable=protected-access field = field.reshape(field.shape[1:]) fields[k] = field @@ -873,7 +872,7 @@ def _make_data_warp( for k in mjxw.types.DataWarp.__annotations__.keys(): field = _get_nested_attr(dw, k, split='__') field = _wp_to_np_type(field) - if mjxw.types.BATCH_DIM['Data'][k]: + if mjxw.types._BATCH_DIM['Data'][k]: # pylint: disable=protected-access field = field.reshape(field.shape[1:]) impl_fields[k] = field diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index b2fc6e11..dba9c429 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -311,10 +311,10 @@ class ModelIOTest(parameterized.TestCase): def check_ndim(path, x): k = _get_name_from_path(path) - if k not in mjxw_types.NDIM['Model']: + if k not in mjxw_types._NDIM['Model']: return - is_batched = mjxw_types.BATCH_DIM['Model'][k] - expected_ndim = mjxw_types.NDIM['Model'][k] - is_batched + is_batched = mjxw_types._BATCH_DIM['Model'][k] + expected_ndim = mjxw_types._NDIM['Model'][k] - is_batched if not hasattr(x, 'ndim'): return msg = f'Field {k} has ndim {x.ndim} but expected {expected_ndim}' @@ -745,10 +745,10 @@ class DataIOTest(parameterized.TestCase): def check_ndim(path, x): k = _get_name_from_path(path) - if k not in mjxw_types.NDIM['Data']: + if k not in mjxw_types._NDIM['Data']: return - is_batched = mjxw_types.BATCH_DIM['Data'][k] - expected_ndim = mjxw_types.NDIM['Data'][k] - is_batched + is_batched = mjxw_types._BATCH_DIM['Data'][k] + expected_ndim = mjxw_types._NDIM['Data'][k] - is_batched if not hasattr(x, 'ndim'): return msg = f'Field {k} has ndim {x.ndim} but expected {expected_ndim}' diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index ecc7c0d4..798d9f8b 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -299,7 +299,6 @@ def _collision_jax_impl(m: types.Model, d: types.Data): num_outputs=37, output_dims=output_dims, vmap_method=None, - graph_compatible=True, in_out_argnames={ 'collision_hftri_index', 'collision_pair', diff --git a/mjx/mujoco/mjx/warp/ffi.py b/mjx/mujoco/mjx/warp/ffi.py index e6db9595..f406c408 100644 --- a/mjx/mujoco/mjx/warp/ffi.py +++ b/mjx/mujoco/mjx/warp/ffi.py @@ -97,7 +97,7 @@ def flatten_signature(signature: inspect.Signature, args: Tuple[Any, ...]): def jax_callable_variadic_tuple( func: Callable, # pylint: disable=g-bare-generic num_outputs: int = 1, - graph_compatible: bool = True, + graph_mode: ffi.GraphMode = ffi.GraphMode.WARP, vmap_method: Optional[str] = None, output_dims: Optional[dict[str, tuple[int, ...]]] = None, in_out_argnames: Optional[Sequence[str]] = None, @@ -116,7 +116,7 @@ def jax_callable_variadic_tuple( my_callable = ffi.jax_callable( func_wrapper, num_outputs=num_outputs, - graph_compatible=graph_compatible, + graph_mode=graph_mode, vmap_method=vmap_method, output_dims=output_dims, in_out_argnames=in_out_argnames, @@ -154,7 +154,7 @@ def _format_arg(arg: Any, name: str, annotation: Any, verbose: bool): # Add stride 0 to first axis in case the underlying argument should be # batched. # NB: the outer marshalling does an "expand_dims" on Model fields. - is_batch_field = mjx_warp_types.BATCH_DIM['Model'].get(name, False) + is_batch_field = mjx_warp_types._BATCH_DIM['Model'].get(name, False) # pylint: disable=protected-access if arg.shape[0] == 1 and is_batch_field: old_strides = arg.strides arg.strides = (0,) + arg.strides[1:] @@ -243,13 +243,13 @@ def marshal_jax_warp_callable(func): # function. m_expanded = jax.tree.map_with_path( lambda path, x: _expand_dim_from_path( - path, x, mjx_warp_types.NDIM['Model'] + path, x, mjx_warp_types._NDIM['Model'] # pylint: disable=protected-access ), m, ) d_expanded = jax.tree.map_with_path( lambda path, x: _expand_dim_from_path( - path, x, mjx_warp_types.NDIM['Data'] + path, x, mjx_warp_types._NDIM['Data'] # pylint: disable=protected-access ), d, ) @@ -292,9 +292,9 @@ def _maybe_broadcast_to( cls_str: str, ) -> Any: """Broadcasts fields that are used in MuJoCo Warp.""" - ndim = _get_mapping_from_tree_path(path, mjx_warp_types.NDIM[cls_str]) + ndim = _get_mapping_from_tree_path(path, mjx_warp_types._NDIM[cls_str]) needs_batch_dim = _get_mapping_from_tree_path( - path, mjx_warp_types.BATCH_DIM[cls_str] + path, mjx_warp_types._BATCH_DIM[cls_str] # pylint: disable=protected-access ) needs_batch_dim = bool(needs_batch_dim) and (ndim is not None and ndim > 0) if needs_batch_dim and not is_batched: @@ -319,13 +319,13 @@ def marshal_custom_vmap(vmap_func): # Flatten batch dims into the first axis if the vmap was nested. m_flat = jax.tree.map_with_path( lambda path, x: _flatten_batch_dim( - path, x, mjx_warp_types.NDIM['Model'] + path, x, mjx_warp_types._NDIM['Model'] # pylint: disable=protected-access ), m, ) d_broadcast_flat = jax.tree.map_with_path( lambda path, x: _flatten_batch_dim( - path, x, mjx_warp_types.NDIM['Data'] + path, x, mjx_warp_types._NDIM['Data'] # pylint: disable=protected-access ), d_broadcast, ) @@ -336,7 +336,7 @@ def marshal_custom_vmap(vmap_func): out_batched = jax.tree.map_with_path( # NB: if a field is not in MuJoCo Warp, we let JAX do its magic. lambda path, x: _get_mapping_from_tree_path( - path, mjx_warp_types.BATCH_DIM['Data'] + path, mjx_warp_types._BATCH_DIM['Data'] # pylint: disable=protected-access ) or x, out_batched, diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index aa54840f..d3a2a4b5 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -1207,7 +1207,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): num_outputs=182, output_dims=output_dims, vmap_method=None, - graph_compatible=True, in_out_argnames={ 'act', 'act_dot', @@ -3298,7 +3297,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): num_outputs=194, output_dims=output_dims, vmap_method=None, - graph_compatible=True, in_out_argnames={ 'act', 'act_dot', diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index 9bb37dfb..cd0e4138 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -179,7 +179,6 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): num_outputs=19, output_dims=output_dims, vmap_method=None, - graph_compatible=True, in_out_argnames={ 'flexedge_length', 'flexedge_velocity', diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 52bc0f45..387e7945 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -439,7 +439,7 @@ def _from_elt(cont, axis_size, d, axis_dest): batching.register_vmappable(DataWarp, int, int, _to_elt, _from_elt, None) -NDIM = { +_NDIM = { 'Data': { 'act': 2, 'act_dot': 2, @@ -1028,7 +1028,7 @@ NDIM = { }, 'Statistic': {'meaninertia': 0}, } -BATCH_DIM = { +_BATCH_DIM = { 'Data': { 'act': True, 'act_dot': True,