diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index 9658ad53..9396d7f3 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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. diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 4c7b0065..a16a8780 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -71,6 +71,7 @@ _MULTIPLE_CONSTRAINTS = """ + @@ -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