From 67960543a010ff6b1deb20df933c3ff5f82faf33 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Tue, 14 Oct 2025 11:14:02 -0700 Subject: [PATCH] Updates for MJX viewer to help avoid graph recapture when impl='warp'. PiperOrigin-RevId: 819312078 Change-Id: I2f94eff9720fac8d06017123a51770d39743f2d6 --- mjx/mujoco/mjx/viewer.py | 57 ++++++++++++++++++++++++++-------------- 1 file changed, 38 insertions(+), 19 deletions(-) diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py index 934611e4..cc110a89 100644 --- a/mjx/mujoco/mjx/viewer.py +++ b/mjx/mujoco/mjx/viewer.py @@ -14,13 +14,13 @@ # ============================================================================== """An example integration of MJX with the MuJoCo viewer.""" +import copy import logging import os os.environ['XLA_FLAGS'] = '--xla_gpu_graph_min_graph_size=1' import time # pylint: disable=g-import-not-at-top from typing import Sequence -import warnings from absl import app from absl import flags @@ -71,11 +71,9 @@ def _main(argv: Sequence[str]) -> None: # TODO(robotic-simulation): improved warp backend performance with MJX viewer if _IMPL.value == 'warp': - warnings.warn( + logging.info( 'The native MuJoCo Warp viewer is currently recommended for best' ' performance.', - UserWarning, - stacklevel=2, ) if _WP_KERNEL_CACHE_DIR.value: @@ -102,33 +100,54 @@ def _main(argv: Sequence[str]) -> None: print(f'Default backend: {jax.default_backend()}') step_fn = mjx.step + + def set_model_fn(mx, gravity, tolerance, ls_tolerance, timestep): + return mx.tree_replace({ + 'opt.gravity': jp.array(gravity), + 'opt.tolerance': jp.array(tolerance), + 'opt.ls_tolerance': jp.array(ls_tolerance), + 'opt.timestep': jp.array(timestep), + }) + + def set_data_fn(dx, ctrl, act, xfrc_applied, qpos, qvel, time_): + return dx.tree_replace({ + 'ctrl': jp.array(ctrl), + 'act': jp.array(act), + 'xfrc_applied': jp.array(xfrc_applied), + 'qpos': jp.array(qpos), + 'qvel': jp.array(qvel), + 'time': jp.array(time_), + }) + if _JIT.value: print('JIT-compiling the model physics step...') start = time.time() - step_fn = jax.jit(step_fn, donate_argnums=(1,)).lower(mx, dx).compile() + step_fn = jax.jit(step_fn, donate_argnums=(1,), keep_unused=True).lower(mx, dx).compile() elapsed = time.time() - start print(f'Compilation took {elapsed}s.') + set_model_fn = ( + jax.jit(set_model_fn, donate_argnums=(0,), keep_unused=True) + .lower(mx, m.opt.gravity, m.opt.tolerance, m.opt.ls_tolerance, m.opt.timestep) + .compile() + ) + set_data_fn = ( + jax.jit(set_data_fn, donate_argnums=(0,), keep_unused=True) + .lower(dx, d.ctrl, d.act, d.xfrc_applied, d.qpos, d.qvel, d.time) + .compile() + ) viewer = mujoco.viewer.launch_passive(m, d, key_callback=key_callback) with viewer: + opt = copy.copy(m.opt) while True: start = time.time() # TODO(robotics-simulation): recompile when changing disable flags, etc. - dx = dx.replace( - ctrl=jp.array(d.ctrl), - act=jp.array(d.act), - xfrc_applied=jp.array(d.xfrc_applied), - ) - dx = dx.replace( - qpos=jp.array(d.qpos), qvel=jp.array(d.qvel), time=jp.array(d.time) - ) # handle resets - mx = mx.tree_replace({ - 'opt.gravity': m.opt.gravity, - 'opt.tolerance': m.opt.tolerance, - 'opt.ls_tolerance': m.opt.ls_tolerance, - 'opt.timestep': m.opt.timestep, - }) + dx = set_data_fn(dx, d.ctrl, d.act, d.xfrc_applied, d.qpos, d.qvel, d.time) + + if m.opt != opt: + opt = copy.copy(m.opt) + mx = set_model_fn(mx, m.opt.gravity, m.opt.tolerance, m.opt.ls_tolerance, m.opt.timestep) if _VIEWER_GLOBAL_STATE['running']: dx = step_fn(mx, dx)