From 1f1d5287fbafb3ebc2ab4b53818587d0be758ba9 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Wed, 19 Mar 2025 14:43:21 -0700 Subject: [PATCH] Fix bind.set when using a single joint instead of a list. PiperOrigin-RevId: 738546056 Change-Id: I26ae045be9e983cb0760132b9d47eac018ca33cd --- mjx/mujoco/mjx/_src/support.py | 21 +++++++++++++-------- mjx/mujoco/mjx/_src/support_test.py | 4 ++++ 2 files changed, 17 insertions(+), 8 deletions(-) diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index aa1b0768..591ee0ee 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -507,14 +507,19 @@ class BindData(object): iter(value) except TypeError: value = [value] - if name == 'qpos': - adr = self.model.jnt_qposadr[self.id] - typ = self.model.jnt_type[self.id] - num = sum((typ == jt) * jt.qpos_width() for jt in JointType) - elif name == 'qvel' or name == 'qacc': - adr = self.model.jnt_dofadr[self.id] - typ = self.model.jnt_type[self.id] - num = sum((typ == jt) * jt.dof_width() for jt in JointType) + if name in ('qpos', 'qvel', 'qacc'): + adr = num = 0 + if name == 'qpos': + adr = self.model.jnt_qposadr[self.id] + typ = self.model.jnt_type[self.id] + num = sum((typ == jt) * jt.qpos_width() for jt in JointType) + elif name == 'qvel' or name == 'qacc': + adr = self.model.jnt_dofadr[self.id] + typ = self.model.jnt_type[self.id] + num = sum((typ == jt) * jt.dof_width() for jt in JointType) + if not isinstance(self.id, list): + adr = [adr] + num = [num] elif isinstance(self.id, list): adr = self.id * dim num = [dim for _ in range(len(self.id))] diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index 5bba0df1..5948f9e1 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -287,6 +287,10 @@ class SupportTest(parameterized.TestCase): dx6 = dx.bind(mx, s.joints[::2]).set('qpos', [1, 0, 0, 0, 8]) np.testing.assert_array_equal(dx6.bind(mx, s.joints).qpos, qpos_desired) np.testing.assert_array_almost_equal(dx.bind(mx, s.joints).qpos, d.qpos) + dx6a = dx.bind(mx, s.joints[0]).set('qpos', qpos_desired[:4]) + np.testing.assert_array_equal( + dx6a.bind(mx, s.joints[0]).qpos, qpos_desired[:4] + ) dx7 = dx.bind(mx, s.joints[::2]).set('qvel', [2.0, -1.2, 0.5, 0.3]) np.testing.assert_array_almost_equal( dx7.bind(mx, s.joints).qvel, [2.0, -1.2, 0.5, 0.0, 0.3], decimal=6