Elliptic friction in MJX.
PiperOrigin-RevId: 641384367 Change-Id: I510c565940324fbbf392ce537ce27e0cb9af3eb1
This commit is contained in:
committed by
Copybara-Service
parent
5b6800116a
commit
4c3d9461ae
@@ -10,6 +10,13 @@ General
|
||||
|
||||
1. Added :ref:`maxhullvert<asset-mesh-maxhullvert>`, the maximum number of vertices in a mesh's convex hull.
|
||||
|
||||
MJX
|
||||
~~~
|
||||
|
||||
2. Added support for :ref:`elliptic friction cones<option-cone>`.
|
||||
3. Fixed a bug that resulted in less-optimal linesearch solutions for some difficult constraint settings.
|
||||
4. Fixed a bug in the Newton solver that sometimes resulted in less-optimal gradients.
|
||||
|
||||
|
||||
Version 3.1.6 (Jun 3, 2024)
|
||||
---------------------------
|
||||
|
||||
+3
-5
@@ -198,13 +198,13 @@ The following features are **fully supported** in MJX:
|
||||
* - :ref:`Geom <mjtGeom>`
|
||||
- ``PLANE``, ``HFIELD``, ``SPHERE``, ``CAPSULE``, ``BOX``, ``MESH`` are fully implemented. ``ELLIPSOID`` and ``CYLINDER`` are implemented but only collide with other primitives.
|
||||
* - :ref:`Constraint <mjtConstraint>`
|
||||
- ``EQUALITY``, ``LIMIT_JOINT``, ``CONTACT_FRICTIONLESS``, ``CONTACT_PYRAMIDAL``
|
||||
- ``EQUALITY``, ``LIMIT_JOINT``, ``CONTACT_FRICTIONLESS``, ``CONTACT_PYRAMIDAL``, ``CONTACT_ELLIPTIC``
|
||||
* - :ref:`Equality <mjtEq>`
|
||||
- ``CONNECT``, ``WELD``, ``JOINT``
|
||||
* - :ref:`Integrator <mjtIntegrator>`
|
||||
- ``EULER``, ``RK4``
|
||||
* - :ref:`Cone <mjtCone>`
|
||||
- ``PYRAMIDAL``
|
||||
- ``PYRAMIDAL``, ``ELLIPTIC``
|
||||
* - :ref:`Condim <coContact>`
|
||||
- 1, 3, 4, 6
|
||||
* - :ref:`Solver <mjtSolver>`
|
||||
@@ -225,7 +225,7 @@ The following features are **in development** and coming soon:
|
||||
* - :ref:`Geom <mjtGeom>`
|
||||
- ``SDF``. Collisions between (``SPHERE``, ``BOX``, ``MESH``, ``HFIELD``) and ``CYLINDER``. Collisions between (``BOX``, ``MESH``, ``HFIELD``) and ``ELLIPSOID``.
|
||||
* - :ref:`Constraint <mjtConstraint>`
|
||||
- :ref:`Frictionloss <coFriction>`, ``CONTACT_ELLIPTIC``, ``FRICTION_DOF``
|
||||
- :ref:`Frictionloss <coFriction>`, ``FRICTION_DOF``
|
||||
* - :ref:`Integrator <mjtIntegrator>`
|
||||
- ``IMPLICIT``, ``IMPLICITFAST``
|
||||
* - Dynamics
|
||||
@@ -240,8 +240,6 @@ The following features are **in development** and coming soon:
|
||||
- ``MUSCLE``
|
||||
* - :ref:`Tendon Wrapping <mjtWrap>`
|
||||
- ``NONE``, ``JOINT``, ``PULLEY``, ``SITE``, ``SPHERE``, ``CYLINDER``
|
||||
* - :ref:`Cone <mjtCone>`
|
||||
- ``ELLIPTIC``
|
||||
* - Fluid Model
|
||||
- :ref:`flEllipsoid`
|
||||
* - :ref:`Tendons <tendon>`
|
||||
|
||||
@@ -239,7 +239,8 @@ def plane_convex(plane: GeomInfo, convex: ConvexInfo) -> Collision:
|
||||
plane_pos = convex.mat.T @ (plane.pos - convex.pos)
|
||||
n = convex.mat.T @ plane.mat[:, 2]
|
||||
support = (plane_pos - vert) @ n
|
||||
idx = _manifold_points(vert, support > jp.maximum(0, support.max() - 1e-4), n)
|
||||
# search for manifold points within a 1mm skin depth
|
||||
idx = _manifold_points(vert, support > jp.maximum(0, support.max() - 1e-3), n)
|
||||
pos = vert[idx]
|
||||
|
||||
# convert to world frame
|
||||
@@ -970,6 +971,7 @@ def _box_box(b1: ConvexInfo, b2: ConvexInfo) -> Collision:
|
||||
# Go back to world frame.
|
||||
pos = b2.pos + pos @ b2.mat.T
|
||||
n = normal @ b2.mat.T
|
||||
dist = jp.where(jp.isinf(dist), jp.finfo(float).max, dist)
|
||||
|
||||
return dist, pos, n
|
||||
|
||||
@@ -1029,6 +1031,7 @@ def _convex_convex(c1: ConvexInfo, c2: ConvexInfo) -> Collision:
|
||||
pos = c2.pos + pos @ c2.mat.T
|
||||
n = normal @ c2.mat.T
|
||||
n = -n if swapped else n
|
||||
dist = jp.where(jp.isinf(dist), jp.finfo(float).max, dist)
|
||||
|
||||
return dist, pos, n
|
||||
|
||||
|
||||
+240
-187
@@ -24,6 +24,7 @@ from mujoco.mjx._src import math
|
||||
from mujoco.mjx._src import support
|
||||
# pylint: disable=g-importing-member
|
||||
from mujoco.mjx._src.dataclasses import PyTreeNode
|
||||
from mujoco.mjx._src.types import ConeType
|
||||
from mujoco.mjx._src.types import ConstraintType
|
||||
from mujoco.mjx._src.types import Contact
|
||||
from mujoco.mjx._src.types import Data
|
||||
@@ -35,18 +36,14 @@ from mujoco.mjx._src.types import Model
|
||||
import numpy as np
|
||||
|
||||
|
||||
_CONDIM_EFC_COUNT = {1: 1, 3: 4, 4: 6, 6: 10}
|
||||
|
||||
|
||||
class _Efc(PyTreeNode):
|
||||
"""Support data for creating constraint matrices."""
|
||||
J: jax.Array
|
||||
pos: jax.Array
|
||||
pos_norm: jax.Array
|
||||
pos_aref: jax.Array
|
||||
pos_imp: jax.Array
|
||||
invweight: jax.Array
|
||||
solref: jax.Array
|
||||
solimp: jax.Array
|
||||
frictionloss: jax.Array
|
||||
|
||||
|
||||
def _kbi(
|
||||
@@ -59,13 +56,13 @@ def _kbi(
|
||||
timeconst, dampratio = solref
|
||||
|
||||
if not m.opt.disableflags & DisableBit.REFSAFE:
|
||||
timeconst = jp.maximum(timeconst, 2 * m.opt.timestep) * (timeconst > 0)
|
||||
timeconst = jp.maximum(timeconst, 2 * m.opt.timestep)
|
||||
|
||||
dmin, dmax, width, mid, power = solimp
|
||||
|
||||
dmin = jp.clip(dmin, mujoco.mjMINIMP, mujoco.mjMAXIMP)
|
||||
dmax = jp.clip(dmax, mujoco.mjMINIMP, mujoco.mjMAXIMP)
|
||||
width = jp.maximum(0, width)
|
||||
width = jp.maximum(mujoco.mjMINVAL, width)
|
||||
mid = jp.clip(mid, mujoco.mjMINIMP, mujoco.mjMAXIMP)
|
||||
power = jp.maximum(1, power)
|
||||
|
||||
@@ -73,8 +70,8 @@ def _kbi(
|
||||
k = 1 / (dmax * dmax * timeconst * timeconst * dampratio * dampratio)
|
||||
b = 2 / (dmax * timeconst)
|
||||
# TODO(robotics-simulation): check various solparam settings in model gen test
|
||||
k = jp.where(dampratio <= 0, -solref[0] / (dmax * dmax), k)
|
||||
b = jp.where(timeconst <= 0, -solref[1] / dmax, b)
|
||||
k = jp.where(solref[0] <= 0, -solref[0] / (dmax * dmax), k)
|
||||
b = jp.where(solref[1] <= 0, -solref[1] / dmax, b)
|
||||
|
||||
imp_x = jp.abs(pos) / width
|
||||
imp_a = (1.0 / jp.power(mid, power - 1)) * jp.power(imp_x, power)
|
||||
@@ -87,254 +84,280 @@ def _kbi(
|
||||
return k, b, imp # corresponds to K, B, I of efc_KBIP
|
||||
|
||||
|
||||
def _instantiate_equality_connect(m: Model, d: Data) -> Optional[_Efc]:
|
||||
def _row(j: jax.Array, *args) -> _Efc:
|
||||
"""Creates an efc row, ensuring args all have same row count."""
|
||||
if len(j.shape) < 2:
|
||||
return _Efc(j, *args) # if j isn't batched, ignore
|
||||
|
||||
args = list(args)
|
||||
for i, arg in enumerate(args):
|
||||
if not arg.shape or arg.shape[0] != j.shape[0]:
|
||||
args[i] = jp.tile(arg, (j.shape[0],) + (1,) * (len(arg.shape)))
|
||||
return _Efc(j, *args)
|
||||
|
||||
|
||||
def _efc_equality_connect(m: Model, d: Data) -> Optional[_Efc]:
|
||||
"""Calculates constraint rows for connect equality constraints."""
|
||||
|
||||
ids = np.nonzero(m.eq_type == EqType.CONNECT)[0]
|
||||
|
||||
if (m.opt.disableflags & DisableBit.EQUALITY) or ids.size == 0:
|
||||
eq_id = np.nonzero(m.eq_type == EqType.CONNECT)[0]
|
||||
if (m.opt.disableflags & DisableBit.EQUALITY) or eq_id.size == 0:
|
||||
return None
|
||||
|
||||
id1, id2, data = m.eq_obj1id[ids], m.eq_obj2id[ids], m.eq_data[ids]
|
||||
|
||||
@jax.vmap
|
||||
def fn(data, id1, id2):
|
||||
def rows(obj1id, obj2id, data, solref, solimp):
|
||||
anchor1, anchor2 = data[0:3], data[3:6]
|
||||
# find global points
|
||||
pos1 = d.xmat[id1] @ anchor1 + d.xpos[id1]
|
||||
pos2 = d.xmat[id2] @ anchor2 + d.xpos[id2]
|
||||
|
||||
# compute position error
|
||||
cpos = pos1 - pos2
|
||||
# error is difference in global positions
|
||||
pos1 = d.xmat[obj1id] @ anchor1 + d.xpos[obj1id]
|
||||
pos2 = d.xmat[obj2id] @ anchor2 + d.xpos[obj2id]
|
||||
pos = pos1 - pos2
|
||||
|
||||
# compute Jacobian difference (opposite of contact: 0 - 1)
|
||||
jacp1, _ = support.jac(m, d, pos1, id1)
|
||||
jacp2, _ = support.jac(m, d, pos2, id2)
|
||||
jacp1, _ = support.jac(m, d, pos1, obj1id)
|
||||
jacp2, _ = support.jac(m, d, pos2, obj2id)
|
||||
j = (jacp1 - jacp2).T
|
||||
pos_imp = math.norm(pos)
|
||||
invweight = m.body_invweight0[obj1id, 0] + m.body_invweight0[obj2id, 0]
|
||||
|
||||
return j, cpos, jp.repeat(math.norm(cpos), 3)
|
||||
return _row(j, pos, pos_imp, invweight, solref, solimp)
|
||||
|
||||
# concatenate to drop connect grouping dimension
|
||||
j, pos, pos_norm = jax.tree_util.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)
|
||||
|
||||
return _Efc(j, pos, pos_norm, invweight, solref, solimp, frictionloss)
|
||||
args = (m.eq_obj1id, m.eq_obj2id, m.eq_data, m.eq_solref, m.eq_solimp)
|
||||
args = jax.tree_util.tree_map(lambda x: x[eq_id], args)
|
||||
# concatenate to drop row grouping
|
||||
return jax.tree_util.tree_map(jp.concatenate, rows(*args))
|
||||
|
||||
|
||||
def _instantiate_equality_weld(m: Model, d: Data) -> Optional[_Efc]:
|
||||
def _efc_equality_weld(m: Model, d: Data) -> Optional[_Efc]:
|
||||
"""Calculates constraint rows for weld equality constraints."""
|
||||
|
||||
ids = np.nonzero(m.eq_type == EqType.WELD)[0]
|
||||
|
||||
if (m.opt.disableflags & DisableBit.EQUALITY) or ids.size == 0:
|
||||
eq_id = np.nonzero(m.eq_type == EqType.WELD)[0]
|
||||
if (m.opt.disableflags & DisableBit.EQUALITY) or eq_id.size == 0:
|
||||
return None
|
||||
|
||||
id1, id2, data = m.eq_obj1id[ids], m.eq_obj2id[ids], m.eq_data[ids]
|
||||
|
||||
@jax.vmap
|
||||
def fn(data, id1, id2):
|
||||
def rows(obj1id, obj2id, data, solref, solimp):
|
||||
anchor1, anchor2 = data[0:3], data[3:6]
|
||||
relpose, torquescale = data[6:10], data[10]
|
||||
|
||||
# find global points
|
||||
pos1 = d.xmat[id1] @ anchor2 + d.xpos[id1]
|
||||
pos2 = d.xmat[id2] @ anchor1 + d.xpos[id2]
|
||||
|
||||
# compute position error
|
||||
# error is difference in global position and orientation
|
||||
pos1 = d.xmat[obj1id] @ anchor2 + d.xpos[obj1id]
|
||||
pos2 = d.xmat[obj2id] @ anchor1 + d.xpos[obj2id]
|
||||
cpos = pos1 - pos2
|
||||
|
||||
# compute Jacobian difference (opposite of contact: 0 - 1)
|
||||
jacp1, jacr1 = support.jac(m, d, pos1, id1)
|
||||
jacp2, jacr2 = support.jac(m, d, pos2, id2)
|
||||
jacp1, jacr1 = support.jac(m, d, pos1, obj1id)
|
||||
jacp2, jacr2 = support.jac(m, d, pos2, obj2id)
|
||||
jacdifp = jacp1 - jacp2
|
||||
jacdifr = (jacr1 - jacr2) * torquescale
|
||||
|
||||
# compute orientation error: neg(q1) * q0 * relpose (axis components only)
|
||||
quat = math.quat_mul(d.xquat[id1], relpose)
|
||||
quat1 = math.quat_inv(d.xquat[id2])
|
||||
quat = math.quat_mul(d.xquat[obj1id], relpose)
|
||||
quat1 = math.quat_inv(d.xquat[obj2id])
|
||||
crot = math.quat_mul(quat1, quat)[1:] # copy axis components
|
||||
|
||||
pos = jp.concatenate((cpos, crot * torquescale))
|
||||
|
||||
# correct rotation Jacobian: 0.5 * neg(q1) * (jac0-jac1) * q0 * relpose
|
||||
jac_fn = lambda j: math.quat_mul(math.quat_mul_axis(quat1, j), quat)[1:]
|
||||
jacdifr = 0.5 * jax.vmap(jac_fn)(jacdifr)
|
||||
|
||||
j = jp.concatenate((jacdifp.T, jacdifr.T))
|
||||
pos = jp.concatenate((cpos, crot * torquescale))
|
||||
pos_imp = math.norm(pos)
|
||||
invweight = m.body_invweight0[obj1id] + m.body_invweight0[obj2id]
|
||||
invweight = jp.repeat(invweight, 3, axis=0)
|
||||
|
||||
return j, pos, jp.repeat(math.norm(pos), 6)
|
||||
return _row(j, pos, pos_imp, invweight, solref, solimp)
|
||||
|
||||
# concatenate to drop weld grouping dimension
|
||||
j, pos, pos_norm = jax.tree_util.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)
|
||||
|
||||
return _Efc(j, pos, pos_norm, invweight, solref, solimp, frictionloss)
|
||||
args = (m.eq_obj1id, m.eq_obj2id, m.eq_data, m.eq_solref, m.eq_solimp)
|
||||
args = jax.tree_util.tree_map(lambda x: x[eq_id], args)
|
||||
# concatenate to drop row grouping
|
||||
return jax.tree_util.tree_map(jp.concatenate, rows(*args))
|
||||
|
||||
|
||||
def _instantiate_equality_joint(m: Model, d: Data) -> Optional[_Efc]:
|
||||
def _efc_equality_joint(m: Model, d: Data) -> Optional[_Efc]:
|
||||
"""Calculates constraint rows for joint equality constraints."""
|
||||
|
||||
ids = np.nonzero(m.eq_type == EqType.JOINT)[0]
|
||||
eq_id = np.nonzero(m.eq_type == EqType.JOINT)[0]
|
||||
|
||||
if (m.opt.disableflags & DisableBit.EQUALITY) or ids.size == 0:
|
||||
if (m.opt.disableflags & DisableBit.EQUALITY) or eq_id.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):
|
||||
def rows(obj2id, data, solref, solimp, 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 = (pos2 - ref2) * (obj2id > -1)
|
||||
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
|
||||
deriv = jp.dot(data[1:5], dif_power[:4] * jp.arange(1, 5)) * (obj2id > -1)
|
||||
|
||||
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)
|
||||
j = jp.zeros((m.nv)).at[dofadr2].set(-deriv).at[dofadr1].set(1.0)
|
||||
invweight = m.dof_invweight0[dofadr1]
|
||||
invweight += m.dof_invweight0[dofadr2] * (obj2id > -1)
|
||||
|
||||
return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss)
|
||||
return _row(j, pos, pos, invweight, solref, solimp)
|
||||
|
||||
args = (m.eq_obj1id, m.eq_obj2id, m.eq_data, m.eq_solref, m.eq_solimp)
|
||||
args = jax.tree_util.tree_map(lambda x: x[eq_id], args)
|
||||
dofadr1, dofadr2 = m.jnt_dofadr[args[0]], m.jnt_dofadr[args[1]]
|
||||
qposadr1, qposadr2 = m.jnt_qposadr[args[0]], m.jnt_qposadr[args[1]]
|
||||
args = args[1:] + (dofadr1, dofadr2, qposadr1, qposadr2)
|
||||
|
||||
return rows(*args)
|
||||
|
||||
|
||||
def _instantiate_friction(m: Model, d: Data) -> Optional[_Efc]:
|
||||
def _efc_friction(m: Model, d: Data) -> Optional[_Efc]:
|
||||
# TODO(robotics-team): implement _instantiate_friction
|
||||
del m, d
|
||||
return None
|
||||
|
||||
|
||||
def _instantiate_limit_ball(m: Model, d: Data) -> Optional[_Efc]:
|
||||
def _efc_limit_ball(m: Model, d: Data) -> Optional[_Efc]:
|
||||
"""Calculates constraint rows for ball joint limits."""
|
||||
|
||||
ids = np.nonzero((m.jnt_type == JointType.BALL) & m.jnt_limited)[0]
|
||||
jnt_id = np.nonzero((m.jnt_type == JointType.BALL) & m.jnt_limited)[0]
|
||||
|
||||
if (m.opt.disableflags & DisableBit.LIMIT) or ids.size == 0:
|
||||
if (m.opt.disableflags & DisableBit.LIMIT) or jnt_id.size == 0:
|
||||
return None
|
||||
|
||||
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]])
|
||||
|
||||
@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)
|
||||
def rows(qposadr, dofadr, jnt_range, jnt_margin, solref, solimp):
|
||||
axis, angle = math.quat_to_axis_angle(d.qpos[jp.arange(4) + qposadr])
|
||||
pos = jp.amax(jnt_range) - angle - jnt_margin
|
||||
active = pos < 0
|
||||
return j * active, pos * active
|
||||
j = jp.zeros(m.nv).at[jp.arange(3) + dofadr].set(-axis)
|
||||
invweight = m.dof_invweight0[dofadr]
|
||||
|
||||
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 _row(j * active, pos * active, pos, invweight, solref, solimp)
|
||||
|
||||
return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss)
|
||||
args = (m.jnt_qposadr, m.jnt_dofadr, m.jnt_range, m.jnt_margin, m.jnt_solref)
|
||||
args += (m.jnt_solimp,)
|
||||
args = jax.tree_util.tree_map(lambda x: x[jnt_id], args)
|
||||
|
||||
return rows(*args)
|
||||
|
||||
|
||||
def _instantiate_limit_slide_hinge(m: Model, d: Data) -> Optional[_Efc]:
|
||||
def _efc_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]
|
||||
jnt_id = np.nonzero(slide_hinge & m.jnt_limited)[0]
|
||||
|
||||
if (m.opt.disableflags & DisableBit.LIMIT) or ids.size == 0:
|
||||
if (m.opt.disableflags & DisableBit.LIMIT) or jnt_id.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)
|
||||
def rows(qposadr, dofadr, jnt_range, jnt_margin, solref, solimp):
|
||||
qpos = d.qpos[qposadr]
|
||||
dist_min, dist_max = qpos - jnt_range[0], jnt_range[1] - qpos
|
||||
pos = jp.minimum(dist_min, dist_max) - jnt_margin
|
||||
active = pos < 0
|
||||
return j * active, pos * active
|
||||
j = jp.zeros(m.nv).at[dofadr].set((dist_min < dist_max) * 2 - 1)
|
||||
invweight = m.dof_invweight0[dofadr]
|
||||
|
||||
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 _row(j * active, pos * active, pos, invweight, solref, solimp)
|
||||
|
||||
return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss)
|
||||
args = (m.jnt_qposadr, m.jnt_dofadr, m.jnt_range, m.jnt_margin, m.jnt_solref)
|
||||
args += (m.jnt_solimp,)
|
||||
args = jax.tree_util.tree_map(lambda x: x[jnt_id], args)
|
||||
|
||||
return rows(*args)
|
||||
|
||||
|
||||
def _instantiate_contact(m: Model, d: Data) -> Optional[_Efc]:
|
||||
"""Calculates constraint rows for contacts."""
|
||||
def _efc_contact_frictionless(m: Model, d: Data) -> Optional[_Efc]:
|
||||
"""Calculates constraint rows for frictionless contacts."""
|
||||
|
||||
if d.ncon == 0:
|
||||
con_id = np.nonzero(d.contact.dim == 1)[0]
|
||||
|
||||
if con_id.size == 0:
|
||||
return None
|
||||
|
||||
def contact_efc(c: Contact, condim: int):
|
||||
@jax.vmap
|
||||
def rows(c: Contact):
|
||||
pos = c.dist - c.includemargin
|
||||
active = pos < 0
|
||||
body1, body2 = jp.array(m.geom_bodyid)[c.geom]
|
||||
jac1p, _ = support.jac(m, d, c.pos, body1)
|
||||
jac2p, _ = support.jac(m, d, c.pos, body2)
|
||||
j = (c.frame @ (jac2p - jac1p).T)[0]
|
||||
invweight = m.body_invweight0[body1, 0] + m.body_invweight0[body2, 0]
|
||||
|
||||
@jax.vmap
|
||||
def fn(c: Contact):
|
||||
dist = c.dist - c.includemargin
|
||||
active = dist < 0
|
||||
body1, body2 = jp.array(m.geom_bodyid)[c.geom]
|
||||
jac1p, jac1r = support.jac(m, d, c.pos, body1)
|
||||
jac2p, jac2r = support.jac(m, d, c.pos, body2)
|
||||
diff = c.frame @ (jac2p - jac1p).T
|
||||
if condim > 3: # only calculate rotational diff if needed
|
||||
diff = jp.concatenate((diff, c.frame @ (jac2r - jac1r).T), axis=0)
|
||||
tran = m.body_invweight0[body1, 0] + m.body_invweight0[body2, 0]
|
||||
return _row(j * active, pos * active, pos, invweight, c.solref, c.solimp)
|
||||
|
||||
if condim == 1:
|
||||
return diff[0] * active, tran, dist * active, c.solref, c.solimp
|
||||
contact = jax.tree_util.tree_map(lambda x: x[con_id], d.contact)
|
||||
|
||||
# a pair of opposing pyramid edges per friction dimension
|
||||
# repeat friction directions with positive and negative sign
|
||||
fri = jp.repeat(c.friction[: condim - 1], 2, axis=0).at[1::2].mul(-1)
|
||||
# repeat condims of jacdiff to match +/- friction directions
|
||||
j = diff[0] + jp.repeat(diff[1:condim], 2, axis=0) * fri[:, None]
|
||||
# pyramidal has common invweight across all edges
|
||||
diag_approx = tran + fri[0] * fri[0] * tran
|
||||
inv_w = diag_approx * 2 * fri[0] * fri[0] / m.opt.impratio
|
||||
repeat_fn = lambda x: jp.repeat(x[None], (condim - 1) * 2, axis=0)
|
||||
inv_w, pos, solref, solimp = jax.tree_util.tree_map(
|
||||
repeat_fn, (inv_w, dist, c.solref, c.solimp)
|
||||
)
|
||||
return j * active, inv_w, pos * active, solref, solimp
|
||||
return rows(contact)
|
||||
|
||||
return fn(c)
|
||||
|
||||
# group efc calculations by condim
|
||||
dims, begs = np.unique(d.contact.dim, return_index=True)
|
||||
efcs = []
|
||||
for i in range(len(dims)):
|
||||
dim, beg = dims[i], begs[i]
|
||||
end = begs[i + 1] if i < len(dims) - 1 else None
|
||||
c = jax.tree_util.tree_map(lambda x, b=beg, e=end: x[b:e], d.contact)
|
||||
efc = contact_efc(c, dim)
|
||||
if dim > 1:
|
||||
# remove efc grouping dimension
|
||||
efc = jax.tree_util.tree_map(jp.concatenate, efc)
|
||||
efcs.append(efc)
|
||||
def _efc_contact_pyramidal(m: Model, d: Data, condim: int) -> Optional[_Efc]:
|
||||
"""Calculates constraint rows for frictional pyramidal contacts."""
|
||||
|
||||
efc = jax.tree_util.tree_map(lambda *x: jp.concatenate(x), *efcs)
|
||||
j, invweight, pos, solref, solimp = efc
|
||||
frictionloss = jp.zeros_like(pos)
|
||||
con_id = np.nonzero(d.contact.dim == condim)[0]
|
||||
|
||||
return _Efc(j, pos, pos, invweight, solref, solimp, frictionloss)
|
||||
if con_id.size == 0:
|
||||
return None
|
||||
|
||||
@jax.vmap
|
||||
def rows(c: Contact):
|
||||
pos = c.dist - c.includemargin
|
||||
active = pos < 0
|
||||
body1, body2 = jp.array(m.geom_bodyid)[c.geom]
|
||||
jac1p, jac1r = support.jac(m, d, c.pos, body1)
|
||||
jac2p, jac2r = support.jac(m, d, c.pos, body2)
|
||||
diff = c.frame @ (jac2p - jac1p).T
|
||||
if condim > 3:
|
||||
diff = jp.concatenate((diff, (c.frame @ (jac2r - jac1r).T)), axis=0)
|
||||
# a pair of opposing pyramid edges per friction dimension
|
||||
# repeat friction directions with positive and negative sign
|
||||
fri = jp.repeat(c.friction[: condim - 1], 2, axis=0).at[1::2].mul(-1)
|
||||
# repeat condims of jacdiff to match +/- friction directions
|
||||
j = diff[0] + jp.repeat(diff[1:condim], 2, axis=0) * fri[:, None]
|
||||
|
||||
# pyramidal has common invweight across all edges
|
||||
invweight = m.body_invweight0[body1, 0] + m.body_invweight0[body2, 0]
|
||||
invweight = invweight + fri[0] * fri[0] * invweight
|
||||
invweight = invweight * 2 * fri[0] * fri[0] / m.opt.impratio
|
||||
|
||||
return _row(j * active, pos * active, pos, invweight, c.solref, c.solimp)
|
||||
|
||||
contact = jax.tree_util.tree_map(lambda x: x[con_id], d.contact)
|
||||
# concatenate to drop row grouping
|
||||
return jax.tree_util.tree_map(jp.concatenate, rows(contact))
|
||||
|
||||
|
||||
def _efc_contact_elliptic(m: Model, d: Data, condim: int) -> Optional[_Efc]:
|
||||
"""Calculates constraint rows for frictional elliptic contacts."""
|
||||
|
||||
con_id = np.nonzero(d.contact.dim == condim)[0]
|
||||
|
||||
if con_id.size == 0:
|
||||
return None
|
||||
|
||||
@jax.vmap
|
||||
def rows(c: Contact):
|
||||
pos = c.dist - c.includemargin
|
||||
active = pos < 0
|
||||
obj1id, obj2id = jp.array(m.geom_bodyid)[c.geom]
|
||||
jac1p, jac1r = support.jac(m, d, c.pos, obj1id)
|
||||
jac2p, jac2r = support.jac(m, d, c.pos, obj2id)
|
||||
j = c.frame @ (jac2p - jac1p).T
|
||||
if condim > 3:
|
||||
j = jp.concatenate((j, (c.frame @ (jac2r - jac1r).T)[: condim - 3]))
|
||||
invweight = m.body_invweight0[obj1id, 0] + m.body_invweight0[obj2id, 0]
|
||||
|
||||
# normal row comes from solref, remaining rows from solreffriction
|
||||
solreffriction = c.solreffriction + c.solref * ~c.solreffriction.any()
|
||||
solreffriction = jp.tile(solreffriction, (condim - 1, 1))
|
||||
solref = jp.concatenate((c.solref[None], solreffriction))
|
||||
fri = jp.square(c.friction[0]) / jp.square(c.friction[1 : condim - 1])
|
||||
invweight = jp.array([invweight, invweight / m.opt.impratio])
|
||||
invweight = jp.concatenate((invweight, invweight[1] * fri))
|
||||
pos_aref = jp.zeros(condim).at[0].set(pos)
|
||||
|
||||
return _row(j * active, pos_aref * active, pos, invweight, solref, c.solimp)
|
||||
|
||||
contact = jax.tree_util.tree_map(lambda x: x[con_id], d.contact)
|
||||
# concatenate to drop row grouping
|
||||
return jax.tree_util.tree_map(jp.concatenate, rows(contact))
|
||||
|
||||
|
||||
def counts(efc_type: np.ndarray) -> Tuple[int, int, int, int]:
|
||||
@@ -344,7 +367,8 @@ def counts(efc_type: np.ndarray) -> Tuple[int, int, int, int]:
|
||||
nl = (efc_type == ConstraintType.LIMIT_JOINT).sum()
|
||||
nc_f = (efc_type == ConstraintType.CONTACT_FRICTIONLESS).sum()
|
||||
nc_p = (efc_type == ConstraintType.CONTACT_PYRAMIDAL).sum()
|
||||
nc = nc_f + nc_p
|
||||
nc_e = (efc_type == ConstraintType.CONTACT_ELLIPTIC).sum()
|
||||
nc = nc_f + nc_p + nc_e
|
||||
|
||||
return ne, nf, nl, nc
|
||||
|
||||
@@ -363,25 +387,48 @@ def make_efc_type(
|
||||
num_rows = (m.eq_type == EqType.CONNECT).sum() * 3
|
||||
num_rows += (m.eq_type == EqType.WELD).sum() * 6
|
||||
num_rows += (m.eq_type == EqType.JOINT).sum()
|
||||
efc_types.extend([ConstraintType.EQUALITY] * num_rows)
|
||||
efc_types += [ConstraintType.EQUALITY] * num_rows
|
||||
|
||||
if not m.opt.disableflags & DisableBit.LIMIT:
|
||||
efc_types.extend([ConstraintType.LIMIT_JOINT] * m.jnt_limited.sum())
|
||||
efc_types += [ConstraintType.LIMIT_JOINT] * m.jnt_limited.sum()
|
||||
|
||||
if not m.opt.disableflags & DisableBit.CONTACT:
|
||||
num_rows = sum(_CONDIM_EFC_COUNT[d] for d in dim)
|
||||
efc_types.extend([ConstraintType.CONTACT_PYRAMIDAL] * num_rows)
|
||||
for condim in (1, 3, 4, 6):
|
||||
n = (dim == condim).sum()
|
||||
if condim == 1:
|
||||
efc_types += [ConstraintType.CONTACT_FRICTIONLESS] * n
|
||||
elif m.opt.cone == ConeType.PYRAMIDAL:
|
||||
efc_types += [ConstraintType.CONTACT_PYRAMIDAL] * (condim - 1) * 2 * n
|
||||
elif m.opt.cone == ConeType.ELLIPTIC:
|
||||
efc_types += [ConstraintType.CONTACT_ELLIPTIC] * condim * n
|
||||
else:
|
||||
raise ValueError(f'Unknown cone: {m.opt.cone}')
|
||||
|
||||
return np.array(efc_types)
|
||||
|
||||
|
||||
def make_efc_address(efc_type: np.ndarray, dim: np.ndarray) -> np.ndarray:
|
||||
def make_efc_address(
|
||||
m: Union[Model, mujoco.MjModel], dim: np.ndarray, efc_type: np.ndarray
|
||||
) -> np.ndarray:
|
||||
"""Returns efc_address that maps contacts to constraint row address."""
|
||||
nc = (efc_type == ConstraintType.CONTACT_PYRAMIDAL).sum()
|
||||
nc_start = efc_type.size - nc
|
||||
offsets = np.cumsum([0] + [_CONDIM_EFC_COUNT[d] for d in dim])[:-1]
|
||||
offsets = np.array([0], dtype=int)
|
||||
for condim in (1, 3, 4, 6):
|
||||
n = (dim == condim).sum()
|
||||
if n == 0:
|
||||
continue
|
||||
if condim == 1:
|
||||
offsets = np.concatenate((offsets, [1] * n))
|
||||
elif m.opt.cone == ConeType.PYRAMIDAL:
|
||||
offsets = np.concatenate((offsets, [(condim - 1) * 2] * n))
|
||||
elif m.opt.cone == ConeType.ELLIPTIC:
|
||||
offsets = np.concatenate((offsets, [condim] * n))
|
||||
else:
|
||||
raise ValueError(f'Unknown cone: {m.opt.cone}')
|
||||
|
||||
return nc_start + offsets
|
||||
_, _, _, nc = counts(efc_type)
|
||||
address = efc_type.size - nc + np.cumsum(offsets)[:-1]
|
||||
|
||||
return address
|
||||
|
||||
|
||||
def make_constraint(m: Model, d: Data) -> Data:
|
||||
@@ -390,15 +437,21 @@ def make_constraint(m: Model, d: Data) -> Data:
|
||||
if m.opt.disableflags & DisableBit.CONSTRAINT:
|
||||
efcs = ()
|
||||
else:
|
||||
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_ball(m, d),
|
||||
_instantiate_limit_slide_hinge(m, d),
|
||||
_instantiate_contact(m, d),
|
||||
) if efc is not None)
|
||||
efcs = (
|
||||
_efc_equality_connect(m, d),
|
||||
_efc_equality_weld(m, d),
|
||||
_efc_equality_joint(m, d),
|
||||
_efc_friction(m, d),
|
||||
_efc_limit_ball(m, d),
|
||||
_efc_limit_slide_hinge(m, d),
|
||||
_efc_contact_frictionless(m, d),
|
||||
)
|
||||
if m.opt.cone == ConeType.ELLIPTIC:
|
||||
con_fn = _efc_contact_elliptic
|
||||
else:
|
||||
con_fn = _efc_contact_pyramidal
|
||||
efcs += tuple(con_fn(m, d, dim) for dim in (3, 4, 6))
|
||||
efcs = tuple(efc for efc in efcs if efc is not None)
|
||||
|
||||
if not efcs:
|
||||
z = jp.empty(0)
|
||||
@@ -410,13 +463,13 @@ def make_constraint(m: Model, d: Data) -> Data:
|
||||
|
||||
@jax.vmap
|
||||
def fn(efc):
|
||||
k, b, imp = _kbi(m, efc.solref, efc.solimp, efc.pos_norm)
|
||||
k, b, imp = _kbi(m, efc.solref, efc.solimp, efc.pos_imp)
|
||||
r = jp.maximum(efc.invweight * (1 - imp) / imp, mujoco.mjMINVAL)
|
||||
aref = -b * (efc.J @ d.qvel) - k * imp * efc.pos
|
||||
aref = -b * (efc.J @ d.qvel) - k * imp * efc.pos_aref
|
||||
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)
|
||||
d = d.replace(efc_frictionloss=jp.zeros_like(r))
|
||||
|
||||
return d
|
||||
|
||||
@@ -15,6 +15,7 @@
|
||||
"""Tests for constraint functions."""
|
||||
|
||||
from absl.testing import absltest
|
||||
from absl.testing import parameterized
|
||||
from jax import numpy as jp
|
||||
import mujoco
|
||||
from mujoco import mjx
|
||||
@@ -38,41 +39,32 @@ def _assert_attr_eq(a, b, attr):
|
||||
_assert_eq(getattr(a, attr), getattr(b, attr), attr)
|
||||
|
||||
|
||||
class ConstraintTest(absltest.TestCase):
|
||||
class ConstraintTest(parameterized.TestCase):
|
||||
|
||||
def test_constraints(self):
|
||||
@parameterized.parameters(
|
||||
mujoco.mjtCone.mjCONE_PYRAMIDAL, mujoco.mjtCone.mjCONE_ELLIPTIC
|
||||
)
|
||||
def test_constraints(self, cone):
|
||||
"""Test constraints."""
|
||||
m = test_util.load_test_file('constraints.xml')
|
||||
m.opt.cone = cone
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_step(m, d, 100) # at 100 steps mix of active/inactive constraints
|
||||
mujoco.mj_forward(m, d)
|
||||
mx = mjx.put_model(m)
|
||||
dx = mjx.put_data(m, d)
|
||||
dx = mjx.make_constraint(mx, dx)
|
||||
d_efc_j = d.efc_J.reshape((-1, m.nv))
|
||||
|
||||
# ne, nf, nl order matches
|
||||
efl = d.ne + d.nf + d.nl
|
||||
_assert_eq(d_efc_j[:efl], dx.efc_J[:efl], 'efc_J')
|
||||
_assert_eq(d.efc_D[:efl], dx.efc_D[:efl], 'efc_D')
|
||||
_assert_eq(d.efc_aref[:efl], dx.efc_aref[:efl], 'efc_aref')
|
||||
_assert_eq(dx.efc_frictionloss, 0, 'efc_frictionloss')
|
||||
# sample a mix of active/inactive constraints at different timesteps
|
||||
for key in range(3):
|
||||
mujoco.mj_resetDataKeyframe(m, d, key)
|
||||
mujoco.mj_forward(m, d)
|
||||
mx = mjx.put_model(m)
|
||||
dx = mjx.put_data(m, d)
|
||||
dx = mjx.make_constraint(mx, dx)
|
||||
|
||||
# contact order might not match, so check efcs contact by contact
|
||||
for i in range(d.ncon):
|
||||
geom_match = (dx.contact.geom == d.contact.geom[i]).all(axis=-1)
|
||||
geom_match &= (dx.contact.pos == d.contact.pos[i]).all(axis=-1)
|
||||
self.assertTrue(geom_match.any(), f'contact {i} not found in MJX contact')
|
||||
j = np.nonzero(geom_match)[0][0]
|
||||
self.assertEqual(d.contact.dim[i], dx.contact.dim[j])
|
||||
nc = max(1, (d.contact.dim[i] - 1) * 2)
|
||||
d_beg, dx_beg = d.contact.efc_address[i], dx.contact.efc_address[j]
|
||||
d_end, dx_end = d_beg + nc, dx_beg + nc
|
||||
_assert_eq(d_efc_j[d_beg:d_end], dx.efc_J[dx_beg:dx_end], 'efc_J')
|
||||
_assert_eq(d.efc_D[d_beg:d_end], dx.efc_D[dx_beg:dx_end], 'efc_D')
|
||||
d_efc_aref = d.efc_aref[d_beg:d_end]
|
||||
dx_efc_aref = dx.efc_aref[dx_beg:dx_end]
|
||||
_assert_eq(d_efc_aref, dx_efc_aref, 'efc_aref')
|
||||
order = test_util.efc_order(m, d, dx)
|
||||
d_efc_j = d.efc_J.reshape((-1, m.nv))
|
||||
_assert_eq(d_efc_j, dx.efc_J[order][:d.nefc], 'efc_J')
|
||||
_assert_eq(0, dx.efc_J[order][d.nefc:], 'efc_J')
|
||||
_assert_eq(d.efc_aref, dx.efc_aref[order][:d.nefc], 'efc_aref')
|
||||
_assert_eq(0, dx.efc_aref[order][d.nefc:], 'efc_aref')
|
||||
_assert_eq(d.efc_D, dx.efc_D[order][:d.nefc], 'efc_D')
|
||||
|
||||
def test_disable_refsafe(self):
|
||||
m = test_util.load_test_file('constraints.xml')
|
||||
|
||||
+25
-11
@@ -15,7 +15,7 @@
|
||||
"""Functions to initialize, load, or save data."""
|
||||
|
||||
import copy
|
||||
from typing import List, Union
|
||||
from typing import List, Tuple, Union
|
||||
|
||||
import jax
|
||||
from jax import numpy as jp
|
||||
@@ -133,7 +133,7 @@ def make_data(m: Union[types.Model, mujoco.MjModel]) -> types.Data:
|
||||
"""Allocate and initialize Data."""
|
||||
dim = collision_driver.make_condim(m)
|
||||
efc_type = constraint.make_efc_type(m, dim)
|
||||
efc_address = constraint.make_efc_address(efc_type, dim)
|
||||
efc_address = constraint.make_efc_address(m, dim, efc_type)
|
||||
ne, nf, nl, nc = constraint.counts(efc_type)
|
||||
ncon, nefc = dim.size, ne + nf + nl + nc
|
||||
|
||||
@@ -330,7 +330,7 @@ def _make_contact(
|
||||
c: mujoco._structs._MjContactList,
|
||||
dim: np.ndarray,
|
||||
efc_address: np.ndarray,
|
||||
) -> types.Contact:
|
||||
) -> Tuple[types.Contact, np.ndarray]:
|
||||
"""Converts mujoco.structs._MjContactList into mjx.Contact."""
|
||||
fields = {f.name: getattr(c, f.name) for f in types.Contact.fields()}
|
||||
fields['frame'] = fields['frame'].reshape((-1, 3, 3))
|
||||
@@ -351,21 +351,21 @@ def _make_contact(
|
||||
zero = jax.tree_util.tree_map(
|
||||
lambda x: np.zeros((1,) + x.shape[1:], dtype=x.dtype), fields
|
||||
)
|
||||
zero['dist'][:] = np.finfo(float).max
|
||||
zero['dist'][:] = 1e10
|
||||
fields = jax.tree_util.tree_map(lambda *x: np.concatenate(x), fields, zero)
|
||||
fields = jax.tree_util.tree_map(lambda x: x[contact_map], fields)
|
||||
|
||||
fields['dim'] = dim
|
||||
fields['efc_address'] = efc_address
|
||||
|
||||
return types.Contact(**fields)
|
||||
return types.Contact(**fields), contact_map
|
||||
|
||||
|
||||
def put_data(m: mujoco.MjModel, d: mujoco.MjData, device=None) -> types.Data:
|
||||
"""Puts mujoco.MjData onto a device, resulting in mjx.Data."""
|
||||
dim = collision_driver.make_condim(m)
|
||||
efc_type = constraint.make_efc_type(m, dim)
|
||||
efc_address = constraint.make_efc_address(efc_type, dim)
|
||||
efc_address = constraint.make_efc_address(m, dim, efc_type)
|
||||
ne, nf, nl, nc = constraint.counts(efc_type)
|
||||
ncon, nefc = dim.size, ne + nf + nl + nc
|
||||
|
||||
@@ -388,6 +388,8 @@ def put_data(m: mujoco.MjModel, d: mujoco.MjData, device=None) -> types.Data:
|
||||
# MJX does not support islanding, so only transfer the first solver_niter
|
||||
fields['solver_niter'] = fields['solver_niter'][0]
|
||||
|
||||
contact, contact_map = _make_contact(d.contact, dim, efc_address)
|
||||
|
||||
# pad efc fields: MuJoCo efc arrays are sparse for inactive constraints.
|
||||
# efc_J is also optionally column-sparse (typically for large nv). MJX is
|
||||
# neither: it contains zeros for inactive constraints, and efc_J is always
|
||||
@@ -403,13 +405,25 @@ def put_data(m: mujoco.MjModel, d: mujoco.MjData, device=None) -> types.Data:
|
||||
else:
|
||||
fields['efc_J'] = fields['efc_J'].reshape((-1 if m.nv else 0, m.nv))
|
||||
|
||||
# move efc rows to their correct offsets
|
||||
for fname in ('efc_J', 'efc_frictionloss', 'efc_D', 'efc_aref', 'efc_force'):
|
||||
value = np.zeros((nefc, m.nv)) if fname == 'efc_J' else np.zeros(nefc)
|
||||
for i in range(4):
|
||||
value_beg = sum([ne, nf, nl][:i])
|
||||
d_beg = sum([d.ne, d.nf, d.nl][:i])
|
||||
size = [d.ne, d.nf, d.nl, d.nefc - d.nl - d.nf - d.ne][i]
|
||||
for i in range(3):
|
||||
value_beg = sum([ne, nf][:i])
|
||||
d_beg = sum([d.ne, d.nf][:i])
|
||||
size = [d.ne, d.nf, d.nl][i]
|
||||
value[value_beg : value_beg + size] = fields[fname][d_beg : d_beg + size]
|
||||
|
||||
# for nc, we may reorder contacts so they match MJX order: group by dim
|
||||
for id_to, id_from in enumerate(contact_map):
|
||||
if id_from == -1:
|
||||
continue
|
||||
num_rows = dim[id_to]
|
||||
if num_rows > 1 and m.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL:
|
||||
num_rows = (num_rows - 1) * 2
|
||||
efc_i, efc_o = d.contact.efc_address[id_from], efc_address[id_to]
|
||||
value[efc_o:efc_o + num_rows] = fields[fname][efc_i:efc_i + num_rows]
|
||||
|
||||
fields[fname] = value
|
||||
|
||||
# convert qM and qLD if jacobian is dense
|
||||
@@ -424,7 +438,7 @@ def put_data(m: mujoco.MjModel, d: mujoco.MjData, device=None) -> types.Data:
|
||||
fields['qLD'] = np.zeros((m.nv, m.nv))
|
||||
fields['qLDiagInv'] = np.zeros(0)
|
||||
|
||||
fields['contact'] = _make_contact(d.contact, dim, efc_address)
|
||||
fields['contact'] = contact
|
||||
fields.update(ne=ne, nf=nf, nl=nl, nefc=nefc, ncon=ncon, efc_type=efc_type)
|
||||
|
||||
# copy because device_put is async:
|
||||
|
||||
@@ -140,14 +140,6 @@ class ModelIOTest(parameterized.TestCase):
|
||||
)
|
||||
)
|
||||
|
||||
def test_cone_not_implemented(self):
|
||||
with self.assertRaises(NotImplementedError):
|
||||
mjx.put_model(
|
||||
mujoco.MjModel.from_xml_string(
|
||||
'<mujoco><option cone="elliptic"/><worldbody/></mujoco>'
|
||||
)
|
||||
)
|
||||
|
||||
def test_pgs_not_implemented(self):
|
||||
with self.assertRaises(NotImplementedError):
|
||||
mjx.put_model(
|
||||
@@ -299,7 +291,7 @@ class DataIOTest(parameterized.TestCase):
|
||||
self.assertEqual(dx.contact.dist.shape, (4,))
|
||||
self.assertEqual(d.ncon, 1) # however only 1 contact in this step
|
||||
np.testing.assert_allclose(dx.contact.dist[0], d.contact.dist[0])
|
||||
self.assertTrue(np.isinf(dx.contact.dist[1:]).all())
|
||||
self.assertTrue((dx.contact.dist[1:] > 0).all())
|
||||
self.assertEqual(dx.contact.frame.shape, (4, 3, 3))
|
||||
np.testing.assert_allclose(
|
||||
dx.contact.frame[0].reshape(9), d.contact.frame[0]
|
||||
|
||||
+214
-40
@@ -22,6 +22,7 @@ from mujoco.mjx._src import smooth
|
||||
from mujoco.mjx._src import support
|
||||
# pylint: disable=g-importing-member
|
||||
from mujoco.mjx._src.dataclasses import PyTreeNode
|
||||
from mujoco.mjx._src.types import ConeType
|
||||
from mujoco.mjx._src.types import Data
|
||||
from mujoco.mjx._src.types import DisableBit
|
||||
from mujoco.mjx._src.types import Model
|
||||
@@ -45,8 +46,12 @@ class _Context(PyTreeNode):
|
||||
cost: constraint + Gauss cost
|
||||
prev_cost: cost from previous iter
|
||||
solver_niter: number of solver iterations
|
||||
active: active (quadratic) constraints (nefc,)
|
||||
fri: friction of regularized cone (num(con.dim > 1), 6)
|
||||
dm: regularized constraint mass (num(con.dim > 1))
|
||||
u: friction cone (normal and tangents) (num(con.dim > 1), 6)
|
||||
h: cone hessian (num(con.dim > 1), 6, 6)
|
||||
"""
|
||||
|
||||
qacc: jax.Array
|
||||
qfrc_constraint: jax.Array
|
||||
Jaref: jax.Array # pylint: disable=invalid-name
|
||||
@@ -59,6 +64,11 @@ class _Context(PyTreeNode):
|
||||
cost: jax.Array
|
||||
prev_cost: jax.Array
|
||||
solver_niter: jax.Array
|
||||
active: jax.Array
|
||||
fri: jax.Array
|
||||
dm: jax.Array
|
||||
u: jax.Array
|
||||
h: jax.Array
|
||||
|
||||
@classmethod
|
||||
def create(cls, m: Model, d: Data, grad: bool = True) -> '_Context':
|
||||
@@ -66,6 +76,15 @@ class _Context(PyTreeNode):
|
||||
# TODO(robotics-team): determine nv at which sparse mul is faster
|
||||
ma = support.mul_m(m, d, d.qacc)
|
||||
nv_0 = jp.zeros(m.nv)
|
||||
fri = 0.0
|
||||
if m.opt.cone == ConeType.ELLIPTIC:
|
||||
friction = d.contact.friction[d.contact.dim > 1]
|
||||
dim = d.contact.dim[d.contact.dim > 1]
|
||||
mu = friction[:, 0] / jp.sqrt(m.opt.impratio)
|
||||
fri = jp.concatenate((mu[:, None], friction), axis=1)
|
||||
for condim in (3, 4, 6):
|
||||
fri = fri.at[dim == condim, condim:].set(0)
|
||||
|
||||
ctx = _Context(
|
||||
qacc=d.qacc,
|
||||
qfrc_constraint=d.qfrc_constraint,
|
||||
@@ -79,8 +98,13 @@ class _Context(PyTreeNode):
|
||||
cost=jp.inf,
|
||||
prev_cost=0.0,
|
||||
solver_niter=0,
|
||||
active=0.0,
|
||||
fri=fri,
|
||||
dm=0.0,
|
||||
u=0.0,
|
||||
h=0.0,
|
||||
)
|
||||
ctx = _update_constraint(d, ctx)
|
||||
ctx = _update_constraint(m, d, ctx)
|
||||
if grad:
|
||||
ctx = _update_gradient(m, d, ctx)
|
||||
ctx = ctx.replace(search=-ctx.Mgrad) # start with preconditioned gradient
|
||||
@@ -106,24 +130,68 @@ class _LSPoint(PyTreeNode):
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
m: Model,
|
||||
d: Data,
|
||||
ctx: _Context,
|
||||
alpha: jax.Array,
|
||||
jv: jax.Array,
|
||||
quad: jax.Array,
|
||||
quad_gauss: jax.Array,
|
||||
uu: jax.Array,
|
||||
v0: jax.Array,
|
||||
uv: jax.Array,
|
||||
vv: jax.Array,
|
||||
) -> '_LSPoint':
|
||||
"""Creates a linesearch point with first and second derivatives."""
|
||||
# roughly corresponds to CGEval in mujoco/src/engine/engine_solver.c
|
||||
|
||||
# TODO(robotics-team): change this to support friction constraints
|
||||
active = ((ctx.Jaref + alpha * jv) < 0).at[:d.ne + d.nf].set(True)
|
||||
quad = jax.vmap(jp.multiply)(quad, active) # only active
|
||||
quad_total = quad_gauss + jp.sum(quad, axis=0)
|
||||
cost, deriv_0, deriv_1 = 0.0, 0.0, 0.0
|
||||
quad_total = quad_gauss
|
||||
|
||||
if m.opt.cone == ConeType.ELLIPTIC:
|
||||
mu, u0 = ctx.fri[:, 0], ctx.u[:, 0]
|
||||
n = u0 + alpha * v0
|
||||
tsqr = uu + alpha * (2 * uv + alpha * vv)
|
||||
t = jp.sqrt(tsqr) # tangential force
|
||||
|
||||
bottom_zone = ((tsqr <= 0) & (n < 0)) | ((tsqr > 0) & ((mu * n + t) <= 0))
|
||||
middle_zone = (tsqr > 0) & (n < (mu * t)) & ((mu * n + t) > 0)
|
||||
|
||||
# quadratic cost for equality, friction, limits, frictionless contacts
|
||||
dim1 = d.contact.efc_address[d.contact.dim == 1]
|
||||
nefl = d.ne + d.nf + d.nl
|
||||
active = ((ctx.Jaref + alpha * jv) < 0).at[:d.ne + d.nf].set(True)
|
||||
active = active.at[nefl:].set(False).at[dim1].set(active[dim1])
|
||||
quad_efld = jax.vmap(jp.multiply)(quad, active)
|
||||
quad_total += jp.sum(quad_efld, axis=0)
|
||||
# elliptic bottom zone: quadratic cost
|
||||
efc_elliptic = d.contact.efc_address[d.contact.dim > 1]
|
||||
quad_c = jax.vmap(jp.multiply)(quad[efc_elliptic], bottom_zone)
|
||||
quad_total += jp.sum(quad_c, axis=0)
|
||||
# elliptic middle zone
|
||||
t += (t == 0) * mujoco.mjMINVAL
|
||||
tsqr += (tsqr == 0) * mujoco.mjMINVAL
|
||||
n1 = v0
|
||||
t1 = (uv + alpha * vv) / t
|
||||
t2 = vv / t - (uv + alpha * vv) * t1 / tsqr
|
||||
dm = ctx.dm * middle_zone
|
||||
nmt = n - mu * t
|
||||
cost = 0.5 * jp.sum(dm * jp.square(nmt))
|
||||
deriv_0 = jp.sum(dm * nmt * (n1 - mu * t1))
|
||||
deriv_1 = jp.sum(dm * (jp.square(n1 - mu * t1) - nmt * mu * t2))
|
||||
elif m.opt.cone == ConeType.PYRAMIDAL:
|
||||
active = ((ctx.Jaref + alpha * jv) < 0).at[:d.ne + d.nf].set(True)
|
||||
quad = jax.vmap(jp.multiply)(quad, active) # only active
|
||||
quad_total += jp.sum(quad, axis=0)
|
||||
else:
|
||||
raise NotImplementedError(f'unsupported cone type: {m.opt.cone}')
|
||||
|
||||
alpha_sq = alpha * alpha
|
||||
cost += alpha_sq * quad_total[2] + alpha * quad_total[1] + quad_total[0]
|
||||
deriv_0 += 2 * alpha * quad_total[2] + quad_total[1]
|
||||
deriv_1 += 2 * quad_total[2] + (quad_total[2] == 0) * mujoco.mjMINVAL
|
||||
|
||||
cost = alpha * alpha * quad_total[2] + alpha * quad_total[1] + quad_total[0]
|
||||
deriv_0 = 2 * alpha * quad_total[2] + quad_total[1]
|
||||
deriv_1 = 2 * quad_total[2] + (quad_total[2] == 0) * mujoco.mjMINVAL
|
||||
return _LSPoint(alpha=alpha, cost=cost, deriv_0=deriv_0, deriv_1=deriv_1)
|
||||
|
||||
|
||||
@@ -159,34 +227,95 @@ def _while_loop_scan(cond_fun, body_fun, init_val, max_iter):
|
||||
return jax.lax.scan(_fun, init, None, length=max_iter)[0][0]
|
||||
|
||||
|
||||
def _update_constraint(d: Data, ctx: _Context) -> _Context:
|
||||
def _update_constraint(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
"""Updates constraint force and resulting cost given latst solver iteration.
|
||||
|
||||
Corresponds to CGupdateConstraint in mujoco/src/engine/engine_solver.c
|
||||
|
||||
Args:
|
||||
m: model defining constraints
|
||||
d: data which contains latest qacc and smooth terms
|
||||
ctx: current solver context
|
||||
|
||||
Returns:
|
||||
context with new constraint force and costs
|
||||
"""
|
||||
# TODO(robotics-team): add friction constraints
|
||||
if m.opt.cone == ConeType.PYRAMIDAL:
|
||||
# ne/nf constraints are always active, rest are non-negative constraints
|
||||
active = (ctx.Jaref < 0).at[:d.ne + d.nf].set(True)
|
||||
efc_force = d.efc_D * -ctx.Jaref * active
|
||||
cost = 0.5 * jp.sum(d.efc_D * ctx.Jaref * ctx.Jaref * active)
|
||||
dm, u, h = 0.0, 0.0, 0.0
|
||||
elif m.opt.cone == ConeType.ELLIPTIC:
|
||||
friction = d.contact.friction[d.contact.dim > 1]
|
||||
efc_address = d.contact.efc_address[d.contact.dim > 1]
|
||||
dim = d.contact.dim[d.contact.dim > 1]
|
||||
slice_fn = jax.vmap(lambda x: jax.lax.dynamic_slice(ctx.Jaref, (x,), (6,)))
|
||||
u = slice_fn(efc_address) * ctx.fri
|
||||
mu, n, t = ctx.fri[:, 0], u[:, 0], jax.vmap(math.norm)(u[:, 1:])
|
||||
|
||||
# only count active constraints
|
||||
active = (ctx.Jaref < 0).at[:d.ne + d.nf].set(True)
|
||||
# bottom zone: quadratic
|
||||
bottom_zone = ((t <= 0) & (n < 0)) | ((t > 0) & ((mu * n + t) <= 0))
|
||||
active = (ctx.Jaref < 0).at[:d.ne + d.nf].set(True)
|
||||
adr_i, adr_j = [], []
|
||||
for i, (condim, addr) in enumerate(zip(dim, efc_address)):
|
||||
adr_i.extend(range(addr, addr + condim))
|
||||
adr_j.extend([i] * condim)
|
||||
active = active.at[jp.array(adr_i)].set(bottom_zone[jp.array(adr_j)])
|
||||
efc_force = d.efc_D * -ctx.Jaref * active
|
||||
cost = 0.5 * jp.sum(d.efc_D * ctx.Jaref * ctx.Jaref * active)
|
||||
|
||||
# middle zone: cone
|
||||
middle_zone = (t > 0) & (n < (mu * t)) & ((mu * n + t) > 0)
|
||||
dm = d.efc_D[efc_address] / jp.maximum(
|
||||
mu * mu * (1 + mu * mu), mujoco.mjMINVAL
|
||||
)
|
||||
nmt = n - mu * t
|
||||
cost += 0.5 * jp.sum(dm * nmt * nmt * middle_zone)
|
||||
# tangent and friction for middle zone:
|
||||
force = -dm * nmt * mu * middle_zone
|
||||
force_fri = -force / (t + ~middle_zone * mujoco.mjMINVAL)
|
||||
force_fri = force_fri[:, None] * u[:, 1:] * friction
|
||||
efc_force = efc_force.at[efc_address].add(force)
|
||||
efc_adr, adr_i, adr_j = [], [], []
|
||||
for i, (condim, addr) in enumerate(zip(dim, efc_address)):
|
||||
efc_adr.extend(range(addr + 1, addr + condim))
|
||||
adr_i.extend([i] * (condim - 1))
|
||||
adr_j.extend(range(condim - 1))
|
||||
efc_adr, adr_i, adr_j = jp.array(efc_adr), jp.array(adr_i), jp.array(adr_j)
|
||||
efc_force = efc_force.at[efc_adr].add(force_fri[(adr_i, adr_j)])
|
||||
|
||||
# cone hessian
|
||||
h = 0.0
|
||||
if m.opt.solver == SolverType.NEWTON:
|
||||
t = jp.maximum(t, mujoco.mjMINVAL)
|
||||
# h = mu*N/T^3 * U*U'
|
||||
ttt = jp.maximum(t * t * t, mujoco.mjMINVAL)
|
||||
h = jax.vmap(lambda x, y: x * jp.outer(y, y.T))(mu * n / ttt, u)
|
||||
# add to diagonal: (mu^2 - mu*N/T) * I
|
||||
h += jax.vmap(lambda x: x * jp.eye(6, 6))(mu * mu - mu * n / t)
|
||||
# set first row: (1, -mu/T * U)
|
||||
h_0 = jax.vmap(lambda mu, t, u: jp.append(1, -mu / t * u[1:]))(mu, t, u)
|
||||
h = h.at[:, 0].set(h_0).at[:, :, 0].set(h_0)
|
||||
# pre and post multiply by diag(mu, friction), scale by Dm
|
||||
h *= jax.vmap(lambda d, f: d * jp.outer(f, f.T))(dm, ctx.fri)
|
||||
# only cone constraints
|
||||
h = jax.vmap(jp.multiply)(h, middle_zone)
|
||||
else:
|
||||
raise NotImplementedError(f'unsupported cone type: {m.opt.cone}')
|
||||
|
||||
efc_force = d.efc_D * -ctx.Jaref * active
|
||||
qfrc_constraint = d.efc_J.T @ efc_force
|
||||
gauss = 0.5 * jp.dot(ctx.Ma - d.qfrc_smooth, ctx.qacc - d.qacc_smooth)
|
||||
cost = 0.5 * jp.sum(d.efc_D * ctx.Jaref * ctx.Jaref * active) + gauss
|
||||
|
||||
ctx = ctx.replace(
|
||||
qfrc_constraint=qfrc_constraint,
|
||||
gauss=gauss,
|
||||
cost=cost,
|
||||
cost=cost + gauss,
|
||||
prev_cost=ctx.cost,
|
||||
efc_force=efc_force,
|
||||
active=active,
|
||||
dm=dm,
|
||||
u=u,
|
||||
h=h,
|
||||
)
|
||||
|
||||
return ctx
|
||||
@@ -213,8 +342,17 @@ def _update_gradient(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
if m.opt.solver == SolverType.CG:
|
||||
mgrad = smooth.solve_m(m, d, grad)
|
||||
elif m.opt.solver == SolverType.NEWTON:
|
||||
active = (ctx.Jaref < 0).at[: d.ne + d.nf].set(True)
|
||||
h = (d.efc_J.T * d.efc_D * active) @ d.efc_J
|
||||
if m.opt.cone == ConeType.ELLIPTIC:
|
||||
cm = jp.diag(d.efc_D * ctx.active)
|
||||
efc_address = d.contact.efc_address[d.contact.dim > 1]
|
||||
dim = d.contact.dim[d.contact.dim > 1]
|
||||
# set efc of cone H along diagonal
|
||||
for i, (condim, addr) in enumerate(zip(dim, efc_address)):
|
||||
h_cone = ctx.h[i, :condim, :condim]
|
||||
cm = cm.at[addr:addr+condim, addr:addr+condim].add(h_cone)
|
||||
h = d.efc_J.T @ cm @ d.efc_J
|
||||
else:
|
||||
h = (d.efc_J.T * d.efc_D * ctx.active) @ d.efc_J
|
||||
h = support.full_m(m, d) + h
|
||||
h_ = jax.scipy.linalg.cho_factor(h)
|
||||
mgrad = jax.scipy.linalg.cho_solve(h_, grad)
|
||||
@@ -256,8 +394,28 @@ def _linesearch(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
))
|
||||
quad = jp.stack((0.5 * ctx.Jaref * ctx.Jaref, jv * ctx.Jaref, 0.5 * jv * jv))
|
||||
quad = (quad * d.efc_D).T
|
||||
uu, v0, uv, vv = 0.0, 0.0, 0.0, 0.0
|
||||
if m.opt.cone == ConeType.ELLIPTIC:
|
||||
mask = d.contact.dim > 1
|
||||
# complete vector quadratic (for bottom zone)
|
||||
efc_con, efc_fri = [], []
|
||||
for condim, addr in zip(d.contact.dim[mask], d.contact.efc_address[mask]):
|
||||
efc_con.extend([addr] * (condim - 1))
|
||||
efc_fri.extend(range(addr + 1, addr + condim))
|
||||
quad = quad.at[jp.array(efc_con)].add(quad[jp.array(efc_fri)])
|
||||
|
||||
point_fn = lambda a: _LSPoint.create(d, ctx, a, jv, quad, quad_gauss)
|
||||
# rescale to make primal cone circular
|
||||
jv_fn = jax.vmap(lambda x: jax.lax.dynamic_slice(jv, (x,), (6,)))
|
||||
efc_elliptic = d.contact.efc_address[mask]
|
||||
v = jv_fn(efc_elliptic) * ctx.fri
|
||||
uu = jp.sum(ctx.u[:, 1:] * ctx.u[:, 1:], axis=1)
|
||||
v0 = v[:, 0]
|
||||
uv = jp.sum(ctx.u[:, 1:] * v[:, 1:], axis=1)
|
||||
vv = jp.sum(v[:, 1:] * v[:, 1:], axis=1)
|
||||
|
||||
point_fn = lambda a: _LSPoint.create(
|
||||
m, d, ctx, a, jv, quad, quad_gauss, uu, v0, uv, vv
|
||||
)
|
||||
|
||||
def cond(ctx: _LSContext) -> jax.Array:
|
||||
done = ctx.ls_iter >= m.opt.ls_iterations
|
||||
@@ -274,21 +432,34 @@ def _linesearch(m: Model, d: Data, ctx: _Context) -> _Context:
|
||||
hi_next = point_fn(hi.alpha - hi.deriv_0 / hi.deriv_1)
|
||||
mid = point_fn(0.5 * (lo.alpha + hi.alpha))
|
||||
|
||||
# we swap lo/hi if:
|
||||
# 1) they are not correctly at a bracket boundary (e.g. lo.deriv_0 > 0), OR
|
||||
# 2) if moving to next or mid narrows the bracket
|
||||
swap_lo_next = (lo.deriv_0 > 0) | (lo.deriv_0 < lo_next.deriv_0)
|
||||
lo = jax.tree_util.tree_map(lambda x, y: jp.where(swap_lo_next, y, x), lo, lo_next)
|
||||
swap_lo_mid = (mid.deriv_0 < 0) & (lo.deriv_0 < mid.deriv_0)
|
||||
lo = jax.tree_util.tree_map(lambda x, y: jp.where(swap_lo_mid, y, x), lo, mid)
|
||||
|
||||
swap_hi_next = (hi.deriv_0 < 0) | (hi.deriv_0 > hi_next.deriv_0)
|
||||
hi = jax.tree_util.tree_map(lambda x, y: jp.where(swap_hi_next, y, x), hi, hi_next)
|
||||
swap_hi_mid = (mid.deriv_0 > 0) & (hi.deriv_0 > mid.deriv_0)
|
||||
hi = jax.tree_util.tree_map(lambda x, y: jp.where(swap_hi_mid, y, x), hi, mid)
|
||||
|
||||
swap = swap_lo_next | swap_lo_mid | swap_hi_next | swap_hi_mid
|
||||
|
||||
# swap lo/hi if the derivative points to a narrower bracket width
|
||||
in_bracket = lambda x, y: ((x < y) & (y < 0) | (x > y) & (y > 0))
|
||||
swap_lo_next = in_bracket(lo.deriv_0, lo_next.deriv_0)
|
||||
lo = jax.tree_util.tree_map(
|
||||
lambda x, y: jp.where(swap_lo_next, y, x), lo, lo_next
|
||||
)
|
||||
swap_lo_mid = in_bracket(lo.deriv_0, mid.deriv_0)
|
||||
lo = jax.tree_util.tree_map(
|
||||
lambda x, y: jp.where(swap_lo_mid, y, x), lo, mid
|
||||
)
|
||||
swap_lo_hi_next = in_bracket(lo.deriv_0, hi_next.deriv_0)
|
||||
lo = jax.tree_util.tree_map(
|
||||
lambda x, y: jp.where(swap_lo_hi_next, y, x), lo, hi_next
|
||||
)
|
||||
swap_hi_next = in_bracket(hi.deriv_0, hi_next.deriv_0)
|
||||
hi = jax.tree_util.tree_map(
|
||||
lambda x, y: jp.where(swap_hi_next, y, x), hi, hi_next
|
||||
)
|
||||
swap_hi_mid = in_bracket(hi.deriv_0, mid.deriv_0)
|
||||
hi = jax.tree_util.tree_map(
|
||||
lambda x, y: jp.where(swap_hi_mid, y, x), hi, mid
|
||||
)
|
||||
swap_hi_lo_next = in_bracket(hi.deriv_0, lo_next.deriv_0)
|
||||
hi = jax.tree_util.tree_map(
|
||||
lambda x, y: jp.where(swap_hi_lo_next, y, x), hi, lo_next
|
||||
)
|
||||
swap = swap_lo_next | swap_lo_mid | swap_lo_hi_next
|
||||
swap = swap | swap_hi_next | swap_hi_mid | swap_hi_lo_next
|
||||
ctx = ctx.replace(lo=lo, hi=hi, swap=swap, ls_iter=ctx.ls_iter + 1)
|
||||
|
||||
return ctx
|
||||
@@ -331,14 +502,17 @@ def solve(m: Model, d: Data) -> Data:
|
||||
def body(ctx: _Context) -> _Context:
|
||||
ctx = _linesearch(m, d, ctx)
|
||||
prev_grad, prev_Mgrad = ctx.grad, ctx.Mgrad # pylint: disable=invalid-name
|
||||
ctx = _update_constraint(d, ctx)
|
||||
ctx = _update_constraint(m, d, ctx)
|
||||
ctx = _update_gradient(m, d, ctx)
|
||||
|
||||
# polak-ribiere:
|
||||
beta = jp.dot(ctx.grad, ctx.Mgrad - prev_Mgrad)
|
||||
beta = beta / jp.maximum(mujoco.mjMINVAL, jp.dot(prev_grad, prev_Mgrad))
|
||||
beta = jp.maximum(0, beta)
|
||||
search = -ctx.Mgrad + beta * ctx.search
|
||||
if m.opt.solver == SolverType.NEWTON:
|
||||
search = -ctx.Mgrad
|
||||
else:
|
||||
# polak-ribiere:
|
||||
beta = jp.dot(ctx.grad, ctx.Mgrad - prev_Mgrad)
|
||||
beta = beta / jp.maximum(mujoco.mjMINVAL, jp.dot(prev_grad, prev_Mgrad))
|
||||
beta = jp.maximum(0, beta)
|
||||
search = -ctx.Mgrad + beta * ctx.search
|
||||
ctx = ctx.replace(search=search, solver_niter=ctx.solver_niter + 1)
|
||||
|
||||
return ctx
|
||||
|
||||
@@ -15,16 +15,18 @@
|
||||
"""Tests for constraint functions."""
|
||||
|
||||
from absl.testing import absltest
|
||||
from absl.testing import parameterized
|
||||
import jax
|
||||
import mujoco
|
||||
from mujoco import mjx
|
||||
from mujoco.mjx._src import solver
|
||||
from mujoco.mjx._src import test_util
|
||||
import numpy as np
|
||||
|
||||
|
||||
# tolerance for difference between MuJoCo and MJX constraint calculations,
|
||||
# tolerance for difference between MuJoCo and MJX solver calculations,
|
||||
# mostly due to float precision
|
||||
_TOLERANCE = 5e-5
|
||||
_TOLERANCE = 5e-3
|
||||
|
||||
|
||||
def _assert_eq(a, b, name, tol=_TOLERANCE):
|
||||
@@ -37,72 +39,85 @@ def _assert_attr_eq(a, b, attr, tol=_TOLERANCE):
|
||||
_assert_eq(getattr(a, attr), getattr(b, attr), attr, tol=tol)
|
||||
|
||||
|
||||
class SolverTest(absltest.TestCase):
|
||||
class SolverTest(parameterized.TestCase):
|
||||
|
||||
def test_newton(self):
|
||||
"""Test newton solver."""
|
||||
@parameterized.parameters(
|
||||
# these scene challenges the solver, with CG you need to crank up
|
||||
# the iterations, otherwise it diverges
|
||||
(mujoco.mjtSolver.mjSOL_CG, mujoco.mjtCone.mjCONE_PYRAMIDAL, 100),
|
||||
(mujoco.mjtSolver.mjSOL_CG, mujoco.mjtCone.mjCONE_ELLIPTIC, 100),
|
||||
# Newton converges much more quickly, lower iterations to demonstrate
|
||||
# mgrad is being calculated optimally
|
||||
(mujoco.mjtSolver.mjSOL_NEWTON, mujoco.mjtCone.mjCONE_PYRAMIDAL, 2),
|
||||
(mujoco.mjtSolver.mjSOL_NEWTON, mujoco.mjtCone.mjCONE_ELLIPTIC, 2),
|
||||
)
|
||||
def test_solver(self, solver_, cone, iterations):
|
||||
"""Test newton, CG solver with pyramidal, elliptic cones."""
|
||||
m = test_util.load_test_file('constraints.xml')
|
||||
# it's critical that mgrad is optimally calculated, so lower iterations
|
||||
# to be sure that MJX is converging as quickly as MuJoCo
|
||||
m.opt.iterations = 1
|
||||
m.opt.solver = solver_
|
||||
m.opt.cone = cone
|
||||
m.opt.iterations = iterations
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_step(m, d, 20) # significant constraint forces at 20 steps
|
||||
|
||||
# mj_forward overwrites qacc_warmstart, so let's restore it to what it was
|
||||
# at the beginning of the step so that MJX does not have a trivial solution
|
||||
warmstart = d.qacc_warmstart.copy()
|
||||
mujoco.mj_forward(m, d)
|
||||
d.qacc_warmstart = warmstart
|
||||
def cost(qacc):
|
||||
jaref = np.zeros(d.nefc, dtype=float)
|
||||
cost = np.zeros(1)
|
||||
mujoco.mj_mulJacVec(m, d, jaref, qacc)
|
||||
mujoco.mj_constraintUpdate(m, d, jaref - d.efc_aref, cost, 0)
|
||||
return cost
|
||||
|
||||
dx = jax.jit(mjx.solve)(mjx.put_model(m), mjx.put_data(m, d))
|
||||
# sample a mix of active/inactive constraints at different timesteps
|
||||
for key in range(0, 3):
|
||||
mujoco.mj_resetDataKeyframe(m, d, key)
|
||||
mujoco.mj_step(m, d) # step to generate warmstart
|
||||
|
||||
_assert_attr_eq(d, dx, 'qacc')
|
||||
_assert_attr_eq(d, dx, 'qfrc_constraint')
|
||||
nnz = dx.efc_J.any(axis=1)
|
||||
_assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force')
|
||||
# compare costs
|
||||
mj_cost = cost(d.qacc)
|
||||
ctx = solver._Context.create(mjx.put_model(m), mjx.put_data(m, d))
|
||||
mjx_cost = ctx.cost - ctx.gauss
|
||||
_assert_eq(mj_cost, mjx_cost, 'cost')
|
||||
|
||||
def test_cg(self):
|
||||
"""Test CG solver."""
|
||||
m = test_util.load_test_file('constraints.xml')
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_step(m, d, 20) # significant constraint forces at 20 steps
|
||||
# mj_forward overwrites qacc_warmstart, so let's restore it to what it was
|
||||
# before the step so that MJX does not have a trivial solution
|
||||
warmstart = d.qacc_warmstart.copy()
|
||||
mujoco.mj_forward(m, d)
|
||||
d.qacc_warmstart = warmstart
|
||||
dx = jax.jit(mjx.solve)(mjx.put_model(m), mjx.put_data(m, d))
|
||||
|
||||
# CG does not converge as quickly as Newton but is cheaper to calculate
|
||||
m.opt.solver = mujoco.mjtSolver.mjSOL_CG
|
||||
m.opt.iterations = 8
|
||||
# MJX finds very similar solutions with the newton solver
|
||||
if solver_ == mujoco.mjtSolver.mjSOL_NEWTON:
|
||||
nnz = dx.efc_J.any(axis=1)
|
||||
_assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force')
|
||||
_assert_attr_eq(d, dx, 'qfrc_constraint')
|
||||
_assert_attr_eq(d, dx, 'qacc')
|
||||
|
||||
# mj_forward overwrites qacc_warmstart, so let's restore it to what it was
|
||||
# at the beginning of the step so that MJX does not have a trivial solution
|
||||
warmstart = d.qacc_warmstart.copy()
|
||||
mujoco.mj_forward(m, d)
|
||||
d.qacc_warmstart = warmstart
|
||||
|
||||
dx = jax.jit(mjx.solve)(mjx.put_model(m), mjx.put_data(m, d))
|
||||
|
||||
_assert_attr_eq(d, dx, 'qacc')
|
||||
_assert_attr_eq(d, dx, 'qfrc_constraint', tol=8e-4)
|
||||
nnz = dx.efc_J.any(axis=1)
|
||||
_assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force', tol=5e-4)
|
||||
# both CG and Newton find costs that are nearly the same as MuJoCo, often
|
||||
# lower (due to slight differences in the MJX linsearch algorithm)
|
||||
mj_cost = cost(d.qacc)
|
||||
mjx_cost = cost(dx.qacc)
|
||||
self.assertLess(mjx_cost, mj_cost * 1.01)
|
||||
|
||||
def test_no_warmstart(self):
|
||||
"""Test no warmstart."""
|
||||
m = test_util.load_test_file('constraints.xml')
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_step(m, d, 20) # significant constraint forces at 20 steps
|
||||
# significant constraint forces keyframe 2
|
||||
mujoco.mj_resetDataKeyframe(m, d, 2)
|
||||
m.opt.disableflags |= mujoco.mjtDisableBit.mjDSBL_WARMSTART
|
||||
mujoco.mj_forward(m, d)
|
||||
mx = mjx.put_model(m)
|
||||
dx = jax.jit(mjx.solve)(mx, mjx.put_data(m, d))
|
||||
nnz = dx.efc_J.any(axis=1)
|
||||
# without warmstart, the solution is not as close
|
||||
_assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force', tol=2e-2)
|
||||
# even without warmstart, newton converges quickly
|
||||
_assert_eq(d.efc_force, dx.efc_force[nnz], 'efc_force', tol=2e-4)
|
||||
|
||||
def test_sparse(self):
|
||||
"""Test solver works with sparse mass matrices."""
|
||||
m = test_util.load_test_file('constraints.xml')
|
||||
m.opt.jacobian = mujoco.mjtJacobian.mjJAC_SPARSE
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_step(m, d, 20) # significant constraint forces at 20 steps
|
||||
# significant constraint forces keyframe 2
|
||||
mujoco.mj_resetDataKeyframe(m, d, 2)
|
||||
|
||||
# mj_forward overwrites qacc_warmstart, so let's restore it to what it was
|
||||
# at the beginning of the step so that MJX does not have a trivial solution
|
||||
|
||||
@@ -26,6 +26,7 @@ import mujoco
|
||||
# pylint: disable=g-importing-member
|
||||
from mujoco.mjx._src import forward
|
||||
from mujoco.mjx._src import io
|
||||
from mujoco.mjx._src.types import Data
|
||||
# pylint: enable=g-importing-member
|
||||
import numpy as np
|
||||
|
||||
@@ -104,6 +105,28 @@ def benchmark(
|
||||
return jit_time, run_time, steps
|
||||
|
||||
|
||||
def efc_order(m: mujoco.MjModel, d: mujoco.MjData, dx: Data) -> np.ndarray:
|
||||
"""Returns a sort order such that dx.efc_*[order][:d.nefc] == d.efc_*."""
|
||||
# reorder efc rows to skip inactive constraints and match contact order
|
||||
efl = dx.ne + dx.nf + dx.nl
|
||||
order = np.arange(efl)
|
||||
order[(dx.efc_J[:efl] == 0).all(axis=1)] = 2**16 # move empty rows to end
|
||||
for i in range(dx.ncon):
|
||||
num_rows = dx.contact.dim[i]
|
||||
if dx.contact.dim[i] > 1 and m.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL:
|
||||
num_rows = (dx.contact.dim[i] - 1) * 2
|
||||
if dx.contact.dist[i] > 0: # move empty contacts to end
|
||||
order = np.append(order, np.repeat(2 ** 16, num_rows))
|
||||
continue
|
||||
contact_match = (d.contact.geom == dx.contact.geom[i]).all(axis=-1)
|
||||
contact_match &= (d.contact.pos == dx.contact.pos[i]).all(axis=-1)
|
||||
assert contact_match.any(), f'contact {i} not found'
|
||||
contact_id = np.nonzero(contact_match)[0][0]
|
||||
order = np.append(order, np.repeat(efl + contact_id, num_rows))
|
||||
|
||||
return np.argsort(order, kind='stable')
|
||||
|
||||
|
||||
_ACTUATOR_TYPES = ['motor', 'velocity', 'position', 'general', 'intvelocity']
|
||||
_DYN_TYPES = ['none', 'integrator', 'filter', 'filterexact']
|
||||
_DYN_PRMS = ['0.189', '2.1']
|
||||
|
||||
@@ -135,9 +135,10 @@ class ConeType(enum.IntEnum):
|
||||
|
||||
Attributes:
|
||||
PYRAMIDAL: pyramidal
|
||||
ELLIPTIC: elliptic
|
||||
"""
|
||||
PYRAMIDAL = mujoco.mjtCone.mjCONE_PYRAMIDAL
|
||||
# unsupported: ELLIPTIC
|
||||
ELLIPTIC = mujoco.mjtCone.mjCONE_ELLIPTIC
|
||||
|
||||
|
||||
class JacobianType(enum.IntEnum):
|
||||
@@ -245,7 +246,7 @@ class ConstraintType(enum.IntEnum):
|
||||
# unsupported: LIMIT_TENDON
|
||||
CONTACT_FRICTIONLESS = mujoco.mjtConstraint.mjCNSTR_CONTACT_FRICTIONLESS
|
||||
CONTACT_PYRAMIDAL = mujoco.mjtConstraint.mjCNSTR_CONTACT_PYRAMIDAL
|
||||
# unsupported: CONTACT_ELLIPTIC
|
||||
CONTACT_ELLIPTIC = mujoco.mjtConstraint.mjCNSTR_CONTACT_ELLIPTIC
|
||||
|
||||
|
||||
class CamLightType(enum.IntEnum):
|
||||
@@ -703,7 +704,7 @@ class Contact(PyTreeNode):
|
||||
solref: jax.Array
|
||||
solreffriction: jax.Array
|
||||
solimp: jax.Array
|
||||
# unsupported: mu, H
|
||||
# unsupported: mu, H (calculated locally in solver.py)
|
||||
dim: np.ndarray
|
||||
geom1: jax.Array
|
||||
geom2: jax.Array
|
||||
|
||||
@@ -58,7 +58,7 @@
|
||||
<geom class="box" condim="4"/>
|
||||
</body>
|
||||
|
||||
<body name="box_condim6" pos="6 0 0">
|
||||
<body name="box_condim6" pos="7 0 0">
|
||||
<freejoint/>
|
||||
<geom class="box" condim="6"/>
|
||||
</body>
|
||||
@@ -75,4 +75,13 @@
|
||||
<motor gear="50000" joint="joint3" name="act2"/>
|
||||
<motor gear="75000" joint="joint4" name="act3"/>
|
||||
</actuator>
|
||||
|
||||
<keyframe>
|
||||
<!-- keyframe 0: default position with some motion, zero contacts -->
|
||||
<key qpos='1 0 0 0 -1 0 0 1 0 0 0 0 0 1 0 0 0 4 0 0 1 0 0 0 5 0 0 1 0 0 0 6 0 0 1 0 0 0 7 0 0 1 0 0 0' qvel='1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1 1'/>
|
||||
<!-- keyframe 1: some contacts but constraints are quadratic -->
|
||||
<key qpos='1 0.0011 -8.2e-07 -0.0011 -0.82 1.9 -0.58 0.97 -0.14 0.17 -0.069 0.22 0.29 0.93 0.2 0.15 0.28 4.3 0.23 -0.26 0.97 0.12 0.11 0.18 5.4 0.25 -0.26 0.97 0.11 0.16 0.17 6.4 0.25 -0.26 0.97 0.11 0.16 0.17 7.4 0.25 -0.26 0.97 0.11 0.16 0.17' qvel='-0.14 3.3e-05 0.14 -0.14 -0.93 -3.6 -1.8 1 -0.57 -0.16 0.21 1.1 1 2.2 0.65 1.3 -4.2 -2.6 -5.6 0.54 1.6 2.8 -3.7 -2.3 0.7 -1.8 1.6 2.8 -3.7 -2.3 0.62 -1.8 1.6 2.8 -3.7 -2.3 0.61 -1.8'/>
|
||||
<!-- keyframe 2: some contacts and some constraints are in cone state (for elliptic) -->
|
||||
<key qpos='1 0.0087 2.4e-07 -0.0086 -0.89 1.8 -0.77 0.98 -0.2 -0.0022 -0.026 0.19 0.33 0.86 0.32 0.064 0.38 4.4 0.36 -0.81 0.97 -0.0013 -0.0011 0.25 5.6 0.52 -0.75 0.98 -0.018 0.17 0.094 6.6 0.52 -0.75 0.98 -0.017 0.16 0.094 7.6 0.52 -0.76 0.98 -0.017 0.16 0.094' qvel='0.2 -1.8e-05 -0.2 -0.72 0.072 0.025 0.015 -4.9 0.35 -0.22 0.26 0.99 -4.6 1.7 1.1 0.52 -0.73 0.16 7.9 1 0.043 -0.042 1.3 1.6 -1.6 0.98 0.027 -0.024 1.3 1.6 -1.8 0.96 0.025 -0.022 1.3 1.6 -1.8 0.96'/>
|
||||
</keyframe>
|
||||
</mujoco>
|
||||
|
||||
Reference in New Issue
Block a user