From 706997655735adaa988ae10538ebc1a8c351a981 Mon Sep 17 00:00:00 2001
From: Erik Frey
Date: Mon, 30 Oct 2023 11:37:18 -0700
Subject: [PATCH] Implement joint equality constraint. Streamline constraint
functions.
PiperOrigin-RevId: 577905570
Change-Id: I70580d21e3c612606b967cdd84ad22794c672397
---
doc/changelog.rst | 17 +-
doc/mjx.rst | 8 +-
mjx/mujoco/mjx/_src/constraint.py | 326 ++++++++++--------
mjx/mujoco/mjx/_src/constraint_test.py | 7 +-
mjx/mujoco/mjx/_src/forward_test.py | 6 +-
mjx/mujoco/mjx/_src/smooth_test.py | 2 +-
mjx/mujoco/mjx/_src/test_util.py | 2 +-
mjx/mujoco/mjx/_src/types.py | 21 +-
.../mjx/test_data/{weld.xml => equality.xml} | 21 +-
9 files changed, 233 insertions(+), 177 deletions(-)
rename mjx/mujoco/mjx/test_data/{weld.xml => equality.xml} (59%)
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 @@
+
+
+
+
+
+