Updates for MJX viewer to help avoid graph recapture when impl='warp'.
PiperOrigin-RevId: 819312078 Change-Id: I2f94eff9720fac8d06017123a51770d39743f2d6
This commit is contained in:
committed by
Copybara-Service
parent
f1d45bd542
commit
67960543a0
+38
-19
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user