Correctly transform site_xmat in MJX get_data/put_data.

PiperOrigin-RevId: 590750273
Change-Id: I836e1f5d20725ed9b93bf86fe3e8e61954a56b47
This commit is contained in:
Erik Frey
2023-12-13 16:33:44 -08:00
committed by Copybara-Service
parent da211ddf29
commit 0915d69c3f
2 changed files with 9 additions and 4 deletions
+3 -3
View File
@@ -217,7 +217,7 @@ def _get_contact(
value = value.reshape((-1, 9))
getattr(c, field.name)[:] = value
ncon = con_id.shape[0]
ncon = cx.dist.shape[0]
c.efc_address[:] = np.arange(efc_start, efc_start + ncon * 4, 4)[con_id]
@@ -257,7 +257,7 @@ def get_data(
value = getattr(dx_i, field.name)
if field.name in ('xmat', 'ximat', 'geom_xmat'):
if field.name in ('xmat', 'ximat', 'geom_xmat', 'site_xmat'):
value = value.reshape((-1, 9))
if field.name in ('efc_frictionloss', 'efc_D', 'efc_aref', 'efc_force'):
@@ -318,7 +318,7 @@ def put_data(m: mujoco.MjModel, d: mujoco.MjData, device=None) -> types.Data:
if f.type is jax.Array
}
for fname in ('xmat', 'ximat', 'geom_xmat'):
for fname in ('xmat', 'ximat', 'geom_xmat', 'site_xmat'):
fields[fname] = fields[fname].reshape((-1, 3, 3))
# pad efc fields: MuJoCo efc arrays are sparse for inactive constraints.
+6 -1
View File
@@ -71,6 +71,7 @@ _MULTIPLE_CONSTRAINTS = """
<joint axis="0 1 0" type="hinge" range="-45 45"/>
<joint axis="1 0 0" type="hinge" range="-0.001 0.001"/>
<geom type="capsule" size=".2 .05"/>
<site pos="-0.214 -0.078 0" quat="0.664 0.664 -0.242 -0.242"/>
</body>
</body>
</worldbody>
@@ -318,9 +319,11 @@ class IoTest(parameterized.TestCase):
self.assertEqual(dx.xmat.shape, (3, 3, 3))
self.assertEqual(dx.ximat.shape, (3, 3, 3))
self.assertEqual(dx.geom_xmat.shape, (3, 3, 3))
self.assertEqual(dx.site_xmat.shape, (1, 3, 3))
np.testing.assert_allclose(dx.xmat.reshape((3, 9)), d.xmat)
np.testing.assert_allclose(dx.ximat.reshape((3, 9)), d.ximat)
np.testing.assert_allclose(dx.geom_xmat.reshape((3, 9)), d.geom_xmat)
np.testing.assert_allclose(dx.site_xmat.reshape((1, 9)), d.site_xmat)
# efc_ are also shape transformed and padded
self.assertEqual(dx.efc_J.shape, (21, 8)) # nefc, nv
@@ -369,13 +372,15 @@ class IoTest(parameterized.TestCase):
self.assertEqual(d_2.contact.frame.shape, (1, 9))
np.testing.assert_allclose(d_2.contact.frame, d.contact.frame)
# xmat, ximat, geom_xmat are all shape transformed
# xmat, ximat, geom_xmat, site_xmat are all shape transformed
self.assertEqual(d_2.xmat.shape, (3, 9))
self.assertEqual(d_2.ximat.shape, (3, 9))
self.assertEqual(d_2.geom_xmat.shape, (3, 9))
self.assertEqual(d_2.site_xmat.shape, (1, 9))
np.testing.assert_allclose(d_2.xmat, d.xmat)
np.testing.assert_allclose(d_2.ximat, d.ximat)
np.testing.assert_allclose(d_2.geom_xmat, d.geom_xmat)
np.testing.assert_allclose(d_2.site_xmat, d.site_xmat)
# efc_* are also shape transformed and filtered
self.assertEqual(d_2.efc_J.shape, (64,)) # nefc * nv