diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py
index c919f7f4..c5b94513 100644
--- a/mjx/mujoco/mjx/_src/io.py
+++ b/mjx/mujoco/mjx/_src/io.py
@@ -1110,12 +1110,10 @@ def _get_data_into(
ncon = (d_i._impl.contact.dist <= 0).sum()
efc_active = (d_i._impl.efc_J != 0).any(axis=1)
nefc = int(efc_active.sum())
- result_i.nJ = nefc * m.nv
- if ncon != result_i.ncon or nefc != result_i.nefc:
- mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc) # pylint: disable=protected-access
- result_i.efc_J_rownnz[:] = np.repeat(m.nv, nefc)
- result_i.efc_J_rowadr[:] = np.arange(0, nefc * m.nv, m.nv)
- result_i.efc_J_colind[:] = np.tile(np.arange(m.nv), nefc)
+ nj = (d_i._impl.efc_J != 0).sum() if support.is_sparse(m) else nefc * m.nv
+
+ if ncon != result_i.ncon or nefc != result_i.nefc or nj != result_i.nJ:
+ mujoco._functions._realloc_con_efc(result_i, ncon=ncon, nefc=nefc, nJ=nj) # pylint: disable=protected-access
if d.backend_impl == types.BackendImpl.JAX:
all_fields = types.Data.fields() + types.DataJAX.fields()
@@ -1167,16 +1165,33 @@ def _get_data_into(
value = {'nefc': nefc, 'ncon': ncon}[field.name]
elif field.name.endswith('xmat') or field.name == 'ximat':
value = value.reshape((-1, 9))
+ elif field.name == 'efc_J':
+ value = value[efc_active]
+ if support.is_sparse(m):
+ efc_J_rownnz = np.zeros(nefc, dtype=np.int32)
+ efc_J_rowadr = np.zeros(nefc, dtype=np.int32)
+ efc_J_colind = np.zeros(nj, dtype=np.int32)
+ efc_J = np.zeros(nj)
+ mujoco.mju_dense2sparse(
+ efc_J,
+ value,
+ efc_J_rownnz,
+ efc_J_rowadr,
+ efc_J_colind,
+ )
+ result_i.efc_J_rownnz[:] = efc_J_rownnz
+ result_i.efc_J_rowadr[:] = efc_J_rowadr
+ result_i.efc_J_colind[:] = efc_J_colind
+ value = efc_J
+ else:
+ value = value.reshape(-1)
elif field.name.startswith('efc_'):
value = value[efc_active]
- if field.name == 'efc_J':
- value = value.reshape(-1)
if d.backend_impl == types.BackendImpl.JAX:
if field.name == 'qM' and not support.is_sparse(m):
value = value[dof_i, dof_j]
elif field.name == 'qLD' and not support.is_sparse(m):
- # TODO(erikfrey): provide correct qLDs
- value = np.zeros(m.nM)
+ value = np.zeros(m.nC)
elif field.name == 'qLDiagInv' and not support.is_sparse(m):
value = np.ones(m.nv)
@@ -1191,6 +1206,10 @@ def _get_data_into(
else:
setattr(result_i, field.name, value)
+ # recalculate qLD and qLDiagInv as MJX and MuJoCo have different
+ # representations of the Cholesky decomposition.
+ mujoco.mj_factorM(m, result_i)
+
def get_data_into(
result: Union[mujoco.MjData, List[mujoco.MjData]],
diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py
index 78f81e3f..33565999 100644
--- a/mjx/mujoco/mjx/_src/io_test.py
+++ b/mjx/mujoco/mjx/_src/io_test.py
@@ -96,6 +96,18 @@ _MULTIPLE_CONSTRAINTS = """
"""
+_SIMPLE_BODY = """
+
+
+
+
+
+
+
+
+
+"""
+
class ModelIOTest(parameterized.TestCase):
"""IO tests for mjx.Model."""
@@ -427,11 +439,14 @@ class DataIOTest(parameterized.TestCase):
elif backend_impl == 'c':
np.testing.assert_allclose(dx_from_dense._impl.qM, d.qM, atol=1e-8)
- @parameterized.parameters('jax', 'c')
- def test_get_data(self, backend_impl: str):
+ @parameterized.parameters(
+ ('jax', False), ('jax', True), ('c', False), ('c', True)
+ )
+ def test_get_data(self, backend_impl: str, sparse: bool):
"""Test that get_data makes correct MjData."""
-
m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONSTRAINTS)
+ if sparse:
+ m.opt.jacobian = mujoco.mjtJacobian.mjJAC_SPARSE
d = mujoco.MjData(m)
mujoco.mj_step(m, d, 2)
dx = mjx.put_data(m, d, backend_impl=backend_impl)
@@ -443,6 +458,8 @@ class DataIOTest(parameterized.TestCase):
np.testing.assert_allclose(d_2.cvel, d.cvel)
np.testing.assert_allclose(d_2.cdof_dot, d.cdof_dot)
np.testing.assert_allclose(d_2.qM, d.qM)
+ np.testing.assert_allclose(d_2.qLD, d.qLD, atol=1e-6)
+ np.testing.assert_allclose(d_2.qLDiagInv, d.qLDiagInv, atol=1e-6)
# only 1 contact active
self.assertEqual(d_2.contact.dist.shape, (1,))
@@ -463,8 +480,27 @@ class DataIOTest(parameterized.TestCase):
# efc_* are also shape transformed and filtered
self.assertEqual(d_2.nefc, 14)
- self.assertEqual(d_2.efc_J.shape, (112,)) # nefc * nv
- np.testing.assert_allclose(d_2.efc_J, d.efc_J)
+ if sparse:
+ efc_j = np.zeros((d.nefc, m.nv))
+ mujoco.mju_sparse2dense(
+ efc_j,
+ d.efc_J,
+ d.efc_J_rownnz,
+ d.efc_J_rowadr,
+ d.efc_J_colind,
+ )
+ efc_j2 = np.zeros((d_2.nefc, m.nv))
+ mujoco.mju_sparse2dense(
+ efc_j2,
+ d_2.efc_J,
+ d_2.efc_J_rownnz,
+ d_2.efc_J_rowadr,
+ d_2.efc_J_colind,
+ )
+ np.testing.assert_allclose(efc_j, efc_j2)
+ else:
+ self.assertEqual(d_2.efc_J.shape, (112,)) # nefc * nv
+ np.testing.assert_allclose(d_2.efc_J, d.efc_J)
self.assertEqual(d_2.efc_aref.shape, (14,)) # nefc
np.testing.assert_allclose(d_2.efc_aref, d.efc_aref)
np.testing.assert_allclose(d_2.contact.efc_address, d.contact.efc_address)
@@ -473,6 +509,16 @@ class DataIOTest(parameterized.TestCase):
# check fields specific to the C implementation
np.testing.assert_allclose(d_2.bvh_active, d.bvh_active)
+ def test_get_data_simplebody(self):
+ """Test that get_data works with simple bodies where nC < nM."""
+ m = mujoco.MjModel.from_xml_string(_SIMPLE_BODY)
+ d = mujoco.MjData(m)
+ mujoco.mj_step(m, d, 2)
+ dx = mjx.put_data(m, d)
+ d_2: mujoco.MjData = mjx.get_data(m, dx)
+ np.testing.assert_allclose(d_2.qLD, d.qLD, atol=1e-6)
+ np.testing.assert_allclose(d_2.qLDiagInv, d.qLDiagInv, atol=1e-6)
+
def test_get_data_runs(self):
xml = """