Import google-deepmind/mujoco_warp from GitHub.

PiperOrigin-RevId: 794774483
Change-Id: I2d8d623bbdfe9f28d4ab96bfc3c9128caa9fe86c
This commit is contained in:
Baruch Tabanpour
2025-08-13 16:07:18 -07:00
committed by Copybara-Service
parent eff4dda189
commit 1a7ec97b07
24 changed files with 1578 additions and 1363 deletions
+2
View File
@@ -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
@@ -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):
@@ -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"""
<mujoco>
<worldbody>
<geom type="plane" size="10 10 .001" friction=".01 .02 .03" priority="{priority1}" condim="1" margin=".002"/>
<body>
<geom type="sphere" size=".1" friction=".123 .456 .789" priority="{priority2}" condim="3" margin=".004"/>
<freejoint/>
</body>
</worldbody>
<keyframe>
<key qpos="0 0 .075 1 0 0 0"/>
</keyframe>
</mujoco>
""",
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__":
+38 -18
View File
@@ -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:
@@ -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):
</mujoco>
"""
)
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):
</mujoco>
"""
)
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"""
<mujoco>
<worldbody>
<geom pos="0 0 0" type="cylinder" size="1 .5"/>
<geom pos="1.999 0 0" type="cylinder" size="1 .5"/>
</worldbody>
</mujoco>
"""
)
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"""
<mujoco>
<worldbody>
<geom pos="0 0 2" type="box" name="box2" size="1 1 1"/>
<geom pos="0 0 4.4" euler="0 90 40" type="box" name="box3" size="1 1 1"/>
</worldbody>
</mujoco>"""
)
_, 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"""
<mujoco>
<worldbody>
<geom name="geom1" type="box" pos="0 0 1.9" size="1 1 1"/>
<geom name="geom2" type="box" pos="0 0 0" size="10 10 1"/>
</worldbody>
</mujoco>
"""
)
_, 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"""
<mujoco>
<asset>
<mesh name="smallbox"
vertex="-1 -1 -1 1 -1 -1 1 1 -1 1 1 1 1 -1 1 -1 1 -1 -1 1 1 -1 -1 1"/>
</asset>
<worldbody>
<geom pos="0 0 2" type="mesh" name="box1" mesh="smallbox"/>
<geom pos="0 1 3.99" euler="0 0 40" type="mesh" name="box2" mesh="smallbox"/>
</worldbody>
</mujoco>
"""
)
_, 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"""
<mujoco>
<worldbody>
<geom size="1 1 1" pos="0 0 2" type="box"/>
<geom size="1 1 1" pos="0 1 3.99" euler="0 0 40" type="box"/>
</worldbody>
</mujoco>
"""
)
_, ncon, _, _ = _geom_dist(m, d, 0, 1, MAX_ITERATIONS, multiccd=True)
self.assertEqual(ncon, 5)
if __name__ == "__main__":
wp.init()
@@ -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]
@@ -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]
@@ -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__":
+44
View File
@@ -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)
@@ -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()
+21 -25
View File
@@ -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),
+30 -18
View File
@@ -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)
+1 -1
View File
@@ -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)
+14 -23
View File
@@ -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]
+5 -12
View File
@@ -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,
@@ -328,6 +328,7 @@ class SensorTest(parameterized.TestCase):
</keyframe>
</mujoco>
""",
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 [
+11 -6
View File
@@ -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
File diff suppressed because it is too large Load Diff
+95 -7
View File
@@ -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."""
+24 -20
View File
@@ -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)
+14 -14
View File
@@ -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])
@@ -58,7 +58,7 @@
<body name="box_condim1" pos="4 0 0">
<freejoint align="false"/>
<geom class="box" condim="3"/>
<geom class="box" condim="1"/>
</body>
<body name="box_condim3" pos="5 0 0">
@@ -68,12 +68,12 @@
<body name="box_condim4" pos="6 0 0">
<freejoint align="false"/>
<geom class="box" condim="3"/>
<geom class="box" condim="4"/>
</body>
<body name="box_condim6" pos="7 0 0">
<freejoint align="false"/>
<geom class="box" condim="3"/>
<geom class="box" condim="6"/>
</body>
</worldbody>
+108 -162
View File
@@ -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
+6 -22
View File
@@ -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,