Add support for vector-valued joints in MJX bind.

PiperOrigin-RevId: 721743855
Change-Id: I68d7f6220ade89a5c0c507ea955e45657541d4eb
This commit is contained in:
Alessio Quaglino
2025-01-31 05:52:17 -08:00
committed by Copybara-Service
parent 3d9946cff3
commit 709b8d6b83
2 changed files with 27 additions and 10 deletions
+18 -4
View File
@@ -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})
+9 -6
View File
@@ -162,8 +162,8 @@ class SupportTest(parameterized.TestCase):
<mujoco model="test_bind_model">
<worldbody>
<body pos="10 20 30" name="body1">
<joint axis="1 0 0" type="slide" name="joint1"/>
<geom size="1 2 3" type="box" name="geom1"/>
<joint axis="1 0 0" type="ball" name="joint1"/>
<geom size="1 2 3" type="box" name="geom1" pos="0 1 0"/>
</body>
<body pos="40 50 60" name="body2">
<joint axis="0 1 0" type="slide" name="joint2"/>
@@ -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