From 1a9d3070428dace53bb22d492f4d8ad578059807 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Tue, 1 Oct 2024 13:10:38 -0700 Subject: [PATCH] Use mocap pos/quat in smooth.kinematics. PiperOrigin-RevId: 681136974 Change-Id: I0ece009ed1e68b61b0b96da2e685e86000f7f6b8 --- doc/changelog.rst | 5 +++++ mjx/mujoco/mjx/_src/smooth.py | 8 ++++++++ mjx/mujoco/mjx/_src/smooth_test.py | 3 +++ mjx/mujoco/mjx/_src/types.py | 4 ++-- mjx/mujoco/mjx/test_data/pendula.xml | 7 +++++++ 5 files changed, 25 insertions(+), 2 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index f2cf0f62..1447bd7b 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -14,6 +14,11 @@ General that previously required these plugins. - Replaced the function ``mjs_setActivePlugins`` with :ref:`mjs_activatePlugin`. +MJX +^^^ + +- Added ``mocap_pos`` and ``mocap_quat`` in kinematics. + Bug fixes ^^^^^^^^^ - Fixed a bug where ``actuator_force`` was not set in MJX (:github:issue:`2068`). diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index 6020a1ad..7e70c880 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -98,6 +98,14 @@ def kinematics(m: Model, d: Data) -> Data: m.body_quat, ) + if m.nmocap: + xpos = xpos.at[m.body_mocapid >= 0].set(d.mocap_pos) + mocap_quat = jax.vmap(math.normalize)(d.mocap_quat) + xquat = xquat.at[m.body_mocapid >= 0].set(mocap_quat) + xmat = xmat.at[m.body_mocapid >= 0].set( + jax.vmap(math.quat_to_mat)(mocap_quat) + ) + v_local_to_global = jax.vmap(support.local_to_global) # TODO(erikfrey): confirm that quats are more performant for mjx than mats diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index 782c5d16..b3dcfe97 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -55,6 +55,9 @@ class SmoothTest(absltest.TestCase): # give the system a little kick to ensure we have non-identity rotations d.qvel = np.random.random(m.nv) mujoco.mj_step(m, d, 10) # let dynamics get state significantly non-zero + # randomize mocap + d.mocap_pos = np.random.random(d.mocap_pos.shape) + d.mocap_quat = np.random.random(d.mocap_quat.shape) mujoco.mj_forward(m, d) mx = mjx.put_model(m) diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 6688bd9a..309b7338 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -1265,8 +1265,8 @@ class Data(PyTreeNode): xfrc_applied: jax.Array eq_active: jax.Array # mocap data: - mocap_pos: jax.Array = _restricted_to('mujoco') - mocap_quat: jax.Array = _restricted_to('mujoco') + mocap_pos: jax.Array + mocap_quat: jax.Array # dynamics: qacc: jax.Array act_dot: jax.Array diff --git a/mjx/mujoco/mjx/test_data/pendula.xml b/mjx/mujoco/mjx/test_data/pendula.xml index 82d2f23f..d183e5e9 100644 --- a/mjx/mujoco/mjx/test_data/pendula.xml +++ b/mjx/mujoco/mjx/test_data/pendula.xml @@ -124,6 +124,13 @@ + + + + + + +