Implement joint equality constraint. Streamline constraint functions.

PiperOrigin-RevId: 577905570
Change-Id: I70580d21e3c612606b967cdd84ad22794c672397
This commit is contained in:
Erik Frey
2023-10-30 11:37:18 -07:00
committed by Copybara-Service
parent 50843cad30
commit 7069976557
9 changed files with 233 additions and 177 deletions
+9 -8
View File
@@ -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<saSimulate>`: correct handling of "Pause update", "Fullscreen" and "VSync" buttons.
7. :ref:`simulate<saSimulate>`: 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<equality-weld-torquescale>`.
10. Fixed bug relating to welds modified with :ref:`torquescale<equality-weld-torquescale>`.
Version 3.0.0 (October 18, 2023)
--------------------------------
+5 -3
View File
@@ -188,9 +188,9 @@ The following features are **fully supported** in MJX:
* - :ref:`Geom <mjtGeom>`
- ``PLANE``, ``SPHERE``, ``CAPSULE``, ``BOX``, ``MESH``
* - :ref:`Constraint <mjtConstraint>`
- ``EQUALITY``, ``FRICTION_DOF``, ``LIMIT_JOINT``, ``CONTACT_PYRAMIDAL``
- ``EQUALITY``, ``LIMIT_JOINT``, ``CONTACT_PYRAMIDAL``
* - :ref:`Equality <mjtEq>`
- ``CONNECT``, ``WELD``
- ``CONNECT``, ``WELD``, ``JOINT``
* - :ref:`Integrator <mjtIntegrator>`
- ``EULER``, ``RK4``
* - :ref:`Cone <mjtCone>`
@@ -218,6 +218,8 @@ The following features are **in development** and coming soon:
- ``TRN_TENDON``
* - :ref:`Geom <mjtGeom>`
- ``HFIELD``, ``ELLIPSOID``, ``CYLINDER``, ``SDF``
* - :ref:`Constraint <mjtConstraint>`
- ``CONTACT_FRICTIONLESS``, ``CONTACT_ELLIPTIC``, ``FRICTION_DOF``
* - :ref:`Integrator <mjtIntegrator>`
- ``IMPLICIT``, ``IMPLICITFAST``
* - :ref:`Cone <mjtCone>`
@@ -231,7 +233,7 @@ The following features are **in development** and coming soon:
* - :ref:`Tendons <tendon>`
- :ref:`Spatial <tendon-spatial>`, :ref:`Fixed <tendon-fixed>`
* - :ref:`Equality <mjtEq>`
- ``JOINT``, ``TENDON``
- ``TENDON``
* - :ref:`Sensors <mjtSensor>`
- All except ``mjSENS_PLUGIN``, ``mjSENS_USER``
+182 -144
View File
@@ -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
+4 -3
View File
@@ -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__':
+3 -3
View File
@@ -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)
+1 -1
View File
@@ -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)
+1 -1
View File
@@ -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']
+11 -10
View File
@@ -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
@@ -47,12 +47,25 @@
<freejoint/>
<geom class="free"/>
</body>
<body name="box5" pos="4 0 0">
<geom class="free"/>
<joint name="joint1" axis="1 0 0" type="hinge" />
</body>
<body name="box6" pos="4 0 0">
<geom class="free"/>
<joint name="joint2" axis="1 0 0" type="hinge" />
</body>
</worldbody>
<equality>
<connect name="connect anchor" body1="box1" body2="beam1" anchor="0 0 -1" />
<weld name="weld anchor weak torques" body1="box2" body2="beam2" torquescale="0.002" anchor="0 -2 0"/>
<weld name="weld relpose" body1="box3" body2="beam3" relpose="0 0 0 1 -.3 0 0"/>
<weld name="weld relpose+anchor" body1="box4" body2="beam4" relpose="0 0 0 1 -.3 0 0" anchor="0 0 -1"/>
<connect name="connect anchor" body1="box1" body2="beam1" anchor="0 0 -1" />
<weld name="weld anchor weak torques" body1="box2" body2="beam2" torquescale="0.002" anchor="0 -2 0"/>
<weld name="weld relpose" body1="box3" body2="beam3" relpose="0 0 0 1 -.3 0 0"/>
<weld name="weld relpose+anchor" body1="box4" body2="beam4" relpose="0 0 0 1 -.3 0 0" anchor="0 0 -1"/>
<joint name="joint" joint1="joint1" joint2="joint2" polycoef="0 -1 0.1 0.15 0.2" />
</equality>
</mujoco>