From 33a30480f89a0362a1d70a014ae481e3bc802ac7 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Fri, 3 Oct 2025 02:24:32 -0700 Subject: [PATCH] Updates for MJX viewer related to Warp backend. PiperOrigin-RevId: 814604809 Change-Id: I1cdb135f8f01b399163ece7415b09fd977ce7ad8 --- mjx/mujoco/mjx/viewer.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py index b8bbff9d..934611e4 100644 --- a/mjx/mujoco/mjx/viewer.py +++ b/mjx/mujoco/mjx/viewer.py @@ -20,6 +20,7 @@ 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 @@ -68,6 +69,15 @@ def _main(argv: Sequence[str]) -> None: if len(argv) > 1: raise app.UsageError('Too many command-line arguments.') + # TODO(robotic-simulation): improved warp backend performance with MJX viewer + if _IMPL.value == 'warp': + warnings.warn( + 'The native MuJoCo Warp viewer is currently recommended for best' + ' performance.', + UserWarning, + stacklevel=2, + ) + if _WP_KERNEL_CACHE_DIR.value: wp.config.kernel_cache_dir = _WP_KERNEL_CACHE_DIR.value @@ -95,7 +105,7 @@ def _main(argv: Sequence[str]) -> None: if _JIT.value: print('JIT-compiling the model physics step...') start = time.time() - step_fn = jax.jit(step_fn).lower(mx, dx).compile() + step_fn = jax.jit(step_fn, donate_argnums=(1,)).lower(mx, dx).compile() elapsed = time.time() - start print(f'Compilation took {elapsed}s.')