diff --git a/doc/changelog.rst b/doc/changelog.rst index 95e437ca..9015ccce 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -13,29 +13,30 @@ General MJX ^^^ -2. Fixed bug where mixed ``jnt_limited`` joints were not being constrained correctly. -3. Made ``device_put`` type validation more verbose (fixes :github:issue:`1113`). -4. Removed empty EFC rows from `MJX`, for joints with no limits (fixes :github:issue:`1117`). +2. Added support for joint equality constraints (``mjEQ_JOINT`` in :ref:`mjtEq`). +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`). Python bindings ^^^^^^^^^^^^^^^ -5. Fix the macOS ``mjpython`` launcher to work with the Python interpreter from Apple Command Line +6. Fix the macOS ``mjpython`` launcher to work with the Python interpreter from Apple Command Line Tools. Simulate ^^^^^^^^ -6. :ref:`simulate`: correct handling of "Pause update", "Fullscreen" and "VSync" buttons. +7. :ref:`simulate`: correct handling of "Pause update", "Fullscreen" and "VSync" buttons. Documentation ^^^^^^^^^^^^^ -7. Added documentation for the :ref:`UI` framework. -8. Fixed typos and supported fields in docs (fixes :github:issue:`1105` and :github:issue:`1106`). +8. Added documentation for the :ref:`UI` framework. +9. Fixed typos and supported fields in docs (fixes :github:issue:`1105` and :github:issue:`1106`). Bug fixes ^^^^^^^^^ -9. Fixed bug relating to welds modified with :ref:`torquescale`. +10. Fixed bug relating to welds modified with :ref:`torquescale`. Version 3.0.0 (October 18, 2023) -------------------------------- diff --git a/doc/mjx.rst b/doc/mjx.rst index 5bac29b3..7d1b937a 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -188,9 +188,9 @@ The following features are **fully supported** in MJX: * - :ref:`Geom ` - ``PLANE``, ``SPHERE``, ``CAPSULE``, ``BOX``, ``MESH`` * - :ref:`Constraint ` - - ``EQUALITY``, ``FRICTION_DOF``, ``LIMIT_JOINT``, ``CONTACT_PYRAMIDAL`` + - ``EQUALITY``, ``LIMIT_JOINT``, ``CONTACT_PYRAMIDAL`` * - :ref:`Equality ` - - ``CONNECT``, ``WELD`` + - ``CONNECT``, ``WELD``, ``JOINT`` * - :ref:`Integrator ` - ``EULER``, ``RK4`` * - :ref:`Cone ` @@ -218,6 +218,8 @@ The following features are **in development** and coming soon: - ``TRN_TENDON`` * - :ref:`Geom ` - ``HFIELD``, ``ELLIPSOID``, ``CYLINDER``, ``SDF`` + * - :ref:`Constraint ` + - ``CONTACT_FRICTIONLESS``, ``CONTACT_ELLIPTIC``, ``FRICTION_DOF`` * - :ref:`Integrator ` - ``IMPLICIT``, ``IMPLICITFAST`` * - :ref:`Cone ` @@ -231,7 +233,7 @@ The following features are **in development** and coming soon: * - :ref:`Tendons ` - :ref:`Spatial `, :ref:`Fixed ` * - :ref:`Equality ` - - ``JOINT``, ``TENDON`` + - ``TENDON`` * - :ref:`Sensors ` - All except ``mjSENS_PLUGIN``, ``mjSENS_USER`` diff --git a/mjx/mujoco/mjx/_src/constraint.py b/mjx/mujoco/mjx/_src/constraint.py index 34263053..166b1892 100644 --- a/mjx/mujoco/mjx/_src/constraint.py +++ b/mjx/mujoco/mjx/_src/constraint.py @@ -14,13 +14,12 @@ # ============================================================================== """Core non-smooth constraint functions.""" -from typing import Tuple +from typing import Optional, Tuple import jax from jax import numpy as jp import mujoco from mujoco.mjx._src import math -from mujoco.mjx._src import scan from mujoco.mjx._src import support # pylint: disable=g-importing-member from mujoco.mjx._src.dataclasses import PyTreeNode @@ -35,16 +34,15 @@ import numpy as np class _Efc(PyTreeNode): + """Support data for creating constraint matrices.""" J: jax.Array - R: jax.Array - aref: jax.Array + pos: jax.Array + pos_norm: jax.Array + invweight: jax.Array + solref: jax.Array + solimp: jax.Array frictionloss: jax.Array - @classmethod - def zero(cls, m: Model) -> '_Efc': - z = jp.empty((0,)) - return _Efc(J=jp.empty((0, m.nv)), R=z, aref=z, frictionloss=z) - def _kbi( m: Model, @@ -84,22 +82,18 @@ def _kbi( return k, b, imp # corresponds to K, B, I of efc_KBIP -def _instantiate_connect(m: Model, d: Data) -> _Efc: - """Returns jacobians and supporting data for connect equality constraints.""" +def _instantiate_equality_connect(m: Model, d: Data) -> Optional[_Efc]: + """Calculates constraint rows for connect equality constraints.""" - if (m.opt.disableflags & DisableBit.EQUALITY) or m.neq == 0: - return _Efc.zero(m) + ids = np.nonzero(m.eq_type == EqType.CONNECT)[0] - connect_id = np.nonzero(m.eq_type == EqType.CONNECT)[0] + if (m.opt.disableflags & DisableBit.EQUALITY) or ids.size == 0: + return None - if connect_id.size == 0: - return _Efc.zero(m) + id1, id2, data = m.eq_obj1id[ids], m.eq_obj2id[ids], m.eq_data[ids] - body1id, body2id = m.eq_obj1id[connect_id], m.eq_obj2id[connect_id] - data = m.eq_data[connect_id] - solref, solimp = m.eq_solref[connect_id], m.eq_solimp[connect_id] - - def fn(data, id1, id2, solref, solimp): + @jax.vmap + def fn(data, id1, id2): anchor1, anchor2 = data[0:3], data[3:6] # find global points pos1 = d.xmat[id1] @ anchor1 + d.xpos[id1] @@ -113,35 +107,31 @@ def _instantiate_connect(m: Model, d: Data) -> _Efc: jacp2, _ = support.jac(m, d, pos2, id2) j = (jacp1 - jacp2).T - # impedance, inverse constraint mass, reference acceleration - k, b, imp = _kbi(m, solref, solimp, math.norm(cpos)) - invweight = m.body_invweight0[id1, 0] + m.body_invweight0[id2, 0] - r = jp.maximum(invweight * (1 - imp) / imp, mujoco.mjMINVAL).repeat(3) - aref = -b * (j @ d.qvel) - k * imp * cpos + return j, cpos, jp.repeat(math.norm(cpos), 3) - return _Efc(J=j, R=r, aref=aref, frictionloss=jp.zeros_like(r)) + # concatenate to drop connect grouping dimension + j, pos, pos_norm = jax.tree_map(jp.concatenate, fn(data, id1, id2)) + invweight = m.body_invweight0[id1, 0] + m.body_invweight0[id2, 0] + invweight = jp.repeat(invweight, 3) + solref = jp.tile(m.eq_solref[ids], (3, 1)) + solimp = jp.tile(m.eq_solimp[ids], (3, 1)) + frictionloss = jp.zeros_like(pos_norm) - efcs = jax.vmap(fn)(data, body1id, body2id, solref, solimp) - - return jax.tree_map(jp.concatenate, efcs) + return _Efc(j, pos, pos_norm, invweight, solref, solimp, frictionloss) -def _instantiate_weld(m: Model, d: Data) -> _Efc: - """Returns jacobians and supporting data for connect weld constraints.""" +def _instantiate_equality_weld(m: Model, d: Data) -> Optional[_Efc]: + """Calculates constraint rows for weld equality constraints.""" - if (m.opt.disableflags & DisableBit.EQUALITY) or m.neq == 0: - return _Efc.zero(m) + ids = np.nonzero(m.eq_type == EqType.WELD)[0] - weld_id = np.nonzero(m.eq_type == EqType.WELD)[0] + if (m.opt.disableflags & DisableBit.EQUALITY) or ids.size == 0: + return None - if weld_id.size == 0: - return _Efc.zero(m) + id1, id2, data = m.eq_obj1id[ids], m.eq_obj2id[ids], m.eq_data[ids] - body1id, body2id = m.eq_obj1id[weld_id], m.eq_obj2id[weld_id] - data = m.eq_data[weld_id] - solref, solimp = m.eq_solref[weld_id], m.eq_solimp[weld_id] - - def fn(data, id1, id2, solref, solimp): + @jax.vmap + def fn(data, id1, id2): anchor1, anchor2 = data[0:3], data[3:6] relpose, torquescale = data[6:10], data[10] @@ -168,124 +158,159 @@ def _instantiate_weld(m: Model, d: Data) -> _Efc: jacdifr = 0.5 * jax.vmap(jac_fn)(jacdifr) j = jp.concatenate((jacdifp.T, jacdifr.T)) - pos = jp.concatenate((cpos, crot)).at[3:].mul(torquescale) + pos = jp.concatenate((cpos, crot * torquescale)) - # impedance, inverse constraint mass, reference acceleration - k, b, imp = _kbi(m, solref, solimp, math.norm(pos)) - invweight = m.body_invweight0[id1] + m.body_invweight0[id2] - r = jp.maximum(invweight * (1 - imp) / imp, mujoco.mjMINVAL).repeat(3) - aref = -b * (j @ d.qvel) - k * imp * pos + return j, pos, jp.repeat(math.norm(pos), 6) - return _Efc(J=j, R=r, aref=aref, frictionloss=jp.zeros_like(r)) + # concatenate to drop weld grouping dimension + j, pos, pos_norm = jax.tree_map(jp.concatenate, fn(data, id1, id2)) + invweight = m.body_invweight0[id1] + m.body_invweight0[id2] + invweight = jp.repeat(invweight, 3) + solref = jp.tile(m.eq_solref[ids], (6, 1)) + solimp = jp.tile(m.eq_solimp[ids], (6, 1)) + frictionloss = jp.zeros_like(pos_norm) - efcs = jax.vmap(fn)(data, body1id, body2id, solref, solimp) - - return jax.tree_map(jp.concatenate, efcs) + return _Efc(j, pos, pos_norm, invweight, solref, solimp, frictionloss) -def _instantiate_friction(m: Model, d: Data) -> _Efc: +def _instantiate_equality_joint(m: Model, d: Data) -> Optional[_Efc]: + """Calculates constraint rows for joint equality constraints.""" + + ids = np.nonzero(m.eq_type == EqType.JOINT)[0] + + if (m.opt.disableflags & DisableBit.EQUALITY) or ids.size == 0: + return None + + id1, id2, data = m.eq_obj1id[ids], m.eq_obj2id[ids], m.eq_data[ids] + dofadr1, dofadr2 = m.jnt_dofadr[id1], m.jnt_dofadr[id2] + qposadr1, qposadr2 = m.jnt_qposadr[id1], m.jnt_qposadr[id2] + + @jax.vmap + def fn(data, id2, dofadr1, dofadr2, qposadr1, qposadr2): + pos1, pos2 = d.qpos[qposadr1], d.qpos[qposadr2] + ref1, ref2 = m.qpos0[qposadr1], m.qpos0[qposadr2] + pos2, ref2 = pos2 * (id2 > -1), ref2 * (id2 > -1) + + dif = pos2 - ref2 + dif_power = jp.power(dif, jp.arange(0, 5)) + + deriv = jp.dot(data[1:5], dif_power[:4] * jp.arange(1, 5)) + j = jp.zeros((m.nv)).at[dofadr1].set(1.0).at[dofadr2].set(-deriv) + pos = pos1 - ref1 - jp.dot(data[:5], dif_power) + return j, pos + + j, pos = fn(data, id2, dofadr1, dofadr2, qposadr1, qposadr2) + invweight = m.dof_invweight0[dofadr1] + m.dof_invweight0[dofadr2] * (id2 > -1) + solref, solimp = m.eq_solref[ids], m.eq_solimp[ids] + frictionloss = jp.zeros_like(pos) + + return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss) + + +def _instantiate_friction(m: Model, d: Data) -> Optional[_Efc]: # TODO(robotics-team): implement _instantiate_friction - del d - return _Efc.zero(m) + del m, d + return None -def _instantiate_limit(m: Model, d: Data) -> _Efc: - """Returns jacobians and supporting data for joint limits.""" +def _instantiate_limit_ball(m: Model, d: Data) -> Optional[_Efc]: + """Calculates constraint rows for ball joint limits.""" - if (m.opt.disableflags & DisableBit.LIMIT) or not m.jnt_limited.any(): - return _Efc.zero(m) + ids = np.nonzero((m.jnt_type == JointType.BALL) & m.jnt_limited)[0] - def fn(jnt_typs, jnt_range, solref, solimp, margin, qpos, dofs, invweight0): - js, rs, arefs = [], [], [] - qpos_i, dof_i = 0, 0 + if (m.opt.disableflags & DisableBit.LIMIT) or ids.size == 0: + return None - for i in range(len(jnt_typs)): - jnt_typ = JointType(jnt_typs[i]) + jnt_range = m.jnt_range[ids] + jnt_margin = m.jnt_margin[ids] + qposadr = np.array([np.arange(q, q + 4) for q in m.jnt_qposadr[ids]]) + dofadr = np.array([np.arange(d, d + 3) for d in m.jnt_dofadr[ids]]) - if jnt_typ == JointType.FREE: - # 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 - j = jp.sum( - jax.vmap(jp.multiply)(dofs[dof_i : dof_i + 3], -axis), axis=0 - ) - elif jnt_typ in (JointType.HINGE, JointType.SLIDE): - dist_min = qpos[qpos_i] - jnt_range[i, 0] - dist_max = jnt_range[i, 1] - qpos[qpos_i] - dist = jp.minimum(dist_min, dist_max) - j = dofs[dof_i] * ((dist_min < dist_max) * 2 - 1) - else: - raise RuntimeError(f'unrecognized joint type: {jnt_typ}') + @jax.vmap + def fn(jnt_range, jnt_margin, qposadr, dofadr): + axis, angle = math.quat_to_axis_angle(d.qpos[qposadr]) + j = jp.zeros(m.nv).at[dofadr].set(-axis) + pos = jp.amax(jnt_range) - angle - jnt_margin + active = pos < 0 + return j * active, pos * active - dist = dist - margin[i] - k, b, imp = _kbi(m, solref[i], solimp[i], dist) - r = jp.maximum(invweight0[dof_i] * (1 - imp) / imp, mujoco.mjMINVAL) - aref = -b * (j @ d.qvel) - k * imp * dist - j, aref = j * (dist < 0), aref * (dist < 0) - js, rs, arefs = js + [j], rs + [r], arefs + [aref] - dof_i, qpos_i = dof_i + jnt_typ.dof_width(), qpos_i + jnt_typ.qpos_width() + j, pos = fn(jnt_range, jnt_margin, qposadr, dofadr) + invweight = m.dof_invweight0[m.jnt_dofadr[ids]] + solref, solimp = m.jnt_solref[ids], m.jnt_solimp[ids] + frictionloss = jp.zeros_like(pos) - return jp.stack(js), jp.stack(rs), jp.stack(arefs) - - j, r, aref = scan.flat( - m, - fn, - 'jjjjjqvv', - 'jjj', - m.jnt_type, - m.jnt_range, - m.jnt_solref, - m.jnt_solimp, - m.jnt_margin, - d.qpos, - jp.eye(m.nv), - 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)) + return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss) -def _instantiate_contact(m: Model, d: Data) -> _Efc: - """Returns jacobians and supporitng data for contacts.""" +def _instantiate_limit_slide_hinge(m: Model, d: Data) -> Optional[_Efc]: + """Calculates constraint rows for slide and hinge joint limits.""" + + slide_hinge = np.isin(m.jnt_type, (JointType.SLIDE, JointType.HINGE)) + ids = np.nonzero(slide_hinge & m.jnt_limited)[0] + + if (m.opt.disableflags & DisableBit.LIMIT) or ids.size == 0: + return None + + jnt_range = m.jnt_range[ids] + jnt_margin = m.jnt_margin[ids] + qposadr = m.jnt_qposadr[ids] + dofadr = m.jnt_dofadr[ids] + + @jax.vmap + def fn(jnt_range, jnt_margin, qposadr, dofadr): + dist_min = d.qpos[qposadr] - jnt_range[0] + dist_max = jnt_range[1] - d.qpos[qposadr] + j = jp.zeros(m.nv).at[dofadr].set((dist_min < dist_max) * 2 - 1) + pos = jp.minimum(dist_min, dist_max) - jnt_margin + active = pos < 0 + return j * active, pos * active + + j, pos = fn(jnt_range, jnt_margin, qposadr, dofadr) + invweight = m.dof_invweight0[dofadr] + solref, solimp = m.jnt_solref[ids], m.jnt_solimp[ids] + frictionloss = jp.zeros_like(pos) + + return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss) + + +def _instantiate_contact(m: Model, d: Data) -> Optional[_Efc]: + """Calculates constraint rows for contacts.""" if (m.opt.disableflags & DisableBit.CONTACT) or d.ncon == 0: - return _Efc.zero(m) - - def fn(contact: Contact): - dist = contact.dist - contact.includemargin - k, b, imp = _kbi(m, contact.solref, contact.solimp, dist) + return None + @jax.vmap + def fn(c: Contact): + dist = c.dist - c.includemargin geom_bodyid = jp.array(m.geom_bodyid) - body1, body2 = geom_bodyid[contact.geom1], geom_bodyid[contact.geom2] - diff = support.jac_dif_pair(m, d, contact.pos, body1, body2) + body1, body2 = geom_bodyid[c.geom1], geom_bodyid[c.geom2] + diff = support.jac_dif_pair(m, d, c.pos, body1, body2) t = m.body_invweight0[body1, 0] + m.body_invweight0[body2, 0] # rotate Jacobian differences to contact frame - diff_con = contact.frame @ diff.T + diff_con = c.frame @ diff.T # TODO(robotics-simulation): add support for other friction dimensions # 4 pyramidal friction directions - js, rs = [], [] - for diff_tan, friction in zip(diff_con[1:], contact.friction[:2]): + js, invweights = [], [] + for diff_tan, friction in zip(diff_con[1:], c.friction[:2]): for f in (friction, -friction): js.append(diff_con[0] + diff_tan * f) - rs.append((t + f * f * t) * 2 * f * f * (1 - imp) / imp) + invweights.append((t + f * f * t) * 2 * f * f) - j, r = jp.stack(js), jp.stack(rs) - r = jp.maximum(r, mujoco.mjMINVAL) - aref = -b * (j @ d.qvel) - k * imp * dist - mask_fn = jax.vmap(lambda x, mask=(dist < 0): x * mask) - j, aref = jax.tree_map(mask_fn, (j, aref)) + active = dist < 0 + j, invweight = jp.stack(js) * active, jp.stack(invweights) + pos = jp.repeat(dist, 4) * active + solref, solimp = jp.tile(c.solref, (4, 1)), jp.tile(c.solimp, (4, 1)) - return _Efc(J=j, R=r, aref=aref, frictionloss=jp.zeros_like(r)) + return j, invweight, pos, solref, solimp - return jax.tree_map(jp.concatenate, jax.vmap(fn)(d.contact)) + res = fn(d.contact) + # remove contact grouping dimension: + j, invweight, pos, solref, solimp = jax.tree_map(jp.concatenate, res) + frictionloss = jp.zeros_like(pos) + + return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss) def count_constraints(m: Model, d: Data) -> Tuple[int, int, int, int]: @@ -298,7 +323,8 @@ def count_constraints(m: Model, d: Data) -> Tuple[int, int, int, int]: else: ne_weld = (m.eq_type == EqType.WELD).sum() ne_connect = (m.eq_type == EqType.CONNECT).sum() - ne = ne_weld * 6 + ne_connect * 3 + ne_joint = (m.eq_type == EqType.JOINT).sum() + ne = ne_weld * 6 + ne_connect * 3 + ne_joint nf = 0 @@ -323,23 +349,35 @@ def make_constraint(m: Model, d: Data) -> Data: d = d.tree_replace({'contact.efc_address': np.arange(ns, ns + d.ncon * 4, 4)}) if m.opt.disableflags & DisableBit.CONSTRAINT: - efc = _Efc.zero(m) + efcs = () else: - efcs = ( - _instantiate_connect(m, d), - _instantiate_weld(m, d), + efcs = tuple(efc for efc in ( + _instantiate_equality_connect(m, d), + _instantiate_equality_weld(m, d), + _instantiate_equality_joint(m, d), _instantiate_friction(m, d), - _instantiate_limit(m, d), + _instantiate_limit_ball(m, d), + _instantiate_limit_slide_hinge(m, d), _instantiate_contact(m, d), - ) - efc = jax.tree_map(lambda *x: jp.concatenate(x), *efcs) + ) if efc is not None) - d = d.replace( - efc_J=efc.J, - efc_D=1 / efc.R, - efc_aref=efc.aref, - efc_frictionloss=efc.frictionloss, - nefc=efc.aref.shape[0], - ) + if not efcs: + z = jp.empty(0) + d = d.replace(efc_J=jp.empty((0, m.nv))) + d = d.replace(efc_D=z, efc_aref=z, efc_frictionloss=z, nefc=0) + return d + + efc = jax.tree_map(lambda *x: jp.concatenate(x), *efcs) + + @jax.vmap + def fn(efc): + k, b, imp = _kbi(m, efc.solref, efc.solimp, efc.pos_norm) + r = jp.maximum(efc.invweight * (1 - imp) / imp, mujoco.mjMINVAL) + aref = -b * (efc.J @ d.qvel) - k * imp * efc.pos + return aref, r + + aref, r = fn(efc) + d = d.replace(efc_J=efc.J, efc_D=1 / r, efc_aref=aref) + d = d.replace(efc_frictionloss=efc.frictionloss, nefc=r.shape[0]) return d diff --git a/mjx/mujoco/mjx/_src/constraint_test.py b/mjx/mujoco/mjx/_src/constraint_test.py index 3b912ce9..98448a7d 100644 --- a/mjx/mujoco/mjx/_src/constraint_test.py +++ b/mjx/mujoco/mjx/_src/constraint_test.py @@ -98,6 +98,7 @@ class ConstraintTest(parameterized.TestCase): def test_jnt_range(self): """Tests that mixed joint ranges are respected.""" + # TODO(robotics-simulation): also test ball m = mujoco.MjModel.from_xml_string(self._JNT_RANGE) m.opt.solver = SolverType.CG.value d = mujoco.MjData(m) @@ -105,7 +106,7 @@ class ConstraintTest(parameterized.TestCase): mx = mjx.device_put(m) dx = mjx.device_put(d) - efc = jax.jit(constraint._instantiate_limit)(mx, dx) + efc = jax.jit(constraint._instantiate_limit_slide_hinge)(mx, dx) # first joint is outside the joint range np.testing.assert_array_almost_equal(efc.J[0, 0], -1.0) @@ -146,7 +147,7 @@ class ConstraintTest(parameterized.TestCase): self.assertEqual(dx.efc_J.shape[0], 0) def test_disable_equality(self): - m = test_util.load_test_file('weld.xml') + m = test_util.load_test_file('equality.xml') d = mujoco.MjData(m) m.opt.disableflags = m.opt.disableflags | DisableBit.EQUALITY @@ -171,7 +172,7 @@ class ConstraintTest(parameterized.TestCase): m.opt.disableflags = m.opt.disableflags | DisableBit.CONTACT mx, dx = mjx.device_put(m), mjx.device_put(d) efc = constraint._instantiate_contact(mx, dx) - self.assertEqual(efc.J.shape[0], 0) + self.assertIsNone(efc) if __name__ == '__main__': diff --git a/mjx/mujoco/mjx/_src/forward_test.py b/mjx/mujoco/mjx/_src/forward_test.py index 0857db2f..d5bdf07b 100644 --- a/mjx/mujoco/mjx/_src/forward_test.py +++ b/mjx/mujoco/mjx/_src/forward_test.py @@ -41,7 +41,7 @@ class ForwardTest(parameterized.TestCase): @parameterized.parameters(enumerate(test_util.TEST_FILES)) def test_forward(self, seed, fname): """Test mujoco mj forward function matches mujoco_mjx forward function.""" - if fname in ('weld.xml',): + if fname in ('equality.xml',): return np.random.seed(seed) @@ -71,7 +71,7 @@ class ForwardTest(parameterized.TestCase): 'convex.xml', 'humanoid.xml', 'triple_pendulum.xml', # TODO(b/301485081) - 'weld.xml', + 'equality.xml', ): # skip models with big constraint violations at step 0 or too slow to run return @@ -101,8 +101,8 @@ class ForwardTest(parameterized.TestCase): mujoco.mj_step(m, d) dx = step_jit_fn(mx, dx) - _assert_attr_eq(d, dx, 'qpos', i, test_name, atol=1e-2) _assert_attr_eq(d, dx, 'qvel', i, test_name, atol=1e-2) + _assert_attr_eq(d, dx, 'qpos', i, test_name, atol=1e-2) _assert_attr_eq(d, dx, 'act', i, test_name) _assert_attr_eq(d, dx, 'time', i, test_name) diff --git a/mjx/mujoco/mjx/_src/smooth_test.py b/mjx/mujoco/mjx/_src/smooth_test.py index 4cdfc6ce..71ff1d75 100644 --- a/mjx/mujoco/mjx/_src/smooth_test.py +++ b/mjx/mujoco/mjx/_src/smooth_test.py @@ -43,7 +43,7 @@ class SmoothTest(parameterized.TestCase): @parameterized.parameters(enumerate(test_util.TEST_FILES)) def test_smooth(self, seed, fname): """Tests mujoco mj smooth functions match mujoco_mjx smooth functions.""" - if fname in ('convex.xml', 'weld.xml'): + if fname in ('convex.xml', 'equality.xml'): return np.random.seed(seed) diff --git a/mjx/mujoco/mjx/_src/test_util.py b/mjx/mujoco/mjx/_src/test_util.py index 5ac50816..9f677ff0 100644 --- a/mjx/mujoco/mjx/_src/test_util.py +++ b/mjx/mujoco/mjx/_src/test_util.py @@ -27,13 +27,13 @@ TEST_FILES: List[str] = [ 'ball_pendulum.xml', 'cherry_pendulum.xml', 'convex.xml', + 'equality.xml', 'humanoid.xml', 'mixed_joint_pendulum.xml', 'single_pendulum.xml', 'slide_pendulum.xml', 'triple_pendulum.xml', 'triple_pendulum_free.xml', - 'weld.xml', ] _ACTUATOR_TYPES = ['motor', 'velocity', 'position', 'general', 'intvelocity'] diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index bd3f9fe1..d43e36b1 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -145,7 +145,8 @@ class EqType(enum.IntEnum): """ CONNECT = mujoco.mjtEq.mjEQ_CONNECT WELD = mujoco.mjtEq.mjEQ_WELD - # unsupported: JOINT, TENDON, DISTANCE + JOINT = mujoco.mjtEq.mjEQ_JOINT + # unsupported: TENDON, DISTANCE class TrnType(enum.IntEnum): @@ -380,7 +381,7 @@ class Model(PyTreeNode): nexclude: int neq: int nnumeric: int - nM: int + nM: int # pylint:disable=invalid-name opt: Option stat: Statistic qpos0: jax.Array @@ -419,14 +420,14 @@ class Model(PyTreeNode): dof_bodyid: np.ndarray dof_jntid: np.ndarray dof_parentid: np.ndarray - dof_Madr: np.ndarray + dof_Madr: np.ndarray # pylint:disable=invalid-name dof_solref: jax.Array dof_solimp: jax.Array dof_frictionloss: jax.Array dof_armature: jax.Array dof_damping: jax.Array dof_invweight0: jax.Array - dof_M0: jax.Array + dof_M0: jax.Array # pylint:disable=invalid-name geom_type: np.ndarray geom_contype: np.ndarray geom_conaffinity: np.ndarray @@ -635,14 +636,14 @@ class Data(PyTreeNode): crb: jax.Array actuator_length: jax.Array actuator_moment: jax.Array - qM: jax.Array - qLD: jax.Array - qLDiagInv: jax.Array - qLDiagSqrtInv: jax.Array + qM: jax.Array # pylint:disable=invalid-name + qLD: jax.Array # pylint:disable=invalid-name + qLDiagInv: jax.Array # pylint:disable=invalid-name + qLDiagSqrtInv: jax.Array # pylint:disable=invalid-name contact: Contact - efc_J: jax.Array + efc_J: jax.Array # pylint:disable=invalid-name efc_frictionloss: jax.Array - efc_D: jax.Array + efc_D: jax.Array # pylint:disable=invalid-name # position, velocity dependent: actuator_velocity: jax.Array cvel: jax.Array diff --git a/mjx/mujoco/mjx/test_data/weld.xml b/mjx/mujoco/mjx/test_data/equality.xml similarity index 59% rename from mjx/mujoco/mjx/test_data/weld.xml rename to mjx/mujoco/mjx/test_data/equality.xml index 0e4c418a..e5c9d184 100644 --- a/mjx/mujoco/mjx/test_data/weld.xml +++ b/mjx/mujoco/mjx/test_data/equality.xml @@ -47,12 +47,25 @@ + + + + + + + + + + + + - - - - + + + + +