From e77c3cb2581dd1f45e30d50b31cb0a88de26a626 Mon Sep 17 00:00:00 2001 From: Erik Frey Date: Thu, 8 Feb 2024 16:31:26 -0800 Subject: [PATCH] Update MJX viewer to use mjx.get_data_into and fix bad function parameters. PiperOrigin-RevId: 605461208 Change-Id: I7a4f6a1ed92382e522d50224e7b5cd73d8842750 --- doc/changelog.rst | 2 ++ mjx/mujoco/mjx/viewer.py | 4 ++-- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index 9e6dc75d..55cf4d2d 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -14,6 +14,8 @@ MJX - Avoid reallocating host ``mjData`` arrays when array shapes are unchanged. - Speed up calculation of ``mjx.ncon`` for models with many geoms. - Avoid calling ``mjx.ncon`` in ``mjx.get_data_into`` when ``nc`` can be derived from ``mjx.Data``. +2. Fixed a bug in ``mjx-viewer`` that prevented it from running. Updated ``mjx-viewer`` to use newer + ``mjx.get_data_into`` function call. Version 3.1.2 (February 05, 2024) ----------------------------------- diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py index 72af4649..c29357eb 100644 --- a/mjx/mujoco/mjx/viewer.py +++ b/mjx/mujoco/mjx/viewer.py @@ -44,7 +44,7 @@ def _main(argv: Sequence[str]) -> None: print(f'Default backend: {jax.default_backend()}') print('JIT-compiling the model physics step...') start = time.time() - step_fn = jax.jit(mjx.step).lower(dx).compile() + step_fn = jax.jit(mjx.step).lower(mx, dx).compile() elapsed = time.time() - start print(f'Compilation took {elapsed}s.') @@ -64,7 +64,7 @@ def _main(argv: Sequence[str]) -> None: }) dx = step_fn(mx, dx) - mjx.device_get_into(d, dx) + mjx.get_data_into(d, m, dx) v.sync() elapsed = time.time() - start