Fix #3435. Add token to ensure sequential calls for mjx-warp refit and render.

PiperOrigin-RevId: 960553589
Change-Id: I76caca5a82dd7f96b39e51c5b67b7382ebd1726d
This commit is contained in:
Baruch Tabanpour
2026-08-06 16:05:05 -07:00
committed by Copybara-Service
parent a1d772c9ad
commit 5e3464f475
13 changed files with 245 additions and 66 deletions
+13
View File
@@ -90,6 +90,19 @@ Rendering
**Migration:** Set :at:`softness` to 1 to reproduce the previous appearance of existing models.
MJX
^^^
.. admonition:: Breaking API changes
:class: attention
- :func:`mjx.render` and :func:`mjx.render_with_segmentation` now return the updated :class:`mjx.Data` as the last
element in their return tuple (i.e. ``(rgb, depth, d)`` and ``(rgb, depth, seg, d)``). This ensures JAX/XLA
strictly enforces causal scheduling between sequential ``refit_bvh`` and ``render`` calls.
**Migration:** Update unpacking calls from ``pixels, depth = mjx.render(mx, d, rc)`` to
``pixels, depth, d = mjx.render(mx, d, rc)``.
Bug fixes
^^^^^^^^^
+7 -1
View File
@@ -267,7 +267,7 @@ volume hierarchy (BVH) and executing the raycaster:
d = mjx.refit_bvh(mx, d, rc_pytree)
# 2. Render all configured cameras
pixels, _ = mjx.render(mx, d, rc_pytree)
pixels, _, d = mjx.render(mx, d, rc_pytree)
# 3. Extract the RGB tensor for the first camera (index 0)
rgb = get_rgb(rc_pytree, 0, pixels)
@@ -276,6 +276,12 @@ volume hierarchy (BVH) and executing the raycaster:
rgb, d = render_fn(mx, d, rc.pytree())
.. NOTE::
:func:`~mujoco.mjx.refit_bvh` and :func:`~mujoco.mjx.render` update an internal execution token
(``d._impl._jax_token``) within :class:`~mujoco.mjx.Data`. Passing ``d`` sequentially through
``refit_bvh`` and ``render`` creates an explicit data dependency, preventing XLA from reordering BVH
updates and raycasting passes across iterations or unrolled loops.
.. WARNING::
The batch dimension ``nworld`` is fixed when the render context is created via
:func:`~mujoco.mjx.create_render_context` since the underlying Warp render context allocates