diff --git a/mjx/mujoco/mjx/_src/smooth.py b/mjx/mujoco/mjx/_src/smooth.py index 3a168221..f2e3b721 100644 --- a/mjx/mujoco/mjx/_src/smooth.py +++ b/mjx/mujoco/mjx/_src/smooth.py @@ -291,7 +291,7 @@ def crb(m: Model, d: Data) -> Data: crb_cdof = jax.vmap(math.inert_mul)(crb_dof, d.cdof) qm = support.make_m(m, crb_cdof, d.cdof, m.dof_armature) d = d.replace(qM=qm) - if support.is_sparse(m): + if support.is_sparse(m) and d._qM_sparse.size > 0: # pylint: disable=protected-access d = d.replace(_qM_sparse=qm) return d @@ -353,7 +353,10 @@ def factor_m(m: Model, d: Data) -> Data: qld = (qld / qld[jp.array(madr_ds)]).at[m.dof_Madr].set(qld_diag) d = d.replace(qLD=qld, qLDiagInv=1 / qld_diag) - d = d.replace(_qLD_sparse=d.qLD, _qLDiagInv_sparse=d.qLDiagInv) + if d._qLD_sparse.size > 0: # pylint: disable=protected-access + d = d.replace(_qLD_sparse=d.qLD) + if d._qLDiagInv_sparse.size > 0: # pylint: disable=protected-access + d = d.replace(_qLDiagInv_sparse=d.qLDiagInv) return d diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index 46f1fa2e..8b5218e5 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -87,10 +87,13 @@ class SmoothTest(absltest.TestCase): dx = jax.jit(mjx.crb)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'crb') _assert_attr_eq(d, dx, 'qM') + _assert_eq(dx._qM_sparse, np.zeros(0), '_qM_sparse') # factor_m dx = jax.jit(mjx.factor_m)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'qLD') _assert_attr_eq(d, dx, 'qLDiagInv') + _assert_eq(dx._qLD_sparse, np.zeros(0), '_qLD_sparse') + _assert_eq(dx._qLDiagInv_sparse, np.zeros(0), '_qLDiagInv_sparse') # com_vel dx = jax.jit(mjx.com_vel)(mx, mjx.put_data(m, d)) _assert_attr_eq(d, dx, 'cvel')