Import google-deepmind/mujoco_warp from GitHub.
PiperOrigin-RevId: 794774483 Change-Id: I2d8d623bbdfe9f28d4ab96bfc3c9128caa9fe86c
This commit is contained in:
committed by
Copybara-Service
parent
eff4dda189
commit
1a7ec97b07
@@ -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
@@ -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:
|
||||
|
||||
+145
-17
@@ -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]
|
||||
|
||||
+66
-23
@@ -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__":
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
+671
-948
File diff suppressed because it is too large
Load Diff
+95
-7
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user