diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py
index 238141eb..deb06509 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py
@@ -36,6 +36,8 @@ from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_position as fwd_
from mujoco.mjx.third_party.mujoco_warp._src.forward import fwd_velocity as fwd_velocity
from mujoco.mjx.third_party.mujoco_warp._src.forward import implicit as implicit
from mujoco.mjx.third_party.mujoco_warp._src.forward import rungekutta4 as rungekutta4
+from mujoco.mjx.third_party.mujoco_warp._src.forward import step1 as step1
+from mujoco.mjx.third_party.mujoco_warp._src.forward import step2 as step2
from mujoco.mjx.third_party.mujoco_warp._src.inverse import inverse as inverse
from mujoco.mjx.third_party.mujoco_warp._src.io import get_data_into as get_data_into
from mujoco.mjx.third_party.mujoco_warp._src.io import make_data as make_data
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py
index b29f6d83..871652a3 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py
@@ -38,7 +38,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_k
# TODO(team): improve compile time to enable backward pass
wp.config.enable_backward = False
-MULTI_CONTACT_COUNT = 4
+MULTI_CONTACT_COUNT = 8
mat3c = wp.types.matrix(shape=(MULTI_CONTACT_COUNT, 3), dtype=float)
_CONVEX_COLLISION_PAIRS = [
@@ -288,6 +288,7 @@ def ccd_kernel_builder(
points = mat3c()
+ # TODO(kbayes): remove legacy GJK once multicontact can be enabled
if default_gjk:
simplex, normal = gjk_legacy(
gjk_iterations,
@@ -349,10 +350,11 @@ def ccd_kernel_builder(
count = 0
return
- for i in range(count):
- points[i] = 0.5 * (witness1[i] + witness2[i])
- normal = witness1[0] - witness2[0]
- frame = make_frame(normal)
+ for i in range(count):
+ points[i] = 0.5 * (witness1[i] + witness2[i])
+ normal = witness1[0] - witness2[0]
+ frame = make_frame(normal)
+
for i in range(count):
# limit maximum number of contacts with height field
if _max_contacts_height_field(ngeom, geom_type, geompair2hfgeompair, g1, g2, worldid, ncon_hfield_out):
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py
index 964a673a..7df19a84 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py
@@ -26,6 +26,14 @@ from mujoco.mjx.third_party.mujoco_warp.test_data.collision_sdf.utils import reg
from mujoco.mjx.third_party.mujoco_warp._src import test_util
from mujoco.mjx.third_party.mujoco_warp._src import types
+_TOLERANCE = 5e-5
+
+
+def _assert_eq(a, b, name):
+ tol = _TOLERANCE * 10
+ err_msg = f"mismatch: {name}"
+ np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+
class CollisionTest(parameterized.TestCase):
"""Tests the collision contact functions."""
@@ -863,7 +871,35 @@ class CollisionTest(parameterized.TestCase):
self.assertEqual(d.ncon.numpy()[0], 1)
np.testing.assert_allclose(d.contact.friction.numpy()[0], types.MJ_MINMU)
- # TODO(team): test contact parameter mixing
+ @parameterized.parameters(("1", "1"), ("1", "2"), ("2", "1"))
+ def test_contact_parameter_mixing(self, priority1, priority2):
+ _, mjd, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+
+
+
+
+
+
+ """,
+ keyframe=0,
+ )
+
+ mjwarp.collision(m, d)
+
+ ncon = d.ncon.numpy()[0]
+ _assert_eq(ncon, 1, "ncon")
+ _assert_eq(d.contact.friction.numpy()[0], mjd.contact.friction[0], "friction")
+ _assert_eq(d.contact.solref.numpy()[0], mjd.contact.solref[0], "solref")
+ _assert_eq(d.contact.solimp.numpy()[0], mjd.contact.solimp[0], "solimp")
+ _assert_eq(d.contact.includemargin.numpy()[0], mjd.contact.includemargin[0], "includemargin")
+ _assert_eq(d.contact.dim.numpy()[0], mjd.contact.dim[0], "dim")
if __name__ == "__main__":
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py
index 81cb0748..ff999d2d 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py
@@ -41,7 +41,7 @@ polyindices = wp.types.vector(MAX_POLYVERT, dtype=int)
mat43 = wp.types.matrix(shape=(4, 3), dtype=float)
mat63 = wp.types.matrix(shape=(6, 3), dtype=float)
-MULTI_CONTACT_COUNT = 4
+MULTI_CONTACT_COUNT = 8
mat3c = wp.types.matrix(shape=(MULTI_CONTACT_COUNT, 3), dtype=float)
@@ -93,6 +93,13 @@ class SupportPoint:
vertex_index: int
+@wp.func
+def discrete_geoms(g1: int, g2: int):
+ return (g1 == int(GeomType.MESH.value) or g1 == int(GeomType.BOX.value) or g1 == int(GeomType.HFIELD.value)) and (
+ g2 == int(GeomType.MESH.value) or g2 == int(GeomType.BOX.value) or g2 == int(GeomType.HFIELD.value)
+ )
+
+
@wp.func
def _support_margin(geom: Geom, geomtype: int, dir: wp.vec3):
sp = SupportPoint()
@@ -590,6 +597,7 @@ def _gjk(
use_margin: bool,
):
"""Find distance within a tolerance between two geoms."""
+ is_discrete = discrete_geoms(geomtype1, geomtype2)
cutoff2 = cutoff * cutoff
simplex = mat43()
simplex1 = mat43()
@@ -598,7 +606,7 @@ def _gjk(
simplex_index2 = wp.vec4i()
n = int(0)
coordinates = wp.vec4() # barycentric coordinates
- epsilon = 0.5 * tolerance * tolerance
+ epsilon = wp.where(is_discrete, 0.0, 0.5 * tolerance * tolerance)
# set initial guess
x_k = x1_0 - x2_0
@@ -1193,12 +1201,14 @@ def _polytope4(
@wp.func
-def _epa(tolerance2: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom2: Geom, geomtype1: int, geomtype2: int):
+def _epa(tolerance: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom2: Geom, geomtype1: int, geomtype2: int):
"""Recover penetration data from two geoms in contact given an initial polytope."""
+ is_discrete = discrete_geoms(geomtype1, geomtype2)
upper = FLOAT_MAX
upper2 = FLOAT_MAX
idx = int(-1)
pidx = int(-1)
+ epsilon = wp.where(is_discrete, 1e-15, tolerance * tolerance)
for k in range(epa_iterations):
pidx = int(idx)
@@ -1235,9 +1245,19 @@ def _epa(tolerance2: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom
upper = upper_k
upper2 = upper * upper
- if upper - lower < tolerance2:
+ if upper - lower < epsilon:
break
+ # check if vertex wi is a repeated support point
+ if is_discrete:
+ found_repeated = bool(False)
+ for i in range(pt.nvert - 1):
+ if pt.vert_index1[i] == pt.vert_index1[wi] and pt.vert_index2[i] == pt.vert_index2[wi]:
+ found_repeated = True
+ break
+ if found_repeated:
+ break
+
pt.nmap = _delete_face(pt, idx)
pt.nhorizon = _add_edge(pt, pt.face[idx][0], pt.face[idx][1])
pt.nhorizon = _add_edge(pt, pt.face[idx][1], pt.face[idx][2])
@@ -1251,7 +1271,7 @@ def _epa(tolerance2: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom
if pt.face_index[i] == -2:
continue
- if wp.dot(pt.face_pr[i], pt.vert[wi]) - pt.face_norm2[i] > MJ_MINVAL:
+ if wp.dot(pt.face_pr[i], pt.vert[wi]) - pt.face_norm2[i] > 1e-10:
pt.nmap = _delete_face(pt, i)
pt.nhorizon = _add_edge(pt, pt.face[i][0], pt.face[i][1])
pt.nhorizon = _add_edge(pt, pt.face[i][1], pt.face[i][2])
@@ -1694,7 +1714,7 @@ def plane_normal(v1: wp.vec3, v2: wp.vec3, n: wp.vec3):
@wp.func
def halfspace(a: wp.vec3, n: wp.vec3, p: wp.vec3):
- return wp.dot(p - a, n) > -MJ_MINVAL
+ return wp.dot(p - a, n) > -1e-10
@wp.func
@@ -1725,7 +1745,7 @@ def polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, n
# compute plane normal and distance to plane for each vertex
pn = polyverts()
pd = polyvec()
- for i in range(nface1):
+ for i in range(nface1 - 1):
pdi, pni = plane_normal(face1[i], face1[i + 1], n)
pd[i] = pdi
pn[i] = pni
@@ -1734,14 +1754,11 @@ def polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, n
pn[nface1 - 1] = pni
# reserve 2 * max_sides as max sides for a clipped polygon
- polygon1 = polyclip()
- polygon2 = polyclip()
+ polygon = polyclip()
+ clipped = polyclip()
npolygon = nface2
nclipped = int(0)
- polygon = polygon1
- clipped = polygon2
-
for i in range(nface2):
polygon[i] = face2[i]
@@ -1768,7 +1785,7 @@ def polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, n
# add new vertex to clipped polygon where PQ intersects the clipping edge
t, res = plane_intersect(pn[e], pd[e], P, Q)
- if t < 0.0 or t > 1.0:
+ if t >= 0.0 and t <= 1.0:
clipped[nclipped] = res
nclipped += 1
@@ -1790,8 +1807,8 @@ def polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, n
# no pruning needed
for i in range(npolygon):
witness2[i] = polygon[i]
- witness1[i] = witness2[i] + dir
- return npolygon, witness2, witness1
+ witness1[i] = witness2[i] - dir
+ return npolygon, witness1, witness2
# recover multiple contacts from EPA polytope
@@ -2125,11 +2142,14 @@ def ccd(
witness2[0] = result.x2
return result.dist, 1, witness1, witness2
- dist, x1, x2, idx = _epa(tolerance * tolerance, epa_iterations, pt, geom1, geom2, geomtype1, geomtype2)
+ dist, x1, x2, idx = _epa(tolerance, epa_iterations, pt, geom1, geom2, geomtype1, geomtype2)
+ if idx == -1:
+ return FLOAT_MAX, 0, witness1, witness2
+
if (
multiccd
- and (geomtype1 == int(GeomType.BOX.value) or geomtype1 == int(GeomType.MESH.value))
- and (geomtype2 == int(GeomType.BOX.value) or geomtype2 == int(GeomType.MESH.value))
+ and (geomtype1 == int(GeomType.BOX.value) or (geomtype1 == int(GeomType.MESH.value) and geom1.mesh_polyadr > -1))
+ and (geomtype2 == int(GeomType.BOX.value) or (geomtype2 == int(GeomType.MESH.value) and geom2.mesh_polyadr > -1))
):
num, w1, w2 = multicontact(pt, pt.face[idx], x1, x2, geom1, geom2, geomtype1, geomtype2)
if num > 0:
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py
index a490b862..f47f7e01 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py
@@ -26,10 +26,10 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
-MAX_ITERATIONS = 10
+MAX_ITERATIONS = 20
-def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
+def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int, multiccd=False):
@nested_kernel
def _gjk_kernel(
# Model:
@@ -39,6 +39,15 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
mesh_vertadr: wp.array(dtype=int),
mesh_vertnum: wp.array(dtype=int),
mesh_vert: wp.array(dtype=wp.vec3),
+ mesh_polynum: wp.array(dtype=int),
+ mesh_polyadr: wp.array(dtype=int),
+ mesh_polynormal: wp.array(dtype=wp.vec3),
+ mesh_polyvertadr: wp.array(dtype=int),
+ mesh_polyvertnum: wp.array(dtype=int),
+ mesh_polyvert: wp.array(dtype=int),
+ mesh_polymapadr: wp.array(dtype=int),
+ mesh_polymapnum: wp.array(dtype=int),
+ mesh_polymap: wp.array(dtype=int),
# Data in:
geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.array2d(dtype=wp.mat33),
@@ -59,6 +68,7 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
horizon: wp.array(dtype=int),
# Out:
dist_out: wp.array(dtype=float),
+ ncon_out: wp.array(dtype=int),
pos_out: wp.array(dtype=wp.vec3),
):
MESHGEOM = int(GeomType.MESH.value)
@@ -70,12 +80,22 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
geom1.rot = geom_xmat_in[0, gid1]
geom1.size = geom_size[0, gid1]
geom1.graphadr = -1
+ geom1.mesh_polyadr = -1
if geom_dataid[gid1] >= 0 and geom_type[gid1] == MESHGEOM:
dataid = geom_dataid[gid1]
geom1.vertadr = mesh_vertadr[dataid]
geom1.vertnum = mesh_vertnum[dataid]
+ geom1.mesh_polynum = mesh_polynum[dataid]
+ geom1.mesh_polyadr = mesh_polyadr[dataid]
geom1.vert = mesh_vert
+ geom1.mesh_polynormal = mesh_polynormal
+ geom1.mesh_polyvertadr = mesh_polyvertadr
+ geom1.mesh_polyvertnum = mesh_polyvertnum
+ geom1.mesh_polyvert = mesh_polyvert
+ geom1.mesh_polymapadr = mesh_polymapadr
+ geom1.mesh_polymapnum = mesh_polymapnum
+ geom1.mesh_polymap = mesh_polymap
geom2 = Geom()
geom2.index = -1
@@ -84,23 +104,33 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
geom2.rot = geom_xmat_in[0, gid2]
geom2.size = geom_size[0, gid2]
geom2.graphadr = -1
+ geom2.mesh_polyadr = -1
if geom_dataid[gid2] >= 0 and geom_type[gid2] == MESHGEOM:
dataid = geom_dataid[gid2]
geom2.vertadr = mesh_vertadr[dataid]
geom2.vertnum = mesh_vertnum[dataid]
+ geom2.mesh_polynum = mesh_polynum[dataid]
+ geom2.mesh_polyadr = mesh_polyadr[dataid]
geom2.vert = mesh_vert
+ geom2.mesh_polynormal = mesh_polynormal
+ geom2.mesh_polyvertadr = mesh_polyvertadr
+ geom2.mesh_polyvertnum = mesh_polyvertnum
+ geom2.mesh_polyvert = mesh_polyvert
+ geom2.mesh_polymapadr = mesh_polymapadr
+ geom2.mesh_polymapnum = mesh_polymapnum
+ geom2.mesh_polymap = mesh_polymap
x_1 = geom_xpos_in[0, gid1]
x_2 = geom_xpos_in[0, gid2]
(
dist,
- count,
+ ncon,
x1,
x2,
) = ccd(
- False,
+ multiccd,
1e-6,
1.0e30,
iterations,
@@ -125,6 +155,7 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
)
dist_out[0] = dist
+ ncon_out[0] = ncon
pos_out[0] = x1[0]
pos_out[1] = x2[0]
@@ -133,13 +164,14 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
vert2 = wp.array(shape=(iterations,), dtype=wp.vec3)
vert_index1 = wp.array(shape=(iterations,), dtype=int)
vert_index2 = wp.array(shape=(iterations,), dtype=int)
- face = wp.array(shape=(2 * iterations,), dtype=wp.vec3i)
- face_pr = wp.array(shape=(2 * iterations,), dtype=wp.vec3)
- face_norm2 = wp.array(shape=(2 * iterations,), dtype=float)
- face_index = wp.array(shape=(2 * iterations,), dtype=int)
- face_map = wp.array(shape=(2 * iterations,), dtype=int)
- horizon = wp.array(shape=(2 * iterations,), dtype=int)
+ face = wp.array(shape=(6 * iterations,), dtype=wp.vec3i)
+ face_pr = wp.array(shape=(6 * iterations,), dtype=wp.vec3)
+ face_norm2 = wp.array(shape=(6 * iterations,), dtype=float)
+ face_index = wp.array(shape=(6 * iterations,), dtype=int)
+ face_map = wp.array(shape=(6 * iterations,), dtype=int)
+ horizon = wp.array(shape=(6 * iterations,), dtype=int)
dist_out = wp.array(shape=(1,), dtype=float)
+ ncon_out = wp.array(shape=(1,), dtype=int)
pos_out = wp.array(shape=(2,), dtype=wp.vec3)
wp.launch(
_gjk_kernel,
@@ -151,6 +183,15 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
m.mesh_vertadr,
m.mesh_vertnum,
m.mesh_vert,
+ m.mesh_polynum,
+ m.mesh_polyadr,
+ m.mesh_polynormal,
+ m.mesh_polyvertadr,
+ m.mesh_polyvertnum,
+ m.mesh_polyvert,
+ m.mesh_polymapadr,
+ m.mesh_polymapnum,
+ m.mesh_polymap,
d.geom_xpos,
d.geom_xmat,
gid1,
@@ -170,10 +211,11 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
],
outputs=[
dist_out,
+ ncon_out,
pos_out,
],
)
- return dist_out.numpy()[0], pos_out.numpy()[0], pos_out.numpy()[1]
+ return dist_out.numpy()[0], ncon_out.numpy()[0], pos_out.numpy()[0], pos_out.numpy()[1]
class GJKTest(absltest.TestCase):
@@ -193,7 +235,7 @@ class GJKTest(absltest.TestCase):
"""
)
- dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ dist, _, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
self.assertEqual(1.0, dist)
def test_spheres_touching(self):
@@ -210,7 +252,7 @@ class GJKTest(absltest.TestCase):
"""
)
- dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ dist, _, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
self.assertEqual(0.0, dist)
def test_box_mesh_distance(self):
@@ -238,7 +280,7 @@ class GJKTest(absltest.TestCase):
"""
)
- dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ dist, _, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
self.assertAlmostEqual(0.1, dist)
def test_sphere_sphere_contact(self):
@@ -255,7 +297,7 @@ class GJKTest(absltest.TestCase):
"""
)
- dist, _, _ = _geom_dist(m, d, 0, 1, 0)
+ dist, _, _, _ = _geom_dist(m, d, 0, 1, 0)
self.assertAlmostEqual(-2, dist)
def test_box_box_contact(self):
@@ -271,7 +313,7 @@ class GJKTest(absltest.TestCase):
"""
)
- dist, x1, x2 = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ dist, _, x1, x2 = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
self.assertAlmostEqual(-1, dist)
normal = wp.normalize(x1 - x2)
self.assertAlmostEqual(normal[0], 1)
@@ -312,9 +354,95 @@ class GJKTest(absltest.TestCase):
"""
)
- dist, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
+ dist, _, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS)
self.assertAlmostEqual(-0.01, dist)
+ def test_cylinder_cylinder_contact(self):
+ """Test penetration between two cylinder."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+ """
+ )
+
+ dist, _, _, _ = _geom_dist(m, d, 0, 1, 50)
+ self.assertAlmostEqual(-0.001, dist)
+
+ def test_box_edge(self):
+ """Test box edge."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+ """
+ )
+ _, ncon, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS, multiccd=True)
+ self.assertEqual(ncon, 2)
+
+ def test_box_box_ccd(self):
+ """Test box box."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+ """
+ )
+ _, ncon, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS, multiccd=True)
+ self.assertEqual(ncon, 4)
+
+ def test_mesh_mesh_ccd(self):
+ """Test mesh-mesh multiccd."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+
+
+
+ """
+ )
+
+ _, ncon, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS, multiccd=True)
+ self.assertEqual(ncon, 5)
+
+ def test_box_box_ccd2(self):
+ """Test box-box multiccd 2."""
+
+ _, _, m, d = test_util.fixture(
+ xml=f"""
+
+
+
+
+
+
+ """
+ )
+
+ _, ncon, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS, multiccd=True)
+ self.assertEqual(ncon, 5)
+
if __name__ == "__main__":
wp.init()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py
index b0e6de66..6ab7236e 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py
@@ -20,6 +20,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_point
from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_to_segment_points
from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame
from mujoco.mjx.third_party.mujoco_warp._src.math import normalize_with_norm
+from mujoco.mjx.third_party.mujoco_warp._src.math import safe_div
from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
@@ -1275,7 +1276,7 @@ def sphere_cylinder(
return
# Corner collision
- inv_len = 1.0 / wp.sqrt(p_proj_sqr)
+ inv_len = safe_div(1.0, wp.sqrt(p_proj_sqr))
p_proj = p_proj * (cylinder.size[0] * inv_len)
cap_offset = axis * (wp.sign(x) * cylinder.size[1])
@@ -1365,7 +1366,7 @@ def plane_cylinder(
# Otherwise use cylinder's x-axis scaled by radius
vec = wp.where(
len_sqr >= 1e-12,
- vec * (cylinder.size[0] / wp.sqrt(len_sqr)),
+ vec * safe_div(cylinder.size[0], wp.sqrt(len_sqr)),
wp.vec3(cylinder.rot[0, 0], cylinder.rot[1, 0], cylinder.rot[2, 0]) * cylinder.size[0],
)
@@ -1554,26 +1555,32 @@ def contact_params(
g1 = geoms[0]
g2 = geoms[1]
- p1 = geom_priority[g1]
- p2 = geom_priority[g2]
-
solmix1 = geom_solmix[worldid, g1]
solmix2 = geom_solmix[worldid, g2]
- mix = solmix1 / (solmix1 + solmix2)
- mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 < MJ_MINVAL), 0.5, mix)
- mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 >= MJ_MINVAL), 0.0, mix)
- mix = wp.where((solmix1 >= MJ_MINVAL) and (solmix2 < MJ_MINVAL), 1.0, mix)
- mix = wp.where(p1 == p2, mix, wp.where(p1 > p2, 1.0, 0.0))
-
- margin = wp.max(geom_margin[worldid, g1], geom_margin[worldid, g2])
- gap = wp.max(geom_gap[worldid, g1], geom_gap[worldid, g2])
-
condim1 = geom_condim[g1]
condim2 = geom_condim[g2]
- condim = wp.where(p1 == p2, wp.max(condim1, condim2), wp.where(p1 > p2, condim1, condim2))
- max_geom_friction = wp.max(geom_friction[worldid, g1], geom_friction[worldid, g2])
+ # priority
+ p1 = geom_priority[g1]
+ p2 = geom_priority[g2]
+
+ if p1 > p2:
+ mix = 1.0
+ condim = condim1
+ max_geom_friction = geom_friction[worldid, g1]
+ elif p2 > p1:
+ mix = 0.0
+ condim = condim2
+ max_geom_friction = geom_friction[worldid, g2]
+ else:
+ mix = safe_div(solmix1, solmix1 + solmix2)
+ mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 < MJ_MINVAL), 0.5, mix)
+ mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 >= MJ_MINVAL), 0.0, mix)
+ mix = wp.where((solmix1 >= MJ_MINVAL) and (solmix2 < MJ_MINVAL), 1.0, mix)
+ condim = wp.max(condim1, condim2)
+ max_geom_friction = wp.max(geom_friction[worldid, g1], geom_friction[worldid, g2])
+
friction = vec5(
wp.max(MJ_MINMU, max_geom_friction[0]),
wp.max(MJ_MINMU, max_geom_friction[0]),
@@ -1582,7 +1589,7 @@ def contact_params(
wp.max(MJ_MINMU, max_geom_friction[2]),
)
- if geom_solref[worldid, g1].x > 0.0 and geom_solref[worldid, g2].x > 0.0:
+ if geom_solref[worldid, g1][0] > 0.0 and geom_solref[worldid, g2][0] > 0.0:
solref = mix * geom_solref[worldid, g1] + (1.0 - mix) * geom_solref[worldid, g2]
else:
solref = wp.min(geom_solref[worldid, g1], geom_solref[worldid, g2])
@@ -1591,6 +1598,10 @@ def contact_params(
solimp = mix * geom_solimp[worldid, g1] + (1.0 - mix) * geom_solimp[worldid, g2]
+ # geom priority is ignored
+ margin = wp.max(geom_margin[worldid, g1], geom_margin[worldid, g2])
+ gap = wp.max(geom_gap[worldid, g1], geom_gap[worldid, g2])
+
return geoms, margin, gap, condim, friction, solref, solreffriction, solimp
@@ -1640,13 +1651,13 @@ def _sphere_box(
closest = 2.0 * (box_size[0] + box_size[1] + box_size[2])
k = wp.int32(0)
for i in range(6):
- face_dist = wp.abs(wp.where(i % 2, 1.0, -1.0) * box_size[i / 2] - center[i / 2])
+ face_dist = wp.abs(wp.where(i % 2, 1.0, -1.0) * box_size[i // 2] - center[i // 2])
if closest > face_dist:
closest = face_dist
k = i
nearest = wp.vec3(0.0)
- nearest[k / 2] = wp.where(k % 2, -1.0, 1.0)
+ nearest[k // 2] = wp.where(k % 2, -1.0, 1.0)
pos = center + nearest * (sphere_size - closest) / 2.0
contact_normal = box_rot @ nearest
contact_dist = -closest - sphere_size
@@ -1882,22 +1893,22 @@ def capsule_box(
if x1 > 1:
x1 = 1.0
s1 = 2
- x2 = (v - mb) / mc
+ x2 = safe_div(v - mb, mc)
elif x1 < -1:
x1 = -1.0
s1 = 0
- x2 = (v + mb) / mc
+ x2 = safe_div(v + mb, mc)
x2_over = x2 > 1.0
if x2_over or x2 < -1.0:
if x2_over:
x2 = 1.0
s2 = 2
- x1 = (u - mb) / ma
+ x1 = safe_div(u - mb, ma)
else:
x2 = -1.0
s2 = 0
- x1 = (u + mb) / ma
+ x1 = safe_div(u + mb, ma)
if x1 > 1:
x1 = 1.0
@@ -1918,7 +1929,7 @@ def capsule_box(
bestsegmentpos = x2
bestboxpos = x1
# ct<6 means closest point on box is at lower end or middle of edge
- c2 = ct / 6
+ c2 = ct // 6
clcorner = i + (1 << j) * c2 # index of closest box corner
cledge = j # axis index of closest box edge
@@ -1951,7 +1962,7 @@ def capsule_box(
if cltype == -4: # invalid type
return
- if cltype >= 0 and cltype / 3 != 1: # closest to a corner of the box
+ if cltype >= 0 and cltype // 3 != 1: # closest to a corner of the box
c1 = axisdir ^ clcorner
# Calculate relative orientation between capsule and corner
# There are two possible configurations:
@@ -1981,18 +1992,18 @@ def capsule_box(
ax2 = 1
if axis[ax] * axis[ax] > 0.5: # second point along the edge of the box
- m = 2.0 * box.size[ax] / wp.abs(halfaxis[ax])
+ m = 2.0 * safe_div(box.size[ax], wp.abs(halfaxis[ax]))
secondpos = min(1.0 - wp.float32(mul) * bestsegmentpos, m)
else: # second point along a face of the box
# check for overshoot again
m = 2.0 * min(
- box.size[ax1] / wp.abs(halfaxis[ax1]),
- box.size[ax2] / wp.abs(halfaxis[ax2]),
+ safe_div(box.size[ax1], wp.abs(halfaxis[ax1])),
+ safe_div(box.size[ax2], wp.abs(halfaxis[ax2])),
)
secondpos = -min(1.0 + wp.float32(mul) * bestsegmentpos, m)
secondpos *= wp.float32(mul)
- elif cltype >= 0 and cltype / 3 == 1: # we are on box's edge
+ elif cltype >= 0 and cltype // 3 == 1: # we are on box's edge
# Calculate relative orientation between capsule and edge
# Two possible configurations:
# - T configuration: c1 = 2^n (no additional contacts)
@@ -2028,7 +2039,7 @@ def capsule_box(
# now find out whether we point towards the opposite side or towards one of the sides
# and also find the farthest point along the capsule that is above the box
- e1 = 2.0 * box.size[ax2] / wp.abs(halfaxis[ax2])
+ e1 = 2.0 * safe_div(box.size[ax2], wp.abs(halfaxis[ax2]))
secondpos = min(e1, secondpos)
if ((axisdir & (1 << ax)) != 0) == ((c1 & (1 << ax2)) != 0):
@@ -2036,7 +2047,7 @@ def capsule_box(
else:
e2 = 1.0 + bestboxpos
- e1 = box.size[ax] * e2 / wp.abs(halfaxis[ax])
+ e1 = box.size[ax] * safe_div(e2, wp.abs(halfaxis[ax]))
secondpos = min(e1, secondpos)
secondpos *= wp.float32(mul)
@@ -2055,7 +2066,7 @@ def capsule_box(
for i in range(3):
if i != clface:
- ha_r = wp.float32(mul) / halfaxis[i]
+ ha_r = safe_div(wp.float32(mul), halfaxis[i])
e1 = (box.size[i] - tmp1[i]) * ha_r
if 0 < e1 and e1 < secondpos:
secondpos = e1
@@ -2300,7 +2311,7 @@ def box_box(
if axis_code < 12:
# Handle face-vertex collision
face_idx = axis_code % 6
- box_idx = axis_code / 6
+ box_idx = axis_code // 6
rotmore = _compute_rotmore(face_idx)
r = rotmore @ wp.where(box_idx, rot12, rot21)
@@ -2368,10 +2379,10 @@ def box_box(
bx = cn2[0]
ay = cn1[1]
by = cn2[1]
- C = 1.0 / (ax * by - bx * ay)
+ C = safe_div(1.0, ax * by - bx * ay)
for i in range(4):
- llx = wp.where(i / 2, lx, -lx)
+ llx = wp.where(i // 2, lx, -lx)
lly = wp.where(i % 2, ly, -ly)
x = llx - lp[0]
@@ -2410,7 +2421,7 @@ def box_box(
else:
# Handle edge-edge collision
- edge1 = (axis_code - 12) / 3
+ edge1 = (axis_code - 12) // 3
edge2 = (axis_code - 12) % 3
# Set up non-contacting edges ax1, ax2 for box2 and pax1, pax2 for box 1
@@ -2523,12 +2534,12 @@ def box_box(
bx = pts_cn2[0]
ay = pts_cn1[1]
by = pts_cn2[1]
- C = 1.0 / (ax * by - bx * ay)
+ C = safe_div(1.0, ax * by - bx * ay)
for i in range(4):
if n == max_con_pair:
break
- llx = wp.where(i / 2, lx, -lx)
+ llx = wp.where(i // 2, lx, -lx)
lly = wp.where(i % 2, ly, -ly)
x = llx - pts_lp[0]
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py
index 9ab71173..6f99949f 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py
@@ -1162,6 +1162,7 @@ def _efc_contact_pyramidal(
efcid = wp.atomic_add(nefc_out, worldid, 1)
if efcid >= njmax_in:
+ contact_efc_address_out[conid, dimid] = -1
return
timestep = opt_timestep[worldid]
@@ -1325,6 +1326,7 @@ def _efc_contact_elliptic(
efcid = wp.atomic_add(nefc_out, worldid, 1)
if efcid >= njmax_in:
+ contact_efc_address_out[conid, dimid] = -1
return
timestep = opt_timestep[worldid]
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py
index 036d2052..2b8a3c26 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint_test.py
@@ -37,6 +37,67 @@ def _assert_eq(a, b, name):
np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
+def _assert_efc_eq(d, mjd, nefc, name):
+ """Assert equality of efc fields after sorting both sides."""
+ # Get the ordering indices based on efc_type, efc_pos, efc_vel, efc_aref, efc_d for MJWarp
+ efc_type = d.efc.type.numpy()[0, :nefc]
+ efc_pos = d.efc.pos.numpy()[0, :nefc]
+ efc_vel = d.efc.vel.numpy()[0, :nefc]
+ efc_aref = d.efc.aref.numpy()[0, :nefc]
+ efc_d = d.efc.D.numpy()[0, :nefc]
+ # Get the ordering indices based on efc_type, efc_pos, efc_vel, efc_aref, efc_d for MuJoCo
+ mjd_efc_type = mjd.efc_type[:nefc]
+ mjd_efc_pos = mjd.efc_pos[:nefc]
+ mjd_efc_vel = mjd.efc_vel[:nefc]
+ mjd_efc_aref = mjd.efc_aref[:nefc]
+ mjd_efc_d = mjd.efc_D[:nefc]
+
+ # Create sorting keys using lexsort (more efficient for multiple keys)
+ d_sort_indices = np.lexsort((efc_pos, efc_type, efc_vel, efc_aref, efc_d))
+ mjd_sort_indices = np.lexsort((mjd_efc_pos, mjd_efc_type, mjd_efc_vel, mjd_efc_aref, mjd_efc_d))
+
+ # Sort MJWarp efc fields
+ d_sorted = d.efc.J.numpy()[0, d_sort_indices, :].reshape(-1)
+
+ # Sort MuJoCo efc fields
+ # For J matrix, need to reshape to 2D, sort rows, then flatten
+ nefc = len(mjd_sort_indices)
+ nv = mjd.efc_J.shape[0] // nefc if nefc > 0 else 0
+ if nv > 0:
+ mjd_J_2d = mjd.efc_J.reshape(nefc, nv)
+ mjd_sorted_J = mjd_J_2d[mjd_sort_indices].reshape(-1)
+ else:
+ mjd_sorted_J = mjd.efc_J
+
+ mjd_sorted_D = mjd.efc_D[mjd_sort_indices]
+ mjd_sorted_vel = mjd.efc_vel[mjd_sort_indices]
+ mjd_sorted_aref = mjd.efc_aref[mjd_sort_indices]
+ mjd_sorted_pos = mjd.efc_pos[mjd_sort_indices]
+ mjd_sorted_margin = mjd.efc_margin[mjd_sort_indices]
+ mjd_sorted_type = mjd.efc_type[mjd_sort_indices]
+
+ # Compare sorted data
+ _assert_eq(d_sorted, mjd_sorted_J, f"{name}_J")
+
+ d_sorted = d.efc.D.numpy()[0, d_sort_indices]
+ _assert_eq(d_sorted, mjd_sorted_D, f"{name}_D")
+
+ d_sorted = d.efc.vel.numpy()[0, d_sort_indices]
+ _assert_eq(d_sorted, mjd_sorted_vel, f"{name}_vel")
+
+ d_sorted = d.efc.aref.numpy()[0, d_sort_indices]
+ _assert_eq(d_sorted, mjd_sorted_aref, f"{name}_aref")
+
+ d_sorted = d.efc.pos.numpy()[0, d_sort_indices]
+ _assert_eq(d_sorted, mjd_sorted_pos, f"{name}_pos")
+
+ d_sorted = d.efc.margin.numpy()[0, d_sort_indices]
+ _assert_eq(d_sorted, mjd_sorted_margin, f"{name}_margin")
+
+ d_sorted = d.efc.type.numpy()[0, d_sort_indices]
+ _assert_eq(d_sorted, mjd_sorted_type, f"{name}_type")
+
+
class ConstraintTest(parameterized.TestCase):
@parameterized.parameters(
(ConeType.PYRAMIDAL, 1, 1),
@@ -108,7 +169,7 @@ class ConstraintTest(parameterized.TestCase):
for arr in (d.ne, d.nefc, d.nf, d.nl, d.efc.type):
arr.fill_(-1)
- for arr in (d.efc.J, d.efc.D, d.efc.aref, d.efc.pos, d.efc.margin):
+ for arr in (d.efc.J, d.efc.D, d.efc.vel, d.efc.aref, d.efc.pos, d.efc.margin):
arr.fill_(wp.nan)
mjwarp.make_constraint(m, d)
@@ -117,13 +178,7 @@ class ConstraintTest(parameterized.TestCase):
_assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc")
_assert_eq(d.nf.numpy()[0], mjd.nf, "nf")
_assert_eq(d.nl.numpy()[0], mjd.nl, "nl")
- _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J")
- _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D")
- _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel")
- _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref")
- _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos")
- _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin")
- _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type")
+ _assert_efc_eq(d, mjd, mjd.nefc, "efc")
def test_limit_tendon(self):
"""Test limit tendon constraints."""
@@ -132,20 +187,14 @@ class ConstraintTest(parameterized.TestCase):
for arr in (d.nefc, d.nl, d.efc.type):
arr.fill_(-1)
- for arr in (d.efc.J, d.efc.D, d.efc.aref, d.efc.pos, d.efc.margin):
+ for arr in (d.efc.J, d.efc.D, d.efc.vel, d.efc.aref, d.efc.pos, d.efc.margin):
arr.fill_(wp.nan)
mjwarp.make_constraint(m, d)
_assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc")
_assert_eq(d.nl.numpy()[0], mjd.nl, "nl")
- _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J")
- _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D")
- _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel")
- _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref")
- _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos")
- _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin")
- _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type")
+ _assert_efc_eq(d, mjd, mjd.nefc, "efc")
def test_equality_tendon(self):
"""Test equality tendon constraints."""
@@ -202,13 +251,7 @@ class ConstraintTest(parameterized.TestCase):
_assert_eq(d.nefc.numpy()[0], mjd.nefc, "nefc")
_assert_eq(d.ne.numpy()[0], mjd.ne, "ne")
- _assert_eq(d.efc.J.numpy()[0, : mjd.nefc, :].reshape(-1), mjd.efc_J, "efc_J")
- _assert_eq(d.efc.D.numpy()[0, : mjd.nefc], mjd.efc_D, "efc_D")
- _assert_eq(d.efc.vel.numpy()[0, : mjd.nefc], mjd.efc_vel, "efc_vel")
- _assert_eq(d.efc.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref")
- _assert_eq(d.efc.pos.numpy()[0, : mjd.nefc], mjd.efc_pos, "efc_pos")
- _assert_eq(d.efc.margin.numpy()[0, : mjd.nefc], mjd.efc_margin, "efc_margin")
- _assert_eq(d.efc.type.numpy()[0, : mjd.nefc], mjd.efc_type, "efc_type")
+ _assert_efc_eq(d, mjd, mjd.nefc, "efc")
if __name__ == "__main__":
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py
index ad07b5c4..3ae65b12 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py
@@ -1025,7 +1025,10 @@ def forward(m: Model, d: Data):
@event_scope
def step(m: Model, d: Data):
"""Advance simulation."""
+ # TODO(team): mj_checkPos
+ # TODO(team): mj_checkVel
forward(m, d)
+ # TODO(team): mj_checkAcc
if m.opt.integrator == IntegratorType.EULER:
euler(m, d)
@@ -1035,3 +1038,44 @@ def step(m: Model, d: Data):
implicit(m, d)
else:
raise NotImplementedError(f"integrator {m.opt.integrator} not implemented.")
+
+
+@event_scope
+def step1(m: Model, d: Data):
+ """Advance simulation in two phases: before input is set by user."""
+ energy = m.opt.enableflags & EnableBit.ENERGY
+ # TODO(team): mj_checkPos
+ # TODO(team): mj_checkVel
+ fwd_position(m, d)
+ sensor.sensor_pos(m, d)
+
+ if energy:
+ if m.sensor_e_potential == 0: # not computed by sensor
+ sensor.energy_pos(m, d)
+ else:
+ wp.launch(_zero_energy, dim=d.nworld, inputs=[d.energy])
+
+ fwd_velocity(m, d)
+ sensor.sensor_vel(m, d)
+
+ if energy:
+ if m.sensor_e_kinetic == 0: # not computed by sensor
+ sensor.energy_vel(m, d)
+
+
+@event_scope
+def step2(m: Model, d: Data):
+ """Advance simulation in two phases: after input is set by user."""
+ fwd_actuation(m, d)
+ fwd_acceleration(m, d)
+ solver.solve(m, d)
+ sensor.sensor_acc(m, d)
+ # TODO(team): mj_checkAcc
+
+ # integrate with Euler or implicitfast
+ # TODO(team): implicit
+ if m.opt.integrator == IntegratorType.IMPLICITFAST:
+ implicit(m, d)
+ else:
+ # note: RK4 defaults to Euler
+ euler(m, d)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py
index f18d4150..fa428f46 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py
@@ -279,6 +279,189 @@ class ForwardTest(parameterized.TestCase):
_assert_eq(d.actuator_force.numpy()[0], mjd.actuator_force, "actuator_force")
+ @parameterized.parameters(("humanoid/humanoid.xml", True), ("humanoid/humanoid.xml", False))
+ def test_step1(self, xml, energy):
+ # TODO(team): test more mjcfs
+ mjm, mjd, m, d = test_util.fixture(xml, kick=True, energy=energy)
+
+ # some of the fields updated by step1
+ step1_field = [
+ "xpos",
+ "xquat",
+ "xmat",
+ "xipos",
+ "ximat",
+ "xanchor",
+ "xaxis",
+ "geom_xpos",
+ "geom_xmat",
+ "site_xmat",
+ "subtree_com",
+ "cinert",
+ "cdof",
+ "cam_xpos",
+ "cam_xmat",
+ "light_xpos",
+ "light_xdir",
+ "ten_length",
+ "ten_J",
+ "ten_wrapadr",
+ "ten_wrapnum",
+ "wrap_obj",
+ "wrap_xpos",
+ "qM",
+ "qLD",
+ "nefc",
+ "efc_type",
+ "efc_id",
+ "efc_J",
+ "efc_pos",
+ "efc_margin",
+ "efc_D",
+ "efc_vel",
+ "efc_aref",
+ "efc_frictionloss",
+ "actuator_length",
+ "actuator_moment",
+ "actuator_velocity",
+ "ten_velocity",
+ "cvel",
+ "cdof_dot",
+ "qfrc_spring",
+ "qfrc_damper",
+ "qfrc_gravcomp",
+ "qfrc_fluid",
+ "qfrc_passive",
+ "qfrc_bias",
+ "energy",
+ ]
+ if m.nflexvert:
+ step1_field += ["flexvert_xpos"]
+ if m.nflexedge:
+ step1_field += ["flexedge_length", "flexedge_velocity"]
+
+ def _getattr(arr):
+ if (len(arr) >= 4) & (arr[:4] == "efc_"):
+ return getattr(d.efc, arr[4:]), True
+ return getattr(d, arr), False
+
+ for arr in step1_field:
+ attr, _ = _getattr(arr)
+ if attr.dtype == float:
+ attr.fill_(wp.nan)
+ elif attr.dtype == int:
+ attr.fill_(-1)
+ else:
+ attr.zero_()
+
+ mujoco.mj_step1(mjm, mjd)
+ mjwarp.step1(m, d)
+
+ for arr in step1_field:
+ d_arr, is_nefc = _getattr(arr)
+ d_arr = d_arr.numpy()[0]
+ mjd_arr = getattr(mjd, arr)
+ if arr in ["xmat", "ximat", "geom_xmat", "site_xmat", "cam_xmat"]:
+ mjd_arr = mjd_arr.reshape(-1)
+ d_arr = d_arr.reshape(-1)
+ elif arr == "qM":
+ qM = np.zeros((mjm.nv, mjm.nv))
+ mujoco.mj_fullM(mjm, qM, mjd.qM)
+ mjd_arr = qM
+ elif arr == "actuator_moment":
+ actuator_moment = np.zeros((mjm.nu, mjm.nv))
+ mujoco.mju_sparse2dense(actuator_moment, mjd.actuator_moment, mjd.moment_rownnz, mjd.moment_rowadr, mjd.moment_colind)
+ mjd_arr = actuator_moment
+ elif arr == "ten_J" and mjm.ntendon:
+ ten_J = np.zeros((mjm.ntendon, mjm.nv))
+ mujoco.mju_sparse2dense(ten_J, mjd.ten_J, mjd.ten_J_rownnz, mjd.ten_J_rowadr, mjd.ten_J_colind)
+ mjd_arr = ten_J
+ elif arr == "efc_J":
+ if mjd.efc_J.shape[0] != mjd.nefc * mjm.nv:
+ efc_J = np.zeros((mjd.nefc, mjm.nv))
+ mujoco.mju_sparse2dense(efc_J, mjd.efc_J, mjd.efc_J_rownnz, mjd.efc_J_rowadr, mjd.efc_J_colind)
+ mjd_arr = efc_J
+ else:
+ mjd_arr = mjd_arr.reshape((mjd.nefc, mjm.nv))
+ elif arr == "qLD":
+ vec = np.ones((1, mjm.nv))
+ res = np.zeros((1, mjm.nv))
+ mujoco.mj_solveM(mjm, mjd, res, vec)
+
+ vec_wp = wp.array(vec, dtype=float)
+ res_wp = wp.zeros((1, mjm.nv), dtype=float)
+ mjwarp.solve_m(m, d, res_wp, vec_wp)
+
+ d_arr = res_wp.numpy()[0]
+ mjd_arr = res[0]
+ if is_nefc:
+ d_arr = d_arr[: d.nefc.numpy()[0]]
+
+ _assert_eq(d_arr, mjd_arr, arr)
+
+ # TODO(team): sensor_pos
+ # TODO(team): sensor_vel
+
+ @parameterized.parameters(
+ ("humanoid/humanoid.xml", IntegratorType.EULER),
+ ("humanoid/humanoid.xml", IntegratorType.IMPLICITFAST),
+ ("humanoid/humanoid.xml", IntegratorType.RK4),
+ )
+ def test_step2(self, xml, integrator):
+ mjm, mjd, m, _ = test_util.fixture(xml, kick=True, integrator=integrator)
+
+ # some of the fields updated by step2
+ step2_field = [
+ "act_dot",
+ "actuator_force",
+ "qfrc_actuator",
+ "qfrc_smooth",
+ "qacc",
+ "qvel",
+ "qpos",
+ "efc_force",
+ "qfrc_constraint",
+ ]
+
+ def _getattr(arr):
+ if (len(arr) >= 4) & (arr[:4] == "efc_"):
+ return getattr(d.efc, arr[4:]), True
+ return getattr(d, arr), False
+
+ mujoco.mj_step1(mjm, mjd)
+
+ # input
+ ctrl = 0.1 * np.random.rand(mjm.nu)
+ qfrc_applied = 0.1 * np.random.rand(mjm.nv)
+ xfrc_applied = 0.1 * np.random.rand(mjm.nbody, 6)
+
+ mjd.ctrl = ctrl
+ mjd.qfrc_applied = qfrc_applied
+ mjd.xfrc_applied = xfrc_applied
+
+ d = mjwarp.put_data(mjm, mjd)
+
+ for arr in step2_field:
+ if arr in ["qpos", "qvel"]:
+ continue
+ attr, _ = _getattr(arr)
+ if attr.dtype == float:
+ attr.fill_(wp.nan)
+ elif attr.dtype == int:
+ attr.fill_(-1)
+ else:
+ attr.zero_()
+
+ mujoco.mj_step2(mjm, mjd)
+ mjwarp.step2(m, d)
+
+ for arr in step2_field:
+ d_arr, is_efc = _getattr(arr)
+ d_arr = d_arr.numpy()[0]
+ if is_efc:
+ d_arr = d_arr[: d.nefc.numpy()[0]]
+ _assert_eq(d_arr, getattr(mjd, arr), arr)
+
if __name__ == "__main__":
wp.init()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py
index 0c9dbb25..d786de4c 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py
@@ -396,6 +396,9 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
rangefinder_sensor_adr = np.full(mjm.nsensor, -1)
rangefinder_sensor_adr[sensor_rangefinder_adr] = np.arange(len(sensor_rangefinder_adr))
+ # contact sensor
+ sensor_adr_to_contact_adr = np.clip(np.cumsum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT) - 1, a_min=0, a_max=None)
+
# TODO(team): improve heuristic for selecting broadphase routine
if mjm.ngeom > 1000:
broadphase = types.BroadphaseType.SAP_SEGMENTED
@@ -800,6 +803,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
[mujoco.mjtSensor.mjSENS_SUBTREELINVEL, mujoco.mjtSensor.mjSENS_SUBTREEANGMOM],
).any(),
sensor_contact_adr=wp.array(np.nonzero(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT)[0], dtype=int),
+ sensor_adr_to_contact_adr=wp.array(sensor_adr_to_contact_adr, dtype=int),
sensor_rne_postconstraint=np.isin(
mjm.sensor_type,
[
@@ -870,21 +874,25 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
if nworld < 1 or nworld > MAX_WORLDS:
raise ValueError(f"nworld must be >= 1 and <= {MAX_WORLDS}")
- if nconmax < 1:
- raise ValueError("nconmax must be >= 1")
+ if nconmax < 0:
+ raise ValueError("nconmax must be >= 0")
- if njmax < 1:
- raise ValueError("njmax must be >= 1")
+ if njmax < 0:
+ raise ValueError("njmax must be >= 0")
condim = np.concatenate((mjm.geom_condim, mjm.pair_dim))
condim_max = np.max(condim) if len(condim) > 0 else 0
if mujoco.mj_isSparse(mjm):
qM = wp.zeros((nworld, 1, mjm.nM), dtype=float)
- qLD = wp.zeros((nworld, 1, mjm.nM), dtype=float)
+ qLD = wp.zeros((nworld, 1, mjm.nC), dtype=float)
+ qM_integration = wp.zeros((nworld, 1, mjm.nM), dtype=float)
+ qLD_integration = wp.zeros((nworld, 1, mjm.nM), dtype=float)
else:
qM = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float)
qLD = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float)
+ qM_integration = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float)
+ qLD_integration = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float)
nsensorcontact = np.sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT)
nrangefinder = sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER)
@@ -1009,7 +1017,7 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
gauss=wp.zeros((nworld,), dtype=float),
cost=wp.zeros((nworld,), dtype=float),
prev_cost=wp.zeros((nworld,), dtype=float),
- active=wp.zeros((nworld, njmax), dtype=bool),
+ state=wp.zeros((nworld, njmax), dtype=int),
gtol=wp.zeros((nworld,), dtype=float),
mv=wp.zeros((nworld, mjm.nv), dtype=float),
jv=wp.zeros((nworld, njmax), dtype=float),
@@ -1035,12 +1043,6 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
mid=wp.zeros((nworld,), dtype=wp.vec3),
mid_alpha=wp.zeros((nworld,), dtype=float),
cost_candidate=wp.zeros((nworld, mjm.opt.ls_iterations), dtype=float),
- # elliptic cone
- u=wp.zeros((nconmax,), dtype=types.vec6),
- uu=wp.zeros((nconmax,), dtype=float),
- uv=wp.zeros((nconmax,), dtype=float),
- vv=wp.zeros((nconmax,), dtype=float),
- condim=wp.zeros((nworld, njmax), dtype=int),
),
# RK4
qpos_t0=wp.zeros((nworld, mjm.nq), dtype=float),
@@ -1053,8 +1055,8 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
qfrc_integration=wp.zeros((nworld, mjm.nv), dtype=float),
qacc_integration=wp.zeros((nworld, mjm.nv), dtype=float),
act_vel_integration=wp.zeros((nworld, mjm.nu), dtype=float),
- qM_integration=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float),
- qLD_integration=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float),
+ qM_integration=qM_integration,
+ qLD_integration=qLD_integration,
qLDiagInv_integration=wp.zeros((nworld, mjm.nv), dtype=float),
# sweep-and-prune broadphase
sap_projection_lower=wp.zeros((nworld, mjm.ngeom, 2), dtype=float),
@@ -1153,11 +1155,11 @@ def put_data(
if nworld < 1 or nworld > MAX_WORLDS:
raise ValueError(f"nworld must be >= 1 and <= {MAX_WORLDS}")
- if nconmax < 1:
- raise ValueError("nconmax must be >= 1")
+ if nconmax < 0:
+ raise ValueError("nconmax must be >= 0")
- if njmax < 1:
- raise ValueError("njmax must be >= 1")
+ if njmax < 0:
+ raise ValueError("njmax must be >= 0")
if nworld * mjd.ncon > nconmax:
raise ValueError(f"nconmax overflow (nconmax must be >= {nworld * mjd.ncon})")
@@ -1394,7 +1396,7 @@ def put_data(
gauss=wp.empty(shape=(nworld,), dtype=float),
cost=wp.empty(shape=(nworld,), dtype=float),
prev_cost=wp.empty(shape=(nworld,), dtype=float),
- active=wp.empty(shape=(nworld, njmax), dtype=bool),
+ state=wp.empty(shape=(nworld, njmax), dtype=int),
gtol=wp.empty(shape=(nworld,), dtype=float),
mv=wp.empty(shape=(nworld, mjm.nv), dtype=float),
jv=wp.empty(shape=(nworld, njmax), dtype=float),
@@ -1419,12 +1421,6 @@ def put_data(
mid=wp.empty(shape=(nworld,), dtype=wp.vec3),
mid_alpha=wp.empty(shape=(nworld,), dtype=float),
cost_candidate=wp.empty(shape=(nworld, mjm.opt.ls_iterations), dtype=float),
- # TODO(team): skip allocation if not elliptic
- u=wp.empty((nconmax,), dtype=types.vec6),
- uu=wp.empty((nconmax,), dtype=float),
- uv=wp.empty((nconmax,), dtype=float),
- vv=wp.empty((nconmax,), dtype=float),
- condim=wp.empty((nworld, njmax), dtype=int),
),
# TODO(team): skip allocation if integrator != RK4
qpos_t0=wp.empty((nworld, mjm.nq), dtype=float),
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py
index 8599d025..b197dd9a 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py
@@ -30,6 +30,14 @@ import mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import test_util
from mujoco.mjx.third_party.mujoco_warp._src.io import MAX_WORLDS
+_IO_TEST_MODELS = (
+ "pendula.xml",
+ "collision_sdf/tactile.xml",
+ "flex/cloth.xml",
+ "actuation/tendon_force_limit.xml",
+ "hfield/hfield.xml",
+)
+
def _dims_match(test_obj, d1: Any, d2: Any, prefix: str = ""):
"""Checks that two dataclasses have fields with the same leading dims."""
@@ -43,8 +51,8 @@ def _dims_match(test_obj, d1: Any, d2: Any, prefix: str = ""):
if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array):
s1, s2 = a1.shape, a2.shape
- test_obj.assertEqual(len(s1), len(s2), full_name + f" dims mismatch. Got {s1} and {s2}.")
- test_obj.assertSequenceAlmostEqual(s1, s2, full_name + f" dims mismatch. Got {s1} and {s2}.")
+ test_obj.assertEqual(len(s1), len(s2), f"{full_name} dims mismatch. Got {s1} and {s2}.")
+ test_obj.assertEqual(s1, s2, f"{full_name} dims mismatch. Got {s1} and {s2}.")
def _get_np_scalar_type(val: Any) -> Optional[Union[bool, int, float]]:
@@ -309,18 +317,20 @@ class IOTest(parameterized.TestCase):
self.assertGreater(m1.body_parentid.strides[0], 0)
self.assertLen(m1.body_parentid.strides, m1.body_parentid.ndim)
- def test_put_data_nworld_array(self):
+ @parameterized.parameters(*_IO_TEST_MODELS)
+ def test_put_data_nworld_array(self, xml):
"""Tests that put_data arrays that scale with nworld have leading dim nworld."""
- mjm, mjd, _, _ = test_util.fixture("pendula.xml")
- d1 = mjwarp.put_data(mjm, mjd, nworld=1, nconmax=1_000, njmax=1_000)
- dn = mjwarp.put_data(mjm, mjd, nworld=133, nconmax=1_000, njmax=1_000)
+ mjm, mjd, _, _ = test_util.fixture(xml)
+ d1 = mjwarp.put_data(mjm, mjd, nworld=1, nconmax=3_000, njmax=3_000)
+ dn = mjwarp.put_data(mjm, mjd, nworld=133, nconmax=3_000, njmax=3_000)
_leading_dims_scale_w_nworld(self, d1, dn, 1, 133)
- def test_make_data_nworld_array(self):
+ @parameterized.parameters(*_IO_TEST_MODELS)
+ def test_make_data_nworld_array(self, xml):
"""Tests that make_data arrays that scale with nworld have leading dim nworld."""
- mjm, *_ = test_util.fixture("pendula.xml")
- d1 = mjwarp.make_data(mjm, nworld=1, nconmax=1_000, njmax=1_000)
- dn = mjwarp.make_data(mjm, nworld=133, nconmax=1_000, njmax=1_000)
+ mjm, *_ = test_util.fixture(xml)
+ d1 = mjwarp.make_data(mjm, nworld=1, nconmax=3_000, njmax=3_000)
+ dn = mjwarp.make_data(mjm, nworld=133, nconmax=3_000, njmax=3_000)
_leading_dims_scale_w_nworld(self, d1, dn, 1, 133)
def test_public_api_jax_compat(self):
@@ -328,9 +338,10 @@ class IOTest(parameterized.TestCase):
_check_annotation_compat(mjwarp.Model.__annotations__, "Model.")
_check_annotation_compat(mjwarp.Data.__annotations__, "Data.")
- def test_types_match_annotations(self):
+ @parameterized.parameters(*_IO_TEST_MODELS)
+ def test_types_match_annotations(self, xml):
"""Tests that the types of dataclass fields match the annotations."""
- mjm, _, m, d = test_util.fixture("pendula.xml")
+ mjm, _, m, d = test_util.fixture(xml)
_check_type_matches_annotation(self, m, "Model.")
_check_type_matches_annotation(self, d, "Data.")
@@ -338,14 +349,15 @@ class IOTest(parameterized.TestCase):
d = mjwarp.make_data(mjm, nworld=2)
_check_type_matches_annotation(self, d, "Data.")
- def test_make_put_data_dims_match(self):
+ @parameterized.parameters(*_IO_TEST_MODELS)
+ def test_make_put_data_dims_match(self, xml):
"""Tests that make_data and put_data have matching dimensions."""
- mjm, mjd, _, _ = test_util.fixture("pendula.xml")
- dm2 = mjwarp.make_data(mjm, nworld=2, nconmax=13, njmax=42)
- dm3 = mjwarp.make_data(mjm, nworld=3, nconmax=13, njmax=42)
+ mjm, mjd, _, _ = test_util.fixture(xml)
+ dm2 = mjwarp.make_data(mjm, nworld=2, nconmax=3_000, njmax=4_200)
+ dm3 = mjwarp.make_data(mjm, nworld=3, nconmax=3_000, njmax=4_200)
- dp2 = mjwarp.put_data(mjm, mjd, nworld=2, nconmax=13, njmax=42)
- dp3 = mjwarp.put_data(mjm, mjd, nworld=3, nconmax=13, njmax=42)
+ dp2 = mjwarp.put_data(mjm, mjd, nworld=2, nconmax=3_000, njmax=4_200)
+ dp3 = mjwarp.put_data(mjm, mjd, nworld=3, nconmax=3_000, njmax=4_200)
_dims_match(self, dm2, dp2)
_dims_match(self, dm3, dp3)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py
index 99cc0b20..68c1e29b 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py
@@ -271,7 +271,7 @@ def closest_segment_to_segment_points(a0: wp.vec3, a1: wp.vec3, b0: wp.vec3, b1:
@wp.func
-def safe_div(x: float, y: float) -> float:
+def safe_div(x: Any, y: Any) -> Any:
return x / wp.where(y != 0.0, y, types.MJ_MINVAL)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py
index 114ce4a8..6bc72ce5 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py
@@ -17,6 +17,7 @@ from typing import Tuple
import warp as wp
+from mujoco.mjx.third_party.mujoco_warp._src.math import safe_div
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
@@ -107,7 +108,7 @@ def _ray_quad(a: float, b: float, c: float) -> Tuple[float, wp.vec2]:
det = wp.sqrt(det)
# compute the two solutions
- den = 1.0 / a
+ den = safe_div(1.0, a)
x0 = (-b - det) * den
x1 = (-b + det) * den
x = wp.vec2(x0, x1)
@@ -194,10 +195,7 @@ def _ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp
if x < 0.0:
return wp.inf
- p = wp.vec2(
- lpnt[0] + x * lvec[0],
- lpnt[1] + x * lvec[1],
- )
+ p = wp.vec2(lpnt[0] + x * lvec[0], lpnt[1] + x * lvec[1])
# accept only within rendered rectangle
if (size[0] <= 0.0 or wp.abs(p[0]) <= size[0]) and (size[1] <= 0.0 or wp.abs(p[1]) <= size[1]):
@@ -284,11 +282,7 @@ def _ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec
lpnt, lvec = _ray_map(pos, mat, pnt, vec)
# invert size^2
- s = wp.vec3(
- 1.0 / (size[0] * size[0]),
- 1.0 / (size[1] * size[1]),
- 1.0 / (size[2] * size[2]),
- )
+ s = wp.vec3(safe_div(1.0, size[0] * size[0]), safe_div(1.0, size[1] * size[1]), safe_div(1.0, size[2] * size[2]))
# (x * lvec + lpnt)' * diag(1 / size^2) * (x * lvec + lpnt) = 1
slvec = wp.cw_mul(s, lvec)
@@ -324,10 +318,7 @@ def _ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec:
# process if non-negative
if sol >= 0.0:
# intersection with horizontal face
- p = wp.vec2(
- lpnt[0] + sol * lvec[0],
- lpnt[1] + sol * lvec[1],
- )
+ p = wp.vec2(lpnt[0] + sol * lvec[0], lpnt[1] + sol * lvec[1])
# accept within radius
if wp.dot(p, p) <= size[0] * size[0]:
@@ -392,7 +383,7 @@ def _ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.v
x = sol
# save in all
- all[2 * i + (side + 1) / 2] = sol
+ all[2 * i + (side + 1) // 2] = sol
return x, all
@@ -459,7 +450,7 @@ def _ray_hfield(
b0[1] = 0.0
else:
b0[2] = 0.0
- b1 = b0 + lvec * -wp.dot(lvec, b0) / wp.dot(lvec, lvec)
+ b1 = b0 + lvec * -safe_div(wp.dot(lvec, b0), wp.dot(lvec, lvec))
b1 = wp.normalize(b1)
b2 = wp.cross(b1, lvec)
@@ -473,10 +464,10 @@ def _ray_hfield(
seg[1] = all[i]
# project segment endpoints in horizontal plane, discretize
- dx = (2.0 * size[0]) / float(ncol - 1)
- dy = (2.0 * size[1]) / float(nrow - 1)
- SX = wp.vec2((lpnt[0] * seg[0] * lvec[0] + size[0]) / dx, (lpnt[0] * seg[1] * lvec[0] + size[0]) / dx)
- SY = wp.vec2((lpnt[1] + seg[0] * lvec[1] + size[1]) / dy, (lpnt[1] + seg[1] * lvec[1] + size[1]) / dy)
+ dx = safe_div(2.0 * size[0], float(ncol - 1))
+ dy = safe_div(2.0 * size[1], float(nrow - 1))
+ SX = wp.vec2(safe_div(lpnt[0] * seg[0] * lvec[0] + size[0], dx), safe_div(lpnt[0] * seg[1] * lvec[0] + size[0], dx))
+ SY = wp.vec2(safe_div(lpnt[1] + seg[0] * lvec[1] + size[1], dy), safe_div(lpnt[1] + seg[1] * lvec[1] + size[1], dy))
# compute ranges, with +1 padding
cmin = wp.max(0, int(wp.floor(wp.min(SX[0], SX[1])) - 1.0))
@@ -511,12 +502,12 @@ def _ray_hfield(
for i in range(4):
if all[i] >= 0.0 and (all[i] < x or x < 0.0):
# normalized height of intersection point
- z = (lpnt[2] + all[i] * lvec[2]) / size[2]
+ z = safe_div(lpnt[2] + all[i] * lvec[2], size[2])
# rectangle points: y, y0, z0, z1
# side normal to x-axis
if i < 2:
- y = (lpnt[1] + all[i] * lvec[1] + size[1]) / dy
+ y = safe_div(lpnt[1] + all[i] * lvec[1] + size[1], dy)
y0 = wp.max(0.0, wp.min(float(nrow - 2), wp.floor(y)))
if i == 1:
z0 = hfield_data[adr + int(wp.round(y0)) * nrow + ncol - 1]
@@ -526,7 +517,7 @@ def _ray_hfield(
z1 = hfield_data[adr + int(wp.round(y0 + 1.0)) * nrow]
# side normal to y-axis
else:
- y = (lpnt[0] + all[i] * lvec[0] + size[0]) / dx
+ y = safe_div(lpnt[0] + all[i] * lvec[0] + size[0], dx)
y0 = wp.max(0.0, wp.min(float(ncol - 2), wp.floor(y)))
if i == 3:
z0 = hfield_data[adr + int(wp.round(y0)) + (nrow - 1) * ncol]
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py
index 48175829..bc561af6 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py
@@ -1565,7 +1565,7 @@ def _sensor_acc(
sensor_adr: wp.array(dtype=int),
sensor_cutoff: wp.array(dtype=float),
sensor_acc_adr: wp.array(dtype=int),
- sensor_contact_adr: wp.array(dtype=int),
+ sensor_adr_to_contact_adr: wp.array(dtype=int),
# Data in:
njmax_in: int,
ncon_in: wp.array(dtype=int),
@@ -1643,16 +1643,9 @@ def _sensor_acc(
num = dim // size # number of slots
adr = sensor_adr[sensorid]
-
- # TODO(team): precompute sensorid to contactsensorid mapping
- contactsensorid = int(0)
- for i in range(sensor_contact_adr.size):
- if sensorid == sensor_contact_adr[i]:
- contactsensorid = i
- break
+ contactsensorid = sensor_adr_to_contact_adr[sensorid]
nmatch = sensor_contact_nmatch_in[worldid, contactsensorid]
-
for i in range(wp.min(nmatch, num)):
# sorted contact id
cid = sensor_contact_matchid_in[worldid, contactsensorid, i]
@@ -1808,7 +1801,7 @@ def _sensor_touch(
):
conid, sensortouchadrid = wp.tid()
- if conid > ncon_in[0]:
+ if conid >= ncon_in[0]:
return
sensorid = sensor_touch_adr[sensortouchadrid]
@@ -2114,7 +2107,7 @@ def sensor_acc(m: Model, d: Data):
wp.launch(
_sensor_tactile_zero,
- dim=(d.nworld, m.nsensordata),
+ dim=(d.nworld, m.nsensor),
inputs=[
m.sensor_type,
m.sensor_dim,
@@ -2222,7 +2215,7 @@ def sensor_acc(m: Model, d: Data):
m.sensor_adr,
m.sensor_cutoff,
m.sensor_acc_adr,
- m.sensor_contact_adr,
+ m.sensor_adr_to_contact_adr,
d.njmax,
d.ncon,
d.xpos,
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py
index ff36df1c..853b71c2 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py
@@ -328,6 +328,7 @@ class SensorTest(parameterized.TestCase):
""",
+ keyframe=keyframe,
)
d.sensordata.zero_()
@@ -420,6 +421,7 @@ class SensorTest(parameterized.TestCase):
# data combinations
field = ["found", "force", "torque", "dist", "pos", "normal", "tangent"]
datas = itertools.chain.from_iterable([itertools.combinations(field, i) for i in range(len(field))])
+ datas = list(datas)
for num in [1, 2, 3, 4, 5]:
for geoms in [
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py
index 73276fc3..ea5b10cf 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py
@@ -242,7 +242,7 @@ def _flex_edges(
j = body_dofadr[flex_vertbodyid[vbase + v[1]]]
vel1 = wp.vec3(qvel_in[worldid, i], qvel_in[worldid, i + 1], qvel_in[worldid, i + 2])
vel2 = wp.vec3(qvel_in[worldid, j], qvel_in[worldid, j + 1], qvel_in[worldid, j + 2])
- flexedge_velocity_out[worldid, edgeid] = wp.dot(vel2 - vel1, vec) / vecnorm
+ flexedge_velocity_out[worldid, edgeid] = math.safe_div(wp.dot(vel2 - vel1, vec), vecnorm)
@wp.kernel
@@ -374,11 +374,16 @@ def _subtree_com_acc(
def _subtree_div(
# Model:
subtree_mass: wp.array2d(dtype=float),
+ # Data in:
+ subtree_com_in: wp.array2d(dtype=wp.vec3),
# Data out:
subtree_com_out: wp.array2d(dtype=wp.vec3),
):
worldid, bodyid = wp.tid()
- subtree_com_out[worldid, bodyid] /= subtree_mass[worldid, bodyid]
+ com = subtree_com_in[worldid, bodyid]
+ mass = subtree_mass[worldid, bodyid]
+ if mass != 0.0:
+ subtree_com_out[worldid, bodyid] = com / mass
@wp.kernel
@@ -492,7 +497,7 @@ def com_pos(m: Model, d: Data):
outputs=[d.subtree_com],
)
- wp.launch(_subtree_div, dim=(d.nworld, m.nbody), inputs=[m.subtree_mass], outputs=[d.subtree_com])
+ wp.launch(_subtree_div, dim=(d.nworld, m.nbody), inputs=[m.subtree_mass, d.subtree_com], outputs=[d.subtree_com])
wp.launch(
_cinert,
dim=(d.nworld, m.nbody),
@@ -1546,7 +1551,7 @@ def _tendon_dot(
# chain rule, second term: Jdot += (jac2 - jac1) * d / dt (dpnt)
Jdot += wp.dot(jacdif, dvel)
- ten_Jdot_out[worldid, tenid, i] += Jdot / divisor
+ ten_Jdot_out[worldid, tenid, i] += math.safe_div(Jdot, divisor)
# TODO(team): j += 2 if geom wrapping
j += 1
@@ -1867,8 +1872,8 @@ def _transmission(
# compute derivatives of length w.r.t. vec and axis
if ok == 1:
- scale = 1.0 - av / sdet
- dldv = axis * scale + vec / sdet
+ scale = 1.0 - math.safe_div(av, sdet)
+ dldv = axis * scale + math.safe_div(vec, sdet)
dlda = vec * scale
else:
dldv = axis
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py
index ed74254f..c77de1e1 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py
@@ -28,6 +28,8 @@ from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
+wp.set_module_options({"enable_backward": False})
+
@wp.func
def _rescale(nv: int, stat_meaninertia: float, value: float) -> float:
@@ -39,68 +41,142 @@ def _in_bracket(x: wp.vec3, y: wp.vec3) -> bool:
return (x[1] < y[1] and y[1] < 0.0) or (x[1] > y[1] and y[1] > 0.0)
+@wp.func
+def _eval_cost(quad: wp.vec3, alpha: float) -> float:
+ return alpha * alpha * quad[2] + alpha * quad[1] + quad[0]
+
+
@wp.func
def _eval_pt(quad: wp.vec3, alpha: float) -> wp.vec3:
return wp.vec3(
- alpha * alpha * quad[2] + alpha * quad[1] + quad[0],
+ _eval_cost(quad, alpha),
2.0 * alpha * quad[2] + quad[1],
2.0 * quad[2],
)
@wp.func
-def _eval_pt_elliptic(
+def _eval(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
+ # Data in:
+ ncon_in: wp.array(dtype=int),
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
+ efc_quad_in: wp.array2d(dtype=wp.vec3),
# In:
- impratio: float,
- friction: types.vec5,
- u0: float,
- uu: float,
- uv: float,
- vv: float,
- jv: float,
- D: float,
- quad: wp.vec3,
+ worldid: int,
+ efcid: int,
alpha: float,
-) -> wp.vec3:
- mu = friction[0] / wp.sqrt(impratio)
- v0 = jv * mu
- n = u0 + alpha * v0
- tsqr = uu + alpha * (2.0 * uv + alpha * vv)
- t = wp.sqrt(tsqr) # tangential force
+ # Out:
+ out: wp.array(dtype=wp.vec3),
+):
+ ne = ne_in[worldid]
+ nf = nf_in[worldid]
- bottom_zone = ((tsqr <= 0.0) and (n < 0)) or ((tsqr > 0.0) and ((mu * n + t) <= 0.0))
- middle_zone = (tsqr > 0) and (n < (mu * t)) and ((mu * n + t) > 0.0)
+ # equality
+ if efcid < ne:
+ wp.atomic_add(out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+ # friction
+ elif efcid < ne + nf:
+ # search point, friction loss, bound (rf)
+ start = efc_Jaref_in[worldid, efcid]
+ dir = efc_jv_in[worldid, efcid]
+ x = start + alpha * dir
+ f = efc_frictionloss_in[worldid, efcid]
+ rf = math.safe_div(f, efc_D_in[worldid, efcid])
- # elliptic bottom zone: quadratic cose
- if bottom_zone:
- pt = _eval_pt(quad, alpha)
+ # -bound < x < bound : quadratic
+ if (-rf < x) and (x < rf):
+ quad = efc_quad_in[worldid, efcid]
+ # x < -bound: linear negative
+ elif x <= -rf:
+ quad = wp.vec3(f * (-0.5 * rf - start), -f * dir, 0.0)
+ # bound < x : linear positive
+ else:
+ quad = wp.vec3(f * (-0.5 * rf + start), f * dir, 0.0)
+
+ wp.atomic_add(out, worldid, _eval_pt(quad, alpha))
+ # elliptic friction cone contact
+ elif efc_type_in[worldid, efcid] == int(types.ConstraintType.CONTACT_ELLIPTIC.value):
+ # extract contact info
+ conid = efc_id_in[worldid, efcid]
+
+ if conid >= ncon_in[0]:
+ return
+
+ efcid0 = contact_efc_address_in[conid, 0]
+
+ if efcid != efcid0:
+ return
+
+ friction = contact_friction_in[conid]
+ mu = friction[0] / wp.sqrt(opt_impratio[worldid])
+
+ # unpack quad
+ efcid1 = contact_efc_address_in[conid, 1]
+ efcid2 = contact_efc_address_in[conid, 2]
+ u0 = efc_quad_in[worldid, efcid1][0]
+ v0 = efc_quad_in[worldid, efcid1][1]
+ uu = efc_quad_in[worldid, efcid1][2]
+ uv = efc_quad_in[worldid, efcid2][0]
+ vv = efc_quad_in[worldid, efcid2][1]
+ dm = efc_quad_in[worldid, efcid2][2]
+
+ # compute N, Tsqr
+ N = u0 + alpha * v0
+ Tsqr = uu + alpha * (2.0 * uv + alpha * vv)
+
+ # no tangential force: top or bottom zone
+ if Tsqr <= 0.0:
+ # bottom zone: quadratic cost
+ if N < 0.0:
+ wp.atomic_add(out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+
+ # top zone: nothing to do
+ # otherwise regular processing
+ else:
+ # tangential force
+ T = wp.sqrt(Tsqr)
+
+ # N >= mu * T : top zone
+ if N >= mu * T:
+ # nothing to do
+ pass
+ # mu * N + T <= 0 : bottom zone
+ elif mu * N + T <= 0.0:
+ wp.atomic_add(out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+
+ # otherwise middle zone
+ else:
+ # derivatives
+ N1 = v0
+ T1 = (uv + alpha * vv) / T
+ T2 = vv / T - (uv + alpha * vv) * T1 / (T * T)
+
+ # add to cost
+ cost = wp.vec3(
+ 0.5 * dm * (N - mu * T) * (N - mu * T),
+ dm * (N - mu * T) * (N1 - mu * T1),
+ dm * ((N1 - mu * T1) * (N1 - mu * T1) + (N - mu * T) * (-mu * T2)),
+ )
+
+ wp.atomic_add(out, worldid, cost)
else:
- pt = wp.vec3(0.0)
+ # search point
+ x = efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid]
- # elliptic middle zone
- if t == 0.0:
- t += types.MJ_MINVAL
-
- if tsqr == 0.0:
- tsqr += types.MJ_MINVAL
-
- n1 = v0
- t1 = (uv + alpha * vv) / t
- t2 = vv / t - (uv + alpha * vv) * t1 / tsqr
-
- if middle_zone:
- mu2 = mu * mu
- dm = D / wp.max(mu2 * (1.0 + mu2), types.MJ_MINVAL)
- nmt = n - mu * t
- n1mut1 = n1 - mu * t1
-
- pt += wp.vec3(
- 0.5 * dm * nmt * nmt,
- dm * nmt * n1mut1,
- dm * (n1mut1 * n1mut1 - nmt * mu * t2),
- )
-
- return pt
+ # active
+ if x < 0.0:
+ wp.atomic_add(out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
@wp.kernel
@@ -134,101 +210,23 @@ def linesearch_iterative_init_gtol_p0_gauss(
@wp.kernel
-def linesearch_iterative_init_p0_elliptic0(
- # Data in:
- ne_in: wp.array(dtype=int),
- nf_in: wp.array(dtype=int),
- nl_in: wp.array(dtype=int),
- nefc_in: wp.array(dtype=int),
- efc_Jaref_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_condim_in: wp.array2d(dtype=int),
- # Data out:
- efc_p0_out: wp.array(dtype=wp.vec3),
-):
- worldid, efcid = wp.tid()
-
- if efcid >= nefc_in[worldid]:
- return
-
- if efc_done_in[worldid]:
- return
-
- active = efc_Jaref_in[worldid, efcid] < 0.0
-
- nef = ne_in[worldid] + nf_in[worldid]
- nefl = nef + nl_in[worldid]
- if efcid < nef:
- active = True
- elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
- active = False
-
- if active:
- quad = efc_quad_in[worldid, efcid]
- wp.atomic_add(efc_p0_out, worldid, wp.vec3(quad[0], quad[1], 2.0 * quad[2]))
-
-
-@wp.kernel
-def linesearch_iterative_init_p0_elliptic1(
+def linesearch_iterative_init_p0(
# Model:
opt_impratio: wp.array(dtype=float),
# Data in:
ncon_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_dim_in: wp.array(dtype=int),
- contact_efc_address_in: wp.array2d(dtype=int),
- contact_worldid_in: wp.array(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_u_in: wp.array(dtype=types.vec6),
- efc_uu_in: wp.array(dtype=float),
- efc_uv_in: wp.array(dtype=float),
- efc_vv_in: wp.array(dtype=float),
- # Data out:
- efc_p0_out: wp.array(dtype=wp.vec3),
-):
- conid = wp.tid()
-
- if conid >= ncon_in[0]:
- return
-
- worldid = contact_worldid_in[conid]
- if efc_done_in[worldid]:
- return
-
- if contact_dim_in[conid] < 2:
- return
-
- efcid = contact_efc_address_in[conid, 0]
-
- pt = _eval_pt_elliptic(
- opt_impratio[worldid],
- contact_friction_in[conid],
- efc_u_in[conid][0],
- efc_uu_in[conid],
- efc_uv_in[conid],
- efc_vv_in[conid],
- efc_jv_in[worldid, efcid],
- efc_D_in[worldid, efcid],
- efc_quad_in[worldid, efcid],
- 0.0,
- )
-
- wp.atomic_add(efc_p0_out, worldid, pt)
-
-
-@wp.kernel
-def linesearch_iterative_init_p0_pyramidal(
- # Data in:
ne_in: wp.array(dtype=int),
nf_in: wp.array(dtype=int),
nefc_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
efc_Jaref_in: wp.array2d(dtype=float),
+ efc_jv_in: wp.array2d(dtype=float),
efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
# Data out:
efc_p0_out: wp.array(dtype=wp.vec3),
):
@@ -237,15 +235,25 @@ def linesearch_iterative_init_p0_pyramidal(
if efcid >= nefc_in[worldid]:
return
- if efc_done_in[worldid]:
- return
-
- if efc_Jaref_in[worldid, efcid] >= 0.0 and efcid >= ne_in[worldid] + nf_in[worldid]:
- return
-
- quad = efc_quad_in[worldid, efcid]
-
- wp.atomic_add(efc_p0_out, worldid, wp.vec3(quad[0], quad[1], 2.0 * quad[2]))
+ _eval(
+ opt_impratio,
+ ncon_in,
+ ne_in,
+ nf_in,
+ contact_friction_in,
+ contact_efc_address_in,
+ efc_type_in,
+ efc_id_in,
+ efc_D_in,
+ efc_frictionloss_in,
+ efc_Jaref_in,
+ efc_jv_in,
+ efc_quad_in,
+ worldid,
+ efcid,
+ 0.0,
+ efc_p0_out,
+ )
@wp.kernel
@@ -270,105 +278,23 @@ def linesearch_iterative_init_lo_gauss(
@wp.kernel
-def linesearch_iterative_init_lo_elliptic0(
- # Data in:
- ne_in: wp.array(dtype=int),
- nf_in: wp.array(dtype=int),
- nl_in: wp.array(dtype=int),
- nefc_in: wp.array(dtype=int),
- efc_Jaref_in: wp.array2d(dtype=float),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_lo_alpha_in: wp.array(dtype=float),
- efc_condim_in: wp.array2d(dtype=int),
- # Data out:
- efc_lo_out: wp.array(dtype=wp.vec3),
-):
- worldid, efcid = wp.tid()
-
- if efcid >= nefc_in[worldid]:
- return
-
- if efc_done_in[worldid]:
- return
-
- alpha = efc_lo_alpha_in[worldid]
-
- active = efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid] < 0.0
-
- nef = ne_in[worldid] + nf_in[worldid]
- nefl = nef + nl_in[worldid]
- if efcid < nef:
- active = True
- elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
- active = False
-
- if active:
- wp.atomic_add(efc_lo_out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
-
-
-@wp.kernel
-def linesearch_iterative_init_lo_elliptic1(
+def linesearch_iterative_init_lo(
# Model:
opt_impratio: wp.array(dtype=float),
# Data in:
ncon_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_dim_in: wp.array(dtype=int),
- contact_efc_address_in: wp.array2d(dtype=int),
- contact_worldid_in: wp.array(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_lo_alpha_in: wp.array(dtype=float),
- efc_u_in: wp.array(dtype=types.vec6),
- efc_uu_in: wp.array(dtype=float),
- efc_uv_in: wp.array(dtype=float),
- efc_vv_in: wp.array(dtype=float),
- # Data out:
- efc_lo_out: wp.array(dtype=wp.vec3),
-):
- conid = wp.tid()
-
- if conid >= ncon_in[0]:
- return
-
- worldid = contact_worldid_in[conid]
- if efc_done_in[worldid]:
- return
-
- if contact_dim_in[conid] < 2:
- return
-
- efcid = contact_efc_address_in[conid, 0]
- alpha = efc_lo_alpha_in[worldid]
- pt = _eval_pt_elliptic(
- opt_impratio[worldid],
- contact_friction_in[conid],
- efc_u_in[conid][0],
- efc_uu_in[conid],
- efc_uv_in[conid],
- efc_vv_in[conid],
- efc_jv_in[worldid, efcid],
- efc_D_in[worldid, efcid],
- efc_quad_in[worldid, efcid],
- alpha,
- )
- wp.atomic_add(efc_lo_out, worldid, pt)
-
-
-@wp.kernel
-def linesearch_iterative_init_lo_pyramidal(
- # Data in:
ne_in: wp.array(dtype=int),
nf_in: wp.array(dtype=int),
nefc_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
efc_Jaref_in: wp.array2d(dtype=float),
efc_jv_in: wp.array2d(dtype=float),
efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
efc_lo_alpha_in: wp.array(dtype=float),
# Data out:
efc_lo_out: wp.array(dtype=wp.vec3),
@@ -378,13 +304,25 @@ def linesearch_iterative_init_lo_pyramidal(
if efcid >= nefc_in[worldid]:
return
- if efc_done_in[worldid]:
- return
-
- alpha = efc_lo_alpha_in[worldid]
-
- if efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid] < 0.0 or (efcid < ne_in[worldid] + nf_in[worldid]):
- wp.atomic_add(efc_lo_out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+ _eval(
+ opt_impratio,
+ ncon_in,
+ ne_in,
+ nf_in,
+ contact_friction_in,
+ contact_efc_address_in,
+ efc_type_in,
+ efc_id_in,
+ efc_D_in,
+ efc_frictionloss_in,
+ efc_Jaref_in,
+ efc_jv_in,
+ efc_quad_in,
+ worldid,
+ efcid,
+ efc_lo_alpha_in[worldid],
+ efc_lo_out,
+ )
@wp.kernel
@@ -462,147 +400,20 @@ def linesearch_iterative_next_alpha_gauss(
@wp.kernel
-def linesearch_iterative_next_quad_elliptic0(
- # Data in:
- ne_in: wp.array(dtype=int),
- nf_in: wp.array(dtype=int),
- nl_in: wp.array(dtype=int),
- nefc_in: wp.array(dtype=int),
- efc_Jaref_in: wp.array2d(dtype=float),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_ls_done_in: wp.array(dtype=bool),
- efc_lo_next_alpha_in: wp.array(dtype=float),
- efc_hi_next_alpha_in: wp.array(dtype=float),
- efc_mid_alpha_in: wp.array(dtype=float),
- efc_condim_in: wp.array2d(dtype=int),
- # Data out:
- efc_lo_next_out: wp.array(dtype=wp.vec3),
- efc_hi_next_out: wp.array(dtype=wp.vec3),
- efc_mid_out: wp.array(dtype=wp.vec3),
-):
- worldid, efcid = wp.tid()
-
- if efcid >= nefc_in[worldid]:
- return
-
- if efc_done_in[worldid]:
- return
-
- if efc_ls_done_in[worldid]:
- return
-
- nef = ne_in[worldid] + nf_in[worldid]
- nefl = nef + nl_in[worldid]
-
- quad = efc_quad_in[worldid, efcid]
- jaref = efc_Jaref_in[worldid, efcid]
- jv = efc_jv_in[worldid, efcid]
-
- alpha = efc_lo_next_alpha_in[worldid]
-
- active = jaref + alpha * jv < 0.0
- if efcid < nef:
- active = True
- elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
- active = False
-
- if active:
- wp.atomic_add(efc_lo_next_out, worldid, _eval_pt(quad, alpha))
-
- alpha = efc_hi_next_alpha_in[worldid]
-
- active = jaref + alpha * jv < 0.0
- if efcid < nef:
- active = True
- elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
- active = False
-
- if active:
- wp.atomic_add(efc_hi_next_out, worldid, _eval_pt(quad, alpha))
-
- alpha = efc_mid_alpha_in[worldid]
-
- active = jaref + alpha * jv < 0.0
- if efcid < nef:
- active = True
- elif efcid >= nefl and efc_condim_in[worldid, efcid] > 1:
- active = False
-
- if active:
- wp.atomic_add(efc_mid_out, worldid, _eval_pt(quad, alpha))
-
-
-@wp.kernel
-def linesearch_iterative_next_quad_elliptic1(
+def linesearch_iterative_next_quad(
# Model:
opt_impratio: wp.array(dtype=float),
# Data in:
ncon_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_dim_in: wp.array(dtype=int),
- contact_efc_address_in: wp.array2d(dtype=int),
- contact_worldid_in: wp.array(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_lo_next_alpha_in: wp.array(dtype=float),
- efc_hi_next_alpha_in: wp.array(dtype=float),
- efc_mid_alpha_in: wp.array(dtype=float),
- efc_u_in: wp.array(dtype=types.vec6),
- efc_uu_in: wp.array(dtype=float),
- efc_uv_in: wp.array(dtype=float),
- efc_vv_in: wp.array(dtype=float),
- # Data out:
- efc_lo_next_out: wp.array(dtype=wp.vec3),
- efc_hi_next_out: wp.array(dtype=wp.vec3),
- efc_mid_out: wp.array(dtype=wp.vec3),
-):
- conid = wp.tid()
-
- if conid >= ncon_in[0]:
- return
-
- worldid = contact_worldid_in[conid]
-
- if efc_done_in[worldid]:
- return
-
- if contact_dim_in[conid] < 2:
- return
-
- efcid = contact_efc_address_in[conid, 0]
- impratio = opt_impratio[worldid]
- friction = contact_friction_in[conid]
- u = efc_u_in[conid][0]
- uu = efc_uu_in[conid]
- uv = efc_uv_in[conid]
- vv = efc_vv_in[conid]
- jv = efc_jv_in[worldid, efcid]
- d = efc_D_in[worldid, efcid]
- quad = efc_quad_in[worldid, efcid]
-
- alpha = efc_lo_next_alpha_in[worldid]
- pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha)
- wp.atomic_add(efc_lo_next_out, worldid, pt)
-
- alpha = efc_hi_next_alpha_in[worldid]
- pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha)
- wp.atomic_add(efc_hi_next_out, worldid, pt)
-
- alpha = efc_mid_alpha_in[worldid]
- pt = _eval_pt_elliptic(impratio, friction, u, uu, uv, vv, jv, d, quad, alpha)
- wp.atomic_add(efc_mid_out, worldid, pt)
-
-
-@wp.kernel
-def linesearch_iterative_next_quad_pyramidal(
- # Data in:
ne_in: wp.array(dtype=int),
nf_in: wp.array(dtype=int),
nefc_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
efc_Jaref_in: wp.array2d(dtype=float),
efc_jv_in: wp.array2d(dtype=float),
efc_quad_in: wp.array2d(dtype=wp.vec3),
@@ -627,23 +438,68 @@ def linesearch_iterative_next_quad_pyramidal(
if efc_ls_done_in[worldid]:
return
- nef_active = efcid < ne_in[worldid] + nf_in[worldid]
+ # lo_next
+ _eval(
+ opt_impratio,
+ ncon_in,
+ ne_in,
+ nf_in,
+ contact_friction_in,
+ contact_efc_address_in,
+ efc_type_in,
+ efc_id_in,
+ efc_D_in,
+ efc_frictionloss_in,
+ efc_Jaref_in,
+ efc_jv_in,
+ efc_quad_in,
+ worldid,
+ efcid,
+ efc_lo_next_alpha_in[worldid],
+ efc_lo_next_out,
+ )
- quad = efc_quad_in[worldid, efcid]
- jaref = efc_Jaref_in[worldid, efcid]
- jv = efc_jv_in[worldid, efcid]
+ # hi_next
+ _eval(
+ opt_impratio,
+ ncon_in,
+ ne_in,
+ nf_in,
+ contact_friction_in,
+ contact_efc_address_in,
+ efc_type_in,
+ efc_id_in,
+ efc_D_in,
+ efc_frictionloss_in,
+ efc_Jaref_in,
+ efc_jv_in,
+ efc_quad_in,
+ worldid,
+ efcid,
+ efc_hi_next_alpha_in[worldid],
+ efc_hi_next_out,
+ )
- alpha = efc_lo_next_alpha_in[worldid]
- if jaref + alpha * jv < 0.0 or nef_active:
- wp.atomic_add(efc_lo_next_out, worldid, _eval_pt(quad, alpha))
-
- alpha = efc_hi_next_alpha_in[worldid]
- if jaref + alpha * jv < 0.0 or nef_active:
- wp.atomic_add(efc_hi_next_out, worldid, _eval_pt(quad, alpha))
-
- alpha = efc_mid_alpha_in[worldid]
- if jaref + alpha * jv < 0.0 or nef_active:
- wp.atomic_add(efc_mid_out, worldid, _eval_pt(quad, alpha))
+ # mid
+ _eval(
+ opt_impratio,
+ ncon_in,
+ ne_in,
+ nf_in,
+ contact_friction_in,
+ contact_efc_address_in,
+ efc_type_in,
+ efc_id_in,
+ efc_D_in,
+ efc_frictionloss_in,
+ efc_Jaref_in,
+ efc_jv_in,
+ efc_quad_in,
+ worldid,
+ efcid,
+ efc_mid_alpha_in[worldid],
+ efc_mid_out,
+ )
@wp.kernel
@@ -740,81 +596,68 @@ def _linesearch_iterative(m: types.Model, d: types.Data):
wp.launch(
linesearch_iterative_init_gtol_p0_gauss,
dim=(d.nworld,),
- inputs=[
- m.nv, m.opt.tolerance, m.opt.ls_tolerance, m.stat.meaninertia, d.efc.search_dot,
- d.efc.quad_gauss, d.efc.done
- ],
- outputs=[d.efc.gtol, d.efc.p0]) # fmt: skip
+ inputs=[m.nv, m.opt.tolerance, m.opt.ls_tolerance, m.stat.meaninertia, d.efc.search_dot, d.efc.quad_gauss, d.efc.done],
+ outputs=[d.efc.gtol, d.efc.p0],
+ )
- if m.opt.cone == types.ConeType.ELLIPTIC:
- wp.launch(
- linesearch_iterative_init_p0_elliptic0,
- dim=(d.nworld, d.njmax,),
- inputs=[
- d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.quad, d.efc.done,
- d.efc.condim
- ],
- outputs=[d.efc.p0]) # fmt: skip
- wp.launch(
- linesearch_iterative_init_p0_elliptic1,
- dim=(d.nconmax),
- inputs=[
- m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim,
- d.contact.efc_address, d.contact.worldid, d.efc.D, d.efc.jv, d.efc.quad,
- d.efc.done, d.efc.u, d.efc.uu, d.efc.uv, d.efc.vv
- ],
- outputs=[d.efc.p0]) # fmt: skip
- else:
- wp.launch(
- linesearch_iterative_init_p0_pyramidal,
- dim=(d.nworld, d.njmax,),
- inputs=[
- d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.quad, d.efc.done
- ], outputs=[d.efc.p0]) # fmt: skip
+ wp.launch(
+ linesearch_iterative_init_p0,
+ dim=(d.nworld, d.njmax),
+ inputs=[
+ m.opt.impratio,
+ d.ncon,
+ d.ne,
+ d.nf,
+ d.nefc,
+ d.contact.friction,
+ d.contact.efc_address,
+ d.efc.type,
+ d.efc.id,
+ d.efc.D,
+ d.efc.frictionloss,
+ d.efc.Jaref,
+ d.efc.jv,
+ d.efc.quad,
+ ],
+ outputs=[d.efc.p0],
+ )
wp.launch(
linesearch_iterative_init_lo_gauss,
dim=(d.nworld,),
+ inputs=[d.efc.quad_gauss, d.efc.done, d.efc.p0],
+ outputs=[d.efc.lo, d.efc.lo_alpha],
+ )
+ wp.launch(
+ linesearch_iterative_init_lo,
+ dim=(d.nworld, d.njmax),
inputs=[
- d.efc.quad_gauss, d.efc.done, d.efc.p0
+ m.opt.impratio,
+ d.ncon,
+ d.ne,
+ d.nf,
+ d.nefc,
+ d.contact.friction,
+ d.contact.efc_address,
+ d.efc.type,
+ d.efc.id,
+ d.efc.D,
+ d.efc.frictionloss,
+ d.efc.Jaref,
+ d.efc.jv,
+ d.efc.quad,
+ d.efc.lo_alpha,
],
- outputs=[d.efc.lo, d.efc.lo_alpha]) # fmt: skip
-
- if m.opt.cone == types.ConeType.ELLIPTIC:
- wp.launch(
- linesearch_iterative_init_lo_elliptic0,
- dim=(d.nworld, d.njmax,),
- inputs=[
- d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad,
- d.efc.done, d.efc.lo_alpha, d.efc.condim
- ],
- outputs=[d.efc.lo]) # fmt: skip
- wp.launch(
- linesearch_iterative_init_lo_elliptic1,
- dim=(d.nconmax),
- inputs=[
- m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim,
- d.contact.efc_address, d.contact.worldid, d.efc.D, d.efc.jv, d.efc.quad,
- d.efc.done, d.efc.lo_alpha, d.efc.u, d.efc.uu, d.efc.uv, d.efc.vv
- ],
- outputs=[d.efc.lo]) # fmt: skip
- else:
- wp.launch(
- linesearch_iterative_init_lo_pyramidal,
- dim=(d.nworld, d.njmax,),
- inputs=[
- d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.jv,
- d.efc.quad, d.efc.done, d.efc.lo_alpha
- ],
- outputs=[d.efc.lo]) # fmt: skip
+ outputs=[d.efc.lo],
+ )
# set the lo/hi interval bounds
-
wp.launch(
linesearch_iterative_init_bounds,
dim=(d.nworld,),
inputs=[d.efc.done, d.efc.p0, d.efc.lo, d.efc.lo_alpha],
- outputs=[d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha]) # fmt: skip
+ outputs=[d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha],
+ )
for _ in range(m.opt.ls_iterations):
# NOTE: we always launch ls_iterations kernels, but the kernels may early exit if done
@@ -823,61 +666,77 @@ def _linesearch_iterative(m: types.Model, d: types.Data):
wp.launch(
linesearch_iterative_next_alpha_gauss,
dim=(d.nworld,),
- inputs=[
- d.efc.quad_gauss, d.efc.done, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha
- ],
- outputs=[
- d.efc.lo_next, d.efc.lo_next_alpha, d.efc.hi_next, d.efc.hi_next_alpha, d.efc.mid, d.efc.mid_alpha
- ]) # fmt: skip
+ inputs=[d.efc.quad_gauss, d.efc.done, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha],
+ outputs=[d.efc.lo_next, d.efc.lo_next_alpha, d.efc.hi_next, d.efc.hi_next_alpha, d.efc.mid, d.efc.mid_alpha],
+ )
- if m.opt.cone == types.ConeType.ELLIPTIC:
- wp.launch(
- linesearch_iterative_next_quad_elliptic0,
- dim=(d.nworld, d.njmax,),
- inputs=[
- d.ne, d.nf, d.nl, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad, d.efc.done, d.efc.ls_done,
- d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha, d.efc.condim
- ],
- outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip
- wp.launch(
- linesearch_iterative_next_quad_elliptic1,
- dim=(d.nconmax),
- inputs=[
- m.opt.impratio, d.ncon, d.contact.friction, d.contact.dim, d.contact.efc_address, d.contact.worldid, d.efc.D,
- d.efc.jv, d.efc.quad, d.efc.done, d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha, d.efc.u, d.efc.uu,
- d.efc.uv, d.efc.vv
- ],
- outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip
- else:
- wp.launch(
- linesearch_iterative_next_quad_pyramidal,
- dim=(d.nworld, d.njmax,),
- inputs=[
- d.ne, d.nf, d.nefc, d.efc.Jaref, d.efc.jv, d.efc.quad, d.efc.done, d.efc.ls_done,
- d.efc.lo_next_alpha, d.efc.hi_next_alpha, d.efc.mid_alpha
- ],
- outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid]) # fmt: skip
+ wp.launch(
+ linesearch_iterative_next_quad,
+ dim=(d.nworld, d.njmax),
+ inputs=[
+ m.opt.impratio,
+ d.ncon,
+ d.ne,
+ d.nf,
+ d.nefc,
+ d.contact.friction,
+ d.contact.efc_address,
+ d.efc.type,
+ d.efc.id,
+ d.efc.D,
+ d.efc.frictionloss,
+ d.efc.Jaref,
+ d.efc.jv,
+ d.efc.quad,
+ d.efc.done,
+ d.efc.ls_done,
+ d.efc.lo_next_alpha,
+ d.efc.hi_next_alpha,
+ d.efc.mid_alpha,
+ ],
+ outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid],
+ )
wp.launch(
linesearch_iterative_swap,
dim=(d.nworld,),
inputs=[
- d.efc.gtol, d.efc.done, d.efc.ls_done, d.efc.p0, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha,
- d.efc.lo_next, d.efc.lo_next_alpha, d.efc.hi_next, d.efc.hi_next_alpha, d.efc.mid, d.efc.mid_alpha
+ d.efc.gtol,
+ d.efc.done,
+ d.efc.ls_done,
+ d.efc.p0,
+ d.efc.lo,
+ d.efc.lo_alpha,
+ d.efc.hi,
+ d.efc.hi_alpha,
+ d.efc.lo_next,
+ d.efc.lo_next_alpha,
+ d.efc.hi_next,
+ d.efc.hi_next_alpha,
+ d.efc.mid,
+ d.efc.mid_alpha,
],
- outputs=[
- d.efc.alpha, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha
- ]) # fmt: skip
+ outputs=[d.efc.alpha, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha],
+ )
@wp.kernel
def linesearch_parallel_fused(
# Model:
nlsp: int,
+ opt_impratio: wp.array(dtype=float),
# Data in:
+ njmax_in: int,
+ ncon_in: wp.array(dtype=int),
ne_in: wp.array(dtype=int),
nf_in: wp.array(dtype=int),
nefc_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
efc_Jaref_in: wp.array2d(dtype=float),
efc_jv_in: wp.array2d(dtype=float),
efc_quad_in: wp.array2d(dtype=wp.vec3),
@@ -891,25 +750,96 @@ def linesearch_parallel_fused(
if efc_done_in[worldid]:
return
- efc_quad_total_candidate = efc_quad_gauss_in[worldid]
-
alpha = float(alphaid) / float(nlsp - 1)
+
+ out = _eval_cost(efc_quad_gauss_in[worldid], alpha)
+
ne = ne_in[worldid]
nf = nf_in[worldid]
- for efcid in range(nefc_in[worldid]):
- Jaref = efc_Jaref_in[worldid, efcid]
- jv = efc_jv_in[worldid, efcid]
- quad = efc_quad_in[worldid, efcid]
- if (Jaref + alpha * jv) < 0.0 or (efcid < ne + nf):
- efc_quad_total_candidate += quad
+ # TODO(team): _eval with option to only compute cost
+ for efcid in range(min(njmax_in, nefc_in[worldid])):
+ # equality
+ if efcid < ne:
+ out += _eval_cost(efc_quad_in[worldid, efcid], alpha)
+ # friction
+ elif efcid < ne + nf:
+ # search point, friction loss, bound (rf)
+ start = efc_Jaref_in[worldid, efcid]
+ dir = efc_jv_in[worldid, efcid]
+ x = start + alpha * dir
+ f = efc_frictionloss_in[worldid, efcid]
+ rf = math.safe_div(f, efc_D_in[worldid, efcid])
- alpha_sq = alpha * alpha
- quad_total0 = efc_quad_total_candidate[0]
- quad_total1 = efc_quad_total_candidate[1]
- quad_total2 = efc_quad_total_candidate[2]
+ # -bound < x < bound : quadratic
+ if (-rf < x) and (x < rf):
+ quad = efc_quad_in[worldid, efcid]
+ # x < -bound: linear negative
+ elif x <= -rf:
+ quad = wp.vec3(f * (-0.5 * rf - start), -f * dir, 0.0)
+ # bound < x : linear positive
+ else:
+ quad = wp.vec3(f * (-0.5 * rf + start), f * dir, 0.0)
- efc_cost_candidate_out[worldid, alphaid] = alpha_sq * quad_total2 + alpha * quad_total1 + quad_total0
+ out += _eval_cost(quad, alpha)
+ # limit and contact
+ elif efc_type_in[worldid, efcid] == int(types.ConstraintType.CONTACT_ELLIPTIC.value):
+ # extract contact info
+ conid = efc_id_in[worldid, efcid]
+
+ if conid >= ncon_in[0]:
+ continue
+
+ efcid0 = contact_efc_address_in[conid, 0]
+ if efcid != efcid0:
+ continue
+
+ friction = contact_friction_in[conid]
+ mu = friction[0] / wp.sqrt(opt_impratio[worldid])
+
+ # unpack quad
+ efcid1 = contact_efc_address_in[conid, 1]
+ efcid2 = contact_efc_address_in[conid, 2]
+ u0 = efc_quad_in[worldid, efcid1][0]
+ v0 = efc_quad_in[worldid, efcid1][1]
+ uu = efc_quad_in[worldid, efcid1][2]
+ uv = efc_quad_in[worldid, efcid2][0]
+ vv = efc_quad_in[worldid, efcid2][1]
+ dm = efc_quad_in[worldid, efcid2][2]
+
+ # compute N, Tsqr
+ N = u0 + alpha * v0
+ Tsqr = uu + alpha * (2.0 * uv + alpha * vv)
+
+ # no tangential force: top or bottom zone
+ if Tsqr <= 0.0:
+ # bottom zone: quadratic cost
+ if N < 0.0:
+ out += _eval_cost(efc_quad_in[worldid, efcid], alpha)
+ # otherwise regular processing
+ else:
+ # tangential force
+ T = wp.sqrt(Tsqr)
+
+ # N >= mu * T : top zone
+ if N >= mu * T:
+ # nothing to do
+ pass
+ # mu * N + T <= 0 : bottom zone
+ elif mu * N + T <= 0.0:
+ out += _eval_cost(efc_quad_in[worldid, efcid], alpha)
+ # otherwise middle zone
+ else:
+ out += 0.5 * dm * (N - mu * T) * (N - mu * T)
+ else:
+ # search point
+ x = efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid]
+
+ # active
+ if x < 0.0:
+ out += _eval_cost(efc_quad_in[worldid, efcid], alpha)
+
+ efc_cost_candidate_out[worldid, alphaid] = out
@wp.kernel
@@ -946,9 +876,18 @@ def _linesearch_parallel(m: types.Model, d: types.Data):
dim=(d.nworld, m.nlsp),
inputs=[
m.nlsp,
+ m.opt.impratio,
+ d.njmax,
+ d.ncon,
d.ne,
d.nf,
d.nefc,
+ d.contact.friction,
+ d.contact.efc_address,
+ d.efc.type,
+ d.efc.id,
+ d.efc.D,
+ d.efc.frictionloss,
d.efc.Jaref,
d.efc.jv,
d.efc.quad,
@@ -1023,7 +962,7 @@ def linesearch_jv_fused(nv: int, dofs_per_thread: int):
@wp.kernel
-def linesearch_init_quad_gauss(
+def linesearch_prepare_gauss(
# Model:
nv: int,
# Data in:
@@ -1052,16 +991,21 @@ def linesearch_init_quad_gauss(
@wp.kernel
-def linesearch_init_quad(
+def linesearch_prepare_quad(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
# Data in:
+ ncon_in: wp.array(dtype=int),
nefc_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
efc_D_in: wp.array2d(dtype=float),
- efc_frictionloss_in: wp.array2d(dtype=float),
efc_Jaref_in: wp.array2d(dtype=float),
efc_jv_in: wp.array2d(dtype=float),
efc_done_in: wp.array(dtype=bool),
- # In:
- disable_floss: bool,
# Data out:
efc_quad_out: wp.array2d(dtype=wp.vec3),
):
@@ -1076,63 +1020,69 @@ def linesearch_init_quad(
Jaref = efc_Jaref_in[worldid, efcid]
jv = efc_jv_in[worldid, efcid]
efc_D = efc_D_in[worldid, efcid]
- floss = efc_frictionloss_in[worldid, efcid]
- if floss > 0.0 and not disable_floss:
- rf = math.safe_div(floss, efc_D)
- if Jaref <= -rf:
- efc_quad_out[worldid, efcid] = wp.vec3(floss * (-0.5 * rf - Jaref), -floss * jv, 0.0)
- return
- elif Jaref >= rf:
- efc_quad_out[worldid, efcid] = wp.vec3(floss * (-0.5 * rf + Jaref), floss * jv, 0.0)
+ # init with scalar quadratic
+ quad = wp.vec3(0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D)
+
+ # elliptic cone: extra processing
+ if efc_type_in[worldid, efcid] == int(types.ConstraintType.CONTACT_ELLIPTIC.value):
+ # extract contact info
+ conid = efc_id_in[worldid, efcid]
+
+ if conid >= ncon_in[0]:
return
- efc_quad_out[worldid, efcid] = wp.vec3(0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D)
+ efcid0 = contact_efc_address_in[conid, 0]
+ if efcid != efcid0:
+ return
-@wp.kernel
-def linesearch_quad_elliptic(
- # Data in:
- ncon_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_dim_in: wp.array(dtype=int),
- contact_efc_address_in: wp.array2d(dtype=int),
- contact_worldid_in: wp.array(dtype=int),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_u_in: wp.array(dtype=types.vec6),
- # Data out:
- efc_quad_out: wp.array2d(dtype=wp.vec3),
- efc_uv_out: wp.array(dtype=float),
- efc_vv_out: wp.array(dtype=float),
-):
- conid, dimid = wp.tid()
- dimid += 1
+ dim = contact_dim_in[conid]
+ friction = contact_friction_in[conid]
+ mu = friction[0] / wp.sqrt(opt_impratio[worldid])
- if conid >= ncon_in[0]:
- return
+ u0 = Jaref * mu
+ v0 = jv * mu
- worldid = contact_worldid_in[conid]
- if efc_done_in[worldid]:
- return
+ uu = float(0.0)
+ uv = float(0.0)
+ vv = float(0.0)
+ for j in range(1, dim):
+ # complete vector quadratic (for bottom zone)
+ efcidj = contact_efc_address_in[conid, j]
+ if efcidj < 0:
+ return
+ jvj = efc_jv_in[worldid, efcidj]
+ jarefj = efc_Jaref_in[worldid, efcidj]
+ dj = efc_D_in[worldid, efcidj]
+ DJj = dj * jarefj
- condim = contact_dim_in[conid]
+ quad += wp.vec3(
+ 0.5 * jarefj * DJj,
+ jvj * DJj,
+ 0.5 * jvj * dj * jvj,
+ )
- if condim == 1 or (dimid >= condim):
- return
+ # rescale to make primal cone circular
+ frictionj = friction[j - 1]
+ uj = jarefj * frictionj
+ vj = jvj * frictionj
- efcid0 = contact_efc_address_in[conid, 0]
- efcid = contact_efc_address_in[conid, dimid]
+ # accumulate sums of squares
+ uu += uj * uj
+ uv += uj * vj
+ vv += vj * vj
- # complete vector quadratic (for bottom zone)
- wp.atomic_add(efc_quad_out, worldid, efcid0, efc_quad_in[worldid, efcid])
+ quad1 = wp.vec3(u0, v0, uu)
+ efcid1 = contact_efc_address_in[conid, 1]
+ efc_quad_out[worldid, efcid1] = quad1
- # rescale to make primal cone circular
- u = efc_u_in[conid][dimid]
- v = efc_jv_in[worldid, efcid] * contact_friction_in[conid][dimid - 1]
- wp.atomic_add(efc_uv_out, conid, u * v)
- wp.atomic_add(efc_vv_out, conid, v * v)
+ mu2 = mu * mu
+ quad2 = wp.vec3(uv, vv, efc_D / (mu2 * (1.0 + mu2)))
+ efcid2 = contact_efc_address_in[conid, 2]
+ efc_quad_out[worldid, efcid2] = quad2
+
+ efc_quad_out[worldid, efcid] = quad
@wp.kernel
@@ -1213,50 +1163,33 @@ def _linesearch(m: types.Model, d: types.Data):
# prepare quadratics
# quad_gauss = [gauss, search.T @ Ma - search.T @ qfrc_smooth, 0.5 * search.T @ mv]
wp.launch(
- linesearch_init_quad_gauss,
+ linesearch_prepare_gauss,
dim=(d.nworld),
inputs=[m.nv, d.qfrc_smooth, d.efc.Ma, d.efc.search, d.efc.gauss, d.efc.mv, d.efc.done],
outputs=[d.efc.quad_gauss],
)
# quad = [0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D]
-
- disable_floss = m.opt.disableflags & types.DisableBit.FRICTIONLOSS
wp.launch(
- linesearch_init_quad,
+ linesearch_prepare_quad,
dim=(d.nworld, d.njmax),
inputs=[
+ m.opt.impratio,
+ d.ncon,
d.nefc,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.efc_address,
+ d.efc.type,
+ d.efc.id,
d.efc.D,
- d.efc.frictionloss,
d.efc.Jaref,
d.efc.jv,
d.efc.done,
- disable_floss,
],
outputs=[d.efc.quad],
)
- if m.opt.cone == types.ConeType.ELLIPTIC:
- d.efc.uv.zero_()
- d.efc.vv.zero_()
- wp.launch(
- linesearch_quad_elliptic,
- dim=(d.nconmax, m.condim_max - 1),
- inputs=[
- d.ncon,
- d.contact.friction,
- d.contact.dim,
- d.contact.efc_address,
- d.contact.worldid,
- d.efc.jv,
- d.efc.quad,
- d.efc.done,
- d.efc.u,
- ],
- outputs=[d.efc.quad, d.efc.uv, d.efc.vv],
- )
-
if m.opt.ls_parallel:
_linesearch_parallel(m, d)
else:
@@ -1351,288 +1284,128 @@ def update_constraint_init_cost(
@wp.kernel
-def update_constraint_efc_pyramidal(
+def update_constraint_efc(
+ # Model:
+ opt_impratio: wp.array(dtype=float),
# Data in:
+ ncon_in: wp.array(dtype=int),
ne_in: wp.array(dtype=int),
nf_in: wp.array(dtype=int),
nefc_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_dim_in: wp.array(dtype=int),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
efc_D_in: wp.array2d(dtype=float),
efc_frictionloss_in: wp.array2d(dtype=float),
efc_Jaref_in: wp.array2d(dtype=float),
- # In:
- disable_floss: int,
+ efc_done_in: wp.array(dtype=bool),
# Data out:
efc_force_out: wp.array2d(dtype=float),
efc_cost_out: wp.array(dtype=float),
- efc_active_out: wp.array2d(dtype=bool),
+ efc_state_out: wp.array2d(dtype=int),
):
worldid, efcid = wp.tid()
if efcid >= nefc_in[worldid]:
return
+ if efc_done_in[worldid]:
+ return
+
efc_D = efc_D_in[worldid, efcid]
Jaref = efc_Jaref_in[worldid, efcid]
- cost = 0.5 * efc_D * Jaref * Jaref
- efc_force = -efc_D * Jaref
-
ne = ne_in[worldid]
nf = nf_in[worldid]
if efcid < ne:
# equality
- pass
- elif efcid < ne + nf and not disable_floss:
+ efc_force_out[worldid, efcid] = -efc_D * Jaref
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.QUADRATIC.value)
+ wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref)
+ elif efcid < ne + nf:
# friction
f = efc_frictionloss_in[worldid, efcid]
- if f > 0.0:
- rf = math.safe_div(f, efc_D)
- if Jaref <= -rf:
- efc_force_out[worldid, efcid] = f
- efc_active_out[worldid, efcid] = False
- wp.atomic_add(efc_cost_out, worldid, -0.5 * rf - Jaref)
- return
- elif Jaref >= rf:
- efc_force_out[worldid, efcid] = -f
- efc_active_out[worldid, efcid] = False
- wp.atomic_add(efc_cost_out, worldid, -0.5 * rf + Jaref)
- return
- else:
- # limit, contact
+ rf = math.safe_div(f, efc_D)
+ if Jaref <= -rf:
+ efc_force_out[worldid, efcid] = f
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.LINEARNEG.value)
+ wp.atomic_add(efc_cost_out, worldid, -f * (0.5 * rf + Jaref))
+ elif Jaref >= rf:
+ efc_force_out[worldid, efcid] = -f
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.LINEARPOS.value)
+ wp.atomic_add(efc_cost_out, worldid, -f * (0.5 * rf - Jaref))
+ else:
+ efc_force_out[worldid, efcid] = -efc_D * Jaref
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.QUADRATIC.value)
+ wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref)
+ elif efc_type_in[worldid, efcid] != int(types.ConstraintType.CONTACT_ELLIPTIC.value):
+ # limit, frictionless contact, pyramidal friction cone contact
if Jaref >= 0.0:
efc_force_out[worldid, efcid] = 0.0
- efc_active_out[worldid, efcid] = False
- return
-
- efc_force_out[worldid, efcid] = efc_force
- efc_active_out[worldid, efcid] = True
- wp.atomic_add(efc_cost_out, worldid, cost)
-
-
-@wp.kernel
-def update_constraint_u_elliptic(
- # Model:
- opt_impratio: wp.array(dtype=float),
- # Data in:
- ncon_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_dim_in: wp.array(dtype=int),
- contact_efc_address_in: wp.array2d(dtype=int),
- contact_worldid_in: wp.array(dtype=int),
- efc_Jaref_in: wp.array2d(dtype=float),
- efc_done_in: wp.array(dtype=bool),
- # Data out:
- efc_u_out: wp.array(dtype=types.vec6),
- efc_uu_out: wp.array(dtype=float),
- efc_condim_out: wp.array2d(dtype=int),
-):
- conid, dimid = wp.tid()
-
- if conid >= ncon_in[0]:
- return
-
- worldid = contact_worldid_in[conid]
- if efc_done_in[worldid]:
- return
-
- efcid = contact_efc_address_in[conid, dimid]
-
- condim = contact_dim_in[conid]
- efc_condim_out[worldid, efcid] = condim
-
- if condim == 1:
- return
-
- if dimid < condim:
- if dimid == 0:
- fri = contact_friction_in[conid][0] / wp.sqrt(opt_impratio[worldid])
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.SATISFIED.value)
else:
- fri = contact_friction_in[conid][dimid - 1]
- u = efc_Jaref_in[worldid, efcid] * fri
- efc_u_out[conid][dimid] = u
- if dimid > 0:
- wp.atomic_add(efc_uu_out, conid, u * u)
+ efc_force_out[worldid, efcid] = -efc_D * Jaref
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.QUADRATIC.value)
+ wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref)
+ else: # elliptic friction cone contact
+ conid = efc_id_in[worldid, efcid]
-
-@wp.kernel
-def update_constraint_active_elliptic_bottom_zone(
- # Model:
- opt_impratio: wp.array(dtype=float),
- # Data in:
- ncon_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_dim_in: wp.array(dtype=int),
- contact_efc_address_in: wp.array2d(dtype=int),
- contact_worldid_in: wp.array(dtype=int),
- efc_done_in: wp.array(dtype=bool),
- efc_u_in: wp.array(dtype=types.vec6),
- efc_uu_in: wp.array(dtype=float),
- # Data out:
- efc_active_out: wp.array2d(dtype=bool),
-):
- conid, dimid = wp.tid()
-
- if conid >= ncon_in[0]:
- return
-
- worldid = contact_worldid_in[conid]
- if efc_done_in[worldid]:
- return
-
- condim = contact_dim_in[conid]
- if condim == 1:
- return
-
- mu = contact_friction_in[conid][0] / wp.sqrt(opt_impratio[worldid])
- n = efc_u_in[conid][0]
- tt = efc_uu_in[conid]
- if tt <= 0.0:
- t = 0.0
- else:
- t = wp.sqrt(tt)
-
- # bottom zone: quadratic
- bottom_zone = ((t <= 0.0) and (n < 0.0)) or ((t > 0.0) and ((mu * n + t) <= 0.0))
-
- # update active
- efcid = contact_efc_address_in[conid, dimid]
- efc_active_out[worldid, efcid] = bottom_zone
-
-
-@wp.kernel
-def update_constraint_efc_elliptic0(
- # Data in:
- ne_in: wp.array(dtype=int),
- nf_in: wp.array(dtype=int),
- nl_in: wp.array(dtype=int),
- nefc_in: wp.array(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_frictionloss_in: wp.array2d(dtype=float),
- efc_Jaref_in: wp.array2d(dtype=float),
- efc_active_in: wp.array2d(dtype=bool),
- efc_done_in: wp.array(dtype=bool),
- # In:
- disable_floss: int,
- # Data out:
- efc_force_out: wp.array2d(dtype=float),
- efc_cost_out: wp.array(dtype=float),
- efc_active_out: wp.array2d(dtype=bool),
-):
- worldid, efcid = wp.tid()
-
- if efcid >= nefc_in[worldid]:
- return
-
- if efc_done_in[worldid]:
- return
-
- efc_D = efc_D_in[worldid, efcid]
- Jaref = efc_Jaref_in[worldid, efcid]
-
- ne = ne_in[worldid]
- nf = nf_in[worldid]
- nl = nl_in[worldid]
-
- if efcid < ne:
- # equality
- efc_active_out[worldid, efcid] = True
- elif efcid < ne + nf and not disable_floss:
- # friction
- f = efc_frictionloss_in[worldid, efcid]
- if f > 0.0:
- rf = math.safe_div(f, efc_D)
- if Jaref <= -rf:
- efc_force_out[worldid, efcid] = f
- efc_active_out[worldid, efcid] = False
- wp.atomic_add(efc_cost_out, worldid, -0.5 * rf - Jaref)
- return
- elif Jaref >= rf:
- efc_force_out[worldid, efcid] = -f
- efc_active_out[worldid, efcid] = False
- wp.atomic_add(efc_cost_out, worldid, -0.5 * rf + Jaref)
- return
- elif efcid < ne + nf + nl:
- # limits
- if Jaref < 0.0:
- efc_active_out[worldid, efcid] = True
- else:
- efc_force_out[worldid, efcid] = 0.0
- efc_active_out[worldid, efcid] = False
- return
- else:
- # contact
- if not efc_active_in[worldid, efcid]: # calculated by solve_active_elliptic_bottom_zone
- efc_force_out[worldid, efcid] = 0.0
+ if conid >= ncon_in[0]:
return
- efc_force_out[worldid, efcid] = -efc_D * Jaref
- wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref)
+ dim = contact_dim_in[conid]
+ friction = contact_friction_in[conid]
+ mu = friction[0] / wp.sqrt(opt_impratio[worldid])
-
-@wp.kernel
-def update_constraint_efc_elliptic1(
- # Model:
- opt_impratio: wp.array(dtype=float),
- # Data in:
- ncon_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_dim_in: wp.array(dtype=int),
- contact_efc_address_in: wp.array2d(dtype=int),
- contact_worldid_in: wp.array(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_done_in: wp.array(dtype=bool),
- efc_u_in: wp.array(dtype=types.vec6),
- efc_uu_in: wp.array(dtype=float),
- # Data out:
- efc_force_out: wp.array2d(dtype=float),
- efc_cost_out: wp.array(dtype=float),
-):
- conid, dimid = wp.tid()
-
- if conid >= ncon_in[0]:
- return
-
- worldid = contact_worldid_in[conid]
- if efc_done_in[worldid]:
- return
-
- condim = contact_dim_in[conid]
-
- if condim == 1 or dimid >= condim:
- return
-
- friction = contact_friction_in[conid]
- efcid = contact_efc_address_in[conid, dimid]
-
- mu = friction[0] / wp.sqrt(opt_impratio[worldid])
- n = efc_u_in[conid][0]
- tt = efc_uu_in[conid]
- if tt <= 0.0:
- t = 0.0
- else:
- t = wp.sqrt(tt)
-
- # middle zone: cone
- middle_zone = (t > 0.0) and (n < (mu * t)) and ((mu * n + t) > 0.0)
-
- # tangent and friction for middle zone:
- if middle_zone:
efcid0 = contact_efc_address_in[conid, 0]
- mu2 = mu * mu
- dm = efc_D_in[worldid, efcid0] / wp.max(mu2 * float(1.0 + mu2), types.MJ_MINVAL)
+ if efcid0 < 0:
+ return
- nmt = n - mu * t
+ N = efc_Jaref_in[worldid, efcid0] * mu
- force = -dm * nmt * mu
- if dimid > 0:
- force_fri = -force / t
- force_fri *= efc_u_in[conid][dimid] * friction[dimid - 1]
- efc_force_out[worldid, efcid] += force_fri
+ ufrictionj = float(0.0)
+ TT = float(0.0)
+ for j in range(1, dim):
+ efcidj = contact_efc_address_in[conid, j]
+ if efcidj < 0:
+ return
+ frictionj = friction[j - 1]
+ uj = efc_Jaref_in[worldid, efcidj] * frictionj
+ TT += uj * uj
+ if efcid == efcidj:
+ ufrictionj = uj * frictionj
+
+ if TT <= 0.0:
+ T = 0.0
else:
- efc_force_out[worldid, efcid] += force
- worldid = contact_worldid_in[conid]
- wp.atomic_add(efc_cost_out, worldid, 0.5 * dm * nmt * nmt)
+ T = wp.sqrt(TT)
+
+ # top zone
+ if (N >= mu * T) or ((T <= 0.0) and (N >= 0.0)):
+ efc_force_out[worldid, efcid] = 0.0
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.SATISFIED.value)
+ # bottom zone
+ elif (mu * N + T <= 0.0) or ((T <= 0.0) and (N < 0.0)):
+ efc_force_out[worldid, efcid] = -efc_D * Jaref
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.QUADRATIC.value)
+ wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref)
+ # middle zone
+ else:
+ dm = math.safe_div(efc_D_in[worldid, efcid0], mu * mu * (1.0 + mu * mu))
+ nmt = N - mu * T
+
+ force = -dm * nmt * mu
+
+ if efcid == efcid0:
+ efc_force_out[worldid, efcid] = force
+ wp.atomic_add(efc_cost_out, worldid, 0.5 * dm * nmt * nmt)
+ else:
+ efc_force_out[worldid, efcid] = -math.safe_div(force, T) * ufrictionj
+
+ efc_state_out[worldid, efcid] = int(types.ConstraintState.CONE.value)
@wp.kernel
@@ -1653,6 +1426,7 @@ def update_constraint_zero_qfrc_constraint(
@wp.kernel
def update_constraint_init_qfrc_constraint(
# Data in:
+ njmax_in: int,
nefc_in: wp.array(dtype=int),
efc_J_in: wp.array3d(dtype=float),
efc_force_in: wp.array2d(dtype=float),
@@ -1666,7 +1440,7 @@ def update_constraint_init_qfrc_constraint(
return
sum_qfrc = float(0.0)
- for efcid in range(nefc_in[worldid]):
+ for efcid in range(min(njmax_in, nefc_in[worldid])):
efc_J = efc_J_in[worldid, efcid, dofid]
force = efc_force_in[worldid, efcid]
sum_qfrc += efc_J * force
@@ -1717,8 +1491,6 @@ def update_constraint_gauss_cost(nv: int, dofs_per_thread: int):
def _update_constraint(m: types.Model, d: types.Data):
"""Update constraint arrays after each solve iteration."""
- disable_floss = m.opt.disableflags & types.DisableBit.FRICTIONLOSS
-
wp.launch(
update_constraint_init_cost,
dim=(d.nworld),
@@ -1726,92 +1498,27 @@ def _update_constraint(m: types.Model, d: types.Data):
outputs=[d.efc.gauss, d.efc.cost, d.efc.prev_cost],
)
- if m.opt.cone == types.ConeType.PYRAMIDAL:
- wp.launch(
- update_constraint_efc_pyramidal,
- dim=(d.nworld, d.njmax),
- inputs=[
- d.ne,
- d.nf,
- d.nefc,
- d.efc.D,
- d.efc.frictionloss,
- d.efc.Jaref,
- disable_floss,
- ],
- outputs=[d.efc.force, d.efc.cost, d.efc.active],
- )
- elif m.opt.cone == types.ConeType.ELLIPTIC:
- d.efc.uu.zero_()
- d.efc.active.zero_()
- d.efc.condim.fill_(-1)
- wp.launch(
- update_constraint_u_elliptic,
- dim=(d.nconmax, m.condim_max),
- inputs=[
- m.opt.impratio,
- d.ncon,
- d.contact.friction,
- d.contact.dim,
- d.contact.efc_address,
- d.contact.worldid,
- d.efc.Jaref,
- d.efc.done,
- ],
- outputs=[d.efc.u, d.efc.uu, d.efc.condim],
- )
- wp.launch(
- update_constraint_active_elliptic_bottom_zone,
- dim=(d.nconmax, m.condim_max),
- inputs=[
- m.opt.impratio,
- d.ncon,
- d.contact.friction,
- d.contact.dim,
- d.contact.efc_address,
- d.contact.worldid,
- d.efc.done,
- d.efc.u,
- d.efc.uu,
- ],
- outputs=[d.efc.active],
- )
- wp.launch(
- update_constraint_efc_elliptic0,
- dim=(d.nworld, d.njmax),
- inputs=[
- d.ne,
- d.nf,
- d.nl,
- d.nefc,
- d.efc.D,
- d.efc.frictionloss,
- d.efc.Jaref,
- d.efc.active,
- d.efc.done,
- disable_floss,
- ],
- outputs=[d.efc.force, d.efc.cost, d.efc.active],
- )
- wp.launch(
- update_constraint_efc_elliptic1,
- dim=(d.nconmax, m.condim_max),
- inputs=[
- m.opt.impratio,
- d.ncon,
- d.contact.friction,
- d.contact.dim,
- d.contact.efc_address,
- d.contact.worldid,
- d.efc.D,
- d.efc.done,
- d.efc.u,
- d.efc.uu,
- ],
- outputs=[d.efc.force, d.efc.cost],
- )
- else:
- raise ValueError(f"Unknown cone type: {m.opt.cone}")
+ wp.launch(
+ update_constraint_efc,
+ dim=(d.nworld, d.njmax),
+ inputs=[
+ m.opt.impratio,
+ d.ncon,
+ d.ne,
+ d.nf,
+ d.nefc,
+ d.contact.friction,
+ d.contact.dim,
+ d.contact.efc_address,
+ d.efc.type,
+ d.efc.id,
+ d.efc.D,
+ d.efc.frictionloss,
+ d.efc.Jaref,
+ d.efc.done,
+ ],
+ outputs=[d.efc.force, d.efc.cost, d.efc.state],
+ )
# qfrc_constraint = efc_J.T @ efc_force
wp.launch(
@@ -1824,7 +1531,7 @@ def _update_constraint(m: types.Model, d: types.Data):
wp.launch(
update_constraint_init_qfrc_constraint,
dim=(d.nworld, m.nv),
- inputs=[d.nefc, d.efc.J, d.efc.force, d.efc.done],
+ inputs=[d.njmax, d.nefc, d.efc.J, d.efc.force, d.efc.done],
outputs=[d.qfrc_constraint],
)
@@ -1947,10 +1654,11 @@ def update_gradient_copy_lower_triangle(
@wp.kernel
def update_gradient_JTDAJ(
# Data in:
+ njmax_in: int,
nefc_in: wp.array(dtype=int),
efc_J_in: wp.array3d(dtype=float),
efc_D_in: wp.array2d(dtype=float),
- efc_active_in: wp.array2d(dtype=bool),
+ efc_state_in: wp.array2d(dtype=int),
efc_done_in: wp.array(dtype=bool),
# Data out:
efc_h_out: wp.array3d(dtype=float),
@@ -1967,23 +1675,26 @@ def update_gradient_JTDAJ(
sum_h = float(0.0)
efc_D = efc_D_in[worldid, 0]
- active = efc_active_in[worldid, 0]
+ # TODO(team): sparse efc_J
efc_Ji = efc_J_in[worldid, 0, dofi]
efc_Jj = efc_J_in[worldid, 0, dofj]
- for efcid in range(nefc - 1):
- # TODO(team): sparse efc_J
- sum_h += efc_Ji * efc_Jj * efc_D * float(active)
+ efc_state = efc_state_in[worldid, 0]
+ for efcid in range(min(njmax_in, nefc) - 1):
+ if efc_state == int(types.ConstraintState.QUADRATIC.value) and efc_D != 0.0:
+ sum_h += efc_Ji * efc_Jj * efc_D
jj = efcid + 1
efc_D = efc_D_in[worldid, jj]
- active = efc_active_in[worldid, jj]
efc_Ji = efc_J_in[worldid, jj, dofi]
efc_Jj = efc_J_in[worldid, jj, dofj]
+ efc_state = efc_state_in[worldid, jj]
- sum_h += efc_Ji * efc_Jj * efc_D * float(active)
+ if efc_state == int(types.ConstraintState.QUADRATIC.value) and efc_D != 0.0:
+ sum_h += efc_Ji * efc_Jj * efc_D
efc_h_out[worldid, dofi, dofj] += sum_h
+# TODO(thowell): combine with JTDAJ ?
@wp.kernel
def update_gradient_JTCJ(
# Model:
@@ -1993,15 +1704,17 @@ def update_gradient_JTCJ(
# Data in:
nconmax_in: int,
ncon_in: wp.array(dtype=int),
+ contact_dist_in: wp.array(dtype=float),
+ contact_includemargin_in: wp.array(dtype=float),
contact_friction_in: wp.array(dtype=types.vec5),
contact_dim_in: wp.array(dtype=int),
contact_efc_address_in: wp.array2d(dtype=int),
contact_worldid_in: wp.array(dtype=int),
efc_J_in: wp.array3d(dtype=float),
efc_D_in: wp.array2d(dtype=float),
+ efc_Jaref_in: wp.array2d(dtype=float),
+ efc_state_in: wp.array2d(dtype=int),
efc_done_in: wp.array(dtype=bool),
- efc_u_in: wp.array(dtype=types.vec6),
- efc_uu_in: wp.array(dtype=float),
# In:
nblocks_perblock: int,
dim_block: int,
@@ -2028,37 +1741,45 @@ def update_gradient_JTCJ(
if condim == 1:
continue
- fri = contact_friction_in[conid]
- mu = fri[0] / wp.sqrt(opt_impratio[worldid])
- n = efc_u_in[conid][0]
- tt = efc_uu_in[conid]
- if tt <= 0.0:
- t = 0.0
- else:
- t = wp.sqrt(tt)
-
- middle_zone = (t > 0) and (n < (mu * t)) and ((mu * n + t) > 0.0)
-
- if not middle_zone:
+ # check contact status
+ if contact_dist_in[conid] - contact_includemargin_in[conid] >= 0.0:
continue
- t = wp.max(t, types.MJ_MINVAL)
- ttt = wp.max(t * t * t, types.MJ_MINVAL)
+ efcid0 = contact_efc_address_in[conid, 0]
+ if efc_state_in[worldid, efcid0] != int(types.ConstraintState.CONE.value):
+ continue
+
+ fri = contact_friction_in[conid]
+ mu = math.safe_div(fri[0], wp.sqrt(opt_impratio[worldid]))
mu2 = mu * mu
- efc0 = contact_efc_address_in[conid, 0]
- dm = efc_D_in[worldid, efc0] / wp.max(mu2 * (1.0 + mu2), types.MJ_MINVAL)
+ dm = math.safe_div(efc_D_in[worldid, efcid0], mu2 * (1.0 + mu2))
if dm == 0.0:
continue
- u = efc_u_in[conid]
+ n = efc_Jaref_in[worldid, efcid0] * mu
+ u = types.vec6(n, 0.0, 0.0, 0.0, 0.0, 0.0)
+
+ tt = float(0.0)
+ for j in range(1, condim):
+ efcidj = contact_efc_address_in[conid, j]
+ uj = efc_Jaref_in[worldid, efcidj] * fri[j - 1]
+ tt += uj * uj
+ u[j] = uj
+
+ if tt <= 0.0:
+ t = 0.0
+ else:
+ t = wp.sqrt(tt)
+ t = wp.max(t, types.MJ_MINVAL)
+ ttt = wp.max(t * t * t, types.MJ_MINVAL)
efc_h = float(0.0)
for dim1id in range(condim):
if dim1id == 0:
- efcid1 = efc0
+ efcid1 = efcid0
else:
efcid1 = contact_efc_address_in[conid, dim1id]
@@ -2069,7 +1790,7 @@ def update_gradient_JTCJ(
for dim2id in range(0, dim1id + 1):
if dim2id == 0:
- efcid2 = efc0
+ efcid2 = efcid0
else:
efcid2 = contact_efc_address_in[conid, dim2id]
@@ -2082,15 +1803,15 @@ def update_gradient_JTCJ(
if dim1id == 0 and dim2id == 0:
hcone = 1.0
elif dim1id == 0:
- hcone = -mu / t * uj
+ hcone = -math.safe_div(mu, t) * uj
elif dim2id == 0:
- hcone = -mu / t * ui
+ hcone = -math.safe_div(mu, t) * ui
else:
- hcone = mu * n / ttt * ui * uj
+ hcone = mu * math.safe_div(n, ttt) * ui * uj
# add to diagonal: mu^2 - mu * n / t
if dim1id == dim2id:
- hcone += mu2 - mu * n / t
+ hcone += mu2 - mu * math.safe_div(n, t)
# pre and post multiply by diag(mu, friction) scale by dm
if dim1id == 0:
@@ -2111,7 +1832,6 @@ def update_gradient_JTCJ(
if dim1id != dim2id:
efc_h += hcone * efc_J12 * efc_J21
- worldid = contact_worldid_in[conid]
efc_h_out[worldid, dof1id, dof2id] += efc_h
@@ -2211,10 +1931,11 @@ def _update_gradient(m: types.Model, d: types.Data):
update_gradient_JTDAJ,
dim=(d.nworld, lower_triangle_dim),
inputs=[
+ d.njmax,
d.nefc,
d.efc.J,
d.efc.D,
- d.efc.active,
+ d.efc.state,
d.efc.done,
],
outputs=[d.efc.h],
@@ -2253,15 +1974,17 @@ def _update_gradient(m: types.Model, d: types.Data):
m.dof_tri_col,
d.nconmax,
d.ncon,
+ d.contact.dist,
+ d.contact.includemargin,
d.contact.friction,
d.contact.dim,
d.contact.efc_address,
d.contact.worldid,
d.efc.J,
d.efc.D,
+ d.efc.Jaref,
+ d.efc.state,
d.efc.done,
- d.efc.u,
- d.efc.uu,
nblocks_perblock,
dim_block,
],
@@ -2549,7 +2272,7 @@ def _solve(m: types.Model, d: types.Data):
# This branch is mostly for when JAX is used as it is currently not compatible
# with CUDA graph conditional.
# It should be removed when JAX becomes compatible.
- for i in range(m.opt.iterations):
+ for _ in range(m.opt.iterations):
_solver_iteration(m, d)
wp.copy(d.qacc_warmstart, d.qacc)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py
index c4629d20..610a8f68 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py
@@ -25,6 +25,7 @@ import mujoco_warp as mjwarp
from mujoco.mjx.third_party.mujoco_warp._src import solver
from mujoco.mjx.third_party.mujoco_warp._src import test_util
+from mujoco.mjx.third_party.mujoco_warp._src.math import safe_div
from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType
@@ -34,7 +35,7 @@ _TOLERANCE = 5e-3
def _assert_eq(a, b, name):
- tol = _TOLERANCE * 10 # avoid test noise
+ tol = _TOLERANCE * 20 # avoid test noise
err_msg = f"mismatch: {name}"
np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
@@ -68,6 +69,93 @@ class SolverTest(parameterized.TestCase):
_assert_eq(mjwarp_cost, mj_cost, name="cost")
+ def test_init_linesearch(self):
+ """Test linesearch initialization."""
+ for keyframe in range(3):
+ # TODO(team): Add the case of elliptic cone friction
+ mjm, mjd, m, d = test_util.fixture(
+ "constraints.xml",
+ keyframe=keyframe,
+ cone=ConeType.PYRAMIDAL,
+ iterations=0,
+ ls_iterations=0,
+ )
+
+ # One step to obtain more non-zeros results
+ mjwarp.step(m, d)
+
+ # Calculate target values
+ efc_search_np = d.efc.search.numpy()[0]
+ efc_J_np = d.efc.J.numpy()[0]
+ efc_gauss_np = d.efc.gauss.numpy()[0]
+ efc_Ma_np = d.efc.Ma.numpy()[0]
+ efc_Jaref_np = d.efc.Jaref.numpy()[0]
+ efc_D_np = d.efc.D.numpy()[0]
+ qfrc_smooth_np = d.qfrc_smooth.numpy()[0]
+ nefc = d.nefc.numpy()[0]
+
+ target_mv = np.zeros(mjm.nv)
+ mujoco.mj_mulM(mjm, mjd, target_mv, efc_search_np)
+ target_jv = efc_J_np @ efc_search_np
+ target_quad_gauss = np.array(
+ [
+ efc_gauss_np,
+ np.dot(efc_search_np, efc_Ma_np - qfrc_smooth_np),
+ 0.5 * np.dot(efc_search_np, target_mv),
+ ]
+ )
+ target_quad = np.transpose(
+ np.vstack(
+ [
+ 0.5 * efc_Jaref_np * efc_Jaref_np * efc_D_np,
+ target_jv * efc_Jaref_np * efc_D_np,
+ 0.5 * target_jv * target_jv * efc_D_np,
+ ]
+ )
+ )
+
+ # launch linesearch with 0 iteration just doing the initialization step
+ d.efc.jv.zero_()
+ d.efc.quad.zero_()
+ solver._linesearch(m, d)
+
+ efc_mv = d.efc.mv.numpy()[0]
+ efc_jv = d.efc.jv.numpy()[0]
+ efc_quad_gauss = d.efc.quad_gauss.numpy()[0]
+ efc_quad = d.efc.quad.numpy()[0]
+ _assert_eq(efc_mv, target_mv, "mv")
+ _assert_eq(efc_jv[:nefc], target_jv[:nefc], "jv")
+ _assert_eq(efc_quad_gauss, target_quad_gauss, "quad_gauss")
+ _assert_eq(efc_quad[:nefc], target_quad[:nefc], "quad")
+
+ @parameterized.parameters(
+ (ConeType.PYRAMIDAL, False),
+ (ConeType.ELLIPTIC, False),
+ (ConeType.PYRAMIDAL, True),
+ (ConeType.ELLIPTIC, True),
+ )
+ def test_update_gradient_CG(self, cone, sparse):
+ """Test _update_gradient function is correct for the CG solver."""
+ mjm, mjd, m, d = test_util.fixture(
+ "humanoid/humanoid.xml",
+ cone=cone,
+ solver=SolverType.CG,
+ sparse=sparse,
+ iterations=0,
+ keyframe=0,
+ )
+
+ # Solve with 0 iterations just initializes and exit
+ mjwarp.solve(m, d)
+
+ # Calculate Mgrad with Mujoco C
+ mj_Mgrad = np.zeros(shape=(1, mjm.nv), dtype=float)
+ mj_grad = np.tile(d.efc.grad.numpy(), (1, 1))
+ mujoco.mj_solveM(mjm, mjd, mj_Mgrad, mj_grad)
+
+ efc_Mgrad = d.efc.Mgrad.numpy()[0]
+ _assert_eq(efc_Mgrad, mj_Mgrad[0], name="Mgrad")
+
@parameterized.parameters(ConeType.PYRAMIDAL, ConeType.ELLIPTIC)
def test_parallel_linesearch(self, cone):
"""Test that iterative and parallel linesearch leads to equivalent results."""
@@ -126,12 +214,12 @@ class SolverTest(parameterized.TestCase):
self.assertLessEqual(abs(alpha_iterative - alpha_parallel_50), abs(alpha_iterative - alpha_parallel_10))
@parameterized.parameters(
- (ConeType.PYRAMIDAL, SolverType.CG, 5, 5, False, False),
- (ConeType.ELLIPTIC, SolverType.CG, 5, 5, False, False),
- (ConeType.PYRAMIDAL, SolverType.NEWTON, 2, 4, False, False),
- (ConeType.ELLIPTIC, SolverType.NEWTON, 2, 5, False, False),
- (ConeType.PYRAMIDAL, SolverType.NEWTON, 2, 4, True, True),
- (ConeType.ELLIPTIC, SolverType.NEWTON, 3, 16, True, True),
+ (ConeType.PYRAMIDAL, SolverType.CG, 10, 5, False, False),
+ (ConeType.ELLIPTIC, SolverType.CG, 10, 5, False, False),
+ (ConeType.PYRAMIDAL, SolverType.NEWTON, 5, 10, False, False),
+ (ConeType.ELLIPTIC, SolverType.NEWTON, 5, 10, False, False),
+ (ConeType.PYRAMIDAL, SolverType.NEWTON, 5, 64, True, True),
+ (ConeType.ELLIPTIC, SolverType.NEWTON, 5, 64, True, True),
)
def test_solve(self, cone, solver_, iterations, ls_iterations, sparse, ls_parallel):
"""Tests solve."""
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py
index 8c5197ee..2a1c1613 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py
@@ -254,12 +254,6 @@ class JointType(enum.IntEnum):
SLIDE = mujoco.mjtJoint.mjJNT_SLIDE
HINGE = mujoco.mjtJoint.mjJNT_HINGE
- def dof_width(self) -> int:
- return {0: 6, 1: 3, 2: 1, 3: 1}[self.value]
-
- def qpos_width(self) -> int:
- return {0: 7, 1: 4, 2: 1, 3: 1}[self.value]
-
class ConeType(enum.IntEnum):
"""Type of friction cone.
@@ -328,6 +322,24 @@ class SolverType(enum.IntEnum):
# unsupported: PGS
+class ConstraintState(enum.IntEnum):
+ """State of constraint.
+
+ Attributes:
+ SATISFIED: constraint satisfied, zero cost (limit, contact)
+ QUADRATIC: quadratic cost (equality, friction, limit, contact)
+ LINEARNEG: linear cost, negative side (friction)
+ LINEARPOS: linear cost, positive side (friction)
+ CONE: square distance to cone cost (elliptic contact)
+ """
+
+ SATISFIED = mujoco.mjtConstraintState.mjCNSTRSTATE_SATISFIED
+ QUADRATIC = mujoco.mjtConstraintState.mjCNSTRSTATE_QUADRATIC
+ LINEARNEG = mujoco.mjtConstraintState.mjCNSTRSTATE_LINEARNEG
+ LINEARPOS = mujoco.mjtConstraintState.mjCNSTRSTATE_LINEARPOS
+ CONE = mujoco.mjtConstraintState.mjCNSTRSTATE_CONE
+
+
class ConstraintType(enum.IntEnum):
"""Type of constraint.
@@ -473,6 +485,7 @@ class EqType(enum.IntEnum):
CONNECT: connect two bodies at a point (ball joint)
JOINT: couple the values of two scalar joints with cubic
WELD: fix relative position and orientation of two bodies
+ TENDON: couple the lengths of two tendons with cubic
"""
CONNECT = mujoco.mjtEq.mjEQ_CONNECT
@@ -623,7 +636,7 @@ class Constraint:
gauss: gauss Cost (nworld,)
cost: constraint + Gauss cost (nworld,)
prev_cost: cost from previous iter (nworld,)
- active: active (quadratic) constraints (nworld, njmax)
+ state: constraint state (nworld, njmax)
gtol: linesearch termination tolerance (nworld,)
mv: qM @ search (nworld, nv)
jv: efc_J @ search (nworld, njmax)
@@ -648,11 +661,6 @@ class Constraint:
mid: loss at mid_alpha (nworld, 3)
mid_alpha: midpoint between lo_alpha and hi_alpha (nworld,)
cost_candidate: costs associated with step sizes (nworld, nlsp)
- u: friction cone (normal and tangents) (nconmax, 6)
- uu: elliptic cone variables (nconmax,)
- uv: elliptic cone variables (nconmax,)
- vv: elliptic cone variables (nconmax,)
- condim: if contact: condim, else: -1 (nworld, njmax)
"""
type: wp.array2d(dtype=int)
@@ -677,7 +685,7 @@ class Constraint:
gauss: wp.array(dtype=float)
cost: wp.array(dtype=float)
prev_cost: wp.array(dtype=float)
- active: wp.array2d(dtype=bool)
+ state: wp.array2d(dtype=int)
gtol: wp.array(dtype=float)
mv: wp.array2d(dtype=float)
jv: wp.array2d(dtype=float)
@@ -703,12 +711,6 @@ class Constraint:
mid: wp.array(dtype=wp.vec3)
mid_alpha: wp.array(dtype=float)
cost_candidate: wp.array2d(dtype=float)
- # elliptic cone
- u: wp.array(dtype=vec6)
- uu: wp.array(dtype=float)
- uv: wp.array(dtype=float)
- vv: wp.array(dtype=float)
- condim: wp.array2d(dtype=int)
@dataclasses.dataclass
@@ -1013,7 +1015,8 @@ class Model:
sensor_e_kinetic: evaluate energy_vel
sensor_tendonactfrc_adr: address for tendonactfrc sensor (<=nsensor,)
sensor_subtree_vel: evaluate subtree_vel
- sensor_contact_adr: addresses for contact sensors
+ sensor_contact_adr: addresses for contact sensors (<=nsensor,)
+ sensor_adr_to_contact_adr: map sensor adr to contact adr (nsensor,)
sensor_rne_postconstraint: evaluate rne_postconstraint
sensor_rangefinder_bodyid: bodyid for rangefinder (nrangefinder,)
plugin: globally registered plugin slot number (nplugin,)
@@ -1324,6 +1327,7 @@ class Model:
sensor_tendonactfrc_adr: wp.array(dtype=int) # warp only
sensor_subtree_vel: bool # warp only
sensor_contact_adr: wp.array(dtype=int) # warp only
+ sensor_adr_to_contact_adr: wp.array(dtype=int) # warp only
sensor_rne_postconstraint: bool # warp only
sensor_rangefinder_bodyid: wp.array(dtype=int) # warp only
plugin: wp.array(dtype=int)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py
index 0f863b88..04c43001 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py
@@ -143,21 +143,21 @@ def wrap_circle(end: wp.vec4, side: wp.vec2, radius: float) -> Tuple[float, wp.v
# construct the two solutions, compute goodness
sol00 = wp.vec2(
- (end[0] * sqrad + radius * end[1] * sqrt0) / sqlen0,
- (end[1] * sqrad - radius * end[0] * sqrt0) / sqlen0,
+ math.safe_div(end[0] * sqrad + radius * end[1] * sqrt0, sqlen0),
+ math.safe_div(end[1] * sqrad - radius * end[0] * sqrt0, sqlen0),
)
sol01 = wp.vec2(
- (end[2] * sqrad - radius * end[3] * sqrt1) / sqlen1,
- (end[3] * sqrad + radius * end[2] * sqrt1) / sqlen1,
+ math.safe_div(end[2] * sqrad - radius * end[3] * sqrt1, sqlen1),
+ math.safe_div(end[3] * sqrad + radius * end[2] * sqrt1, sqlen1),
)
sol10 = wp.vec2(
- (end[0] * sqrad - radius * end[1] * sqrt0) / sqlen0,
- (end[1] * sqrad + radius * end[0] * sqrt0) / sqlen0,
+ math.safe_div(end[0] * sqrad - radius * end[1] * sqrt0, sqlen0),
+ math.safe_div(end[1] * sqrad + radius * end[0] * sqrt0, sqlen0),
)
sol11 = wp.vec2(
- (end[2] * sqrad + radius * end[3] * sqrt1) / sqlen1,
- (end[3] * sqrad - radius * end[2] * sqrt1) / sqlen1,
+ math.safe_div(end[2] * sqrad + radius * end[3] * sqrt1, sqlen1),
+ math.safe_div(end[3] * sqrad - radius * end[2] * sqrt1, sqlen1),
)
# goodness: close to sd, or shorter path
@@ -249,11 +249,11 @@ def wrap_inside(
pnt *= radius
# compute function parameters: asin(A * z) + asin(B * z) - 2 * asin(z) + G = 0
- A = radius / len0
- B = radius / len1
+ A = math.safe_div(radius, len0)
+ B = math.safe_div(radius, len1)
sq_A = A * A
sq_B = B * B
- cosG = (len0 * len0 + len1 * len1 - dd) / (2.0 * len0 * len1)
+ cosG = math.safe_div(len0 * len0 + len1 * len1 - dd, 2.0 * len0 * len1)
if cosG < -1.0 + MJ_MINVAL:
return -1.0, pnt, pnt
elif cosG > 1.0 - MJ_MINVAL:
@@ -285,7 +285,7 @@ def wrap_inside(
return 0.0, pnt, pnt
# new point
- z1 = z - f / df
+ z1 = z - math.safe_div(f, df)
# make sure we are moving to the left; SHOULD NOT OCCUR
if z1 > z:
@@ -435,8 +435,8 @@ def wrap(
# set vertical coordinates
L0 = wp.sqrt((p0[0] - res0[0]) * (p0[0] - res0[0]) + (p0[1] - res0[1]) * (p0[1] - res0[1]))
L1 = wp.sqrt((p1[0] - res1[0]) * (p1[0] - res1[0]) + (p1[1] - res1[1]) * (p1[1] - res1[1]))
- res0[2] = p0[2] + (p1[2] - p0[2]) * L0 / (L0 + wlen + L1)
- res1[2] = p0[2] + (p1[2] - p0[2]) * (L0 + wlen) / (L0 + wlen + L1)
+ res0[2] = p0[2] + (p1[2] - p0[2]) * math.safe_div(L0, L0 + wlen + L1)
+ res1[2] = p0[2] + (p1[2] - p0[2]) * math.safe_div(L0 + wlen, L0 + wlen + L1)
# correct wlen for height
height = wp.abs(res1[2] - res0[2])
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml
index 8e5be10f..f6fb68e0 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/constraints.xml
@@ -58,7 +58,7 @@
-
+
@@ -68,12 +68,12 @@
-
+
-
+
diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py
index d3a2a4b5..1daac020 100644
--- a/mjx/mujoco/mjx/warp/forward.py
+++ b/mjx/mujoco/mjx/warp/forward.py
@@ -227,7 +227,7 @@ def _forward_shim(
nlsp: int,
nmeshface: int,
nmocap: int,
- nsensordata: int,
+ nsensor: int,
nsensortaxel: int,
nsite: int,
ntendon: int,
@@ -257,6 +257,7 @@ def _forward_shim(
rangefinder_sensor_adr: wp.array(dtype=int),
sensor_acc_adr: wp.array(dtype=int),
sensor_adr: wp.array(dtype=int),
+ sensor_adr_to_contact_adr: wp.array(dtype=int),
sensor_contact_adr: wp.array(dtype=int),
sensor_cutoff: wp.array(dtype=float),
sensor_datatype: wp.array(dtype=int),
@@ -481,13 +482,11 @@ def _forward_shim(
efc__Jaref: wp.array2d(dtype=float),
efc__Ma: wp.array2d(dtype=float),
efc__Mgrad: wp.array2d(dtype=float),
- efc__active: wp.array2d(dtype=bool),
efc__alpha: wp.array(dtype=float),
efc__aref: wp.array2d(dtype=float),
efc__beta: wp.array(dtype=float),
efc__cholesky_L_tmp: wp.array3d(dtype=float),
efc__cholesky_y_tmp: wp.array2d(dtype=float),
- efc__condim: wp.array2d(dtype=int),
efc__cost: wp.array(dtype=float),
efc__cost_candidate: wp.array2d(dtype=float),
efc__done: wp.array(dtype=bool),
@@ -522,12 +521,9 @@ def _forward_shim(
efc__quad_gauss: wp.array(dtype=wp.vec3),
efc__search: wp.array2d(dtype=float),
efc__search_dot: wp.array(dtype=float),
+ efc__state: wp.array2d(dtype=int),
efc__type: wp.array2d(dtype=int),
- efc__u: wp.array(dtype=mjwp_types.vec6),
- efc__uu: wp.array(dtype=float),
- efc__uv: wp.array(dtype=float),
efc__vel: wp.array2d(dtype=float),
- efc__vv: wp.array(dtype=float),
):
_m.stat = _s
_m.opt = _o
@@ -713,7 +709,7 @@ def _forward_shim(
_m.nlsp = nlsp
_m.nmeshface = nmeshface
_m.nmocap = nmocap
- _m.nsensordata = nsensordata
+ _m.nsensor = nsensor
_m.nsensortaxel = nsensortaxel
_m.nsite = nsite
_m.ntendon = ntendon
@@ -769,6 +765,7 @@ def _forward_shim(
_m.rangefinder_sensor_adr = rangefinder_sensor_adr
_m.sensor_acc_adr = sensor_acc_adr
_m.sensor_adr = sensor_adr
+ _m.sensor_adr_to_contact_adr = sensor_adr_to_contact_adr
_m.sensor_contact_adr = sensor_contact_adr
_m.sensor_cutoff = sensor_cutoff
_m.sensor_datatype = sensor_datatype
@@ -868,13 +865,11 @@ def _forward_shim(
_d.efc.Jaref = efc__Jaref
_d.efc.Ma = efc__Ma
_d.efc.Mgrad = efc__Mgrad
- _d.efc.active = efc__active
_d.efc.alpha = efc__alpha
_d.efc.aref = efc__aref
_d.efc.beta = efc__beta
_d.efc.cholesky_L_tmp = efc__cholesky_L_tmp
_d.efc.cholesky_y_tmp = efc__cholesky_y_tmp
- _d.efc.condim = efc__condim
_d.efc.cost = efc__cost
_d.efc.cost_candidate = efc__cost_candidate
_d.efc.done = efc__done
@@ -909,12 +904,9 @@ def _forward_shim(
_d.efc.quad_gauss = efc__quad_gauss
_d.efc.search = efc__search
_d.efc.search_dot = efc__search_dot
+ _d.efc.state = efc__state
_d.efc.type = efc__type
- _d.efc.u = efc__u
- _d.efc.uu = efc__uu
- _d.efc.uv = efc__uv
_d.efc.vel = efc__vel
- _d.efc.vv = efc__vv
_d.energy = energy
_d.epa_face = epa_face
_d.epa_horizon = epa_horizon
@@ -1154,13 +1146,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'efc__Jaref': d._impl.efc__Jaref.shape,
'efc__Ma': d._impl.efc__Ma.shape,
'efc__Mgrad': d._impl.efc__Mgrad.shape,
- 'efc__active': d._impl.efc__active.shape,
'efc__alpha': d._impl.efc__alpha.shape,
'efc__aref': d._impl.efc__aref.shape,
'efc__beta': d._impl.efc__beta.shape,
'efc__cholesky_L_tmp': d._impl.efc__cholesky_L_tmp.shape,
'efc__cholesky_y_tmp': d._impl.efc__cholesky_y_tmp.shape,
- 'efc__condim': d._impl.efc__condim.shape,
'efc__cost': d._impl.efc__cost.shape,
'efc__cost_candidate': d._impl.efc__cost_candidate.shape,
'efc__done': d._impl.efc__done.shape,
@@ -1195,16 +1185,13 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'efc__quad_gauss': d._impl.efc__quad_gauss.shape,
'efc__search': d._impl.efc__search.shape,
'efc__search_dot': d._impl.efc__search_dot.shape,
+ 'efc__state': d._impl.efc__state.shape,
'efc__type': d._impl.efc__type.shape,
- 'efc__u': d._impl.efc__u.shape,
- 'efc__uu': d._impl.efc__uu.shape,
- 'efc__uv': d._impl.efc__uv.shape,
'efc__vel': d._impl.efc__vel.shape,
- 'efc__vv': d._impl.efc__vv.shape,
}
jf = ffi.jax_callable_variadic_tuple(
_forward_shim,
- num_outputs=182,
+ num_outputs=177,
output_dims=output_dims,
vmap_method=None,
in_out_argnames={
@@ -1343,13 +1330,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'efc__Jaref',
'efc__Ma',
'efc__Mgrad',
- 'efc__active',
'efc__alpha',
'efc__aref',
'efc__beta',
'efc__cholesky_L_tmp',
'efc__cholesky_y_tmp',
- 'efc__condim',
'efc__cost',
'efc__cost_candidate',
'efc__done',
@@ -1384,12 +1369,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'efc__quad_gauss',
'efc__search',
'efc__search_dot',
+ 'efc__state',
'efc__type',
- 'efc__u',
- 'efc__uu',
- 'efc__uv',
'efc__vel',
- 'efc__vv',
},
)
out = jf(
@@ -1574,7 +1556,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m._impl.nlsp,
m.nmeshface,
m.nmocap,
- m.nsensordata,
+ m.nsensor,
m._impl.nsensortaxel,
m.nsite,
m.ntendon,
@@ -1604,6 +1586,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m._impl.rangefinder_sensor_adr,
m._impl.sensor_acc_adr,
m.sensor_adr,
+ m._impl.sensor_adr_to_contact_adr,
m._impl.sensor_contact_adr,
m.sensor_cutoff,
m.sensor_datatype,
@@ -1827,13 +1810,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
d._impl.efc__Jaref,
d._impl.efc__Ma,
d._impl.efc__Mgrad,
- d._impl.efc__active,
d._impl.efc__alpha,
d._impl.efc__aref,
d._impl.efc__beta,
d._impl.efc__cholesky_L_tmp,
d._impl.efc__cholesky_y_tmp,
- d._impl.efc__condim,
d._impl.efc__cost,
d._impl.efc__cost_candidate,
d._impl.efc__done,
@@ -1868,12 +1849,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
d._impl.efc__quad_gauss,
d._impl.efc__search,
d._impl.efc__search_dot,
+ d._impl.efc__state,
d._impl.efc__type,
- d._impl.efc__u,
- d._impl.efc__uu,
- d._impl.efc__uv,
d._impl.efc__vel,
- d._impl.efc__vv,
)
d = d.tree_replace({
'act': out[0],
@@ -2011,53 +1989,48 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'_impl.efc__Jaref': out[132],
'_impl.efc__Ma': out[133],
'_impl.efc__Mgrad': out[134],
- '_impl.efc__active': out[135],
- '_impl.efc__alpha': out[136],
- '_impl.efc__aref': out[137],
- '_impl.efc__beta': out[138],
- '_impl.efc__cholesky_L_tmp': out[139],
- '_impl.efc__cholesky_y_tmp': out[140],
- '_impl.efc__condim': out[141],
- '_impl.efc__cost': out[142],
- '_impl.efc__cost_candidate': out[143],
- '_impl.efc__done': out[144],
- '_impl.efc__force': out[145],
- '_impl.efc__frictionloss': out[146],
- '_impl.efc__gauss': out[147],
- '_impl.efc__grad': out[148],
- '_impl.efc__grad_dot': out[149],
- '_impl.efc__gtol': out[150],
- '_impl.efc__h': out[151],
- '_impl.efc__hi': out[152],
- '_impl.efc__hi_alpha': out[153],
- '_impl.efc__hi_next': out[154],
- '_impl.efc__hi_next_alpha': out[155],
- '_impl.efc__id': out[156],
- '_impl.efc__jv': out[157],
- '_impl.efc__lo': out[158],
- '_impl.efc__lo_alpha': out[159],
- '_impl.efc__lo_next': out[160],
- '_impl.efc__lo_next_alpha': out[161],
- '_impl.efc__ls_done': out[162],
- '_impl.efc__margin': out[163],
- '_impl.efc__mid': out[164],
- '_impl.efc__mid_alpha': out[165],
- '_impl.efc__mv': out[166],
- '_impl.efc__p0': out[167],
- '_impl.efc__pos': out[168],
- '_impl.efc__prev_Mgrad': out[169],
- '_impl.efc__prev_cost': out[170],
- '_impl.efc__prev_grad': out[171],
- '_impl.efc__quad': out[172],
- '_impl.efc__quad_gauss': out[173],
- '_impl.efc__search': out[174],
- '_impl.efc__search_dot': out[175],
- '_impl.efc__type': out[176],
- '_impl.efc__u': out[177],
- '_impl.efc__uu': out[178],
- '_impl.efc__uv': out[179],
- '_impl.efc__vel': out[180],
- '_impl.efc__vv': out[181],
+ '_impl.efc__alpha': out[135],
+ '_impl.efc__aref': out[136],
+ '_impl.efc__beta': out[137],
+ '_impl.efc__cholesky_L_tmp': out[138],
+ '_impl.efc__cholesky_y_tmp': out[139],
+ '_impl.efc__cost': out[140],
+ '_impl.efc__cost_candidate': out[141],
+ '_impl.efc__done': out[142],
+ '_impl.efc__force': out[143],
+ '_impl.efc__frictionloss': out[144],
+ '_impl.efc__gauss': out[145],
+ '_impl.efc__grad': out[146],
+ '_impl.efc__grad_dot': out[147],
+ '_impl.efc__gtol': out[148],
+ '_impl.efc__h': out[149],
+ '_impl.efc__hi': out[150],
+ '_impl.efc__hi_alpha': out[151],
+ '_impl.efc__hi_next': out[152],
+ '_impl.efc__hi_next_alpha': out[153],
+ '_impl.efc__id': out[154],
+ '_impl.efc__jv': out[155],
+ '_impl.efc__lo': out[156],
+ '_impl.efc__lo_alpha': out[157],
+ '_impl.efc__lo_next': out[158],
+ '_impl.efc__lo_next_alpha': out[159],
+ '_impl.efc__ls_done': out[160],
+ '_impl.efc__margin': out[161],
+ '_impl.efc__mid': out[162],
+ '_impl.efc__mid_alpha': out[163],
+ '_impl.efc__mv': out[164],
+ '_impl.efc__p0': out[165],
+ '_impl.efc__pos': out[166],
+ '_impl.efc__prev_Mgrad': out[167],
+ '_impl.efc__prev_cost': out[168],
+ '_impl.efc__prev_grad': out[169],
+ '_impl.efc__quad': out[170],
+ '_impl.efc__quad_gauss': out[171],
+ '_impl.efc__search': out[172],
+ '_impl.efc__search_dot': out[173],
+ '_impl.efc__state': out[174],
+ '_impl.efc__type': out[175],
+ '_impl.efc__vel': out[176],
})
return d
@@ -2278,7 +2251,7 @@ def _step_shim(
nlsp: int,
nmeshface: int,
nmocap: int,
- nsensordata: int,
+ nsensor: int,
nsensortaxel: int,
nsite: int,
ntendon: int,
@@ -2308,6 +2281,7 @@ def _step_shim(
rangefinder_sensor_adr: wp.array(dtype=int),
sensor_acc_adr: wp.array(dtype=int),
sensor_adr: wp.array(dtype=int),
+ sensor_adr_to_contact_adr: wp.array(dtype=int),
sensor_contact_adr: wp.array(dtype=int),
sensor_cutoff: wp.array(dtype=float),
sensor_datatype: wp.array(dtype=int),
@@ -2545,13 +2519,11 @@ def _step_shim(
efc__Jaref: wp.array2d(dtype=float),
efc__Ma: wp.array2d(dtype=float),
efc__Mgrad: wp.array2d(dtype=float),
- efc__active: wp.array2d(dtype=bool),
efc__alpha: wp.array(dtype=float),
efc__aref: wp.array2d(dtype=float),
efc__beta: wp.array(dtype=float),
efc__cholesky_L_tmp: wp.array3d(dtype=float),
efc__cholesky_y_tmp: wp.array2d(dtype=float),
- efc__condim: wp.array2d(dtype=int),
efc__cost: wp.array(dtype=float),
efc__cost_candidate: wp.array2d(dtype=float),
efc__done: wp.array(dtype=bool),
@@ -2586,12 +2558,9 @@ def _step_shim(
efc__quad_gauss: wp.array(dtype=wp.vec3),
efc__search: wp.array2d(dtype=float),
efc__search_dot: wp.array(dtype=float),
+ efc__state: wp.array2d(dtype=int),
efc__type: wp.array2d(dtype=int),
- efc__u: wp.array(dtype=mjwp_types.vec6),
- efc__uu: wp.array(dtype=float),
- efc__uv: wp.array(dtype=float),
efc__vel: wp.array2d(dtype=float),
- efc__vv: wp.array(dtype=float),
):
_m.stat = _s
_m.opt = _o
@@ -2778,7 +2747,7 @@ def _step_shim(
_m.nlsp = nlsp
_m.nmeshface = nmeshface
_m.nmocap = nmocap
- _m.nsensordata = nsensordata
+ _m.nsensor = nsensor
_m.nsensortaxel = nsensortaxel
_m.nsite = nsite
_m.ntendon = ntendon
@@ -2835,6 +2804,7 @@ def _step_shim(
_m.rangefinder_sensor_adr = rangefinder_sensor_adr
_m.sensor_acc_adr = sensor_acc_adr
_m.sensor_adr = sensor_adr
+ _m.sensor_adr_to_contact_adr = sensor_adr_to_contact_adr
_m.sensor_contact_adr = sensor_contact_adr
_m.sensor_cutoff = sensor_cutoff
_m.sensor_datatype = sensor_datatype
@@ -2936,13 +2906,11 @@ def _step_shim(
_d.efc.Jaref = efc__Jaref
_d.efc.Ma = efc__Ma
_d.efc.Mgrad = efc__Mgrad
- _d.efc.active = efc__active
_d.efc.alpha = efc__alpha
_d.efc.aref = efc__aref
_d.efc.beta = efc__beta
_d.efc.cholesky_L_tmp = efc__cholesky_L_tmp
_d.efc.cholesky_y_tmp = efc__cholesky_y_tmp
- _d.efc.condim = efc__condim
_d.efc.cost = efc__cost
_d.efc.cost_candidate = efc__cost_candidate
_d.efc.done = efc__done
@@ -2977,12 +2945,9 @@ def _step_shim(
_d.efc.quad_gauss = efc__quad_gauss
_d.efc.search = efc__search
_d.efc.search_dot = efc__search_dot
+ _d.efc.state = efc__state
_d.efc.type = efc__type
- _d.efc.u = efc__u
- _d.efc.uu = efc__uu
- _d.efc.uv = efc__uv
_d.efc.vel = efc__vel
- _d.efc.vv = efc__vv
_d.energy = energy
_d.epa_face = epa_face
_d.epa_horizon = epa_horizon
@@ -3244,13 +3209,11 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'efc__Jaref': d._impl.efc__Jaref.shape,
'efc__Ma': d._impl.efc__Ma.shape,
'efc__Mgrad': d._impl.efc__Mgrad.shape,
- 'efc__active': d._impl.efc__active.shape,
'efc__alpha': d._impl.efc__alpha.shape,
'efc__aref': d._impl.efc__aref.shape,
'efc__beta': d._impl.efc__beta.shape,
'efc__cholesky_L_tmp': d._impl.efc__cholesky_L_tmp.shape,
'efc__cholesky_y_tmp': d._impl.efc__cholesky_y_tmp.shape,
- 'efc__condim': d._impl.efc__condim.shape,
'efc__cost': d._impl.efc__cost.shape,
'efc__cost_candidate': d._impl.efc__cost_candidate.shape,
'efc__done': d._impl.efc__done.shape,
@@ -3285,16 +3248,13 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'efc__quad_gauss': d._impl.efc__quad_gauss.shape,
'efc__search': d._impl.efc__search.shape,
'efc__search_dot': d._impl.efc__search_dot.shape,
+ 'efc__state': d._impl.efc__state.shape,
'efc__type': d._impl.efc__type.shape,
- 'efc__u': d._impl.efc__u.shape,
- 'efc__uu': d._impl.efc__uu.shape,
- 'efc__uv': d._impl.efc__uv.shape,
'efc__vel': d._impl.efc__vel.shape,
- 'efc__vv': d._impl.efc__vv.shape,
}
jf = ffi.jax_callable_variadic_tuple(
_step_shim,
- num_outputs=194,
+ num_outputs=189,
output_dims=output_dims,
vmap_method=None,
in_out_argnames={
@@ -3445,13 +3405,11 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'efc__Jaref',
'efc__Ma',
'efc__Mgrad',
- 'efc__active',
'efc__alpha',
'efc__aref',
'efc__beta',
'efc__cholesky_L_tmp',
'efc__cholesky_y_tmp',
- 'efc__condim',
'efc__cost',
'efc__cost_candidate',
'efc__done',
@@ -3486,12 +3444,9 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'efc__quad_gauss',
'efc__search',
'efc__search_dot',
+ 'efc__state',
'efc__type',
- 'efc__u',
- 'efc__uu',
- 'efc__uv',
'efc__vel',
- 'efc__vv',
},
)
out = jf(
@@ -3677,7 +3632,7 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m._impl.nlsp,
m.nmeshface,
m.nmocap,
- m.nsensordata,
+ m.nsensor,
m._impl.nsensortaxel,
m.nsite,
m.ntendon,
@@ -3707,6 +3662,7 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m._impl.rangefinder_sensor_adr,
m._impl.sensor_acc_adr,
m.sensor_adr,
+ m._impl.sensor_adr_to_contact_adr,
m._impl.sensor_contact_adr,
m.sensor_cutoff,
m.sensor_datatype,
@@ -3943,13 +3899,11 @@ def _step_jax_impl(m: types.Model, d: types.Data):
d._impl.efc__Jaref,
d._impl.efc__Ma,
d._impl.efc__Mgrad,
- d._impl.efc__active,
d._impl.efc__alpha,
d._impl.efc__aref,
d._impl.efc__beta,
d._impl.efc__cholesky_L_tmp,
d._impl.efc__cholesky_y_tmp,
- d._impl.efc__condim,
d._impl.efc__cost,
d._impl.efc__cost_candidate,
d._impl.efc__done,
@@ -3984,12 +3938,9 @@ def _step_jax_impl(m: types.Model, d: types.Data):
d._impl.efc__quad_gauss,
d._impl.efc__search,
d._impl.efc__search_dot,
+ d._impl.efc__state,
d._impl.efc__type,
- d._impl.efc__u,
- d._impl.efc__uu,
- d._impl.efc__uv,
d._impl.efc__vel,
- d._impl.efc__vv,
)
d = d.tree_replace({
'act': out[0],
@@ -4139,53 +4090,48 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'_impl.efc__Jaref': out[144],
'_impl.efc__Ma': out[145],
'_impl.efc__Mgrad': out[146],
- '_impl.efc__active': out[147],
- '_impl.efc__alpha': out[148],
- '_impl.efc__aref': out[149],
- '_impl.efc__beta': out[150],
- '_impl.efc__cholesky_L_tmp': out[151],
- '_impl.efc__cholesky_y_tmp': out[152],
- '_impl.efc__condim': out[153],
- '_impl.efc__cost': out[154],
- '_impl.efc__cost_candidate': out[155],
- '_impl.efc__done': out[156],
- '_impl.efc__force': out[157],
- '_impl.efc__frictionloss': out[158],
- '_impl.efc__gauss': out[159],
- '_impl.efc__grad': out[160],
- '_impl.efc__grad_dot': out[161],
- '_impl.efc__gtol': out[162],
- '_impl.efc__h': out[163],
- '_impl.efc__hi': out[164],
- '_impl.efc__hi_alpha': out[165],
- '_impl.efc__hi_next': out[166],
- '_impl.efc__hi_next_alpha': out[167],
- '_impl.efc__id': out[168],
- '_impl.efc__jv': out[169],
- '_impl.efc__lo': out[170],
- '_impl.efc__lo_alpha': out[171],
- '_impl.efc__lo_next': out[172],
- '_impl.efc__lo_next_alpha': out[173],
- '_impl.efc__ls_done': out[174],
- '_impl.efc__margin': out[175],
- '_impl.efc__mid': out[176],
- '_impl.efc__mid_alpha': out[177],
- '_impl.efc__mv': out[178],
- '_impl.efc__p0': out[179],
- '_impl.efc__pos': out[180],
- '_impl.efc__prev_Mgrad': out[181],
- '_impl.efc__prev_cost': out[182],
- '_impl.efc__prev_grad': out[183],
- '_impl.efc__quad': out[184],
- '_impl.efc__quad_gauss': out[185],
- '_impl.efc__search': out[186],
- '_impl.efc__search_dot': out[187],
- '_impl.efc__type': out[188],
- '_impl.efc__u': out[189],
- '_impl.efc__uu': out[190],
- '_impl.efc__uv': out[191],
- '_impl.efc__vel': out[192],
- '_impl.efc__vv': out[193],
+ '_impl.efc__alpha': out[147],
+ '_impl.efc__aref': out[148],
+ '_impl.efc__beta': out[149],
+ '_impl.efc__cholesky_L_tmp': out[150],
+ '_impl.efc__cholesky_y_tmp': out[151],
+ '_impl.efc__cost': out[152],
+ '_impl.efc__cost_candidate': out[153],
+ '_impl.efc__done': out[154],
+ '_impl.efc__force': out[155],
+ '_impl.efc__frictionloss': out[156],
+ '_impl.efc__gauss': out[157],
+ '_impl.efc__grad': out[158],
+ '_impl.efc__grad_dot': out[159],
+ '_impl.efc__gtol': out[160],
+ '_impl.efc__h': out[161],
+ '_impl.efc__hi': out[162],
+ '_impl.efc__hi_alpha': out[163],
+ '_impl.efc__hi_next': out[164],
+ '_impl.efc__hi_next_alpha': out[165],
+ '_impl.efc__id': out[166],
+ '_impl.efc__jv': out[167],
+ '_impl.efc__lo': out[168],
+ '_impl.efc__lo_alpha': out[169],
+ '_impl.efc__lo_next': out[170],
+ '_impl.efc__lo_next_alpha': out[171],
+ '_impl.efc__ls_done': out[172],
+ '_impl.efc__margin': out[173],
+ '_impl.efc__mid': out[174],
+ '_impl.efc__mid_alpha': out[175],
+ '_impl.efc__mv': out[176],
+ '_impl.efc__p0': out[177],
+ '_impl.efc__pos': out[178],
+ '_impl.efc__prev_Mgrad': out[179],
+ '_impl.efc__prev_cost': out[180],
+ '_impl.efc__prev_grad': out[181],
+ '_impl.efc__quad': out[182],
+ '_impl.efc__quad_gauss': out[183],
+ '_impl.efc__search': out[184],
+ '_impl.efc__search_dot': out[185],
+ '_impl.efc__state': out[186],
+ '_impl.efc__type': out[187],
+ '_impl.efc__vel': out[188],
})
return d
diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py
index 387e7945..e86b691d 100644
--- a/mjx/mujoco/mjx/warp/types.py
+++ b/mjx/mujoco/mjx/warp/types.py
@@ -175,6 +175,7 @@ class ModelWarp(PyTreeNode):
qM_tiles: Tuple[TileSet, ...]
rangefinder_sensor_adr: np.ndarray
sensor_acc_adr: np.ndarray
+ sensor_adr_to_contact_adr: np.ndarray
sensor_contact_adr: np.ndarray
sensor_e_kinetic: bool
sensor_e_potential: bool
@@ -241,13 +242,11 @@ class DataWarp(PyTreeNode):
efc__Jaref: jax.Array
efc__Ma: jax.Array
efc__Mgrad: jax.Array
- efc__active: jax.Array
efc__alpha: jax.Array
efc__aref: jax.Array
efc__beta: jax.Array
efc__cholesky_L_tmp: jax.Array
efc__cholesky_y_tmp: jax.Array
- efc__condim: jax.Array
efc__cost: jax.Array
efc__cost_candidate: jax.Array
efc__done: jax.Array
@@ -282,12 +281,9 @@ class DataWarp(PyTreeNode):
efc__quad_gauss: jax.Array
efc__search: jax.Array
efc__search_dot: jax.Array
+ efc__state: jax.Array
efc__type: jax.Array
- efc__u: jax.Array
- efc__uu: jax.Array
- efc__uv: jax.Array
efc__vel: jax.Array
- efc__vv: jax.Array
energy: jax.Array
energy_vel_mul_m_skip: jax.Array
epa_face: jax.Array
@@ -390,10 +386,6 @@ DATA_NON_VMAP = {
'contact__solref',
'contact__solreffriction',
'contact__worldid',
- 'efc__u',
- 'efc__uu',
- 'efc__uv',
- 'efc__vv',
'epa_face',
'epa_horizon',
'epa_index',
@@ -483,13 +475,11 @@ _NDIM = {
'efc__Jaref': 2,
'efc__Ma': 2,
'efc__Mgrad': 2,
- 'efc__active': 2,
'efc__alpha': 1,
'efc__aref': 2,
'efc__beta': 1,
'efc__cholesky_L_tmp': 3,
'efc__cholesky_y_tmp': 2,
- 'efc__condim': 2,
'efc__cost': 1,
'efc__cost_candidate': 2,
'efc__done': 1,
@@ -524,12 +514,9 @@ _NDIM = {
'efc__quad_gauss': 2,
'efc__search': 2,
'efc__search_dot': 1,
+ 'efc__state': 2,
'efc__type': 2,
- 'efc__u': 2,
- 'efc__uu': 1,
- 'efc__uv': 1,
'efc__vel': 2,
- 'efc__vv': 1,
'energy': 2,
'energy_vel_mul_m_skip': 1,
'epa_face': 3,
@@ -932,6 +919,7 @@ _NDIM = {
'rangefinder_sensor_adr': 1,
'sensor_acc_adr': 1,
'sensor_adr': 1,
+ 'sensor_adr_to_contact_adr': 1,
'sensor_contact_adr': 1,
'sensor_cutoff': 1,
'sensor_datatype': 1,
@@ -1072,13 +1060,11 @@ _BATCH_DIM = {
'efc__Jaref': True,
'efc__Ma': True,
'efc__Mgrad': True,
- 'efc__active': True,
'efc__alpha': True,
'efc__aref': True,
'efc__beta': True,
'efc__cholesky_L_tmp': True,
'efc__cholesky_y_tmp': True,
- 'efc__condim': True,
'efc__cost': True,
'efc__cost_candidate': True,
'efc__done': True,
@@ -1113,12 +1099,9 @@ _BATCH_DIM = {
'efc__quad_gauss': True,
'efc__search': True,
'efc__search_dot': True,
+ 'efc__state': True,
'efc__type': True,
- 'efc__u': False,
- 'efc__uu': False,
- 'efc__uv': False,
'efc__vel': True,
- 'efc__vv': False,
'energy': True,
'energy_vel_mul_m_skip': True,
'epa_face': False,
@@ -1521,6 +1504,7 @@ _BATCH_DIM = {
'rangefinder_sensor_adr': False,
'sensor_acc_adr': False,
'sensor_adr': False,
+ 'sensor_adr_to_contact_adr': False,
'sensor_contact_adr': False,
'sensor_cutoff': False,
'sensor_datatype': False,