From 827401554e2bdd1c27c160e9288166452b84921d Mon Sep 17 00:00:00 2001 From: Kristian Hartikainen Date: Mon, 23 Sep 2024 11:30:13 +0300 Subject: [PATCH 1/2] Enable binary `.mjb` in `mjx/testspeed.py` --- mjx/mujoco/mjx/testspeed.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) 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( From 6f031f2293dc9339b42cac8a690ec1a16aee8b8e Mon Sep 17 00:00:00 2001 From: Kristian Hartikainen Date: Thu, 26 Sep 2024 12:38:56 +0300 Subject: [PATCH 2/2] Enable binary `.mjb` in `mjx/viewer.py` --- mjx/mujoco/mjx/viewer.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py index 84fd475f..08b24553 100644 --- a/mjx/mujoco/mjx/viewer.py +++ b/mjx/mujoco/mjx/viewer.py @@ -48,7 +48,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)