diff --git a/doc/changelog.rst b/doc/changelog.rst index 6daa2ae9..4ff449a7 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -12,6 +12,7 @@ Bug fixes 2. Fixed typos and supported fields in docs (fixes :github:issue:`1105` and :github:issue:`1106`). 3. Fixed bug where mixed ``jnt_limited`` joints were not being constrained correctly. 4. Made ``device_put`` type validation more verbose (fixes :github:issue:`1113`). +5. Removed empty EFC rows from `MJX`, for joints with no limits (fixes :github:issue:`1117`). Documentation ^^^^^^^^^^^^^ diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index 3f020198..91e83cb9 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -203,7 +203,8 @@ def _instantiate_limit(m: Model, d: Data) -> _Efc: jnt_typ = JointType(jnt_typs[i]) if jnt_typ == JointType.FREE: - return None # omit constraint rows for free joints + # this row gets removed via jnt_limited filter: + dist, j = jp.zeros(()), jp.zeros((m.nv)) elif jnt_typ == JointType.BALL: axis, angle = math.quat_to_axis_angle(qpos[qpos_i : qpos_i + 4]) dist = jp.amax(jnt_range[i]) - angle @@ -228,19 +229,13 @@ def _instantiate_limit(m: Model, d: Data) -> _Efc: return jp.stack(js), jp.stack(rs), jp.stack(arefs) - jnt_range = jp.where( - m.jnt_limited[:, None], - m.jnt_range, - jp.array([-jp.inf, jp.inf]), - ) - j, r, aref = scan.flat( m, fn, 'jjjjjqvv', 'jjj', m.jnt_type, - jnt_range, + m.jnt_range, m.jnt_solref, m.jnt_solimp, m.jnt_margin, @@ -249,6 +244,10 @@ def _instantiate_limit(m: Model, d: Data) -> _Efc: m.dof_invweight0, ) + # ignore rows for joints with no limits + jnt_limited = m.jnt_limited.astype(bool) + j, r, aref = j[jnt_limited], r[jnt_limited], aref[jnt_limited] + return _Efc(J=j, R=r, aref=aref, frictionloss=jp.zeros_like(r)) @@ -303,12 +302,12 @@ def count_constraints(m: Model, d: Data) -> Tuple[int, int, int, int]: nf = 0 - if (m.opt.disableflags & DisableBit.LIMIT) or not m.jnt_limited.any(): + if m.opt.disableflags & DisableBit.LIMIT: nl = 0 else: - nl = (m.jnt_type != JointType.FREE).sum() + nl = int(m.jnt_limited.sum()) - if (m.opt.disableflags & DisableBit.CONTACT): + if m.opt.disableflags & DisableBit.CONTACT: nc = 0 else: nc = d.ncon * 4 diff --git a/mjx/mujoco/mjx/_src/constraint_test.py b/mjx/mujoco/mjx/_src/constraint_test.py index 8bc77042..2c8c1a45 100644 --- a/mjx/mujoco/mjx/_src/constraint_test.py +++ b/mjx/mujoco/mjx/_src/constraint_test.py @@ -110,8 +110,8 @@ class ConstraintTest(parameterized.TestCase): # first joint is outside the joint range np.testing.assert_array_almost_equal(efc.J[0, 0], -1.0) - # second joint does not hit joint range - np.testing.assert_array_almost_equal(efc.aref[1], 0.0) + # second joint has no range, so only one efc row + self.assertEqual(efc.J.shape[0], 1) def test_disable_refsafe(self): m = test_util.load_test_file('ant.xml')