MJX Data.where. Fixes #3377

PiperOrigin-RevId: 943875492
Change-Id: I963e4ad052ce9184a7548a468aa4aea722d10469
This commit is contained in:
Taylor Howell
2026-07-07 07:13:40 -07:00
committed by Copybara-Service
parent 6d02ee4756
commit 6608f1affa
5 changed files with 170 additions and 29 deletions
+14
View File
@@ -132,6 +132,20 @@ Since JAX and Warp diverge in their implementations of contact buffers, contacts
For more details and examples of using MJX-Warp in the wild, see the announcement in MuJoCo Playground
`here <https://github.com/google-deepmind/mujoco_playground/discussions/197>`__.
Batched ``Data`` updates
~~~~~~~~~~~~~~~~~~~~~~~~
With MJX-JAX it is possible to reset a subset of environments in a batch with
`jax.tree.map(jax.numpy.where, done, reset_data, data)`. However, this approach does not work out-of-the-box for
MJX-Warp due to internal implementation details.
To support batched ``Data`` updates for both implementations, MJX provides a unified `where` method on `Data` objects:
.. code-block:: python
data = data.where(done, reset_data)
.. _MjxWarpGraphModes:
Graph Modes
+43
View File
@@ -177,6 +177,49 @@ class ForwardTest(absltest.TestCase):
np.testing.assert_allclose(dx.qvel, 1 + m.opt.timestep)
def test_where(self):
m = mujoco.MjModel.from_xml_string("""
<mujoco>
<worldbody>
<body>
<joint type="slide" axis="1 0 0"/>
<geom size="0.1"/>
</body>
</worldbody>
</mujoco>
""")
d_template = mjx.make_data(m)
d1 = d_template.replace(qpos=jp.array([1.0]))
d2 = d_template.replace(qpos=jp.array([2.0]))
# Test scalar condition (outside vmap)
out_true = d1.where(True, d2)
np.testing.assert_allclose(out_true.qpos, d2.qpos)
out_false = d1.where(False, d2)
np.testing.assert_allclose(out_false.qpos, d1.qpos)
# Test batched condition (inside vmap)
@jax.vmap
def merge_batched(done, r, s):
return s.where(done, r)
done_batch = jp.array([True, False])
r_batch = jax.vmap(lambda x: d_template.replace(qpos=jp.array([x])))(
jp.array([2.0, 3.0])
)
s_batch = jax.vmap(lambda x: d_template.replace(qpos=jp.array([x])))(
jp.array([1.0, 1.0])
)
merged = merge_batched(done_batch, r_batch, s_batch)
# env 0: done=True -> r -> qpos=2.0
# env 1: done=False -> s -> qpos=1.0
np.testing.assert_allclose(merged.qpos, jp.array([[2.0], [1.0]]))
class ActuatorTest(parameterized.TestCase):
+51 -10
View File
@@ -1155,21 +1155,62 @@ class Data(PyTreeNode):
return val
def __getitem__(self, key):
def get_name_from_path(path: jax.tree_util.KeyPath) -> str:
if any(isinstance(p, jax.tree_util.SequenceKey) for p in path):
is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path]
path = path[: is_seq_key.index(True)]
assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path)
path = [p for p in path if p.name != '_impl']
attr = '__'.join(p.name for p in path)
return attr
if self.impl == Impl.WARP:
return jax.tree.map_with_path(
lambda path, x, k=key: x[k]
if get_name_from_path(path) not in mjxw_types.DATA_NON_VMAP
if tree_path_to_attr_str(path) not in mjxw_types.DATA_NON_VMAP
else x,
self,
)
return jax.tree.map(lambda x: x[key], self)
def where(self, done: jax.Array, other: 'Data') -> 'Data':
"""Selectively merge self and other based on done.
Args:
done: Boolean array (or scalar inside vmap) indicating reset status.
other: Data object to select when done is True.
Returns:
Merged Data object.
"""
if self.impl != Impl.JAX and self.impl != Impl.WARP:
raise NotImplementedError(
'where is only supported for JAX and WARP implementations.'
)
if self.impl == Impl.JAX:
return jax.tree.map(
lambda x, y: jax.numpy.where(done, x, y), other, self
)
# Warp impl:
def merge_leaf(path, r_val, s_val):
field_name = tree_path_to_attr_str(path)
is_batched = mjxw_types._BATCH_DIM['Data'].get(field_name, True)
if is_batched:
return jax.numpy.where(done, r_val, s_val)
else:
return s_val
return jax.tree_util.tree_map_with_path(merge_leaf, other, self)
def tree_path_to_attr_str(path: jax.tree_util.KeyPath) -> str:
"""Converts a tree path to a dataclass attribute string."""
if not isinstance(path, tuple):
raise NotImplementedError(
f'Parsing for jax tree path {path} not implemented.'
)
if any(isinstance(p, jax.tree_util.SequenceKey) for p in path):
# get the path up to the first sequence key, we assume variadic sequences
is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path]
path = path[: is_seq_key.index(True)]
assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path)
path = [p for p in path if p.name != '_impl']
return '__'.join(p.name for p in path)
+9 -19
View File
@@ -26,6 +26,7 @@ import numpy as np
import warp as wp
from mujoco.mjx.third_party.warp._src.jax import ffi as warp_ffi
from mujoco.mjx._src.types import tree_path_to_attr_str
from mujoco.mjx.warp import types as mjx_warp_types
@@ -186,29 +187,12 @@ def format_args_for_warp(func, verbose=False):
return wrapper
def _tree_path_to_attr_str(path: jax.tree_util.KeyPath) -> str:
"""Converts a tree path to a dataclass attribute string."""
if not isinstance(path, tuple):
raise NotImplementedError(
f'Parsing for jax tree path {path} not implemented.'
)
if any(isinstance(p, jax.tree_util.SequenceKey) for p in path):
# get the path up to the first sequence key, we assume variadic sequences
is_seq_key = [isinstance(p, jax.tree_util.SequenceKey) for p in path]
path = path[: is_seq_key.index(True)]
assert all(isinstance(p, jax.tree_util.GetAttrKey) for p in path)
path = [p for p in path if p.name != '_impl']
return '__'.join(p.name for p in path)
def _get_mapping_from_tree_path(
path: jax.tree_util.KeyPath,
mapping: dict[str, int],
) -> Optional[int]:
"""Gets the mapped value from a tree path."""
attr = _tree_path_to_attr_str(path)
attr = tree_path_to_attr_str(path)
# None if the MJX public field is not present in the MJX-Warp mapping.
return mapping.get(attr)
@@ -323,11 +307,13 @@ def _check_leading_dim(
has_batch_dim = _get_mapping_from_tree_path(
path, mjx_warp_types._BATCH_DIM['Data']
)
attr = _tree_path_to_attr_str(path)
attr = tree_path_to_attr_str(path)
if has_batch_dim and leaf.shape[0] != expected_batch_dim:
raise ValueError(
f'Leaf node batch size ({leaf.shape[0]}) and expected batch size'
f' ({expected_batch_dim}) do not match for field {attr}.'
' If you are trying to merge Data objects (e.g. for resets),'
' use Data.where instead of jax.tree.map.'
)
if (
not has_batch_dim
@@ -339,6 +325,8 @@ def _check_leading_dim(
raise ValueError(
f'Leaf node leading dim ({leaf.shape[0]}) does not match naconmax'
f' ({expected_naconmax}) for field {attr}.'
' If you are trying to merge Data objects (e.g. for resets),'
' use Data.where instead of jax.tree.map.'
)
if (
not has_batch_dim
@@ -350,6 +338,8 @@ def _check_leading_dim(
raise ValueError(
f'Leaf node leading dim ({leaf.shape[0]}) does not match njmax'
f' ({expected_njmax}) for field {attr}.'
' If you are trying to merge Data objects (e.g. for resets),'
' use Data.where instead of jax.tree.map.'
)
+53
View File
@@ -386,6 +386,59 @@ class StepTest(parameterized.TestCase):
)
_ = jax.jit(jax.vmap(forward.step, in_axes=(None, 0)))(mx, dx_batch)
def test_where_autoreset(self):
if not _FORCE_TEST:
if not mjxw.WARP_INSTALLED:
self.skipTest('Warp not installed.')
if not io.has_cuda_gpu_device():
self.skipTest('No CUDA GPU device available.')
m = test_util.load_test_file('pendula.xml')
mx = mjx.put_model(m, impl='warp')
batch_size = 4
num_steps = 2
data_template = mjx.make_data(
m, impl='warp', naconmax=16 * batch_size, njmax=64
)
def reset_fn(key):
qpos = data_template.qpos.at[0].set(0.1)
d = data_template.replace(qpos=qpos)
return forward.forward(mx, d)
keys = jp.arange(batch_size)
batched_data = jax.vmap(reset_fn)(keys)
def step_fn(key, data):
stepped_data = forward.step(mx, data)
# dummy condition: reset if first joint pos > 0.05
done = stepped_data.qpos[0] > 0.05
reset_data = reset_fn(key)
return stepped_data.where(done, reset_data)
def rollout(data, key):
def body(carry, _):
data, key = carry
key, sk = jax.random.split(key)
step_keys = jax.random.split(sk, batch_size)
data = jax.vmap(step_fn)(step_keys, data)
return (data, key), None
(next_data, key), _ = jax.lax.scan(
body, (data, key), None, length=num_steps
)
return next_data
out = jax.jit(rollout)(batched_data, jax.random.PRNGKey(1))
self.assertEqual(out.qpos.shape, (batch_size, m.nq))
self.assertEqual(
out._impl.contact__type.shape, (data_template._impl.naconmax,)
)
if __name__ == '__main__':
absltest.main()