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