Implement joint equality constraint. Streamline constraint functions.
PiperOrigin-RevId: 577905570 Change-Id: I70580d21e3c612606b967cdd84ad22794c672397
This commit is contained in:
committed by
Copybara-Service
parent
50843cad30
commit
7069976557
+9
-8
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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__':
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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']
|
||||
|
||||
@@ -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>
|
||||
Reference in New Issue
Block a user