Add warp graph capture and allow graph conditional.
PiperOrigin-RevId: 792208987 Change-Id: I139ff1db734b66734d588c3e7fe73426659cb484
This commit is contained in:
committed by
Copybara-Service
parent
4533129103
commit
9e5fdd6983
@@ -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
|
||||
|
||||
|
||||
@@ -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}'
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user