Add warp graph capture and allow graph conditional.

PiperOrigin-RevId: 792208987
Change-Id: I139ff1db734b66734d588c3e7fe73426659cb484
This commit is contained in:
Baruch Tabanpour
2025-08-07 10:09:22 -07:00
committed by Copybara-Service
parent 4533129103
commit 9e5fdd6983
7 changed files with 20 additions and 25 deletions
+2 -3
View File
@@ -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
+6 -6
View File
@@ -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}'
-1
View File
@@ -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',
+10 -10
View File
@@ -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,
-2
View File
@@ -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',
-1
View File
@@ -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',
+2 -2
View File
@@ -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,