From f24de91cc9d6724b838bafc7883a0d463d7f76c3 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Thu, 24 Oct 2024 16:07:28 -0700 Subject: [PATCH] Use eq_active in MJX. Fixes #2173. PiperOrigin-RevId: 689549368 Change-Id: I14d9817c5ca5dfc2a735ab8bacfb11b75c9feadc --- doc/changelog.rst | 1 + mjx/mujoco/mjx/_src/constraint.py | 38 +++++++++++++++++++------- mjx/mujoco/mjx/_src/constraint_test.py | 14 ++++++++-- 3 files changed, 41 insertions(+), 12 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index 8f69bb07..109c54f9 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -15,6 +15,7 @@ MJX ^^^ - Added ``apply_ft``, ``jac``, and ``xfrc_accumulate`` as public functions. - Added ``TOUCH`` sensor. +- Added support for ``eq_active``. Fixes :github:issue:`2173`. Bug fixes ^^^^^^^^^ diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index 26b684e8..158d9157 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -107,7 +107,9 @@ def _efc_equality_connect(m: Model, d: Data) -> Optional[_Efc]: return None @jax.vmap - def rows(is_site, obj1id, obj2id, body1id, body2id, data, solref, solimp): + def rows( + is_site, obj1id, obj2id, body1id, body2id, data, solref, solimp, active + ): anchor1, anchor2 = data[0:3], data[3:6] pos1 = d.xmat[body1id] @ anchor1 + d.xpos[body1id] @@ -128,7 +130,8 @@ def _efc_equality_connect(m: Model, d: Data) -> Optional[_Efc]: invweight = m.body_invweight0[body1id, 0] + m.body_invweight0[body2id, 0] zero = jp.zeros_like(pos) - return _row(j, pos, pos_imp, invweight, solref, solimp, zero, zero) + efc = _row(j, pos, pos_imp, invweight, solref, solimp, zero, zero) + return jax.tree_util.tree_map(lambda x: x * active, efc) is_site = m.eq_objtype == ObjType.SITE body1id = np.copy(m.eq_obj1id) @@ -147,6 +150,7 @@ def _efc_equality_connect(m: Model, d: Data) -> Optional[_Efc]: m.eq_data, m.eq_solref, m.eq_solimp, + d.eq_active, ) args = jax.tree_util.tree_map(lambda x: x[eq_id], args) # concatenate to drop row grouping @@ -161,7 +165,9 @@ def _efc_equality_weld(m: Model, d: Data) -> Optional[_Efc]: return None @jax.vmap - def rows(is_site, obj1id, obj2id, body1id, body2id, data, solref, solimp): + def rows( + is_site, obj1id, obj2id, body1id, body2id, data, solref, solimp, active + ): anchor1, anchor2 = data[0:3], data[3:6] relpose, torquescale = data[6:10], data[10] @@ -208,7 +214,8 @@ def _efc_equality_weld(m: Model, d: Data) -> Optional[_Efc]: invweight = jp.repeat(invweight, 3, axis=0) zero = jp.zeros_like(pos) - return _row(j, pos, pos_imp, invweight, solref, solimp, zero, zero) + efc = _row(j, pos, pos_imp, invweight, solref, solimp, zero, zero) + return jax.tree_util.tree_map(lambda x: x * active, efc) is_site = m.eq_objtype == ObjType.SITE body1id = np.copy(m.eq_obj1id) @@ -227,6 +234,7 @@ def _efc_equality_weld(m: Model, d: Data) -> Optional[_Efc]: m.eq_data, m.eq_solref, m.eq_solimp, + d.eq_active, ) args = jax.tree_util.tree_map(lambda x: x[eq_id], args) # concatenate to drop row grouping @@ -242,7 +250,9 @@ def _efc_equality_joint(m: Model, d: Data) -> Optional[_Efc]: return None @jax.vmap - def rows(obj2id, data, solref, solimp, dofadr1, dofadr2, qposadr1, qposadr2): + def rows( + obj2id, data, solref, solimp, active, dofadr1, dofadr2, qposadr1, qposadr2 + ): pos1, pos2 = d.qpos[qposadr1], d.qpos[qposadr2] ref1, ref2 = m.qpos0[qposadr1], m.qpos0[qposadr2] dif = (pos2 - ref2) * (obj2id > -1) @@ -255,9 +265,11 @@ def _efc_equality_joint(m: Model, d: Data) -> Optional[_Efc]: invweight += m.dof_invweight0[dofadr2] * (obj2id > -1) zero = jp.zeros_like(pos) - return _row(j, pos, pos, invweight, solref, solimp, zero, zero) + efc = _row(j, pos, pos, invweight, solref, solimp, zero, zero) + return jax.tree_util.tree_map(lambda x: x * active, efc) args = (m.eq_obj1id, m.eq_obj2id, m.eq_data, m.eq_solref, m.eq_solimp) + args += (d.eq_active,) args = jax.tree_util.tree_map(lambda x: x[eq_id], args) dofadr1, dofadr2 = m.jnt_dofadr[args[0]], m.jnt_dofadr[args[1]] qposadr1, qposadr2 = m.jnt_qposadr[args[0]], m.jnt_qposadr[args[1]] @@ -274,7 +286,7 @@ def _efc_equality_tendon(m: Model, d: Data) -> Optional[_Efc]: if (m.opt.disableflags & DisableBit.EQUALITY) or eq_id.size == 0: return None - obj1id, obj2id, data, solref, solimp = jax.tree_util.tree_map( + obj1id, obj2id, data, solref, solimp, active = jax.tree_util.tree_map( lambda x: x[eq_id], ( m.eq_obj1id, @@ -282,11 +294,14 @@ def _efc_equality_tendon(m: Model, d: Data) -> Optional[_Efc]: m.eq_data, m.eq_solref, m.eq_solimp, + d.eq_active, ), ) @jax.vmap - def rows(obj2id, data, solref, solimp, invweight, jac1, jac2, pos1, pos2): + def rows( + obj2id, data, solref, solimp, invweight, jac1, jac2, pos1, pos2, active + ): dif = pos2 * (obj2id > -1) dif_power = jp.power(dif, jp.arange(0, 5)) pos = pos1 - jp.dot(data[:5], dif_power) @@ -294,7 +309,8 @@ def _efc_equality_tendon(m: Model, d: Data) -> Optional[_Efc]: j = jac1 + jac2 * -deriv zero = jp.zeros_like(pos) - return _row(j, pos, pos, invweight, solref, solimp, zero, zero) + efc = _row(j, pos, pos, invweight, solref, solimp, zero, zero) + return jax.tree_util.tree_map(lambda x: x * active, efc) inv1, inv2 = m.tendon_invweight0[obj1id], m.tendon_invweight0[obj2id] jac1, jac2 = d.ten_J[obj1id], d.ten_J[obj2id] @@ -302,7 +318,9 @@ def _efc_equality_tendon(m: Model, d: Data) -> Optional[_Efc]: pos2 = d.ten_length[obj2id] - m.tendon_length0[obj2id] invweight = inv1 + inv2 * (obj2id > -1) - return rows(obj2id, data, solref, solimp, invweight, jac1, jac2, pos1, pos2) + return rows( + obj2id, data, solref, solimp, invweight, jac1, jac2, pos1, pos2, active + ) def _efc_friction(m: Model, d: Data) -> Optional[_Efc]: diff --git a/mjx/mujoco/mjx/_src/constraint_test.py b/mjx/mujoco/mjx/_src/constraint_test.py index ceb43147..76d9bc43 100644 --- a/mjx/mujoco/mjx/_src/constraint_test.py +++ b/mjx/mujoco/mjx/_src/constraint_test.py @@ -41,10 +41,17 @@ def _assert_attr_eq(a, b, attr): class ConstraintTest(parameterized.TestCase): + def setUp(self): + super().setUp() + np.random.seed(42) + @parameterized.parameters( - mujoco.mjtCone.mjCONE_PYRAMIDAL, mujoco.mjtCone.mjCONE_ELLIPTIC + {'cone': mujoco.mjtCone.mjCONE_PYRAMIDAL, 'rand_eq_active': False}, + {'cone': mujoco.mjtCone.mjCONE_ELLIPTIC, 'rand_eq_active': False}, + {'cone': mujoco.mjtCone.mjCONE_PYRAMIDAL, 'rand_eq_active': True}, + {'cone': mujoco.mjtCone.mjCONE_ELLIPTIC, 'rand_eq_active': True}, ) - def test_constraints(self, cone): + def test_constraints(self, cone, rand_eq_active): """Test constraints.""" m = test_util.load_test_file('constraints.xml') m.opt.cone = cone @@ -53,6 +60,8 @@ class ConstraintTest(parameterized.TestCase): # sample a mix of active/inactive constraints at different timesteps for key in range(3): mujoco.mj_resetDataKeyframe(m, d, key) + if rand_eq_active: + d.eq_active[:] = np.random.randint(0, 2, size=m.neq) mujoco.mj_forward(m, d) mx = mjx.put_model(m) dx = mjx.put_data(m, d) @@ -66,6 +75,7 @@ class ConstraintTest(parameterized.TestCase): _assert_eq(0, dx.efc_aref[order][d.nefc :], 'efc_aref') _assert_eq(d.efc_D, dx.efc_D[order][: d.nefc], 'efc_D') _assert_eq(d.efc_pos, dx.efc_pos[order][: d.nefc], 'efc_pos') + _assert_eq(dx.efc_pos[order][d.nefc:], 0, 'efc_pos') _assert_eq( d.efc_frictionloss, dx.efc_frictionloss[order][: d.nefc],