Add jit flag to viewer (easier debugging). Make convex.xml stable.
PiperOrigin-RevId: 678312872 Change-Id: I0d65d7c2d05a7b5103fa0be2233cbc5ba232817c
This commit is contained in:
committed by
Copybara-Service
parent
5bca7876c4
commit
b09c1a977e
@@ -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" />
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user