diff --git a/mjx/mujoco/mjx/testspeed.py b/mjx/mujoco/mjx/testspeed.py index 8065b3d5..9ccba845 100644 --- a/mjx/mujoco/mjx/testspeed.py +++ b/mjx/mujoco/mjx/testspeed.py @@ -22,7 +22,7 @@ from etils import epath import mujoco from mujoco import mjx -_MJCF = flags.DEFINE_string('mjcf', None, 'path to model', required=True) +_MJCF = flags.DEFINE_string('mjcf', None, 'path to model `.xml` or `.mjb`', required=True) _BASE_PATH = flags.DEFINE_string( 'base_path', None, 'base path, defaults to mujoco.mjx resource path' ) @@ -50,7 +50,10 @@ def _main(argv: Sequence[str]): path = epath.resource_path('mujoco.mjx') / 'test_data' path = _BASE_PATH.value or path f = epath.Path(path) / _MJCF.value - m = mujoco.MjModel.from_xml_path(f.as_posix()) + if f.suffix == '.mjb': + m = mujoco.MjModel.from_binary_path(f.as_posix()) + else: + m = mujoco.MjModel.from_xml_path(f.as_posix()) print(f'Rolling out {_NSTEP.value} steps at dt = {m.opt.timestep:.3f}...') jit_time, run_time, steps = mjx.benchmark( diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py index 4ee1fda9..7ad564c9 100644 --- a/mjx/mujoco/mjx/viewer.py +++ b/mjx/mujoco/mjx/viewer.py @@ -51,7 +51,10 @@ def _main(argv: Sequence[str]) -> None: jax.config.update('jax_debug_nans', True) print(f'Loading model from: {_MODEL_PATH.value}.') - m = mujoco.MjModel.from_xml_path(_MODEL_PATH.value) + if _MODEL_PATH.value.endswith('.mjb'): + m = mujoco.MjModel.from_binary_path(_MODEL_PATH.value) + else: + m = mujoco.MjModel.from_xml_path(_MODEL_PATH.value) d = mujoco.MjData(m) mx = mjx.put_model(m) dx = mjx.put_data(m, d)