Add jit flag to viewer (easier debugging). Make convex.xml stable.

PiperOrigin-RevId: 678312872
Change-Id: I0d65d7c2d05a7b5103fa0be2233cbc5ba232817c
This commit is contained in:
Baruch Tabanpour
2024-09-24 10:48:40 -07:00
committed by Copybara-Service
parent 5bca7876c4
commit b09c1a977e
2 changed files with 19 additions and 8 deletions
+1 -1
View File
@@ -2,7 +2,7 @@
<default>
<geom friction="0.5 0.0 0.0"/>
</default>
<option solver="CG" iterations="6" ls_iterations="6"/>
<option solver="Newton" iterations="10" ls_iterations="6"/>
<asset>
<mesh name="tetrahedron" file="meshes/tetrahedron.stl" scale="0.4 0.4 0.4" />
<mesh name="dodecahedron" file="meshes/dodecahedron.stl" scale="0.04 0.04 0.04" />
+18 -7
View File
@@ -21,10 +21,13 @@ from typing import Sequence
from absl import app
from absl import flags
import jax
from jax import numpy as jp
import mujoco
from mujoco import mjx
import mujoco.viewer
_JIT = flags.DEFINE_bool('jit', True, 'To jit or not to jit.')
_MODEL_PATH = flags.DEFINE_string('mjcf', None, 'Path to a MuJoCo MJCF file.',
required=True)
@@ -54,11 +57,13 @@ def _main(argv: Sequence[str]) -> None:
dx = mjx.put_data(m, d)
print(f'Default backend: {jax.default_backend()}')
print('JIT-compiling the model physics step...')
start = time.time()
step_fn = jax.jit(mjx.step).lower(mx, dx).compile()
elapsed = time.time() - start
print(f'Compilation took {elapsed}s.')
step_fn = mjx.step
if _JIT.value:
print('JIT-compiling the model physics step...')
start = time.time()
step_fn = jax.jit(step_fn).lower(mx, dx).compile()
elapsed = time.time() - start
print(f'Compilation took {elapsed}s.')
viewer = mujoco.viewer.launch_passive(m, d, key_callback=key_callback)
with viewer:
@@ -66,8 +71,14 @@ def _main(argv: Sequence[str]) -> None:
start = time.time()
# TODO(robotics-simulation): recompile when changing disable flags, etc.
dx = dx.replace(ctrl=d.ctrl, act=d.act, xfrc_applied=d.xfrc_applied)
dx = dx.replace(qpos=d.qpos, qvel=d.qvel, time=d.time) # handle resets
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,