From 3a17be2d24ad8b89272115ef5027b799f233ba19 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Mon, 23 Sep 2024 15:22:27 -0700 Subject: [PATCH] Add pause key to MJX viewer. PiperOrigin-RevId: 677959701 Change-Id: I7f816eb6bf702461f0d749b2aa6cd1664c258f3d --- mjx/mujoco/mjx/viewer.py | 21 ++++++++++++++++++--- 1 file changed, 18 insertions(+), 3 deletions(-) diff --git a/mjx/mujoco/mjx/viewer.py b/mjx/mujoco/mjx/viewer.py index 0a1605a6..84fd475f 100644 --- a/mjx/mujoco/mjx/viewer.py +++ b/mjx/mujoco/mjx/viewer.py @@ -14,6 +14,7 @@ # ============================================================================== """An example integration of MJX with the MuJoCo viewer.""" +import logging import time from typing import Sequence @@ -28,6 +29,17 @@ _MODEL_PATH = flags.DEFINE_string('mjcf', None, 'Path to a MuJoCo MJCF file.', required=True) +_VIEWER_GLOBAL_STATE = { + 'running': True, +} + + +def key_callback(key: int) -> None: + if key == 32: # Space bar + _VIEWER_GLOBAL_STATE['running'] = not _VIEWER_GLOBAL_STATE['running'] + logging.info('RUNNING = %s', _VIEWER_GLOBAL_STATE['running']) + + def _main(argv: Sequence[str]) -> None: """Launches MuJoCo passive viewer fed by MJX.""" if len(argv) > 1: @@ -48,7 +60,8 @@ def _main(argv: Sequence[str]) -> None: elapsed = time.time() - start print(f'Compilation took {elapsed}s.') - with mujoco.viewer.launch_passive(m, d) as v: + viewer = mujoco.viewer.launch_passive(m, d, key_callback=key_callback) + with viewer: while True: start = time.time() @@ -62,9 +75,11 @@ def _main(argv: Sequence[str]) -> None: 'opt.timestep': m.opt.timestep, }) - dx = step_fn(mx, dx) + if _VIEWER_GLOBAL_STATE['running']: + dx = step_fn(mx, dx) + mjx.get_data_into(d, m, dx) - v.sync() + viewer.sync() elapsed = time.time() - start if elapsed < m.opt.timestep: