diff --git a/mjx/mujoco/mjx/_src/support.py b/mjx/mujoco/mjx/_src/support.py index 041942b9..09fc5365 100644 --- a/mjx/mujoco/mjx/_src/support.py +++ b/mjx/mujoco/mjx/_src/support.py @@ -477,11 +477,25 @@ class BindData(object): iter(value) except TypeError: value = [value] - if len(value) == 1: - array = array.at[self.id].set(value[0]) + if 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() + ) + elif isinstance(self.id, list): + adr = self.id + num = [1 for _ in range(len(self.id))] else: - for i, v in enumerate(value): - array = array.at[self.id[i]].set(v) + adr = [self.id] + num = [1] + i = 0 + for a, n in zip(adr, num): + array = array.at[a: a + n].set(value[i: i + n]) + i += n return self.data.replace(**{self.__getname(name): array}) diff --git a/mjx/mujoco/mjx/_src/support_test.py b/mjx/mujoco/mjx/_src/support_test.py index 6fadc843..965f369c 100644 --- a/mjx/mujoco/mjx/_src/support_test.py +++ b/mjx/mujoco/mjx/_src/support_test.py @@ -162,8 +162,8 @@ class SupportTest(parameterized.TestCase): - - + + @@ -220,12 +220,13 @@ class SupportTest(parameterized.TestCase): np.testing.assert_array_equal(mx.bind(s.joints).axis, m.jnt_axis) np.testing.assert_array_equal(mx.bind(s.joints).qposadr, m.jnt_qposadr) + qposnum = [4, 1, 1] 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 + d.qpos[m.jnt_qposadr[i]:m.jnt_qposadr[i] + qposnum[i]], decimal=6 ) np.testing.assert_array_equal(dx.bind(mx, s.actuators).ctrl, d.ctrl) @@ -260,10 +261,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]) + qpos_1step = [1.00000e00, -3.67875e-06, 0, 0, 0, -3.924e-05] + qpos_desired = [1, 0, 0, 0, 0, 8] + np.testing.assert_array_almost_equal(d.qpos, qpos_1step) 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]) + 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) # test invalid name