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 @@
+
+
+
+