From 1a7ec97b07633a98c6a47a713b541ed3f9836bc4 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Wed, 13 Aug 2025 16:07:18 -0700 Subject: [PATCH] Import google-deepmind/mujoco_warp from GitHub. PiperOrigin-RevId: 794774483 Change-Id: I2d8d623bbdfe9f28d4ab96bfc3c9128caa9fe86c --- .../mjx/third_party/mujoco_warp/__init__.py | 2 + .../mujoco_warp/_src/collision_convex.py | 12 +- .../mujoco_warp/_src/collision_driver_test.py | 38 +- .../mujoco_warp/_src/collision_gjk.py | 56 +- .../mujoco_warp/_src/collision_gjk_test.py | 162 +- .../mujoco_warp/_src/collision_primitive.py | 87 +- .../mujoco_warp/_src/constraint.py | 2 + .../mujoco_warp/_src/constraint_test.py | 89 +- .../third_party/mujoco_warp/_src/forward.py | 44 + .../mujoco_warp/_src/forward_test.py | 183 ++ .../mjx/third_party/mujoco_warp/_src/io.py | 46 +- .../third_party/mujoco_warp/_src/io_test.py | 48 +- .../mjx/third_party/mujoco_warp/_src/math.py | 2 +- .../mjx/third_party/mujoco_warp/_src/ray.py | 37 +- .../third_party/mujoco_warp/_src/sensor.py | 17 +- .../mujoco_warp/_src/sensor_test.py | 2 + .../third_party/mujoco_warp/_src/smooth.py | 17 +- .../third_party/mujoco_warp/_src/solver.py | 1619 +++++++---------- .../mujoco_warp/_src/solver_test.py | 102 +- .../mjx/third_party/mujoco_warp/_src/types.py | 44 +- .../third_party/mujoco_warp/_src/util_misc.py | 28 +- .../mujoco_warp/test_data/constraints.xml | 6 +- mjx/mujoco/mjx/warp/forward.py | 270 ++- mjx/mujoco/mjx/warp/types.py | 28 +- 24 files changed, 1578 insertions(+), 1363 deletions(-) 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,