Use eq_active in MJX. Fixes #2173.

PiperOrigin-RevId: 689549368
Change-Id: I14d9817c5ca5dfc2a735ab8bacfb11b75c9feadc
This commit is contained in:
Baruch Tabanpour
2024-10-24 16:07:28 -07:00
committed by Copybara-Service
parent 2de430a61f
commit f24de91cc9
3 changed files with 41 additions and 12 deletions
+1
View File
@@ -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
^^^^^^^^^
+28 -10
View File
@@ -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]:
+12 -2
View File
@@ -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],