From 3dcfad3293522f564eeeb68bc178edb5ffe9c823 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Mon, 27 Jan 2025 14:35:16 -0800 Subject: [PATCH] Add qpos to MJX bindings. PiperOrigin-RevId: 720311208 Change-Id: Iaf5f905e206256f38cd71949bbebd5151b2a2041 --- mjx/mujoco/mjx/_src/support.py | 27 +++++++++++++++++++++++---- mjx/mujoco/mjx/_src/support_test.py | 16 +++++++++++++--- 2 files changed, 36 insertions(+), 7 deletions(-) diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index 0e4e21ec..01326d52 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -25,6 +25,7 @@ from mujoco.mjx._src import scan from mujoco.mjx._src.types import ConeType from mujoco.mjx._src.types import Data from mujoco.mjx._src.types import JacobianType +from mujoco.mjx._src.types import JointType from mujoco.mjx._src.types import Model # pylint: enable=g-importing-member import numpy as np @@ -433,20 +434,38 @@ class BindData(object): return name else: raise AttributeError('ctrl is not available for this type') + if name == 'qpos': + if self.prefix == 'jnt_': + return name + else: + raise AttributeError('qpos is not available for this type') else: return self.prefix + name def __getattr__(self, name: str): - if name == 'sensordata': - adr = self.model.sensor_adr[self.id] - num = self.model.sensor_dim[self.id] + if name == 'sensordata' or name == 'qpos': + adr = num = 0 + if name == 'sensordata': + adr = self.model.sensor_adr[self.id] + num = self.model.sensor_dim[self.id] + elif name == 'qpos': + adr = self.model.jnt_qposadr[self.id] + typ = self.model.jnt_type[self.id] + num = ( + (typ == JointType.FREE) * JointType.FREE.qpos_width() + + (typ == JointType.BALL) * JointType.BALL.qpos_width() + + (typ == JointType.HINGE) * JointType.HINGE.qpos_width() + + (typ == JointType.SLIDE) * JointType.SLIDE.qpos_width() + ) if isinstance(self.id, list): idx = [] for a, n in zip(adr, num): idx.extend(a + j for j in range(n)) return getattr(self.data, name)[idx, ...] - else: + elif num > 1: return getattr(self.data, name)[adr : adr + num, ...] + else: + return getattr(self.data, name)[adr, ...] return getattr(self.data, self.__getname(name))[self.id, ...] def set(self, name: str, value: jax.Array) -> Data: diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index da046c70..eb4c5b20 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -161,15 +161,15 @@ class SupportTest(parameterized.TestCase): xml = """ - + - + - + @@ -223,6 +223,10 @@ class SupportTest(parameterized.TestCase): for i in range(m.njnt): np.testing.assert_array_equal(m.bind(s.joints[i]).axis, m.jnt_axis[i, :]) np.testing.assert_array_equal(mx.bind(s.joints[i]).axis, m.jnt_axis[i, :]) + np.testing.assert_array_almost_equal( + dx.bind(mx, s.joints[i]).qpos, + d.qpos[m.jnt_qposadr[i]], decimal=6 + ) np.testing.assert_array_equal(dx.bind(mx, s.actuators).ctrl, d.ctrl) for i in range(m.nu): @@ -256,6 +260,12 @@ class SupportTest(parameterized.TestCase): np.testing.assert_array_equal(dx5.bind(mx, s.actuators).ctrl, [0, 7, 0]) np.testing.assert_array_equal(dx.bind(mx, s.actuators).ctrl, [0, 0, 0]) + np.testing.assert_array_almost_equal(d.qpos, [0, 0, -3.924e-05]) + np.testing.assert_array_almost_equal(dx.bind(mx, s.joints).qpos, d.qpos) + dx6 = dx.bind(mx, s.joints[1:]).set('qpos', [8, 0]) + np.testing.assert_array_equal(dx6.bind(mx, s.joints).qpos, [0, 8, 0]) + np.testing.assert_array_almost_equal(dx.bind(mx, s.joints).qpos, d.qpos) + # test invalid name with self.assertRaises(AttributeError): print(dx.bind(mx, s.geoms).ctrl)