Import google-deepmind/mujoco_warp from GitHub.

PiperOrigin-RevId: 792201820
Change-Id: Ib7f4a1c28c8e48f91f00878c49e2aafdc169ba9e
This commit is contained in:
Baruch Tabanpour
2025-08-07 09:53:15 -07:00
committed by Copybara-Service
parent bb4aeb44f1
commit 4533129103
27 changed files with 2483 additions and 577 deletions
+3
View File
@@ -64,7 +64,9 @@ from mujoco.mjx.third_party.mujoco_warp._src.solver import solve as solve
from mujoco.mjx.third_party.mujoco_warp._src.support import contact_force as contact_force
from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m as mul_m
from mujoco.mjx.third_party.mujoco_warp._src.support import xfrc_accumulate as xfrc_accumulate
from mujoco.mjx.third_party.mujoco_warp._src.test_util import BenchmarkSuite as BenchmarkSuite
from mujoco.mjx.third_party.mujoco_warp._src.test_util import benchmark as benchmark
from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseFilter as BroadphaseFilter
from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType as BroadphaseType
from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType as ConeType
from mujoco.mjx.third_party.mujoco_warp._src.types import Constraint as Constraint
@@ -72,6 +74,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import Contact as Contact
from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit as DisableBit
from mujoco.mjx.third_party.mujoco_warp._src.types import DynType as DynType
from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit as EnableBit
from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType as IntegratorType
from mujoco.mjx.third_party.mujoco_warp._src.types import JointType as JointType
from mujoco.mjx.third_party.mujoco_warp._src.types import Option as Option
from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType as SolverType
@@ -321,7 +321,8 @@ def ccd_kernel_builder(
x1 += hfield_prism_vertex(geom1.hfprism, i)
x1 = x1 / 6.0
dist, x1, x2 = ccd(
dist, count, witness1, witness2 = ccd(
False,
1e-6,
0.0,
gjk_iterations,
@@ -344,13 +345,13 @@ def ccd_kernel_builder(
epa_map_in[tid],
epa_horizon_in[tid],
)
count = 0
if dist < 0.0:
count = 1
points[0] = 0.5 * (x1 + x2)
normal = x1 - x2
if dist >= 0.0:
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):
# limit maximum number of contacts with height field
@@ -475,7 +475,7 @@ class CollisionTest(parameterized.TestCase):
@classmethod
def setUpClass(cls):
register_sdf_plugins(mjwarp._src.collision_sdf)
register_sdf_plugins(mjwarp)
@parameterized.parameters(_SDF_SDF.keys())
def test_sdf_collision(self, fixture):
File diff suppressed because it is too large Load Diff
@@ -96,9 +96,11 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
(
dist,
count,
x1,
x2,
) = ccd(
False,
1e-6,
1.0e30,
iterations,
@@ -123,8 +125,8 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int):
)
dist_out[0] = dist
pos_out[0] = x1
pos_out[1] = x2
pos_out[0] = x1[0]
pos_out[1] = x2[0]
vert = wp.array(shape=(iterations,), dtype=wp.vec3)
vert1 = wp.array(shape=(iterations,), dtype=wp.vec3)
@@ -253,8 +255,7 @@ class GJKTest(absltest.TestCase):
"""
)
# TODO(kbayes): use margin trick instead of EPA for penetration recovery
dist, _, _ = _geom_dist(m, d, 0, 1, 500)
dist, _, _ = _geom_dist(m, d, 0, 1, 0)
self.assertAlmostEqual(-2, dist)
def test_box_box_contact(self):
+69 -21
View File
@@ -65,11 +65,33 @@ def transform_aabb(aabb_pos: wp.vec3, aabb_size: wp.vec3, pos: wp.vec3, ori: wp.
return aabb
@wp.func
def radial_field(a: wp.vec3, x: wp.vec3, size: wp.vec3) -> wp.vec3:
field = wp.cw_div(-size, a)
field = wp.normalize(field)
field[0] *= wp.sign(x[0])
field[1] *= wp.sign(x[1])
field[2] *= wp.sign(x[2])
return field
@wp.func
def sphere(p: wp.vec3, size: wp.vec3) -> float:
return wp.length(p) - size[0]
@wp.func
def box(p: wp.vec3, size: wp.vec3) -> float:
a = wp.abs(p) - size
if a[0] >= 0 or a[1] >= 0 or a[2] >= 0:
z = wp.vec3(0.0, 0.0, 0.0)
b = wp.max(a, z)
return wp.norm_l2(b) + wp.min(wp.max(a), 0.0)
b = radial_field(a, p, size)
t = -wp.cw_div(a, wp.abs(b))
return -wp.min(t) * wp.norm_l2(b)
@wp.func
def ellipsoid(p: wp.vec3, size: wp.vec3) -> float:
scaled_p = wp.vec3(p[0] / size[0], p[1] / size[1], p[2] / size[2])
@@ -91,6 +113,24 @@ def grad_sphere(p: wp.vec3) -> wp.vec3:
wp.vec3(0.0)
@wp.func
def grad_box(p: wp.vec3, size: wp.vec3) -> wp.vec3:
a = wp.abs(p) - size
if wp.max(a) < 0:
return radial_field(a, p, size)
z = wp.vec3(0.0, 0.0, 0.0)
b = wp.max(a, z)
c = wp.norm_l2(b)
g = wp.cw_mul(wp.div(b, c), wp.cw_div(p, wp.abs(p)))
if a[0] <= 0:
g[0] = 0.0
if a[1] <= 0:
g[1] = 0.0
if a[2] <= 0:
g[2] = 0.0
return g
@wp.func
def grad_ellipsoid(p: wp.vec3, size: wp.vec3) -> wp.vec3:
a = wp.vec3(p[0] / size[0], p[1] / size[1], p[2] / size[2])
@@ -128,8 +168,12 @@ def user_sdf_grad(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3:
@wp.func
def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float:
if type == int(GeomType.SPHERE.value):
if type == int(GeomType.PLANE.value):
return p[2]
elif type == int(GeomType.SPHERE.value):
return sphere(p, attr)
elif type == int(GeomType.BOX.value):
return box(p, attr)
elif type == int(GeomType.ELLIPSOID.value):
return ellipsoid(p, attr)
elif type == int(GeomType.SDF.value):
@@ -140,8 +184,13 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float:
@wp.func
def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3:
if type == int(GeomType.SPHERE.value):
if type == int(GeomType.PLANE.value):
grad = wp.vec3(0.0, 0.0, 1.0)
return grad
elif type == int(GeomType.SPHERE.value):
return grad_sphere(p)
elif type == int(GeomType.BOX.value):
return grad_box(p, attr)
elif type == int(GeomType.ELLIPSOID.value):
return grad_ellipsoid(p, attr)
elif type == int(GeomType.SDF.value):
@@ -444,41 +493,40 @@ def _sdf_narrowphase(
g1_plugin = geom_plugin_index[g1]
g2_plugin = geom_plugin_index[g2]
g2_to_g1_rot = wp.transpose(geom2.rot) * geom1.rot
g2_to_g1_pos = wp.transpose(geom2.rot) * (geom1.pos - geom2.pos)
g1_to_g2_rot = wp.transpose(geom1.rot) * geom2.rot
g1_to_g2_pos = wp.transpose(geom1.rot) * (geom2.pos - geom1.pos)
aabb_pos = geom_aabb[g1, 0]
aabb_size = geom_aabb[g1, 1]
aabb1 = transform_aabb(aabb_pos, aabb_size, g2_to_g1_pos, g2_to_g1_rot)
identity = wp.identity(3, dtype=float)
aabb1 = transform_aabb(aabb_pos, aabb_size, wp.vec3(0.0), identity)
aabb_pos = geom_aabb[g2, 0]
aabb_size = geom_aabb[g2, 1]
aabb2 = transform_aabb(aabb_pos, aabb_size, wp.vec3(0.0), wp.mat33(1.0))
aabb2 = transform_aabb(aabb_pos, aabb_size, g1_to_g2_pos, g1_to_g2_rot)
aabb_intersection = AABB()
aabb_intersection.min = wp.max(aabb1.min, aabb2.min)
aabb_intersection.max = wp.min(aabb1.max, aabb2.max)
geom_pos2 = geom_pos[worldid, g2]
quat2 = geom_quat[worldid, g2]
geom_mat2 = math.quat_to_mat(quat2)
rot2 = math.mul(geom2.rot, math.transpose(geom_mat2))
pos2 = wp.sub(geom2.pos, math.mul(rot2, geom_pos2))
pos2 = geom2.pos
rot2 = geom2.rot
pos1 = geom1.pos
rot1 = geom1.rot
if type1 == int(GeomType.SDF.value):
geom_pos1 = geom_pos[worldid, g1]
quat1 = geom_quat[worldid, g1]
geom_mat1 = math.quat_to_mat(quat1)
rot1 = math.mul(geom1.rot, math.transpose(geom_mat1))
pos1 = wp.sub(geom1.pos, math.mul(rot1, geom_pos1))
attr1 = plugin_attr[g1_plugin]
g1_plugin_id = plugin[g1_plugin]
else:
pos1 = geom1.pos
rot1 = geom1.rot
attr1 = geom1.size
g1_plugin_id = -1
if g2_plugin != -1:
attr2 = plugin_attr[g2_plugin]
g2_plugin_id = plugin[g2_plugin]
else:
attr2 = geom2.size
g2_plugin_id = -1
for i in range(sdf_initpoints):
x_g2 = wp.vec3(
aabb_intersection.min[0] + (aabb_intersection.max[0] - aabb_intersection.min[0]) * halton(i, 2),
@@ -486,11 +534,11 @@ def _sdf_narrowphase(
aabb_intersection.min[2] + (aabb_intersection.max[2] - aabb_intersection.min[2]) * halton(i, 5),
)
x = geom2.rot * x_g2 + geom2.pos
x = geom1.rot * x_g2 + geom1.pos
x0_initial = wp.transpose(rot2) * (x - pos2)
dist, pos, n = gradient_descent(
type1, x0_initial, attr1, plugin_attr[g2_plugin], pos1, rot1, pos2, rot2, g1_plugin_id, plugin[g2_plugin], sdf_iterations
type1, x0_initial, attr1, attr2, pos1, rot1, pos2, rot2, g1_plugin_id, g2_plugin_id, sdf_iterations
)
write_contact(
@@ -65,33 +65,32 @@ class ConstraintTest(parameterized.TestCase):
xml = f"""
<mujoco>
<worldbody>
<body pos="0.0 0 0">
<freejoint/>
<body>
<geom type="sphere" size=".1" condim="{condim1}"/>
</body>
<body pos="0.05 0 0">
<freejoint/>
</body>
<body>
<geom type="sphere" size=".1" condim="{condim2}"/>
</body>
<body>
<geom type="ellipsoid" size=".1 .1 .1" condim="{condim2}"/>
</body>
</worldbody>
<keyframe>
<key qpos=".10 .11 .12 .7071 .7071 0 0" />
</keyframe>
</mujoco>
"""
_, mjd, m, d = test_util.fixture(xml=xml, cone=cone)
for arr in (
d.efc.D,
d.efc.aref,
d.efc.pos,
d.efc.margin,
):
arr.zero_()
_, mjd, m, d = test_util.fixture(xml=xml, cone=cone, keyframe=0)
# fill with nan to check whether we are not reading uninitialized values
d.efc.J.fill_(wp.nan)
for arr in (d.efc.J, d.efc.D, d.efc.aref, d.efc.pos, d.efc.margin):
arr.fill_(wp.nan)
mjwarp.make_constraint(m, d)
_assert_eq(d.ncon.numpy()[0], mjd.ncon, "ncon")
_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.aref.numpy()[0, : mjd.nefc], mjd.efc_aref, "efc_aref")
@@ -105,21 +104,12 @@ class ConstraintTest(parameterized.TestCase):
def test_constraints(self, cone):
"""Test constraints."""
for key in range(3):
mjm, mjd, m, d = test_util.fixture("constraints.xml", sparse=False, cone=cone, keyframe=key)
_, mjd, m, d = test_util.fixture("constraints.xml", sparse=False, cone=cone, keyframe=key)
for arr in (
d.efc.D,
d.efc.aref,
d.efc.pos,
d.efc.margin,
d.ne,
d.nefc,
d.nf,
d.nl,
):
arr.zero_()
d.efc.J.fill_(wp.nan)
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):
arr.fill_(wp.nan)
mjwarp.make_constraint(m, d)
@@ -140,8 +130,10 @@ class ConstraintTest(parameterized.TestCase):
for keyframe in range(-1, 1):
_, mjd, m, d = test_util.fixture("tendon/tendon_limit.xml", sparse=False, keyframe=keyframe)
for arr in (d.nefc, d.nl, d.efc.J, d.efc.D, d.efc.aref, d.efc.pos, d.efc.margin):
arr.zero_()
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):
arr.fill_(wp.nan)
mjwarp.make_constraint(m, d)
@@ -197,9 +189,15 @@ class ConstraintTest(parameterized.TestCase):
<key qpos=".1 .2 .3"/>
</keyframe>
</mujoco>
"""
""",
keyframe=0,
)
for arr in (d.nefc, d.ne, d.efc.type):
arr.fill_(-1)
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")
+1 -1
View File
@@ -1017,9 +1017,9 @@ def forward(m: Model, d: Data):
fwd_actuation(m, d)
fwd_acceleration(m, d, factorize=True)
sensor.sensor_acc(m, d)
solver.solve(m, d)
sensor.sensor_acc(m, d)
@event_scope
+72 -11
View File
@@ -27,6 +27,9 @@ from mujoco.mjx.third_party.mujoco_warp._src import types
# number of max iterations to run GJK/EPA
MJ_CCD_ITERATIONS = 12
# max number of worlds supported
MAX_WORLDS = 2**24
def _hfield_geom_pair(mjm: mujoco.MjModel) -> Tuple[int, np.array]:
geom1, geom2 = np.triu_indices(mjm.ngeom, k=1)
@@ -115,6 +118,20 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
if mjm.opt.noslip_iterations > 0:
raise NotImplementedError(f"noslip solver not implemented.")
# contact sensor
is_contact_sensor = mjm.sensor_type == types.SensorType.CONTACT
if is_contact_sensor.any():
# matching
if (
(mjm.sensor_objtype[is_contact_sensor] != types.ObjType.GEOM)
| (mjm.sensor_reftype[is_contact_sensor] != types.ObjType.GEOM)
).any():
raise NotImplementedError("Contact sensor: only geom1-geom2 matching is implemented.")
# reduction
if (mjm.sensor_intprm[is_contact_sensor, 1] != 1).any():
raise NotImplementedError(f"Contact sensor: only mindist reduction is implemented.")
# TODO(team): remove after _update_gradient for Newton uses tile operations for islands
nv_max = 60
if mjm.nv > nv_max and mjm.opt.jacobian == mujoco.mjtJacobian.mjJAC_DENSE:
@@ -366,10 +383,11 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
array._is_batched = True
if not expand_dim:
array.strides = (0,) + array.strides[1:]
array.shape = (MAX_WORLDS,) + array.shape[1:]
return array
array.strides = (0,) + array.strides
array.ndim += 1
array.shape = (1,) + array.shape
array.shape = (MAX_WORLDS,) + array.shape
return array
# rangefinder
@@ -415,6 +433,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
nwrap=mjm.nwrap,
nsensor=mjm.nsensor,
nsensordata=mjm.nsensordata,
nsensortaxel=sum(mjm.mesh_vertnum[mjm.sensor_objid[mjm.sensor_type == mujoco.mjtSensor.mjSENS_TACTILE]]),
nmeshvert=mjm.nmeshvert,
nmeshface=mjm.nmeshface,
nmeshgraph=mjm.nmeshgraph,
@@ -595,10 +614,13 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
mesh_vertadr=wp.array(mjm.mesh_vertadr, dtype=int),
mesh_vertnum=wp.array(mjm.mesh_vertnum, dtype=int),
mesh_vert=wp.array(mjm.mesh_vert, dtype=wp.vec3),
mesh_normaladr=wp.array(mjm.mesh_normaladr, dtype=int),
mesh_normal=wp.array(mjm.mesh_normal, dtype=wp.vec3),
mesh_faceadr=wp.array(mjm.mesh_faceadr, dtype=int),
mesh_face=wp.array(mjm.mesh_face, dtype=wp.vec3i),
mesh_graphadr=wp.array(mjm.mesh_graphadr, dtype=int),
mesh_graph=wp.array(mjm.mesh_graph, dtype=int),
mesh_quat=wp.array(mjm.mesh_quat, dtype=wp.quat),
mesh_polynum=wp.array(mjm.mesh_polynum, dtype=int),
mesh_polyadr=wp.array(mjm.mesh_polyadr, dtype=int),
mesh_polynormal=wp.array(mjm.mesh_polynormal, dtype=wp.vec3),
@@ -709,6 +731,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
sensor_objid=wp.array(mjm.sensor_objid, dtype=int),
sensor_reftype=wp.array(mjm.sensor_reftype, dtype=int),
sensor_refid=wp.array(mjm.sensor_refid, dtype=int),
sensor_intprm=wp.array(mjm.sensor_intprm, dtype=int),
sensor_dim=wp.array(mjm.sensor_dim, dtype=int),
sensor_adr=wp.array(mjm.sensor_adr, dtype=int),
sensor_cutoff=wp.array(mjm.sensor_cutoff, dtype=float),
@@ -776,6 +799,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
mjm.sensor_type,
[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_rne_postconstraint=np.isin(
mjm.sensor_type,
[
@@ -798,6 +822,24 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
block_dim=types.BlockDim(),
geom_pair_type_count=tuple(geom_type_pair_count),
has_sdf_geom=bool(np.any(mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF)),
taxel_vertadr=wp.array(
[
j + mjm.mesh_vertadr[mjm.sensor_objid[i]]
for i in range(mjm.nsensor)
if mjm.sensor_type[i] == mujoco.mjtSensor.mjSENS_TACTILE
for j in range(mjm.mesh_vertnum[mjm.sensor_objid[i]])
],
dtype=int,
),
taxel_sensorid=wp.array(
[
i
for i in range(mjm.nsensor)
if mjm.sensor_type[i] == mujoco.mjtSensor.mjSENS_TACTILE
for j in range(mjm.mesh_vertnum[mjm.sensor_objid[i]])
],
dtype=int,
),
)
return m
@@ -811,11 +853,12 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host).
nworld (int, optional): Number of worlds. Defaults to 1.
nconmax (int, optional): Maximum number of contacts for all worlds. Defaults to -1.
njmax (int, optional): Maximum number of constraints for all worlds. Defaults to -1.
njmax (int, optional): Maximum number of constraints per world. Defaults to -1.
Returns:
Data: The data object containing the current state and output arrays (device).
"""
# TODO(team): move to Model?
if nconmax == -1:
# TODO(team): heuristic for nconmax
@@ -823,6 +866,16 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
if njmax == -1:
# TODO(team): heuristic for njmax
njmax = 20 * 6
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 njmax < 1:
raise ValueError("njmax must be >= 1")
condim = np.concatenate((mjm.geom_condim, mjm.pair_dim))
condim_max = np.max(condim) if len(condim) > 0 else 0
@@ -833,6 +886,7 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
qM = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float)
qLD = 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)
return types.Data(
@@ -841,6 +895,7 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
njmax=njmax,
solver_niter=wp.zeros(nworld, dtype=int),
ncon=wp.zeros(1, dtype=int),
ncon_world=wp.zeros(nworld, dtype=int),
ncon_hfield=wp.zeros((nworld, _hfield_geom_pair(mjm)[0]), dtype=int), # warp only
ne=wp.zeros(nworld, dtype=int),
ne_connect=wp.zeros(nworld, dtype=int), # warp only
@@ -965,8 +1020,6 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
prev_grad=wp.zeros((nworld, mjm.nv), dtype=float),
prev_Mgrad=wp.zeros((nworld, mjm.nv), dtype=float),
beta=wp.zeros((nworld,), dtype=float),
beta_num=wp.zeros((nworld,), dtype=float),
beta_den=wp.zeros((nworld,), dtype=float),
done=wp.zeros((nworld,), dtype=bool),
# linesearch
ls_done=wp.zeros((nworld,), dtype=bool),
@@ -1051,6 +1104,10 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
sensor_rangefinder_vec=wp.zeros((nworld, nrangefinder), dtype=wp.vec3),
sensor_rangefinder_dist=wp.zeros((nworld, nrangefinder), dtype=float),
sensor_rangefinder_geomid=wp.zeros((nworld, nrangefinder), dtype=int),
sensor_contact_nmatch=wp.zeros((nworld, nsensorcontact), dtype=int),
sensor_contact_matchid=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=int),
sensor_contact_criteria=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=float),
sensor_contact_direction=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=float),
# ray
ray_bodyexclude=wp.zeros(1, dtype=int),
ray_dist=wp.zeros((nworld, 1), dtype=float),
@@ -1078,7 +1135,7 @@ def put_data(
mjd (mujoco.MjData): The data object containing current state and output arrays (host).
nworld (int, optional): The number of worlds. Defaults to 1.
nconmax (int, optional): The maximum number of contacts for all worlds. Defaults to -1.
njmax (int, optional): The maximum number of constraints for all worlds. Defaults to -1.
njmax (int, optional): The maximum number of constraints per world. Defaults to -1.
Returns:
Data: The data object containing the current state and output arrays (device).
@@ -1089,12 +1146,12 @@ def put_data(
nworld = nworld or 1
# TODO(team): better heuristic for nconmax
nconmax = nconmax or max(512, mjd.ncon * nworld)
nconmax = nconmax or max(512, 4 * mjd.ncon * nworld)
# TODO(team): better heuristic for njmax
njmax = njmax or max(5, mjd.nefc)
njmax = njmax or max(5, 4 * mjd.nefc)
if nworld < 1:
raise ValueError("nworld must be >= 1")
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")
@@ -1192,6 +1249,7 @@ def put_data(
efc_force_fill[:, :nefc] = np.tile(mjd.efc_force, (nworld, 1))
efc_margin_fill[:, :nefc] = np.tile(mjd.efc_margin, (nworld, 1))
nsensorcontact = np.sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT)
nrangefinder = sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER)
# some helper functions to simplify the data field definitions below
@@ -1225,6 +1283,7 @@ def put_data(
njmax=njmax,
solver_niter=tile(mjd.solver_niter[0]),
ncon=arr([mjd.ncon * nworld]),
ncon_world=wp.zeros(nworld, dtype=int),
ncon_hfield=wp.zeros((nworld, _hfield_geom_pair(mjm)[0]), dtype=int), # warp only
ne=wp.full(shape=(nworld), value=mjd.ne),
ne_connect=wp.full(shape=(nworld), value=ne_connect),
@@ -1346,8 +1405,6 @@ def put_data(
prev_grad=wp.empty(shape=(nworld, mjm.nv), dtype=float),
prev_Mgrad=wp.empty(shape=(nworld, mjm.nv), dtype=float),
beta=wp.empty(shape=(nworld,), dtype=float),
beta_num=wp.empty(shape=(nworld,), dtype=float),
beta_den=wp.empty(shape=(nworld,), dtype=float),
done=wp.empty(shape=(nworld,), dtype=bool),
ls_done=wp.zeros(shape=(nworld,), dtype=bool),
p0=wp.empty(shape=(nworld,), dtype=wp.vec3),
@@ -1429,6 +1486,10 @@ def put_data(
sensor_rangefinder_vec=wp.zeros((nworld, nrangefinder), dtype=wp.vec3),
sensor_rangefinder_dist=wp.zeros((nworld, nrangefinder), dtype=float),
sensor_rangefinder_geomid=wp.zeros((nworld, nrangefinder), dtype=int),
sensor_contact_nmatch=wp.zeros((nworld, nsensorcontact), dtype=int),
sensor_contact_matchid=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=int),
sensor_contact_criteria=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=float),
sensor_contact_direction=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=float),
# ray
ray_bodyexclude=wp.zeros(1, dtype=int),
ray_dist=wp.zeros((nworld, 1), dtype=float),
+33 -3
View File
@@ -23,10 +23,12 @@ import mujoco
import numpy as np
import warp as wp
from absl.testing import absltest
from absl.testing import parameterized
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
def _dims_match(test_obj, d1: Any, d2: Any, prefix: str = ""):
@@ -173,7 +175,7 @@ def _leading_dims_scale_w_nworld(test_obj, d1: Any, d2: Any, nworld1: int, nworl
test_obj.assertEqual(s1, nworld1, full_name + f" has leading dim {s1} with nworld={nworld1}. {msg}")
class IOTest(absltest.TestCase):
class IOTest(parameterized.TestCase):
def test_make_put_data(self):
"""Tests that make_data and put_data are producing the same shapes for all arrays."""
mjm, _, _, d = test_util.fixture("pendula.xml")
@@ -295,11 +297,11 @@ class IOTest(absltest.TestCase):
m1 = mjwarp.put_model(mjm)
self.assertTrue(hasattr(m1.geom_pos, "_is_batched"))
self.assertEqual(m1.geom_pos.shape[0], 1)
self.assertEqual(m1.geom_pos.shape[0], MAX_WORLDS)
self.assertEqual(m1.geom_pos.strides[0], 0)
self.assertLen(m1.geom_pos.strides, m1.geom_pos.ndim)
self.assertTrue(hasattr(m1.opt.gravity, "_is_batched"))
self.assertEqual(m1.opt.gravity.shape[0], 1)
self.assertEqual(m1.opt.gravity.shape[0], MAX_WORLDS)
self.assertEqual(m1.opt.gravity.strides[0], 0)
self.assertLen(m1.opt.gravity.strides, m1.opt.gravity.ndim)
self.assertFalse(hasattr(m1.body_parentid, "_is_batched"))
@@ -348,6 +350,34 @@ class IOTest(absltest.TestCase):
_dims_match(self, dm2, dp2)
_dims_match(self, dm3, dp3)
@parameterized.parameters(
'<contact geom1="plane"/>',
'<contact geom2="plane"/>',
'<contact site="site"/>',
'<contact reduce="maxforce"/>',
'<contact reduce="netforce"/>',
'<contact geom1="plane" geom2="sphere"/>',
)
def test_contact_sensor(self, contact_sensor):
mjm = mujoco.MjModel.from_xml_string(f"""
<mujoco>
<worldbody>
<site name="site"/>
<geom name="plane" type="plane" size="10 10 .001"/>
<body name="body">
<geom name="sphere" size=".1"/>
<joint type="slide" axis="0 0 1"/>
</body>
</worldbody>
<sensor>
{contact_sensor}
</sensor>
</mujoco>
""")
with self.assertRaises(NotImplementedError):
mjwarp.put_model(mjm)
if __name__ == "__main__":
wp.init()
+10 -10
View File
@@ -26,8 +26,8 @@ from mujoco.mjx.third_party.mujoco_warp._src.test_util import fixture
class JAXTest(parameterized.TestCase):
@parameterized.parameters("humanoid/humanoid.xml", "pendula.xml")
def test_jax(self, xml):
@parameterized.product(xml=("pendula.xml", "humanoid/humanoid.xml"), graph_conditional=(True, False))
def test_jax(self, xml, graph_conditional):
os.environ["XLA_FLAGS"] = "--xla_gpu_graph_min_graph_size=1"
# Force JAX to allocate memory on demand and deallocate when not needed (slow)
os.environ["XLA_PYTHON_CLIENT_ALLOCATOR"] = "platform"
@@ -38,7 +38,7 @@ class JAXTest(parameterized.TestCase):
self.skipTest("JAX not installed")
from jax import numpy as jp
from warp.jax_experimental.ffi import jax_callable
from mujoco.mjx.third_party.warp.jax_experimental import ffi
if jax.default_backend() != "gpu":
self.skipTest("JAX default backend is not GPU")
@@ -51,14 +51,12 @@ class JAXTest(parameterized.TestCase):
xml,
nworld=NWORLDS,
nconmax=NWORLDS * NCONTACTS,
njmax=NWORLDS * NCONTACTS * 4,
njmax=NCONTACTS * 4,
iterations=1,
ls_iterations=4,
kick=True,
)
# Disable CUDA graph conditional
m.opt.graph_conditional = False
m.opt.graph_conditional = graph_conditional
def warp_step(
qpos_in: wp.array(dtype=wp.float32, ndim=2),
@@ -82,14 +80,16 @@ class JAXTest(parameterized.TestCase):
return qpos, qvel
warp_step_fn = jax_callable(
warp_step_fn = ffi.jax_callable(
warp_step,
num_outputs=2,
output_dims={"qpos_out": (NWORLDS, mjm.nq), "qvel_out": (NWORLDS, mjm.nv)},
graph_compatible=True,
graph_mode=ffi.GraphMode.WARP,
)
jax_qpos = jp.tile(jp.array(m.qpos0.numpy()), (NWORLDS, 1))
# temp qpos0 array to get the right numpy shape
qpos0_temp = wp.array(ptr=m.qpos0.ptr, shape=(1,) + m.qpos0.shape[1:], dtype=wp.float32)
jax_qpos = jp.tile(jp.array(qpos0_temp), (NWORLDS, 1))
jax_qvel = jp.zeros((NWORLDS, m.nv))
jax_unroll_fn = jax.jit(unroll).lower(jax_qpos, jax_qvel).compile()
+445 -1
View File
@@ -21,17 +21,21 @@ from mujoco.mjx.third_party.mujoco_warp._src import math
from mujoco.mjx.third_party.mujoco_warp._src import ray
from mujoco.mjx.third_party.mujoco_warp._src import smooth
from mujoco.mjx.third_party.mujoco_warp._src import support
from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType
from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import DataType
from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import JointType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType
from mujoco.mjx.third_party.mujoco_warp._src.types import SensorType
from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType
from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
from mujoco.mjx.third_party.mujoco_warp._src.types import vec6
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
@@ -1546,6 +1550,7 @@ def _frameangacc(
@wp.kernel
def _sensor_acc(
# Model:
opt_cone: int,
body_rootid: wp.array(dtype=int),
jnt_dofadr: wp.array(dtype=int),
geom_bodyid: wp.array(dtype=int),
@@ -1555,10 +1560,15 @@ def _sensor_acc(
sensor_datatype: wp.array(dtype=int),
sensor_objtype: wp.array(dtype=int),
sensor_objid: wp.array(dtype=int),
sensor_intprm: wp.array2d(dtype=int),
sensor_dim: wp.array(dtype=int),
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),
# Data in:
njmax_in: int,
ncon_in: wp.array(dtype=int),
xpos_in: wp.array2d(dtype=wp.vec3),
xipos_in: wp.array2d(dtype=wp.vec3),
geom_xpos_in: wp.array2d(dtype=wp.vec3),
@@ -1569,8 +1579,18 @@ def _sensor_acc(
cvel_in: wp.array2d(dtype=wp.spatial_vector),
actuator_force_in: wp.array2d(dtype=float),
qfrc_actuator_in: wp.array2d(dtype=float),
contact_dist_in: wp.array(dtype=float),
contact_pos_in: wp.array(dtype=wp.vec3),
contact_frame_in: wp.array(dtype=wp.mat33),
contact_friction_in: wp.array(dtype=vec5),
contact_dim_in: wp.array(dtype=int),
contact_efc_address_in: wp.array2d(dtype=int),
efc_force_in: wp.array2d(dtype=float),
cacc_in: wp.array2d(dtype=wp.spatial_vector),
cfrc_int_in: wp.array2d(dtype=wp.spatial_vector),
sensor_contact_nmatch_in: wp.array2d(dtype=int),
sensor_contact_matchid_in: wp.array3d(dtype=int),
sensor_contact_direction_in: wp.array3d(dtype=float),
# Data out:
sensordata_out: wp.array2d(dtype=float),
):
@@ -1580,7 +1600,123 @@ def _sensor_acc(
objid = sensor_objid[sensorid]
out = sensordata_out[worldid]
if sensortype == int(SensorType.ACCELEROMETER.value):
if sensortype == int(SensorType.CONTACT.value):
dataspec = sensor_intprm[sensorid, 0]
dim = sensor_dim[sensorid]
objtype = sensor_objtype[sensorid]
# found, force, torque, dist, pos, normal, tangent
# TODO(thowell): precompute slot size
found = False
force = False
torque = False
dist = False
pos = False
normal = False
tangent = False
size = int(0)
for i in range(7):
if dataspec & (1 << i):
if i == 0:
found = True
size += 1
elif i == 1:
force = True
size += 3
elif i == 2:
torque = True
size += 3
elif i == 3:
dist = True
size += 1
elif i == 4:
pos = True
size += 3
elif i == 5:
normal = True
size += 3
elif i == 6:
tangent = True
size += 3
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
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]
# contact direction
dir = sensor_contact_direction_in[worldid, contactsensorid, i]
adr_slot = adr + i * size
if found:
out[adr_slot] = float(nmatch)
adr_slot += 1
if force or torque:
contact_forcetorque = support.contact_force_fn(
opt_cone,
njmax_in,
ncon_in,
contact_frame_in,
contact_friction_in,
contact_dim_in,
contact_efc_address_in,
efc_force_in,
worldid,
cid,
False,
)
if force:
out[adr_slot + 0] = contact_forcetorque[0]
out[adr_slot + 1] = contact_forcetorque[1]
out[adr_slot + 2] = dir * contact_forcetorque[2]
adr_slot += 3
if torque:
out[adr_slot + 0] = contact_forcetorque[3]
out[adr_slot + 1] = contact_forcetorque[4]
out[adr_slot + 2] = dir * contact_forcetorque[5]
adr_slot += 3
if dist:
out[adr_slot] = contact_dist_in[cid]
adr_slot += 1
if pos:
contact_pos = contact_pos_in[cid]
out[adr_slot + 0] = contact_pos[0]
out[adr_slot + 1] = contact_pos[1]
out[adr_slot + 2] = contact_pos[2]
adr_slot += 3
if normal:
contact_normal = contact_frame_in[cid][0]
out[adr_slot + 0] = dir * contact_normal[0]
out[adr_slot + 1] = dir * contact_normal[1]
out[adr_slot + 2] = dir * contact_normal[2]
adr_slot += 3
if tangent:
contact_tangent = contact_frame_in[cid][1]
out[adr_slot + 0] = dir * contact_tangent[0]
out[adr_slot + 1] = dir * contact_tangent[1]
out[adr_slot + 2] = dir * contact_tangent[2]
adr_slot += 3
# zero remaining slots
for i in range(nmatch, num):
for j in range(size):
out[adr + i * size + j] = 0.0
elif sensortype == int(SensorType.ACCELEROMETER.value):
vec3 = _accelerometer(
body_rootid, site_bodyid, site_xpos_in, site_xmat_in, subtree_com_in, cvel_in, cacc_in, worldid, objid
)
@@ -1726,6 +1862,210 @@ def _sensor_touch(
wp.atomic_add(sensordata_out[worldid], adr, normalforce)
@wp.kernel
def _sensor_tactile_zero(
# Model:
sensor_type: wp.array(dtype=int),
sensor_dim: wp.array(dtype=int),
sensor_adr: wp.array(dtype=int),
# Data out:
sensordata_out: wp.array2d(dtype=float),
):
worldid, sensorid = wp.tid()
if sensor_type[sensorid] != int(SensorType.TACTILE.value):
return
for i in range(sensor_dim[sensorid]):
sensordata_out[worldid, sensor_adr[sensorid] + i] = 0.0
@wp.func
def _transform_spatial(vec: wp.spatial_vector, dif: wp.vec3) -> wp.vec3:
return wp.spatial_bottom(vec) - wp.cross(dif, wp.spatial_top(vec))
@wp.kernel
def _sensor_tactile(
# Model:
body_rootid: wp.array(dtype=int),
body_weldid: wp.array(dtype=int),
geom_bodyid: wp.array(dtype=int),
mesh_vertadr: wp.array(dtype=int),
mesh_vert: wp.array(dtype=wp.vec3),
mesh_normaladr: wp.array(dtype=int),
mesh_normal: wp.array(dtype=wp.vec3),
mesh_quat: wp.array(dtype=wp.quat),
sensor_objid: wp.array(dtype=int),
sensor_refid: wp.array(dtype=int),
sensor_dim: wp.array(dtype=int),
sensor_adr: wp.array(dtype=int),
plugin: wp.array(dtype=int),
plugin_attr: wp.array(dtype=wp.vec3f),
geom_plugin_index: wp.array(dtype=int),
taxel_vertadr: wp.array(dtype=int),
taxel_sensorid: wp.array(dtype=int),
# Data in:
ncon_in: wp.array(dtype=int),
geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.array2d(dtype=wp.mat33),
subtree_com_in: wp.array2d(dtype=wp.vec3),
cvel_in: wp.array2d(dtype=wp.spatial_vector),
contact_geom_in: wp.array(dtype=wp.vec2i),
contact_worldid_in: wp.array(dtype=int),
# Data out:
sensordata_out: wp.array2d(dtype=float),
):
conid, taxelid = wp.tid()
if conid >= ncon_in[0]:
return
worldid = contact_worldid_in[conid]
# get sensor_id
sensor_id = taxel_sensorid[taxelid]
# get parent weld id
mesh_id = sensor_objid[sensor_id]
geom_id = sensor_refid[sensor_id]
parent_body = geom_bodyid[geom_id]
parent_weld = body_weldid[parent_body]
# contact geom
body1 = body_weldid[geom_bodyid[contact_geom_in[conid][0]]]
body2 = body_weldid[geom_bodyid[contact_geom_in[conid][1]]]
if body1 == parent_weld:
geom = contact_geom_in[conid][1]
elif body2 == parent_weld:
geom = contact_geom_in[conid][0]
else:
return
body = geom_bodyid[geom]
# vertex local position
vertid = taxel_vertadr[taxelid] - mesh_vertadr[mesh_id]
pos = mesh_vert[vertid + mesh_vertadr[mesh_id]]
# position in global frame
xpos = geom_xmat_in[worldid, geom_id] @ pos
xpos += geom_xpos_in[worldid, geom_id]
# position in other geom frame
tmp = xpos - geom_xpos_in[worldid, geom]
lpos = wp.transpose(geom_xmat_in[worldid, geom]) @ tmp
# compute distance
plugin_id = geom_plugin_index[geom]
depth = wp.min(sdf(int(GeomType.SDF.value), lpos, plugin_attr[plugin_id], plugin[plugin_id]), 0.0)
if depth >= 0.0:
return
# get velocity in global
vel_sensor = _transform_spatial(cvel_in[worldid, parent_weld], xpos - subtree_com_in[worldid, body_rootid[parent_weld]])
vel_other = _transform_spatial(
cvel_in[worldid, body], geom_xpos_in[worldid, geom] - subtree_com_in[worldid, body_rootid[body]]
)
vel_rel = vel_sensor - vel_other
# get contact force/torque, rotate into node frame
offset = mesh_normaladr[mesh_id] + 3 * vertid
normal = math.rot_vec_quat(mesh_normal[offset], mesh_quat[mesh_id])
tang1 = math.rot_vec_quat(mesh_normal[offset + 1], mesh_quat[mesh_id])
tang2 = math.rot_vec_quat(mesh_normal[offset + 2], mesh_quat[mesh_id])
kMaxDepth = 0.05
pressure = depth / wp.max(kMaxDepth - depth, MJ_MINVAL)
force = wp.mul(normal, pressure)
# one row of mat^T * force
forceT = wp.vec3()
forceT[0] = wp.dot(force, normal)
forceT[1] = wp.abs(wp.dot(vel_rel, tang1))
forceT[2] = wp.abs(wp.dot(vel_rel, tang2))
# add to sensor output
dim = sensor_dim[sensor_id] / 3
wp.atomic_add(sensordata_out[worldid], sensor_adr[sensor_id] + 0 * dim + vertid, forceT[0])
wp.atomic_add(sensordata_out[worldid], sensor_adr[sensor_id] + 1 * dim + vertid, forceT[1])
wp.atomic_add(sensordata_out[worldid], sensor_adr[sensor_id] + 2 * dim + vertid, forceT[2])
@wp.kernel
def _contact_match(
# Model:
sensor_objid: wp.array(dtype=int),
sensor_refid: wp.array(dtype=int),
sensor_contact_adr: wp.array(dtype=int),
# Data in:
ncon_in: wp.array(dtype=int),
contact_dist_in: wp.array(dtype=float),
contact_geom_in: wp.array(dtype=wp.vec2i),
contact_worldid_in: wp.array(dtype=int),
# Data out:
sensor_contact_nmatch_out: wp.array2d(dtype=int),
sensor_contact_matchid_out: wp.array3d(dtype=int),
sensor_contact_criteria_out: wp.array3d(dtype=float),
sensor_contact_direction_out: wp.array3d(dtype=float),
):
contactsensorid, contactid = wp.tid()
sensorid = sensor_contact_adr[contactsensorid]
if contactid >= ncon_in[0]:
return
# sensor information
objid = sensor_objid[sensorid]
refid = sensor_refid[sensorid]
# contact information
geom = contact_geom_in[contactid]
# geom-geom match
geom0geom1 = objid == geom[0] and refid == geom[1]
geom1geom0 = objid == geom[1] and refid == geom[0]
if geom0geom1 or geom1geom0:
worldid = contact_worldid_in[contactid]
contactmatchid = wp.atomic_add(sensor_contact_nmatch_out[worldid], contactsensorid, 1)
sensor_contact_matchid_out[worldid, contactsensorid, contactmatchid] = contactid
# TODO(thowell): alternative criteria
sensor_contact_criteria_out[worldid, contactsensorid, contactmatchid] = contact_dist_in[contactid]
# contact direction
if geom1geom0:
sensor_contact_direction_out[worldid, contactsensorid, contactmatchid] = -1.0
else:
sensor_contact_direction_out[worldid, contactsensorid, contactmatchid] = 1.0
return
# TODO(thowell): alternative matching
@wp.kernel
def _contact_sort(
# Data in:
sensor_contact_nmatch_in: wp.array2d(dtype=int),
sensor_contact_matchid_in: wp.array3d(dtype=int),
sensor_contact_criteria_in: wp.array3d(dtype=float),
# Data out:
sensor_contact_matchid_out: wp.array3d(dtype=int),
):
worldid, contactsensorid = wp.tid()
nmatch = sensor_contact_nmatch_in[worldid, contactsensorid]
# skip sort
if nmatch <= 1:
return
criteria_tile = wp.tile_load(sensor_contact_criteria_in[worldid, contactsensorid], shape=MJ_MAXCONPAIR)
matchid_tile = wp.tile_load(sensor_contact_matchid_in[worldid, contactsensorid], shape=MJ_MAXCONPAIR)
wp.tile_sort(criteria_tile, matchid_tile)
wp.tile_store(sensor_contact_matchid_out[worldid, contactsensorid], matchid_tile)
@event_scope
def sensor_acc(m: Model, d: Data):
"""Compute acceleration-dependent sensor values."""
@@ -1772,6 +2112,94 @@ def sensor_acc(m: Model, d: Data):
],
)
wp.launch(
_sensor_tactile_zero,
dim=(d.nworld, m.nsensordata),
inputs=[
m.sensor_type,
m.sensor_dim,
m.sensor_adr,
],
outputs=[
d.sensordata,
],
)
wp.launch(
_sensor_tactile,
dim=(d.nconmax, m.nsensortaxel),
inputs=[
m.body_rootid,
m.body_weldid,
m.geom_bodyid,
m.mesh_vertadr,
m.mesh_vert,
m.mesh_normaladr,
m.mesh_normal,
m.mesh_quat,
m.sensor_objid,
m.sensor_refid,
m.sensor_dim,
m.sensor_adr,
m.plugin,
m.plugin_attr,
m.geom_plugin_index,
m.taxel_vertadr,
m.taxel_sensorid,
d.ncon,
d.geom_xpos,
d.geom_xmat,
d.subtree_com,
d.cvel,
d.contact.geom,
d.contact.worldid,
],
outputs=[
d.sensordata,
],
)
if m.sensor_contact_adr.size:
# match criteria
d.sensor_contact_nmatch.zero_()
d.sensor_contact_matchid.zero_()
d.sensor_contact_criteria.zero_()
wp.launch(
_contact_match,
dim=(m.sensor_contact_adr.size, d.nconmax),
inputs=[
m.sensor_objid,
m.sensor_refid,
m.sensor_contact_adr,
d.ncon,
d.contact.dist,
d.contact.geom,
d.contact.worldid,
],
outputs=[
d.sensor_contact_nmatch,
d.sensor_contact_matchid,
d.sensor_contact_criteria,
d.sensor_contact_direction,
],
)
# sorting
wp.launch_tiled(
_contact_sort,
dim=(d.nworld, m.sensor_contact_adr.size),
inputs=[
d.sensor_contact_nmatch,
d.sensor_contact_matchid,
d.sensor_contact_criteria,
],
outputs=[
d.sensor_contact_matchid,
],
block_dim=m.block_dim.contact_sort,
)
if m.sensor_rne_postconstraint:
smooth.rne_postconstraint(m, d)
@@ -1779,6 +2207,7 @@ def sensor_acc(m: Model, d: Data):
_sensor_acc,
dim=(d.nworld, m.sensor_acc_adr.size),
inputs=[
m.opt.cone,
m.body_rootid,
m.jnt_dofadr,
m.geom_bodyid,
@@ -1788,9 +2217,14 @@ def sensor_acc(m: Model, d: Data):
m.sensor_datatype,
m.sensor_objtype,
m.sensor_objid,
m.sensor_intprm,
m.sensor_dim,
m.sensor_adr,
m.sensor_cutoff,
m.sensor_acc_adr,
m.sensor_contact_adr,
d.njmax,
d.ncon,
d.xpos,
d.xipos,
d.geom_xpos,
@@ -1801,8 +2235,18 @@ def sensor_acc(m: Model, d: Data):
d.cvel,
d.actuator_force,
d.qfrc_actuator,
d.contact.dist,
d.contact.pos,
d.contact.frame,
d.contact.friction,
d.contact.dim,
d.contact.efc_address,
d.efc.force,
d.cacc,
d.cfrc_int,
d.sensor_contact_nmatch,
d.sensor_contact_matchid,
d.sensor_contact_direction,
],
outputs=[d.sensordata],
)
@@ -15,6 +15,8 @@
"""Tests for sensor functions."""
import itertools
import mujoco
import numpy as np
import warp as wp
@@ -405,6 +407,66 @@ class SensorTest(parameterized.TestCase):
_assert_eq(d.energy.numpy()[0][1], mjd.energy[1], "kinetic energy")
@parameterized.parameters(
'type="sphere" size=".1"',
'type="capsule" size=".1 .1" euler="0 89 89"',
'type="box" size=".1 .1 .1" euler=".02 .05 .1"',
)
def test_contact_sensor(self, geom):
"""Test contact sensor."""
# create contact sensors
contact_sensor = ""
# 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))])
for num in [1, 2, 3, 4, 5]:
for geoms in [
'geom1="plane" geom2="geom"',
'geom1="geom" geom2="plane"',
'geom1="plane" geom2="sphere"',
'geom1="sphere" geom2="plane"',
'geom1="geom" geom2="sphere"',
'geom1="sphere" geom2="geom"',
]:
for data in datas:
data = " ".join(data)
contact_sensor += f'<contact {geoms} num="{num}" reduce="mindist" data="{data}"/>'
_MJCF = f"""
<mujoco>
<compiler angle="degree"/>
<option cone="pyramidal"/>
<worldbody>
<geom name="plane" type="plane" size="10 10 .001"/>
<body>
<geom name="geom" {geom}/>
<joint type="slide" axis="0 0 1"/>
</body>
<body>
<geom name="sphere" type="sphere" size=".1"/>
<joint type="slide" axis="0 0 1"/>
</body>
</worldbody>
<keyframe>
<key qpos=".09 1"/>
</keyframe>
<sensor>
{contact_sensor}
</sensor>
</mujoco>
"""
_, mjd, m, d = test_util.fixture(xml=_MJCF, keyframe=0)
d.sensordata.zero_()
mjwarp.forward(m, d)
sensordata = d.sensordata.numpy()[0]
_assert_eq(sensordata, mjd.sensordata, "sensordata")
self.assertTrue(sensordata.any()) # check that sensordata is not empty
if __name__ == "__main__":
wp.init()
+3
View File
@@ -1252,6 +1252,7 @@ def _cfrc_ext_contact(
body_rootid: wp.array(dtype=int),
geom_bodyid: wp.array(dtype=int),
# Data in:
njmax_in: int,
ncon_in: wp.array(dtype=int),
subtree_com_in: wp.array2d(dtype=wp.vec3),
contact_pos_in: wp.array(dtype=wp.vec3),
@@ -1282,6 +1283,7 @@ def _cfrc_ext_contact(
# contact force in world frame
force = support.contact_force_fn(
opt_cone,
njmax_in,
ncon_in,
contact_frame_in,
contact_friction_in,
@@ -1351,6 +1353,7 @@ def rne_postconstraint(m: Model, d: Data):
m.opt.cone,
m.body_rootid,
m.geom_bodyid,
d.njmax,
d.ncon,
d.subtree_com,
d.contact.pos,
+98 -121
View File
@@ -14,6 +14,7 @@
# ==============================================================================
from math import ceil
from math import sqrt
import warp as wp
@@ -1651,8 +1652,6 @@ def update_constraint_zero_qfrc_constraint(
@wp.kernel
def update_constraint_init_qfrc_constraint(
# Model:
nv: int,
# Data in:
nefc_in: wp.array(dtype=int),
efc_J_in: wp.array3d(dtype=float),
@@ -1660,48 +1659,59 @@ def update_constraint_init_qfrc_constraint(
efc_done_in: wp.array(dtype=bool),
# Data out:
qfrc_constraint_out: wp.array2d(dtype=float),
):
worldid, efcid = wp.tid()
if efcid >= nefc_in[worldid]:
return
if efc_done_in[worldid]:
return
force = efc_force_in[worldid, efcid]
for i in range(nv):
wp.atomic_add(
qfrc_constraint_out[worldid],
i,
efc_J_in[worldid, efcid, i] * force,
)
@wp.kernel
def update_constraint_gauss_cost(
# Data in:
qacc_in: wp.array2d(dtype=float),
qfrc_smooth_in: wp.array2d(dtype=float),
qacc_smooth_in: wp.array2d(dtype=float),
efc_Ma_in: wp.array2d(dtype=float),
efc_done_in: wp.array(dtype=bool),
# Data out:
efc_gauss_out: wp.array(dtype=float),
efc_cost_out: wp.array(dtype=float),
):
worldid, dofid = wp.tid()
if efc_done_in[worldid]:
return
gauss_cost = (
0.5
* (efc_Ma_in[worldid, dofid] - qfrc_smooth_in[worldid, dofid])
* (qacc_in[worldid, dofid] - qacc_smooth_in[worldid, dofid])
)
wp.atomic_add(efc_gauss_out, worldid, gauss_cost)
wp.atomic_add(efc_cost_out, worldid, gauss_cost)
sum_qfrc = float(0.0)
for efcid in range(nefc_in[worldid]):
efc_J = efc_J_in[worldid, efcid, dofid]
force = efc_force_in[worldid, efcid]
sum_qfrc += efc_J * force
qfrc_constraint_out[worldid, dofid] += sum_qfrc
@cache_kernel
def update_constraint_gauss_cost(nv: int, dofs_per_thread: int):
@nested_kernel
def kernel(
# Data in:
qacc_in: wp.array2d(dtype=float),
qfrc_smooth_in: wp.array2d(dtype=float),
qacc_smooth_in: wp.array2d(dtype=float),
efc_Ma_in: wp.array2d(dtype=float),
efc_done_in: wp.array(dtype=bool),
# Data out:
efc_gauss_out: wp.array(dtype=float),
efc_cost_out: wp.array(dtype=float),
):
worldid, dofstart = wp.tid()
if efc_done_in[worldid]:
return
gauss_cost = float(0.0)
if wp.static(dofs_per_thread >= nv):
for i in range(wp.static(min(dofs_per_thread, nv))):
gauss_cost += (efc_Ma_in[worldid, i] - qfrc_smooth_in[worldid, i]) * (qacc_in[worldid, i] - qacc_smooth_in[worldid, i])
efc_gauss_out[worldid] += 0.5 * gauss_cost
efc_cost_out[worldid] += 0.5 * gauss_cost
else:
for i in range(wp.static(dofs_per_thread)):
ii = dofstart * wp.static(dofs_per_thread) + i
if ii < nv:
gauss_cost += (efc_Ma_in[worldid, ii] - qfrc_smooth_in[worldid, ii]) * (
qacc_in[worldid, ii] - qacc_smooth_in[worldid, ii]
)
wp.atomic_add(efc_gauss_out, worldid, gauss_cost)
wp.atomic_add(efc_cost_out, worldid, gauss_cost)
return kernel
def _update_constraint(m: types.Model, d: types.Data):
@@ -1813,15 +1823,24 @@ def _update_constraint(m: types.Model, d: types.Data):
wp.launch(
update_constraint_init_qfrc_constraint,
dim=(d.nworld, d.njmax),
inputs=[m.nv, d.nefc, d.efc.J, d.efc.force, d.efc.done],
dim=(d.nworld, m.nv),
inputs=[d.nefc, d.efc.J, d.efc.force, d.efc.done],
outputs=[d.qfrc_constraint],
)
# if we are only using 1 thread, it makes sense to do more dofs and skip the atomics.
# For more than 1 thread, dofs_per_thread is lower for better load balancing.
if m.nv > 50:
dofs_per_thread = 20
else:
dofs_per_thread = 50
threads_per_efc = ceil(m.nv / dofs_per_thread)
# gauss = 0.5 * (Ma - qfrc_smooth).T @ (qacc - qacc_smooth)
wp.launch(
update_constraint_gauss_cost,
dim=(d.nworld, m.nv),
update_constraint_gauss_cost(m.nv, dofs_per_thread),
dim=(d.nworld, threads_per_efc),
inputs=[d.qacc, d.qfrc_smooth, d.qacc_smooth, d.efc.Ma, d.efc.done],
outputs=[d.efc.gauss, d.efc.cost],
)
@@ -1927,16 +1946,12 @@ def update_gradient_copy_lower_triangle(
@wp.kernel
def update_gradient_JTDAJ(
# Model:
dof_tri_row: wp.array(dtype=int),
dof_tri_col: wp.array(dtype=int),
# Data in:
nefc_in: wp.array(dtype=int),
efc_J_in: wp.array3d(dtype=float),
efc_D_in: wp.array2d(dtype=float),
efc_active_in: wp.array2d(dtype=bool),
efc_done_in: wp.array(dtype=bool),
# In:
# Data out:
efc_h_out: wp.array3d(dtype=float),
):
@@ -1947,20 +1962,26 @@ def update_gradient_JTDAJ(
nefc = nefc_in[worldid]
dofi = dof_tri_row[elementid]
dofj = dof_tri_col[elementid]
for efcid in range(nefc):
efc_D = efc_D_in[worldid, efcid]
active = efc_active_in[worldid, efcid]
if efc_D == 0.0 or not active:
continue
dofi = (int(sqrt(float(1 + 8 * elementid))) - 1) // 2
dofj = elementid - (dofi * (dofi + 1)) // 2
sum_h = float(0.0)
efc_D = efc_D_in[worldid, 0]
active = efc_active_in[worldid, 0]
efc_Ji = efc_J_in[worldid, 0, dofi]
efc_Jj = efc_J_in[worldid, 0, dofj]
for efcid in range(nefc - 1):
# TODO(team): sparse efc_J
value = efc_J_in[worldid, efcid, dofi] * efc_J_in[worldid, efcid, dofj] * efc_D
if value != 0.0:
wp.atomic_add(efc_h_out[worldid, dofi], dofj, value)
sum_h += efc_Ji * efc_Jj * efc_D * float(active)
jj = efcid + 1
efc_D = efc_D_in[worldid, jj]
active = efc_active_in[worldid, jj]
efc_Ji = efc_J_in[worldid, jj, dofi]
efc_Jj = efc_J_in[worldid, jj, dofj]
sum_h += efc_Ji * efc_Jj * efc_D * float(active)
efc_h_out[worldid, dofi, dofj] += sum_h
@wp.kernel
@@ -2184,12 +2205,12 @@ def _update_gradient(m: types.Model, d: types.Data):
outputs=[d.efc.h],
)
lower_triangle_dim = int(m.nv * (m.nv + 1) / 2)
# TODO(team): Investigate whether d.efc.h initialization can be merged into this kernel
wp.launch(
update_gradient_JTDAJ,
dim=(d.nworld, m.dof_tri_row.size),
dim=(d.nworld, lower_triangle_dim),
inputs=[
m.dof_tri_row,
m.dof_tri_col,
d.nefc,
d.efc.J,
d.efc.D,
@@ -2295,24 +2316,9 @@ def solve_prev_grad_Mgrad(
@wp.kernel
def solve_zero_beta_num_den(
# Data in:
efc_done_in: wp.array(dtype=bool),
# Data out:
efc_beta_num_out: wp.array(dtype=float),
efc_beta_den_out: wp.array(dtype=float),
):
worldid = wp.tid()
if efc_done_in[worldid]:
return
efc_beta_num_out[worldid] = 0.0
efc_beta_den_out[worldid] = 0.0
@wp.kernel
def solve_beta_num_den(
def solve_beta(
# Model:
nv: int,
# Data in:
efc_grad_in: wp.array2d(dtype=float),
efc_Mgrad_in: wp.array2d(dtype=float),
@@ -2320,30 +2326,6 @@ def solve_beta_num_den(
efc_prev_Mgrad_in: wp.array2d(dtype=float),
efc_done_in: wp.array(dtype=bool),
# Data out:
efc_beta_num_out: wp.array(dtype=float),
efc_beta_den_out: wp.array(dtype=float),
):
worldid, dofid = wp.tid()
if efc_done_in[worldid]:
return
prev_Mgrad = efc_prev_Mgrad_in[worldid][dofid]
wp.atomic_add(
efc_beta_num_out,
worldid,
efc_grad_in[worldid, dofid] * (efc_Mgrad_in[worldid, dofid] - prev_Mgrad),
)
wp.atomic_add(efc_beta_den_out, worldid, efc_prev_grad_in[worldid, dofid] * prev_Mgrad)
@wp.kernel
def solve_beta(
# Data in:
efc_beta_num_in: wp.array(dtype=float),
efc_beta_den_in: wp.array(dtype=float),
efc_done_in: wp.array(dtype=bool),
# Data out:
efc_beta_out: wp.array(dtype=float),
):
worldid = wp.tid()
@@ -2351,7 +2333,14 @@ def solve_beta(
if efc_done_in[worldid]:
return
efc_beta_out[worldid] = wp.max(0.0, efc_beta_num_in[worldid] / wp.max(types.MJ_MINVAL, efc_beta_den_in[worldid]))
beta_num = float(0.0)
beta_den = float(0.0)
for dofid in range(nv):
prev_Mgrad = efc_prev_Mgrad_in[worldid][dofid]
beta_num += efc_grad_in[worldid, dofid] * (efc_Mgrad_in[worldid, dofid] - prev_Mgrad)
beta_den += efc_prev_grad_in[worldid, dofid] * prev_Mgrad
efc_beta_out[worldid] = wp.max(0.0, beta_num / wp.max(types.MJ_MINVAL, beta_den))
@wp.kernel
@@ -2451,24 +2440,10 @@ def _solver_iteration(
# polak-ribiere
if m.opt.solver == types.SolverType.CG:
wp.launch(
solve_zero_beta_num_den,
dim=(d.nworld),
inputs=[d.efc.done],
outputs=[d.efc.beta_num, d.efc.beta_den],
)
wp.launch(
solve_beta_num_den,
dim=(d.nworld, m.nv),
inputs=[d.efc.grad, d.efc.Mgrad, d.efc.prev_grad, d.efc.prev_Mgrad, d.efc.done],
outputs=[d.efc.beta_num, d.efc.beta_den],
)
wp.launch(
solve_beta,
dim=(d.nworld,),
inputs=[d.efc.beta_num, d.efc.beta_den, d.efc.done],
inputs=[m.nv, d.efc.grad, d.efc.Mgrad, d.efc.prev_grad, d.efc.prev_Mgrad, d.efc.done],
outputs=[d.efc.beta],
)
@@ -2539,8 +2514,10 @@ def solve(m: types.Model, d: types.Data):
def _solve(m: types.Model, d: types.Data):
"""Finds forces that satisfy constraints."""
# warmstart
wp.copy(d.qacc, d.qacc_warmstart)
if not (m.opt.disableflags & types.DisableBit.WARMSTART):
wp.copy(d.qacc, d.qacc_warmstart)
else:
wp.copy(d.qacc, d.qacc_smooth)
# create context
create_context(m, d, grad=True)
@@ -68,6 +68,63 @@ class SolverTest(parameterized.TestCase):
_assert_eq(mjwarp_cost, mj_cost, name="cost")
@parameterized.parameters(ConeType.PYRAMIDAL, ConeType.ELLIPTIC)
def test_parallel_linesearch(self, cone):
"""Test that iterative and parallel linesearch leads to equivalent results."""
# TODO(team): Enable this case when elliptic/parallel linesearch is working
if cone == ConeType.ELLIPTIC:
return
_, _, m, d = test_util.fixture(
"humanoid/humanoid.xml",
cone=cone,
ls_parallel=False,
iterations=50,
ls_iterations=50,
)
# One step to obtain more non-zeros results
mjwarp.step(m, d)
# Preparing for linesearch
m.opt.iterations = 0
mjwarp.fwd_velocity(m, d)
mjwarp.fwd_acceleration(m, d, factorize=True)
solver.solve(m, d)
# Storing some initial values
d_efc_Ma = d.efc.Ma.numpy().copy()
d_efc_Jaref = d.efc.Jaref.numpy().copy()
d_qacc = d.qacc.numpy().copy()
# Launching iterative linesearch
m.opt.ls_parallel = False
solver._linesearch(m, d)
alpha_iterative = d.efc.alpha.numpy().copy()
# Launching parallel linesearch with 10 testing points
m.nlsp = 10
d.efc.Ma = wp.array2d(d_efc_Ma)
d.efc.Jaref = wp.array(d_efc_Jaref)
d.qacc = wp.array2d(d_qacc)
m.opt.ls_parallel = True
solver._linesearch(m, d)
alpha_parallel_10 = d.efc.alpha.numpy().copy()
# Launching parallel linesearch with 50 testing points
m.nlsp = 50
d.efc.Ma = wp.array2d(d_efc_Ma)
d.efc.Jaref = wp.array(d_efc_Jaref)
d.qacc = wp.array2d(d_qacc)
solver._linesearch(m, d)
alpha_parallel_50 = d.efc.alpha.numpy().copy()
# Checking that iterative and parallel linesearch lead to similar results
# and that increasing ls_iterations leads to better results
_assert_eq(alpha_iterative, alpha_parallel_50, name="linesearch alpha")
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),
+19 -4
View File
@@ -298,7 +298,9 @@ def any_different(v0: wp.vec3, v1: wp.vec3) -> wp.bool:
@wp.func
def _decode_pyramid(pyramid: wp.array(dtype=float), efc_address: int, mu: vec5, condim: int) -> wp.spatial_vector:
def _decode_pyramid(
njmax_in: int, pyramid: wp.array(dtype=float), efc_address: int, mu: vec5, condim: int
) -> wp.spatial_vector:
"""Converts pyramid representation to contact force."""
force = wp.spatial_vector()
@@ -308,8 +310,15 @@ def _decode_pyramid(pyramid: wp.array(dtype=float), efc_address: int, mu: vec5,
force[0] = float(0.0)
for i in range(condim - 1):
dir1 = pyramid[2 * i + efc_address]
dir2 = pyramid[2 * i + efc_address + 1]
adr = 2 * i + efc_address
if adr < njmax_in:
dir1 = pyramid[adr]
else:
dir1 = 0.0
if adr + 1 < njmax_in:
dir2 = pyramid[adr + 1]
else:
dir2 = 0.0
force[0] += dir1 + dir2
force[i + 1] = (dir1 - dir2) * mu[i]
@@ -321,6 +330,7 @@ def contact_force_fn(
# Model:
opt_cone: int,
# Data in:
njmax_in: int,
ncon_in: wp.array(dtype=int),
contact_frame_in: wp.array(dtype=wp.mat33),
contact_friction_in: wp.array(dtype=vec5),
@@ -340,6 +350,7 @@ def contact_force_fn(
if contact_id >= 0 and contact_id <= ncon_in[0] and efc_address >= 0:
if opt_cone == int(ConeType.PYRAMIDAL.value):
force = _decode_pyramid(
njmax_in,
efc_force_in[worldid],
efc_address,
contact_friction_in[contact_id],
@@ -347,7 +358,8 @@ def contact_force_fn(
)
else:
for i in range(condim):
force[i] = efc_force_in[worldid, contact_efc_address_in[contact_id, i]]
if contact_efc_address_in[contact_id, i] < njmax_in:
force[i] = efc_force_in[worldid, contact_efc_address_in[contact_id, i]]
if to_world_frame:
# Transform both top and bottom parts of spatial vector by the full contact frame matrix
@@ -363,6 +375,7 @@ def contact_force_kernel(
# Model:
opt_cone: int,
# Data in:
njmax_in: int,
ncon_in: wp.array(dtype=int),
contact_frame_in: wp.array(dtype=wp.mat33),
contact_friction_in: wp.array(dtype=vec5),
@@ -387,6 +400,7 @@ def contact_force_kernel(
out[tid] = contact_force_fn(
opt_cone,
njmax_in,
ncon_in,
contact_frame_in,
contact_friction_in,
@@ -421,6 +435,7 @@ def contact_force(
dim=(contact_ids.size,),
inputs=[
m.opt.cone,
d.njmax,
d.ncon,
d.contact.frame,
d.contact.friction,
+88 -4
View File
@@ -15,6 +15,8 @@
"""Utilities for testing."""
import importlib
import os
import time
from typing import Callable, Optional, Tuple
@@ -23,6 +25,7 @@ import numpy as np
import warp as wp
from etils import epath
from mujoco.mjx.third_party.mujoco_warp._src import forward
from mujoco.mjx.third_party.mujoco_warp._src import io
from mujoco.mjx.third_party.mujoco_warp._src import warp_util
from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
@@ -242,26 +245,107 @@ def benchmark(
m.actuator_ctrllimited, m.actuator_ctrlrange, i, 0.01
],
outputs=[d.ctrl]) # fmt: skip
wp.synchronize()
run_beg = time.perf_counter()
wp.capture_launch(graph)
wp.synchronize()
run_end = time.perf_counter()
run_end = time.perf_counter()
time_vec[i] = run_end - run_beg
if trace:
trace = _sum(trace, tracer.trace())
else:
trace = tracer.trace()
if measure_alloc or measure_solver_niter:
wp.synchronize()
if measure_alloc:
ncon.append(d.ncon.numpy()[0])
nefc.append(np.sum(d.nefc.numpy()))
if measure_solver_niter:
solver_niter.append(d.solver_niter.numpy())
wp.synchronize()
run_duration = np.sum(time_vec)
return jit_duration, run_duration, trace, ncon, nefc, solver_niter
class BenchmarkSuite:
"""Base suite for all model benchmarks."""
path = ""
batch_size = -1
nconmax = -1
njmax = -1
param_names = ("function",)
params = (
"jit_duration",
"solver_niter_mean",
"solver_niter_p95",
"device_memory_allocated",
"step",
"step.forward",
"step.forward.fwd_position",
"step.forward.fwd_position.kinematics",
"step.forward.fwd_position.com_pos",
"step.forward.fwd_position.camlight",
"step.forward.fwd_position.crb",
"step.forward.fwd_position.tendon_armature",
"step.forward.fwd_position.collision",
"step.forward.fwd_position.make_constraint",
"step.forward.fwd_position.transmission",
"step.forward.sensor_pos",
"step.forward.fwd_velocity",
"step.forward.fwd_velocity.com_vel",
"step.forward.fwd_velocity.passive",
"step.forward.fwd_velocity.rne",
"step.forward.fwd_velocity.tendon_bias",
"step.forward.sensor_vel",
"step.forward.fwd_actuation",
"step.forward.fwd_acceleration",
"step.forward.fwd_acceleration.xfrc_accumulate",
"step.forward.sensor_acc",
"step.forward.solve",
)
number = 1
rounds = 1
sample_time = 0
repeat = 1
def setup_cache(self):
module = importlib.import_module(self.__module__)
path = os.path.join(os.path.realpath(os.path.dirname(module.__file__)), self.path)
mjm = mujoco.MjModel.from_xml_path(path)
mjd = mujoco.MjData(mjm)
if mjm.nkey > 0:
mujoco.mj_resetDataKeyframe(mjm, mjd, 0)
# TODO(team): mj_forward call shouldn't be necessary, but it is
mujoco.mj_forward(mjm, mjd)
wp.init()
free_before = wp.get_device().free_memory
m = io.put_model(mjm)
d = io.put_data(mjm, mjd, self.batch_size, self.nconmax, self.njmax)
jit_duration, _, trace, _, _, solver_niter = benchmark(forward.step, m, d, 1000, True, False, True)
metrics = {
"jit_duration": jit_duration,
"solver_niter_mean": np.mean(solver_niter),
"solver_niter_p95": np.quantile(solver_niter, 0.95),
"device_memory_allocated": free_before - wp.get_device().free_memory,
}
def tree_flatten(d, parent_k=""):
ret = {}
steps = self.batch_size * 1000
for k, v in d.items():
k = parent_k + "." + k if parent_k else k
ret = ret | {k: 1e6 * v[0][0] / steps} | tree_flatten(v[1], k)
return ret
metrics = metrics | tree_flatten(trace)
return metrics
def track_metric(self, metrics, fn):
return metrics[fn]
+77 -52
View File
@@ -48,6 +48,7 @@ class BlockDim:
# ray
ray: int = 64
# sensor
contact_sort: int = 64
energy_vel_kinetic: int = 256
# smooth
cholesky_factorize: int = 256
@@ -59,6 +60,20 @@ class BlockDim:
mul_m_dense: int = 256
class BroadphaseType(enum.IntEnum):
"""Type of broadphase algorithm.
Attributes:
NXN: Broad phase checking all pairs
SAP_TILE: Sweep and prune broad phase using tile sort
SAP_SEGMENTED: Sweep and prune broad phase using segment sort
"""
NXN = 0
SAP_TILE = 1
SAP_SEGMENTED = 2
class BroadphaseFilter(enum.IntFlag):
"""Bitmask specifying which collision functions to run during broadphase.
@@ -133,12 +148,13 @@ class DisableBit(enum.IntFlag):
PASSIVE = mujoco.mjtDisableBit.mjDSBL_PASSIVE
GRAVITY = mujoco.mjtDisableBit.mjDSBL_GRAVITY
CLAMPCTRL = mujoco.mjtDisableBit.mjDSBL_CLAMPCTRL
WARMSTART = mujoco.mjtDisableBit.mjDSBL_WARMSTART
ACTUATION = mujoco.mjtDisableBit.mjDSBL_ACTUATION
REFSAFE = mujoco.mjtDisableBit.mjDSBL_REFSAFE
EULERDAMP = mujoco.mjtDisableBit.mjDSBL_EULERDAMP
FILTERPARENT = mujoco.mjtDisableBit.mjDSBL_FILTERPARENT
SENSOR = mujoco.mjtDisableBit.mjDSBL_SENSOR
# unsupported: MIDPHASE, WARMSTART
# unsupported: MIDPHASE
class EnableBit(enum.IntFlag):
@@ -371,6 +387,7 @@ class SensorType(enum.IntEnum):
SUBTREELINVEL: subtree linear velocity
SUBTREEANGMOM: subtree angular momentum
TOUCH: scalar contact normal forces summed over sensor zone
CONTACT: contacts which occurred during the simulation
ACCELEROMETER: accelerometer
FORCE: force
TORQUE: torque
@@ -381,6 +398,7 @@ class SensorType(enum.IntEnum):
TENDONLIMITFRC: tendon limit force
FRAMELINACC: 3D linear acceleration
FRAMEANGACC: 3D angular acceleration
TACTILE: tactile sensor
"""
MAGNETOMETER = mujoco.mjtSensor.mjSENS_MAGNETOMETER
@@ -414,6 +432,7 @@ class SensorType(enum.IntEnum):
SUBTREELINVEL = mujoco.mjtSensor.mjSENS_SUBTREELINVEL
SUBTREEANGMOM = mujoco.mjtSensor.mjSENS_SUBTREEANGMOM
TOUCH = mujoco.mjtSensor.mjSENS_TOUCH
CONTACT = mujoco.mjtSensor.mjSENS_CONTACT
ACCELEROMETER = mujoco.mjtSensor.mjSENS_ACCELEROMETER
FORCE = mujoco.mjtSensor.mjSENS_FORCE
TORQUE = mujoco.mjtSensor.mjSENS_TORQUE
@@ -424,6 +443,7 @@ class SensorType(enum.IntEnum):
TENDONLIMITFRC = mujoco.mjtSensor.mjSENS_TENDONLIMITFRC
FRAMELINACC = mujoco.mjtSensor.mjSENS_FRAMELINACC
FRAMEANGACC = mujoco.mjtSensor.mjSENS_FRAMEANGACC
TACTILE = mujoco.mjtSensor.mjSENS_TACTILE
class ObjType(enum.IntEnum):
@@ -502,20 +522,6 @@ vec10 = vec10f
vec11 = vec11f
class BroadphaseType(enum.IntEnum):
"""Type of broadphase algorithm.
Attributes:
NXN: Broad phase checking all pairs
SAP_TILE: Sweep and prune broad phase using tile sort
SAP_SEGMENTED: Sweep and prune broad phase using segment sort
"""
NXN = 0
SAP_TILE = 1
SAP_SEGMENTED = 2
@dataclasses.dataclass
class Option:
"""Physics options.
@@ -527,9 +533,9 @@ class Option:
ls_tolerance: CG/Newton linesearch tolerance
gravity: gravitational acceleration
magnetic: global magnetic flux
integrator: integration mode (mjtIntegrator)
cone: type of friction cone (mjtCone)
solver: solver algorithm (mjtSolver)
integrator: integration mode (IntegratorType)
cone: type of friction cone (ConeType)
solver: solver algorithm (SolverType)
iterations: number of main solver iterations
ls_iterations: maximum number of CG/Newton linesearch iterations
disableflags: bit flags for disabling standard features
@@ -542,8 +548,8 @@ class Option:
has_fluid: True if wind, density, or viscosity are non-zero at put_model time
density: density of medium
viscosity: viscosity of medium
broadphase: broadphase type, 0: nxn, 1: sap_tile, 2: sap_segmented
broadphase_filter: broadphase filter bitflag
broadphase: broadphase type (BroadphaseType)
broadphase_filter: broadphase filter bitflag (BroadphaseFilter)
graph_conditional: flag to use cuda graph conditional, should be False when JAX is used
sdf_initpoints: number of starting points for gradient descent
sdf_iterations: max number of iterations for gradient descent
@@ -597,7 +603,7 @@ class Constraint:
"""Constraint data.
Attributes:
type: constraint type (mjtConstraint) (nworld, njmax)
type: constraint type (ConstraintType) (nworld, njmax)
id: id of object of specific type (nworld, njmax)
J: constraint Jacobian (nworld, njmax, nv)
pos: constraint position (equality, contact) (nworld, njmax)
@@ -628,8 +634,6 @@ class Constraint:
prev_grad: previous grad (nworld, nv)
prev_Mgrad: previous Mgrad (nworld, nv)
beta: polak-ribiere beta (nworld,)
beta_num: numerator of beta (nworld,)
beta_den: denominator of beta (nworld,)
done: solver done (nworld,)
ls_done: linesearch done (nworld,)
p0: initial point (nworld, 3)
@@ -684,8 +688,6 @@ class Constraint:
prev_grad: wp.array2d(dtype=float)
prev_Mgrad: wp.array2d(dtype=float)
beta: wp.array(dtype=float)
beta_num: wp.array(dtype=float)
beta_den: wp.array(dtype=float)
done: wp.array(dtype=bool)
# linesearch
ls_done: wp.array(dtype=bool)
@@ -752,6 +754,7 @@ class Model:
nwrap: number of wrap objects in all tendon paths
nsensor: number of sensors
nsensordata: number of elements in sensor data vector
nsensortaxel: number of taxels in all tactile sensors
nmeshvert: number of vertices for all meshes
nmeshface: number of faces for all meshes
nmeshgraph: number of ints in mesh auxiliary data
@@ -771,8 +774,6 @@ class Model:
qM_mulm_i: sparse mass matrix addressing
qM_mulm_j: sparse mass matrix addressing
qM_madr_ij: sparse mass matrix addressing
qLD_update_tree: dof tree ordering for qLD updates
qLD_update_treeadr: index of each dof tree level
M_rownnz: number of non-zeros in each row of qM (nv,)
M_rowadr: index of each row in qM (nv,)
M_colind: column indices of non-zeros in qM (nM,)
@@ -801,7 +802,7 @@ class Model:
body_contype: OR over all geom contypes (nbody,)
body_conaffinity: OR over all geom conaffinities (nbody,)
body_gravcomp: antigravity force, units of body weight (nworld, nbody)
jnt_type: type of joint (mjtJoint) (njnt,)
jnt_type: type of joint (JointType) (njnt,)
jnt_qposadr: start addr in 'qpos' for joint's data (njnt,)
jnt_dofadr: start addr in 'qvel' for joint's data (njnt,)
jnt_bodyid: id of joint's body (njnt,)
@@ -830,7 +831,7 @@ class Model:
dof_solref: constraint solver reference: frictionloss (nworld, nv, NREF)
dof_tri_row: np.tril_indices (mjm.nv)[0]
dof_tri_col: np.tril_indices (mjm.nv)[1]
geom_type: geometric type (mjtGeom) (ngeom,)
geom_type: geometric type (GeomType) (ngeom,)
geom_contype: geom contact type (ngeom,)
geom_conaffinity: geom contact affinity (ngeom,)
geom_condim: contact dimensionality (1, 3, 4, 6) (ngeom,)
@@ -856,11 +857,11 @@ class Model:
hfield_ncol: number of columns in grid (nhfield,)
hfield_size: (x, y, z_top, z_bottom) (nhfield, 4)
hfield_data: elevation data (nhfielddata,)
site_type: geom type for rendering (mjtGeom) (nsite,)
site_type: geom type for rendering (GeomType) (nsite,)
site_bodyid: id of site's body (nsite,)
site_pos: local position offset rel. to body (nworld, nsite, 3)
site_quat: local orientation offset rel. to body (nworld, nsite, 4)
cam_mode: camera tracking mode (mjtCamLight) (ncam,)
cam_mode: camera tracking mode (CamLightType) (ncam,)
cam_bodyid: id of camera's body (ncam,)
cam_targetbodyid: id of targeted body; -1: none (ncam,)
cam_pos: position rel. to body frame (nworld, ncam, 3)
@@ -872,7 +873,7 @@ class Model:
cam_resolution: resolution: pixels [width, height] (ncam, 2)
cam_sensorsize: sensor size: length [width, height] (ncam, 2)
cam_intrinsic: [focal length; principal point] (ncam, 4)
light_mode: light tracking mode (mjtCamLight) (nlight,)
light_mode: light tracking mode (CamLightType) (nlight,)
light_bodyid: id of light's body (nlight,)
light_targetbodyid: id of targeted body; -1: none (nlight,)
light_pos: position rel. to body frame (nworld, nlight, 3)
@@ -883,10 +884,14 @@ class Model:
mesh_vertadr: first vertex address (nmesh,)
mesh_vertnum: number of vertices (nmesh,)
mesh_vert: vertex positions for all meshes (nmeshvert, 3)
mesh_normal: normals for all meshes (nmeshnormal, 3)
mesh_faceadr: first face address (nmesh,)
mesh_face: face indices for all meshes (nface, 3)
mesh_normaladr: first normal address (nmesh,)
mesh_normal: normals for all meshes (nmeshnormal x 3)
mesh_graphadr: graph data address; -1: no graph (nmesh,)
mesh_graph: convex graph data (nmeshgraph,)
mesh_quat: rotation applied to asset vertices (nmesh, 4)
mesh_polynum: number of polygons per mesh (nmesh,)
mesh_polyadr: first polygon address per mesh (nmesh,)
mesh_polynormal: all polygon normals (nmeshpoly, 3)
@@ -896,10 +901,10 @@ class Model:
mesh_polymapadr: first polygon address per vertex (nmeshvert,)
mesh_polymapnum: number of polygons per vertex (nmeshvert,)
mesh_polymap: vertex to polygon map (nmeshpolymap,)
eq_type: constraint type (mjtEq) (neq,)
eq_type: constraint type (EqType) (neq,)
eq_obj1id: id of object 1 (neq,)
eq_obj2id: id of object 2 (neq,)
eq_objtype: type of both objects (mjtObj) (neq,)
eq_objtype: type of both objects (ObjType) (neq,)
eq_active0: initial enable/disable constraint state (neq,)
eq_solref: constraint solver reference (nworld, neq, mjNREF)
eq_solimp: constraint solver impedance (nworld, neq, mjNIMP)
@@ -911,10 +916,10 @@ class Model:
actuator_moment_tiles_nv: tiling configuration
actuator_moment_tiles_nu: tiling configuration
actuator_affine_bias_gain: affine bias/gain present
actuator_trntype: transmission type (mjtTrn) (nu,)
actuator_dyntype: dynamics type (mjtDyn) (nu,)
actuator_gaintype: gain type (mjtGain) (nu,)
actuator_biastype: bias type (mjtBias) (nu,)
actuator_trntype: transmission type (TrnType) (nu,)
actuator_dyntype: dynamics type (DynType) (nu,)
actuator_gaintype: gain type (GainType) (nu,)
actuator_biastype: bias type (BiasType) (nu,)
actuator_trnid: transmission id: joint, tendon, site (nu, 2)
actuator_actadr: first activation address; -1: stateless (nu,)
actuator_actnum: number of activation variables (nu,)
@@ -971,7 +976,7 @@ class Model:
tendon_invweight0: inv. weight in qpos0 (nworld, ntendon)
wrap_objid: object id: geom, site, joint (nwrap,)
wrap_prm: divisor, joint coef, or site id (nwrap,)
wrap_type: wrap object type (mjtWrap) (nwrap,)
wrap_type: wrap object type (WrapType) (nwrap,)
tendon_jnt_adr: joint tendon address (<=nwrap,)
tendon_site_pair_adr: site pair tendon address (<=nwrap,)
tendon_geom_adr: geom tendon address (<=nwrap,)
@@ -982,12 +987,13 @@ class Model:
wrap_site_pair_adr: first address for site wrap pair (<=nwrap,)
wrap_geom_adr: addresses for geom tendon wrap object (<=nwrap,)
wrap_pulley_scale: pulley scaling (nwrap,)
sensor_type: sensor type (mjtSensor) (nsensor,)
sensor_datatype: numeric data type (mjtDataType) (nsensor,)
sensor_objtype: type of sensorized object (mjtObj) (nsensor,)
sensor_type: sensor type (SensorType) (nsensor,)
sensor_datatype: numeric data type (DataType) (nsensor,)
sensor_objtype: type of sensorized object (ObjType) (nsensor,)
sensor_objid: id of sensorized object (nsensor,)
sensor_reftype: type of reference frame (mjtObj) (nsensor,)
sensor_reftype: type of reference frame (ObjType) (nsensor,)
sensor_refid: id of reference frame; -1: global frame (nsensor,)
sensor_intprm: sensor parameters (nsensor, mjNSENS)
sensor_dim: number of scalar outputs (nsensor,)
sensor_adr: address in sensor array (nsensor,)
sensor_cutoff: cutoff for real and positive; 0: ignore (nsensor,)
@@ -1007,6 +1013,7 @@ 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_rne_postconstraint: evaluate rne_postconstraint
sensor_rangefinder_bodyid: bodyid for rangefinder (nrangefinder,)
plugin: globally registered plugin slot number (nplugin,)
@@ -1048,6 +1055,7 @@ class Model:
nwrap: int
nsensor: int
nsensordata: int
nsensortaxel: int
nmeshvert: int
nmeshface: int
nmeshgraph: int
@@ -1193,10 +1201,13 @@ class Model:
mesh_vertadr: wp.array(dtype=int)
mesh_vertnum: wp.array(dtype=int)
mesh_vert: wp.array(dtype=wp.vec3)
mesh_normaladr: wp.array(dtype=int)
mesh_normal: wp.array(dtype=wp.vec3)
mesh_faceadr: wp.array(dtype=int)
mesh_face: wp.array(dtype=wp.vec3i)
mesh_graphadr: wp.array(dtype=int)
mesh_graph: wp.array(dtype=int)
mesh_quat: wp.array(dtype=wp.quat)
mesh_polynum: wp.array(dtype=int)
mesh_polyadr: wp.array(dtype=int)
mesh_polynormal: wp.array(dtype=wp.vec3)
@@ -1295,6 +1306,7 @@ class Model:
sensor_objid: wp.array(dtype=int)
sensor_reftype: wp.array(dtype=int)
sensor_refid: wp.array(dtype=int)
sensor_intprm: wp.array2d(dtype=int)
sensor_dim: wp.array(dtype=int)
sensor_adr: wp.array(dtype=int)
sensor_cutoff: wp.array(dtype=float)
@@ -1311,6 +1323,7 @@ class Model:
sensor_e_kinetic: bool # warp only
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_rne_postconstraint: bool # warp only
sensor_rangefinder_bodyid: wp.array(dtype=int) # warp only
plugin: wp.array(dtype=int)
@@ -1323,6 +1336,8 @@ class Model:
block_dim: BlockDim # warp only
geom_pair_type_count: tuple[int, ...] # warp only
has_sdf_geom: bool # warp only
taxel_vertadr: wp.array(dtype=int) # warp only
taxel_sensorid: wp.array(dtype=int) # warp only
@dataclasses.dataclass
@@ -1365,18 +1380,19 @@ class Data:
Attributes:
nworld: number of worlds
nconmax: maximum number of contacts
njmax: maximum number of constraints
njmax: maximum number of constraints per world
solver_niter: number of solver iterations (nworld,)
ncon: number of detected contacts
ncon_world: number of detected contacts per world (nworld,)
ncon_hfield: number of contacts per geom pair with hfield (nworld, nhfieldgeompair)
ne: number of equality constraints
ne_connect: number of equality connect constraints
ne_weld: number of equality weld constraints
ne_jnt: number of equality joint constraints
ne_ten: number of equality tendon constraints
nf: number of friction constraints
nl: number of limit constraints
nefc: number of constraints (1,)
ne: number of equality constraints (nworld,)
ne_connect: number of equality connect constraints (nworld,)
ne_weld: number of equality weld constraints (nworld,)
ne_jnt: number of equality joint constraints (nworld,)
ne_ten: number of equality tendon constraints (nworld,)
nf: number of friction constraints (nworld,)
nl: number of limit constraints (nworld,)
nefc: number of constraints (nworld,)
nsolving: number of unconverged worlds (1,)
time: simulation time (nworld,)
energy: potential, kinetic energy (nworld, 2)
@@ -1498,6 +1514,10 @@ class Data:
sensor_rangefinder_vec: directions for rangefinder (nworld, nrangefinder, 3)
sensor_rangefinder_dist: distances for rangefinder (nworld, nrangefinder)
sensor_rangefinder_geomid: geomids for rangefinder (nworld, nrangefinder)
sensor_contact_nmatch: match count for each world-sensor (nworld, <=nsensor)
sensor_contact_matchid: id for matching contact (nworld, <=nsensor, MJ_MAXCONPAIR)
sensor_contact_criteria: critera for reduction (nworld, <=nsensor, MJ_MAXCONPAIR)
sensor_contact_direction: direction of contact (nworld, <=nsensor, MJ_MAXCONPAIR)
ray_bodyexclude: id of body to exclude from ray computation
ray_dist: ray distance to nearest geom (nworld, 1)
ray_geomid: id of geom that intersects with ray (nworld, 1)
@@ -1510,6 +1530,7 @@ class Data:
njmax: int # warp only
solver_niter: wp.array(dtype=int)
ncon: wp.array(dtype=int)
ncon_world: wp.array(dtype=int) # warp only
ncon_hfield: wp.array2d(dtype=int) # warp only
ne: wp.array(dtype=int)
ne_connect: wp.array(dtype=int) # warp only
@@ -1653,6 +1674,10 @@ class Data:
sensor_rangefinder_vec: wp.array2d(dtype=wp.vec3) # warp only
sensor_rangefinder_dist: wp.array2d(dtype=float) # warp only
sensor_rangefinder_geomid: wp.array2d(dtype=int) # warp only
sensor_contact_nmatch: wp.array2d(dtype=int) # warp only
sensor_contact_matchid: wp.array3d(dtype=int) # warp only
sensor_contact_criteria: wp.array3d(dtype=float) # warp only
sensor_contact_direction: wp.array3d(dtype=float) # warp only
# ray
ray_bodyexclude: wp.array(dtype=int) # warp only
@@ -14,6 +14,7 @@
# ==============================================================================
import functools
import inspect
from typing import Callable, Optional
import warp as wp
@@ -97,6 +98,11 @@ def event_scope(fn, name: str = ""):
global _STACK
if _STACK is None:
return fn(*args, **kwargs)
for frame_info in inspect.stack():
if frame_info.function in ("capture_while", "capture_if"):
return fn(*args, **kwargs)
# push into next level of stack
saved_stack, _STACK = _STACK, {}
beg = wp.Event(enable_timing=True)
@@ -0,0 +1,144 @@
import warp as wp
@wp.func
def Subtraction(a: float, b: float) -> float:
return wp.max(a, -b)
@wp.func
def Intersection(a: float, b: float) -> float:
return wp.max(a, b)
@wp.func
def circle(rho: float, r: float) -> float:
return rho - r
@wp.func
def smoothUnion(a: float, b: float, k: float) -> float:
h = wp.min(wp.max(0.5 + 0.5 * (b - a) / k, 0.0), 1.0)
return b * (1.0 - h) + a * h - k * h * (1.0 - h)
@wp.func
def smoothIntersection(a: float, b: float, k: float) -> float:
return Subtraction(Intersection(a, b), smoothUnion(Subtraction(a, b), Subtraction(b, a), k))
@wp.func
def extrusion(p: wp.vec3, sdf_2d: float, h: float) -> float:
w = wp.vec2()
w[0] = sdf_2d
w[1] = wp.abs(p[2]) - h
w_abs = wp.vec2()
w_abs[0] = wp.max(w[0], 0.0)
w_abs[1] = wp.max(w[1], 0.0)
return wp.min(wp.max(w[0], w[1]), 0.0) + wp.sqrt(w_abs[0] * w_abs[0] + w_abs[1] * w_abs[1])
@wp.func
def mod(x: float, y: float) -> float:
return x - y * wp.floor(x / y)
@wp.func
def distance2D(p: wp.vec3, attributes: wp.vec3) -> float:
# see https://www.shadertoy.com/view/3lG3WR
D = 2.8
N = 25.0
psi = 3.096e-5 * N * N - 6.557e-3 * N + 0.551 # pressure angle
alpha = 0.0
innerdiameter = -1.0
R = D / 2.0
rho = wp.sqrt(p[0] * p[0] + p[1] * p[1])
Pd = N / D # Diametral Pitch: teeth per unit length of diameter
P = wp.PI / Pd # Circular Pitch
a = 1.0 / Pd # Addendum: radial length of a tooth from the pitch
# circle to the tip of the tooth.
Do = D + 2.0 * a # Outside Diameter
Ro = Do / 2.0
h = 2.2 / Pd
innerR = Ro - h - 0.14 * D
if innerdiameter >= 0.0:
innerR = innerdiameter / 2.0
# Early exit
if innerR - rho > 0.0:
return innerR - rho
# Early exit
if Ro - rho < -0.2:
return rho - Ro
Db = D * wp.cos(psi) # Base Diameter
Rb = Db / 2.0
fi = wp.atan2(p[1], p[0]) + alpha
alphaStride = P / R
invAlpha = wp.acos(Rb / R)
invPhi = wp.tan(invAlpha) - invAlpha
shift = alphaStride / 2.0 - 2.0 * invPhi
fia = mod(fi + shift / 2.0, alphaStride) - shift / 2.0
fib = mod(-fi - shift + shift / 2.0, alphaStride) - shift / 2.0
dista = -1.0e6
distb = -1.0e6
if Rb < rho:
acos_rbRho = wp.acos(Rb / rho)
thetaa = fia + acos_rbRho
thetab = fib + acos_rbRho
ta = wp.sqrt(rho * rho - Rb * Rb)
# https://math.stackexchange.com/questions/1266689/distance-from-a-point-to-the-involute-of-a-circle
dista = ta - Rb * thetaa
distb = ta - Rb * thetab
gearOuter = circle(rho, Ro)
gearLowBase = circle(rho, Ro - h)
crownBase = circle(rho, innerR)
cogs = Intersection(dista, distb)
baseWalls = Intersection(fia - (alphaStride - shift), fib - (alphaStride - shift))
cogs = Intersection(baseWalls, cogs)
cogs = smoothIntersection(gearOuter, cogs, 0.0035 * D)
cogs = smoothUnion(gearLowBase, cogs, Rb - Ro + h)
cogs = Subtraction(cogs, crownBase)
return cogs
@wp.func
def gear(p: wp.vec3, attr: wp.vec3) -> float:
thickness = 0.2
return extrusion(p, distance2D(p, attr), thickness / 2.0)
@wp.func
def gear_sdf_grad(p: wp.vec3, attr: wp.vec3) -> wp.vec3:
grad = wp.vec3()
eps = 1e-6
f_original = gear(p, attr)
x_plus = wp.vec3(p[0] + eps, p[1], p[2])
f_plus = gear(x_plus, attr)
grad[0] = (f_plus - f_original) / eps
x_plus = wp.vec3(p[0], p[1] + eps, p[2])
f_plus = gear(x_plus, attr)
grad[1] = (f_plus - f_original) / eps
x_plus = wp.vec3(p[0], p[1], p[2] + eps)
f_plus = gear(x_plus, attr)
grad[2] = (f_plus - f_original) / eps
return grad
@@ -0,0 +1,66 @@
<mujoco model="touchtest">
<compiler autolimits="true"/>
<option sdf_initpoints="20"/>
<extension>
<plugin plugin="mujoco.sdf.gear">
<instance name="gear">
<config key="alpha" value="0"/>
</instance>
</plugin>
</extension>
<asset>
<texture name="grid" type="2d" builtin="checker" rgb1=".1 .2 .3" rgb2=".2 .3 .4"
width="300" height="300" mark="edge" markrgb=".2 .3 .4"/>
<material name="grid" texture="grid" texrepeat="3 1" texuniform="true"/>
<mesh name="gear">
<plugin instance="gear"/>
</mesh>
<mesh name="sphere" builtin="wedge" params="37 37 45 45 0" scale=".3 .3 .3"/>
<mesh name="box" builtin="plate" params="37 37" scale=".3 .5 .2"/>
</asset>
<visual>
<headlight ambient=".7 .7 .7" diffuse=".2 .2 .2" specular="0.1 0.1 0.1"/>
<map znear="0.01"/>
<scale contactwidth=".02" contactheight=".5"/>
</visual>
<default>
<geom friction="0.4" solimp="0 0.95 0.02"/>
</default>
<statistic center="0 0 1" extent="1" meansize=".1"/>
<worldbody>
<light pos="1 0 .3" dir="-1 0 -.3"/>
<light pos="-1 0 .3" dir="1 0 -.3"/>
<geom name="floor" pos="0 0 -0.01" type="plane" size="3 3 .01"/>
<geom type="sdf" name="gear" mesh="gear" rgba="0.4 0.4 0.4 1" pos="-1.3 0 0">
<plugin instance="gear"/>
</geom>
<body name="ball" pos="-1 1 1">
<joint name="x" type="slide" axis="1 0 0" damping="1"/>
<joint name="y" type="slide" axis="0 1 0" damping="1"/>
<joint name="z" type="slide" axis="0 0 1"/>
<joint name="rx" axis="1 0 0" springdamper="0.2 1"/>
<joint name="ry" axis="0 1 0" springdamper="0.2 1"/>
<geom type="sphere" size=".3" mass="0.1" rgba=".5 .5 .5 .3"/>
<geom name="ball" type="mesh" mesh="sphere" mass="0" contype="0" conaffinity="0" rgba=".5 .5 .5 0"/>
</body>
<body name="finger" pos="0 0 1">
<freejoint/>
<geom type="box" size=".3 .5 .2" mass="0.1" rgba=".5 .5 .5 .3"/>
<geom name="finger" type="mesh" mesh="box" mass="0" contype="0" conaffinity="0" rgba=".5 .5 .5 0"/>
</body>
</worldbody>
<sensor>
<tactile geom="ball" mesh="sphere"/>
<tactile geom="finger" mesh="box"/>
</sensor>
</mujoco>
@@ -22,6 +22,8 @@ import warp as wp
from .bolt import bolt
from .bolt import bolt_sdf_grad
from .gear import gear
from .gear import gear_sdf_grad
from .nut import nut
from .nut import nut_sdf_grad
@@ -31,21 +33,25 @@ class SDFType(enum.Enum):
NUT = "NUT"
BOLT = "BOLT"
GEAR = "GEAR"
def register_sdf_plugins(collision_sdf) -> Dict[str, int]:
def register_sdf_plugins(mjwarp) -> Dict[str, int]:
xml = """<mujoco>
<extension>
<plugin plugin="mujoco.sdf.nut"><instance name="n"/></plugin>
<plugin plugin="mujoco.sdf.bolt"><instance name="b"/></plugin>
<plugin plugin="mujoco.sdf.gear"><instance name="g"/></plugin>
</extension>
<asset>
<mesh name="nm"><plugin instance="n"/></mesh>
<mesh name="bm"><plugin instance="b"/></mesh>
<mesh name="gm"><plugin instance="g"/></mesh>
</asset>
<worldbody>
<body><geom type="sdf" name="ng" mesh="nm"><plugin instance="n"/></geom></body>
<body><geom type="sdf" name="bg" mesh="bm"><plugin instance="b"/></geom></body>
<body><geom type="sdf" name="gg" mesh="gm"><plugin instance="g"/></geom></body>
</worldbody>
</mujoco>"""
@@ -62,6 +68,8 @@ def register_sdf_plugins(collision_sdf) -> Dict[str, int]:
sdf_types[SDFType.NUT.value] = int(m.plugin[i])
elif name == "bg":
sdf_types[SDFType.BOLT.value] = int(m.plugin[i])
elif name == "gg":
sdf_types[SDFType.GEAR.value] = int(m.plugin[i])
@wp.func
def user_sdf(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float:
@@ -70,6 +78,8 @@ def register_sdf_plugins(collision_sdf) -> Dict[str, int]:
result = nut(p, attr)
elif sdf_type == wp.static(sdf_types[SDFType.BOLT.value]):
result = bolt(p, attr)
elif sdf_type == wp.static(sdf_types[SDFType.GEAR.value]):
result = gear(p, attr)
return result
@wp.func
@@ -78,9 +88,11 @@ def register_sdf_plugins(collision_sdf) -> Dict[str, int]:
return nut_sdf_grad(p, attr)
elif sdf_type == wp.static(sdf_types[SDFType.BOLT.value]):
return bolt_sdf_grad(p, attr)
elif sdf_type == wp.static(sdf_types[SDFType.GEAR.value]):
return gear_sdf_grad(p, attr)
return wp.vec3()
collision_sdf.user_sdf = user_sdf
collision_sdf.user_sdf_grad = user_sdf_grad
mjwarp._src.collision_sdf.user_sdf = user_sdf
mjwarp._src.collision_sdf.user_sdf_grad = user_sdf_grad
return sdf_types
-3
View File
@@ -42,7 +42,6 @@ _e = mjwarp.Constraint(
**{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init}
)
@ffi.format_args_for_warp
def _collision_shim(
# Model
@@ -489,8 +488,6 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
@ffi.marshal_jax_warp_callable
def collision(m: types.Model, d: types.Data):
return _collision_jax_impl(m, d)
@collision.def_vmap
@ffi.marshal_custom_vmap
def collision_vmap(unused_axis_size, is_batched, m, d):
+302 -222
View File
@@ -195,6 +195,8 @@ def _forward_shim(
mesh_faceadr: wp.array(dtype=int),
mesh_graph: wp.array(dtype=int),
mesh_graphadr: wp.array(dtype=int),
mesh_normal: wp.array(dtype=wp.vec3),
mesh_normaladr: wp.array(dtype=int),
mesh_polyadr: wp.array(dtype=int),
mesh_polymap: wp.array(dtype=int),
mesh_polymapadr: wp.array(dtype=int),
@@ -204,6 +206,7 @@ def _forward_shim(
mesh_polyvert: wp.array(dtype=int),
mesh_polyvertadr: wp.array(dtype=int),
mesh_polyvertnum: wp.array(dtype=int),
mesh_quat: wp.array(dtype=wp.quat),
mesh_vert: wp.array(dtype=wp.vec3),
mesh_vertadr: wp.array(dtype=int),
mesh_vertnum: wp.array(dtype=int),
@@ -224,6 +227,8 @@ def _forward_shim(
nlsp: int,
nmeshface: int,
nmocap: int,
nsensordata: int,
nsensortaxel: int,
nsite: int,
ntendon: int,
nu: int,
@@ -252,10 +257,13 @@ def _forward_shim(
rangefinder_sensor_adr: wp.array(dtype=int),
sensor_acc_adr: wp.array(dtype=int),
sensor_adr: wp.array(dtype=int),
sensor_contact_adr: wp.array(dtype=int),
sensor_cutoff: wp.array(dtype=float),
sensor_datatype: wp.array(dtype=int),
sensor_dim: wp.array(dtype=int),
sensor_e_kinetic: bool,
sensor_e_potential: bool,
sensor_intprm: wp.array2d(dtype=int),
sensor_limitfrc_adr: wp.array(dtype=int),
sensor_limitpos_adr: wp.array(dtype=int),
sensor_limitvel_adr: wp.array(dtype=int),
@@ -278,6 +286,8 @@ def _forward_shim(
site_size: wp.array(dtype=wp.vec3),
site_type: wp.array(dtype=int),
subtree_mass: wp.array2d(dtype=float),
taxel_sensorid: wp.array(dtype=int),
taxel_vertadr: wp.array(dtype=int),
tendon_actfrclimited: wp.array(dtype=bool),
tendon_actfrcrange: wp.array2d(dtype=wp.vec2),
tendon_adr: wp.array(dtype=int),
@@ -418,6 +428,10 @@ def _forward_shim(
sap_range: wp.array2d(dtype=int),
sap_segment_index: wp.array2d(dtype=int),
sap_sort_index: wp.array3d(dtype=int),
sensor_contact_criteria: wp.array3d(dtype=float),
sensor_contact_direction: wp.array3d(dtype=float),
sensor_contact_matchid: wp.array3d(dtype=int),
sensor_contact_nmatch: wp.array2d(dtype=int),
sensor_rangefinder_dist: wp.array2d(dtype=float),
sensor_rangefinder_geomid: wp.array2d(dtype=int),
sensor_rangefinder_pnt: wp.array2d(dtype=wp.vec3),
@@ -471,8 +485,6 @@ def _forward_shim(
efc__alpha: wp.array(dtype=float),
efc__aref: wp.array2d(dtype=float),
efc__beta: wp.array(dtype=float),
efc__beta_den: wp.array(dtype=float),
efc__beta_num: 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),
@@ -669,6 +681,8 @@ def _forward_shim(
_m.mesh_faceadr = mesh_faceadr
_m.mesh_graph = mesh_graph
_m.mesh_graphadr = mesh_graphadr
_m.mesh_normal = mesh_normal
_m.mesh_normaladr = mesh_normaladr
_m.mesh_polyadr = mesh_polyadr
_m.mesh_polymap = mesh_polymap
_m.mesh_polymapadr = mesh_polymapadr
@@ -678,6 +692,7 @@ def _forward_shim(
_m.mesh_polyvert = mesh_polyvert
_m.mesh_polyvertadr = mesh_polyvertadr
_m.mesh_polyvertnum = mesh_polyvertnum
_m.mesh_quat = mesh_quat
_m.mesh_vert = mesh_vert
_m.mesh_vertadr = mesh_vertadr
_m.mesh_vertnum = mesh_vertnum
@@ -698,6 +713,8 @@ def _forward_shim(
_m.nlsp = nlsp
_m.nmeshface = nmeshface
_m.nmocap = nmocap
_m.nsensordata = nsensordata
_m.nsensortaxel = nsensortaxel
_m.nsite = nsite
_m.ntendon = ntendon
_m.nu = nu
@@ -752,10 +769,13 @@ def _forward_shim(
_m.rangefinder_sensor_adr = rangefinder_sensor_adr
_m.sensor_acc_adr = sensor_acc_adr
_m.sensor_adr = sensor_adr
_m.sensor_contact_adr = sensor_contact_adr
_m.sensor_cutoff = sensor_cutoff
_m.sensor_datatype = sensor_datatype
_m.sensor_dim = sensor_dim
_m.sensor_e_kinetic = sensor_e_kinetic
_m.sensor_e_potential = sensor_e_potential
_m.sensor_intprm = sensor_intprm
_m.sensor_limitfrc_adr = sensor_limitfrc_adr
_m.sensor_limitpos_adr = sensor_limitpos_adr
_m.sensor_limitvel_adr = sensor_limitvel_adr
@@ -779,6 +799,8 @@ def _forward_shim(
_m.site_type = site_type
_m.stat.meaninertia = stat__meaninertia
_m.subtree_mass = subtree_mass
_m.taxel_sensorid = taxel_sensorid
_m.taxel_vertadr = taxel_vertadr
_m.tendon_actfrclimited = tendon_actfrclimited
_m.tendon_actfrcrange = tendon_actfrcrange
_m.tendon_adr = tendon_adr
@@ -850,8 +872,6 @@ def _forward_shim(
_d.efc.alpha = efc__alpha
_d.efc.aref = efc__aref
_d.efc.beta = efc__beta
_d.efc.beta_den = efc__beta_den
_d.efc.beta_num = efc__beta_num
_d.efc.cholesky_L_tmp = efc__cholesky_L_tmp
_d.efc.cholesky_y_tmp = efc__cholesky_y_tmp
_d.efc.condim = efc__condim
@@ -957,6 +977,10 @@ def _forward_shim(
_d.sap_range = sap_range
_d.sap_segment_index = sap_segment_index
_d.sap_sort_index = sap_sort_index
_d.sensor_contact_criteria = sensor_contact_criteria
_d.sensor_contact_direction = sensor_contact_direction
_d.sensor_contact_matchid = sensor_contact_matchid
_d.sensor_contact_nmatch = sensor_contact_nmatch
_d.sensor_rangefinder_dist = sensor_rangefinder_dist
_d.sensor_rangefinder_geomid = sensor_rangefinder_geomid
_d.sensor_rangefinder_pnt = sensor_rangefinder_pnt
@@ -1077,6 +1101,10 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'sap_range': d._impl.sap_range.shape,
'sap_segment_index': d._impl.sap_segment_index.shape,
'sap_sort_index': d._impl.sap_sort_index.shape,
'sensor_contact_criteria': d._impl.sensor_contact_criteria.shape,
'sensor_contact_direction': d._impl.sensor_contact_direction.shape,
'sensor_contact_matchid': d._impl.sensor_contact_matchid.shape,
'sensor_contact_nmatch': d._impl.sensor_contact_nmatch.shape,
'sensor_rangefinder_dist': d._impl.sensor_rangefinder_dist.shape,
'sensor_rangefinder_geomid': d._impl.sensor_rangefinder_geomid.shape,
'sensor_rangefinder_pnt': d._impl.sensor_rangefinder_pnt.shape,
@@ -1130,8 +1158,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'efc__alpha': d._impl.efc__alpha.shape,
'efc__aref': d._impl.efc__aref.shape,
'efc__beta': d._impl.efc__beta.shape,
'efc__beta_den': d._impl.efc__beta_den.shape,
'efc__beta_num': d._impl.efc__beta_num.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,
@@ -1178,7 +1204,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
}
jf = ffi.jax_callable_variadic_tuple(
_forward_shim,
num_outputs=180,
num_outputs=182,
output_dims=output_dims,
vmap_method=None,
graph_compatible=True,
@@ -1265,6 +1291,10 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'sap_range',
'sap_segment_index',
'sap_sort_index',
'sensor_contact_criteria',
'sensor_contact_direction',
'sensor_contact_matchid',
'sensor_contact_nmatch',
'sensor_rangefinder_dist',
'sensor_rangefinder_geomid',
'sensor_rangefinder_pnt',
@@ -1318,8 +1348,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'efc__alpha',
'efc__aref',
'efc__beta',
'efc__beta_den',
'efc__beta_num',
'efc__cholesky_L_tmp',
'efc__cholesky_y_tmp',
'efc__condim',
@@ -1515,6 +1543,8 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m.mesh_faceadr,
m.mesh_graph,
m.mesh_graphadr,
m.mesh_normal,
m.mesh_normaladr,
m._impl.mesh_polyadr,
m._impl.mesh_polymap,
m._impl.mesh_polymapadr,
@@ -1524,6 +1554,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m._impl.mesh_polyvert,
m._impl.mesh_polyvertadr,
m._impl.mesh_polyvertnum,
m.mesh_quat,
m.mesh_vert,
m.mesh_vertadr,
m.mesh_vertnum,
@@ -1544,6 +1575,8 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m._impl.nlsp,
m.nmeshface,
m.nmocap,
m.nsensordata,
m._impl.nsensortaxel,
m.nsite,
m.ntendon,
m.nu,
@@ -1572,10 +1605,13 @@ 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_contact_adr,
m.sensor_cutoff,
m.sensor_datatype,
m.sensor_dim,
m._impl.sensor_e_kinetic,
m._impl.sensor_e_potential,
m.sensor_intprm,
m._impl.sensor_limitfrc_adr,
m._impl.sensor_limitpos_adr,
m._impl.sensor_limitvel_adr,
@@ -1598,6 +1634,8 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m.site_size,
m.site_type,
m._impl.subtree_mass,
m._impl.taxel_sensorid,
m._impl.taxel_vertadr,
m.tendon_actfrclimited,
m.tendon_actfrcrange,
m.tendon_adr,
@@ -1737,6 +1775,10 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
d._impl.sap_range,
d._impl.sap_segment_index,
d._impl.sap_sort_index,
d._impl.sensor_contact_criteria,
d._impl.sensor_contact_direction,
d._impl.sensor_contact_matchid,
d._impl.sensor_contact_nmatch,
d._impl.sensor_rangefinder_dist,
d._impl.sensor_rangefinder_geomid,
d._impl.sensor_rangefinder_pnt,
@@ -1790,8 +1832,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
d._impl.efc__alpha,
d._impl.efc__aref,
d._impl.efc__beta,
d._impl.efc__beta_den,
d._impl.efc__beta_num,
d._impl.efc__cholesky_L_tmp,
d._impl.efc__cholesky_y_tmp,
d._impl.efc__condim,
@@ -1919,104 +1959,106 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'_impl.sap_range': out[79],
'_impl.sap_segment_index': out[80],
'_impl.sap_sort_index': out[81],
'_impl.sensor_rangefinder_dist': out[82],
'_impl.sensor_rangefinder_geomid': out[83],
'_impl.sensor_rangefinder_pnt': out[84],
'_impl.sensor_rangefinder_vec': out[85],
'sensordata': out[86],
'site_xmat': out[87],
'site_xpos': out[88],
'_impl.solver_niter': out[89],
'_impl.subtree_angmom': out[90],
'_impl.subtree_bodyvel': out[91],
'subtree_com': out[92],
'_impl.subtree_linvel': out[93],
'_impl.ten_J': out[94],
'_impl.ten_Jdot': out[95],
'_impl.ten_actfrc': out[96],
'_impl.ten_bias_coef': out[97],
'_impl.ten_length': out[98],
'_impl.ten_velocity': out[99],
'_impl.ten_wrapadr': out[100],
'_impl.ten_wrapnum': out[101],
'time': out[102],
'_impl.wrap_geom_xpos': out[103],
'_impl.wrap_obj': out[104],
'_impl.wrap_xpos': out[105],
'xanchor': out[106],
'xaxis': out[107],
'xfrc_applied': out[108],
'ximat': out[109],
'xipos': out[110],
'xmat': out[111],
'xpos': out[112],
'xquat': out[113],
'_impl.contact__dim': out[114],
'_impl.contact__dist': out[115],
'_impl.contact__efc_address': out[116],
'_impl.contact__frame': out[117],
'_impl.contact__friction': out[118],
'_impl.contact__geom': out[119],
'_impl.contact__includemargin': out[120],
'_impl.contact__pos': out[121],
'_impl.contact__solimp': out[122],
'_impl.contact__solref': out[123],
'_impl.contact__solreffriction': out[124],
'_impl.contact__worldid': out[125],
'_impl.efc__D': out[126],
'_impl.efc__J': out[127],
'_impl.efc__Jaref': out[128],
'_impl.efc__Ma': out[129],
'_impl.efc__Mgrad': out[130],
'_impl.efc__active': out[131],
'_impl.efc__alpha': out[132],
'_impl.efc__aref': out[133],
'_impl.efc__beta': out[134],
'_impl.efc__beta_den': out[135],
'_impl.efc__beta_num': out[136],
'_impl.efc__cholesky_L_tmp': out[137],
'_impl.efc__cholesky_y_tmp': out[138],
'_impl.efc__condim': 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__type': out[174],
'_impl.efc__u': out[175],
'_impl.efc__uu': out[176],
'_impl.efc__uv': out[177],
'_impl.efc__vel': out[178],
'_impl.efc__vv': out[179],
'_impl.sensor_contact_criteria': out[82],
'_impl.sensor_contact_direction': out[83],
'_impl.sensor_contact_matchid': out[84],
'_impl.sensor_contact_nmatch': out[85],
'_impl.sensor_rangefinder_dist': out[86],
'_impl.sensor_rangefinder_geomid': out[87],
'_impl.sensor_rangefinder_pnt': out[88],
'_impl.sensor_rangefinder_vec': out[89],
'sensordata': out[90],
'site_xmat': out[91],
'site_xpos': out[92],
'_impl.solver_niter': out[93],
'_impl.subtree_angmom': out[94],
'_impl.subtree_bodyvel': out[95],
'subtree_com': out[96],
'_impl.subtree_linvel': out[97],
'_impl.ten_J': out[98],
'_impl.ten_Jdot': out[99],
'_impl.ten_actfrc': out[100],
'_impl.ten_bias_coef': out[101],
'_impl.ten_length': out[102],
'_impl.ten_velocity': out[103],
'_impl.ten_wrapadr': out[104],
'_impl.ten_wrapnum': out[105],
'time': out[106],
'_impl.wrap_geom_xpos': out[107],
'_impl.wrap_obj': out[108],
'_impl.wrap_xpos': out[109],
'xanchor': out[110],
'xaxis': out[111],
'xfrc_applied': out[112],
'ximat': out[113],
'xipos': out[114],
'xmat': out[115],
'xpos': out[116],
'xquat': out[117],
'_impl.contact__dim': out[118],
'_impl.contact__dist': out[119],
'_impl.contact__efc_address': out[120],
'_impl.contact__frame': out[121],
'_impl.contact__friction': out[122],
'_impl.contact__geom': out[123],
'_impl.contact__includemargin': out[124],
'_impl.contact__pos': out[125],
'_impl.contact__solimp': out[126],
'_impl.contact__solref': out[127],
'_impl.contact__solreffriction': out[128],
'_impl.contact__worldid': out[129],
'_impl.efc__D': out[130],
'_impl.efc__J': out[131],
'_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],
})
return d
@@ -2025,8 +2067,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
@ffi.marshal_jax_warp_callable
def forward(m: types.Model, d: types.Data):
return _forward_jax_impl(m, d)
@forward.def_vmap
@ffi.marshal_custom_vmap
def forward_vmap(unused_axis_size, is_batched, m, d):
@@ -2207,6 +2247,8 @@ def _step_shim(
mesh_faceadr: wp.array(dtype=int),
mesh_graph: wp.array(dtype=int),
mesh_graphadr: wp.array(dtype=int),
mesh_normal: wp.array(dtype=wp.vec3),
mesh_normaladr: wp.array(dtype=int),
mesh_polyadr: wp.array(dtype=int),
mesh_polymap: wp.array(dtype=int),
mesh_polymapadr: wp.array(dtype=int),
@@ -2216,6 +2258,7 @@ def _step_shim(
mesh_polyvert: wp.array(dtype=int),
mesh_polyvertadr: wp.array(dtype=int),
mesh_polyvertnum: wp.array(dtype=int),
mesh_quat: wp.array(dtype=wp.quat),
mesh_vert: wp.array(dtype=wp.vec3),
mesh_vertadr: wp.array(dtype=int),
mesh_vertnum: wp.array(dtype=int),
@@ -2236,6 +2279,8 @@ def _step_shim(
nlsp: int,
nmeshface: int,
nmocap: int,
nsensordata: int,
nsensortaxel: int,
nsite: int,
ntendon: int,
nu: int,
@@ -2264,10 +2309,13 @@ def _step_shim(
rangefinder_sensor_adr: wp.array(dtype=int),
sensor_acc_adr: wp.array(dtype=int),
sensor_adr: wp.array(dtype=int),
sensor_contact_adr: wp.array(dtype=int),
sensor_cutoff: wp.array(dtype=float),
sensor_datatype: wp.array(dtype=int),
sensor_dim: wp.array(dtype=int),
sensor_e_kinetic: bool,
sensor_e_potential: bool,
sensor_intprm: wp.array2d(dtype=int),
sensor_limitfrc_adr: wp.array(dtype=int),
sensor_limitpos_adr: wp.array(dtype=int),
sensor_limitvel_adr: wp.array(dtype=int),
@@ -2290,6 +2338,8 @@ def _step_shim(
site_size: wp.array(dtype=wp.vec3),
site_type: wp.array(dtype=int),
subtree_mass: wp.array2d(dtype=float),
taxel_sensorid: wp.array(dtype=int),
taxel_vertadr: wp.array(dtype=int),
tendon_actfrclimited: wp.array(dtype=bool),
tendon_actfrcrange: wp.array2d(dtype=wp.vec2),
tendon_adr: wp.array(dtype=int),
@@ -2443,6 +2493,10 @@ def _step_shim(
sap_range: wp.array2d(dtype=int),
sap_segment_index: wp.array2d(dtype=int),
sap_sort_index: wp.array3d(dtype=int),
sensor_contact_criteria: wp.array3d(dtype=float),
sensor_contact_direction: wp.array3d(dtype=float),
sensor_contact_matchid: wp.array3d(dtype=int),
sensor_contact_nmatch: wp.array2d(dtype=int),
sensor_rangefinder_dist: wp.array2d(dtype=float),
sensor_rangefinder_geomid: wp.array2d(dtype=int),
sensor_rangefinder_pnt: wp.array2d(dtype=wp.vec3),
@@ -2496,8 +2550,6 @@ def _step_shim(
efc__alpha: wp.array(dtype=float),
efc__aref: wp.array2d(dtype=float),
efc__beta: wp.array(dtype=float),
efc__beta_den: wp.array(dtype=float),
efc__beta_num: 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),
@@ -2695,6 +2747,8 @@ def _step_shim(
_m.mesh_faceadr = mesh_faceadr
_m.mesh_graph = mesh_graph
_m.mesh_graphadr = mesh_graphadr
_m.mesh_normal = mesh_normal
_m.mesh_normaladr = mesh_normaladr
_m.mesh_polyadr = mesh_polyadr
_m.mesh_polymap = mesh_polymap
_m.mesh_polymapadr = mesh_polymapadr
@@ -2704,6 +2758,7 @@ def _step_shim(
_m.mesh_polyvert = mesh_polyvert
_m.mesh_polyvertadr = mesh_polyvertadr
_m.mesh_polyvertnum = mesh_polyvertnum
_m.mesh_quat = mesh_quat
_m.mesh_vert = mesh_vert
_m.mesh_vertadr = mesh_vertadr
_m.mesh_vertnum = mesh_vertnum
@@ -2724,6 +2779,8 @@ def _step_shim(
_m.nlsp = nlsp
_m.nmeshface = nmeshface
_m.nmocap = nmocap
_m.nsensordata = nsensordata
_m.nsensortaxel = nsensortaxel
_m.nsite = nsite
_m.ntendon = ntendon
_m.nu = nu
@@ -2779,10 +2836,13 @@ def _step_shim(
_m.rangefinder_sensor_adr = rangefinder_sensor_adr
_m.sensor_acc_adr = sensor_acc_adr
_m.sensor_adr = sensor_adr
_m.sensor_contact_adr = sensor_contact_adr
_m.sensor_cutoff = sensor_cutoff
_m.sensor_datatype = sensor_datatype
_m.sensor_dim = sensor_dim
_m.sensor_e_kinetic = sensor_e_kinetic
_m.sensor_e_potential = sensor_e_potential
_m.sensor_intprm = sensor_intprm
_m.sensor_limitfrc_adr = sensor_limitfrc_adr
_m.sensor_limitpos_adr = sensor_limitpos_adr
_m.sensor_limitvel_adr = sensor_limitvel_adr
@@ -2806,6 +2866,8 @@ def _step_shim(
_m.site_type = site_type
_m.stat.meaninertia = stat__meaninertia
_m.subtree_mass = subtree_mass
_m.taxel_sensorid = taxel_sensorid
_m.taxel_vertadr = taxel_vertadr
_m.tendon_actfrclimited = tendon_actfrclimited
_m.tendon_actfrcrange = tendon_actfrcrange
_m.tendon_adr = tendon_adr
@@ -2879,8 +2941,6 @@ def _step_shim(
_d.efc.alpha = efc__alpha
_d.efc.aref = efc__aref
_d.efc.beta = efc__beta
_d.efc.beta_den = efc__beta_den
_d.efc.beta_num = efc__beta_num
_d.efc.cholesky_L_tmp = efc__cholesky_L_tmp
_d.efc.cholesky_y_tmp = efc__cholesky_y_tmp
_d.efc.condim = efc__condim
@@ -2996,6 +3056,10 @@ def _step_shim(
_d.sap_range = sap_range
_d.sap_segment_index = sap_segment_index
_d.sap_sort_index = sap_sort_index
_d.sensor_contact_criteria = sensor_contact_criteria
_d.sensor_contact_direction = sensor_contact_direction
_d.sensor_contact_matchid = sensor_contact_matchid
_d.sensor_contact_nmatch = sensor_contact_nmatch
_d.sensor_rangefinder_dist = sensor_rangefinder_dist
_d.sensor_rangefinder_geomid = sensor_rangefinder_geomid
_d.sensor_rangefinder_pnt = sensor_rangefinder_pnt
@@ -3128,6 +3192,10 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'sap_range': d._impl.sap_range.shape,
'sap_segment_index': d._impl.sap_segment_index.shape,
'sap_sort_index': d._impl.sap_sort_index.shape,
'sensor_contact_criteria': d._impl.sensor_contact_criteria.shape,
'sensor_contact_direction': d._impl.sensor_contact_direction.shape,
'sensor_contact_matchid': d._impl.sensor_contact_matchid.shape,
'sensor_contact_nmatch': d._impl.sensor_contact_nmatch.shape,
'sensor_rangefinder_dist': d._impl.sensor_rangefinder_dist.shape,
'sensor_rangefinder_geomid': d._impl.sensor_rangefinder_geomid.shape,
'sensor_rangefinder_pnt': d._impl.sensor_rangefinder_pnt.shape,
@@ -3181,8 +3249,6 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'efc__alpha': d._impl.efc__alpha.shape,
'efc__aref': d._impl.efc__aref.shape,
'efc__beta': d._impl.efc__beta.shape,
'efc__beta_den': d._impl.efc__beta_den.shape,
'efc__beta_num': d._impl.efc__beta_num.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,
@@ -3229,7 +3295,7 @@ def _step_jax_impl(m: types.Model, d: types.Data):
}
jf = ffi.jax_callable_variadic_tuple(
_step_shim,
num_outputs=192,
num_outputs=194,
output_dims=output_dims,
vmap_method=None,
graph_compatible=True,
@@ -3328,6 +3394,10 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'sap_range',
'sap_segment_index',
'sap_sort_index',
'sensor_contact_criteria',
'sensor_contact_direction',
'sensor_contact_matchid',
'sensor_contact_nmatch',
'sensor_rangefinder_dist',
'sensor_rangefinder_geomid',
'sensor_rangefinder_pnt',
@@ -3381,8 +3451,6 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'efc__alpha',
'efc__aref',
'efc__beta',
'efc__beta_den',
'efc__beta_num',
'efc__cholesky_L_tmp',
'efc__cholesky_y_tmp',
'efc__condim',
@@ -3579,6 +3647,8 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m.mesh_faceadr,
m.mesh_graph,
m.mesh_graphadr,
m.mesh_normal,
m.mesh_normaladr,
m._impl.mesh_polyadr,
m._impl.mesh_polymap,
m._impl.mesh_polymapadr,
@@ -3588,6 +3658,7 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m._impl.mesh_polyvert,
m._impl.mesh_polyvertadr,
m._impl.mesh_polyvertnum,
m.mesh_quat,
m.mesh_vert,
m.mesh_vertadr,
m.mesh_vertnum,
@@ -3608,6 +3679,8 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m._impl.nlsp,
m.nmeshface,
m.nmocap,
m.nsensordata,
m._impl.nsensortaxel,
m.nsite,
m.ntendon,
m.nu,
@@ -3636,10 +3709,13 @@ 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_contact_adr,
m.sensor_cutoff,
m.sensor_datatype,
m.sensor_dim,
m._impl.sensor_e_kinetic,
m._impl.sensor_e_potential,
m.sensor_intprm,
m._impl.sensor_limitfrc_adr,
m._impl.sensor_limitpos_adr,
m._impl.sensor_limitvel_adr,
@@ -3662,6 +3738,8 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m.site_size,
m.site_type,
m._impl.subtree_mass,
m._impl.taxel_sensorid,
m._impl.taxel_vertadr,
m.tendon_actfrclimited,
m.tendon_actfrcrange,
m.tendon_adr,
@@ -3814,6 +3892,10 @@ def _step_jax_impl(m: types.Model, d: types.Data):
d._impl.sap_range,
d._impl.sap_segment_index,
d._impl.sap_sort_index,
d._impl.sensor_contact_criteria,
d._impl.sensor_contact_direction,
d._impl.sensor_contact_matchid,
d._impl.sensor_contact_nmatch,
d._impl.sensor_rangefinder_dist,
d._impl.sensor_rangefinder_geomid,
d._impl.sensor_rangefinder_pnt,
@@ -3867,8 +3949,6 @@ def _step_jax_impl(m: types.Model, d: types.Data):
d._impl.efc__alpha,
d._impl.efc__aref,
d._impl.efc__beta,
d._impl.efc__beta_den,
d._impl.efc__beta_num,
d._impl.efc__cholesky_L_tmp,
d._impl.efc__cholesky_y_tmp,
d._impl.efc__condim,
@@ -4008,104 +4088,106 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'_impl.sap_range': out[91],
'_impl.sap_segment_index': out[92],
'_impl.sap_sort_index': out[93],
'_impl.sensor_rangefinder_dist': out[94],
'_impl.sensor_rangefinder_geomid': out[95],
'_impl.sensor_rangefinder_pnt': out[96],
'_impl.sensor_rangefinder_vec': out[97],
'sensordata': out[98],
'site_xmat': out[99],
'site_xpos': out[100],
'_impl.solver_niter': out[101],
'_impl.subtree_angmom': out[102],
'_impl.subtree_bodyvel': out[103],
'subtree_com': out[104],
'_impl.subtree_linvel': out[105],
'_impl.ten_J': out[106],
'_impl.ten_Jdot': out[107],
'_impl.ten_actfrc': out[108],
'_impl.ten_bias_coef': out[109],
'_impl.ten_length': out[110],
'_impl.ten_velocity': out[111],
'_impl.ten_wrapadr': out[112],
'_impl.ten_wrapnum': out[113],
'time': out[114],
'_impl.wrap_geom_xpos': out[115],
'_impl.wrap_obj': out[116],
'_impl.wrap_xpos': out[117],
'xanchor': out[118],
'xaxis': out[119],
'xfrc_applied': out[120],
'ximat': out[121],
'xipos': out[122],
'xmat': out[123],
'xpos': out[124],
'xquat': out[125],
'_impl.contact__dim': out[126],
'_impl.contact__dist': out[127],
'_impl.contact__efc_address': out[128],
'_impl.contact__frame': out[129],
'_impl.contact__friction': out[130],
'_impl.contact__geom': out[131],
'_impl.contact__includemargin': out[132],
'_impl.contact__pos': out[133],
'_impl.contact__solimp': out[134],
'_impl.contact__solref': out[135],
'_impl.contact__solreffriction': out[136],
'_impl.contact__worldid': out[137],
'_impl.efc__D': out[138],
'_impl.efc__J': out[139],
'_impl.efc__Jaref': out[140],
'_impl.efc__Ma': out[141],
'_impl.efc__Mgrad': out[142],
'_impl.efc__active': out[143],
'_impl.efc__alpha': out[144],
'_impl.efc__aref': out[145],
'_impl.efc__beta': out[146],
'_impl.efc__beta_den': out[147],
'_impl.efc__beta_num': out[148],
'_impl.efc__cholesky_L_tmp': out[149],
'_impl.efc__cholesky_y_tmp': out[150],
'_impl.efc__condim': 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__type': out[186],
'_impl.efc__u': out[187],
'_impl.efc__uu': out[188],
'_impl.efc__uv': out[189],
'_impl.efc__vel': out[190],
'_impl.efc__vv': out[191],
'_impl.sensor_contact_criteria': out[94],
'_impl.sensor_contact_direction': out[95],
'_impl.sensor_contact_matchid': out[96],
'_impl.sensor_contact_nmatch': out[97],
'_impl.sensor_rangefinder_dist': out[98],
'_impl.sensor_rangefinder_geomid': out[99],
'_impl.sensor_rangefinder_pnt': out[100],
'_impl.sensor_rangefinder_vec': out[101],
'sensordata': out[102],
'site_xmat': out[103],
'site_xpos': out[104],
'_impl.solver_niter': out[105],
'_impl.subtree_angmom': out[106],
'_impl.subtree_bodyvel': out[107],
'subtree_com': out[108],
'_impl.subtree_linvel': out[109],
'_impl.ten_J': out[110],
'_impl.ten_Jdot': out[111],
'_impl.ten_actfrc': out[112],
'_impl.ten_bias_coef': out[113],
'_impl.ten_length': out[114],
'_impl.ten_velocity': out[115],
'_impl.ten_wrapadr': out[116],
'_impl.ten_wrapnum': out[117],
'time': out[118],
'_impl.wrap_geom_xpos': out[119],
'_impl.wrap_obj': out[120],
'_impl.wrap_xpos': out[121],
'xanchor': out[122],
'xaxis': out[123],
'xfrc_applied': out[124],
'ximat': out[125],
'xipos': out[126],
'xmat': out[127],
'xpos': out[128],
'xquat': out[129],
'_impl.contact__dim': out[130],
'_impl.contact__dist': out[131],
'_impl.contact__efc_address': out[132],
'_impl.contact__frame': out[133],
'_impl.contact__friction': out[134],
'_impl.contact__geom': out[135],
'_impl.contact__includemargin': out[136],
'_impl.contact__pos': out[137],
'_impl.contact__solimp': out[138],
'_impl.contact__solref': out[139],
'_impl.contact__solreffriction': out[140],
'_impl.contact__worldid': out[141],
'_impl.efc__D': out[142],
'_impl.efc__J': out[143],
'_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],
})
return d
@@ -4114,8 +4196,6 @@ def _step_jax_impl(m: types.Model, d: types.Data):
@ffi.marshal_jax_warp_callable
def step(m: types.Model, d: types.Data):
return _step_jax_impl(m, d)
@step.def_vmap
@ffi.marshal_custom_vmap
def step_vmap(unused_axis_size, is_batched, m, d):
-3
View File
@@ -42,7 +42,6 @@ _e = mjwarp.Constraint(
**{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init}
)
@ffi.format_args_for_warp
def _kinematics_shim(
# Model
@@ -282,8 +281,6 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data):
@ffi.marshal_jax_warp_callable
def kinematics(m: types.Model, d: types.Data):
return _kinematics_jax_impl(m, d)
@kinematics.def_vmap
@ffi.marshal_custom_vmap
def kinematics_vmap(unused_axis_size, is_batched, m, d):
+38 -21
View File
@@ -13,7 +13,6 @@
# limitations under the License.
# ==============================================================================
"""MJX Warp types.
DO NOT EDIT. This file is auto-generated.
"""
import dataclasses
@@ -23,10 +22,8 @@ from jax import tree_util
from jax.interpreters import batching
from mujoco.mjx._src import dataclasses as mjx_dataclasses
import numpy as np
PyTreeNode = mjx_dataclasses.PyTreeNode
@dataclasses.dataclass(frozen=True)
@tree_util.register_pytree_node_class
class TileSet:
@@ -38,7 +35,6 @@ class TileSet:
adr: address of each tile in the set
size: size of all the tiles in this set
"""
adr: np.ndarray
size: int
@@ -59,11 +55,11 @@ class BlockDim:
TODO(team): experimental and may be removed
"""
actuator_velocity: int
cholesky_factorize: int
cholesky_factorize_solve: int
cholesky_solve: int
contact_sort: int
energy_vel_kinetic: int
euler_dense: int
mul_m_dense: int
@@ -87,13 +83,10 @@ class BlockDim:
class StatisticWarp(PyTreeNode):
"""Derived fields from Statistic."""
meaninertia: float
class OptionWarp(PyTreeNode):
"""Derived fields from Option."""
broadphase: int
broadphase_filter: int
epa_iterations: int
@@ -106,10 +99,8 @@ class OptionWarp(PyTreeNode):
sdf_initpoints: int
sdf_iterations: int
class ModelWarp(PyTreeNode):
"""Derived fields from Model."""
M_colind: np.ndarray
M_rowadr: np.ndarray
M_rownnz: np.ndarray
@@ -168,6 +159,7 @@ class ModelWarp(PyTreeNode):
nmeshpoly: int
nmeshpolymap: int
nmeshpolyvert: int
nsensortaxel: int
nxn_geom_pair: np.ndarray
nxn_geom_pair_filtered: np.ndarray
nxn_pairid: np.ndarray
@@ -183,6 +175,7 @@ class ModelWarp(PyTreeNode):
qM_tiles: Tuple[TileSet, ...]
rangefinder_sensor_adr: np.ndarray
sensor_acc_adr: np.ndarray
sensor_contact_adr: np.ndarray
sensor_e_kinetic: bool
sensor_e_potential: bool
sensor_limitfrc_adr: np.ndarray
@@ -197,6 +190,8 @@ class ModelWarp(PyTreeNode):
sensor_touch_adr: np.ndarray
sensor_vel_adr: np.ndarray
subtree_mass: jax.Array
taxel_sensorid: np.ndarray
taxel_vertadr: np.ndarray
ten_wrapadr_site: np.ndarray
ten_wrapnum_site: np.ndarray
tendon_geom_adr: np.ndarray
@@ -209,10 +204,8 @@ class ModelWarp(PyTreeNode):
wrap_site_adr: np.ndarray
wrap_site_pair_adr: np.ndarray
class DataWarp(PyTreeNode):
"""Derived fields from Data."""
act_dot_rk: jax.Array
act_t0: jax.Array
act_vel_integration: jax.Array
@@ -252,8 +245,6 @@ class DataWarp(PyTreeNode):
efc__alpha: jax.Array
efc__aref: jax.Array
efc__beta: jax.Array
efc__beta_den: jax.Array
efc__beta_num: jax.Array
efc__cholesky_L_tmp: jax.Array
efc__cholesky_y_tmp: jax.Array
efc__condim: jax.Array
@@ -321,6 +312,7 @@ class DataWarp(PyTreeNode):
ncollision: jax.Array
ncon: jax.Array
ncon_hfield: jax.Array
ncon_world: jax.Array
nconmax: int
ne: jax.Array
ne_connect: jax.Array
@@ -357,6 +349,10 @@ class DataWarp(PyTreeNode):
sap_range: jax.Array
sap_segment_index: jax.Array
sap_sort_index: jax.Array
sensor_contact_criteria: jax.Array
sensor_contact_direction: jax.Array
sensor_contact_matchid: jax.Array
sensor_contact_nmatch: jax.Array
sensor_rangefinder_dist: jax.Array
sensor_rangefinder_geomid: jax.Array
sensor_rangefinder_pnt: jax.Array
@@ -377,8 +373,6 @@ class DataWarp(PyTreeNode):
wrap_obj: jax.Array
wrap_xpos: jax.Array
shape = property(lambda self: self.cacc.shape)
DATA_NON_VMAP = {
'collision_hftri_index',
'collision_pair',
@@ -421,7 +415,6 @@ DATA_NON_VMAP = {
'ray_bodyexclude',
}
def _to_elt(cont, _, d, axis):
return DataWarp(**{
f.name: (
@@ -494,8 +487,6 @@ NDIM = {
'efc__alpha': 1,
'efc__aref': 2,
'efc__beta': 1,
'efc__beta_den': 1,
'efc__beta_num': 1,
'efc__cholesky_L_tmp': 3,
'efc__cholesky_y_tmp': 2,
'efc__condim': 2,
@@ -568,6 +559,7 @@ NDIM = {
'ncollision': 1,
'ncon': 1,
'ncon_hfield': 2,
'ncon_world': 1,
'nconmax': 0,
'ne': 1,
'ne_connect': 1,
@@ -618,6 +610,10 @@ NDIM = {
'sap_range': 2,
'sap_segment_index': 2,
'sap_sort_index': 3,
'sensor_contact_criteria': 3,
'sensor_contact_direction': 3,
'sensor_contact_matchid': 3,
'sensor_contact_nmatch': 2,
'sensor_rangefinder_dist': 2,
'sensor_rangefinder_geomid': 2,
'sensor_rangefinder_pnt': 3,
@@ -684,6 +680,7 @@ NDIM = {
'block_dim__cholesky_factorize': 0,
'block_dim__cholesky_factorize_solve': 0,
'block_dim__cholesky_solve': 0,
'block_dim__contact_sort': 0,
'block_dim__energy_vel_kinetic': 0,
'block_dim__euler_dense': 0,
'block_dim__mul_m_dense': 0,
@@ -829,6 +826,8 @@ NDIM = {
'mesh_faceadr': 1,
'mesh_graph': 1,
'mesh_graphadr': 1,
'mesh_normal': 2,
'mesh_normaladr': 1,
'mesh_polyadr': 1,
'mesh_polymap': 1,
'mesh_polymapadr': 1,
@@ -838,6 +837,7 @@ NDIM = {
'mesh_polyvert': 1,
'mesh_polyvertadr': 1,
'mesh_polyvertnum': 1,
'mesh_quat': 2,
'mesh_vert': 2,
'mesh_vertadr': 1,
'mesh_vertnum': 1,
@@ -872,6 +872,7 @@ NDIM = {
'nq': 0,
'nsensor': 0,
'nsensordata': 0,
'nsensortaxel': 0,
'nsite': 0,
'ntendon': 0,
'nu': 0,
@@ -931,11 +932,13 @@ NDIM = {
'rangefinder_sensor_adr': 1,
'sensor_acc_adr': 1,
'sensor_adr': 1,
'sensor_contact_adr': 1,
'sensor_cutoff': 1,
'sensor_datatype': 1,
'sensor_dim': 1,
'sensor_e_kinetic': 0,
'sensor_e_potential': 0,
'sensor_intprm': 2,
'sensor_limitfrc_adr': 1,
'sensor_limitpos_adr': 1,
'sensor_limitvel_adr': 1,
@@ -959,6 +962,8 @@ NDIM = {
'site_type': 1,
'stat__meaninertia': 0,
'subtree_mass': 2,
'taxel_sensorid': 1,
'taxel_vertadr': 1,
'ten_wrapadr_site': 1,
'ten_wrapnum_site': 1,
'tendon_actfrclimited': 1,
@@ -1071,8 +1076,6 @@ BATCH_DIM = {
'efc__alpha': True,
'efc__aref': True,
'efc__beta': True,
'efc__beta_den': True,
'efc__beta_num': True,
'efc__cholesky_L_tmp': True,
'efc__cholesky_y_tmp': True,
'efc__condim': True,
@@ -1145,6 +1148,7 @@ BATCH_DIM = {
'ncollision': False,
'ncon': False,
'ncon_hfield': True,
'ncon_world': True,
'nconmax': False,
'ne': True,
'ne_connect': True,
@@ -1195,6 +1199,10 @@ BATCH_DIM = {
'sap_range': True,
'sap_segment_index': True,
'sap_sort_index': True,
'sensor_contact_criteria': True,
'sensor_contact_direction': True,
'sensor_contact_matchid': True,
'sensor_contact_nmatch': True,
'sensor_rangefinder_dist': True,
'sensor_rangefinder_geomid': True,
'sensor_rangefinder_pnt': True,
@@ -1261,6 +1269,7 @@ BATCH_DIM = {
'block_dim__cholesky_factorize': False,
'block_dim__cholesky_factorize_solve': False,
'block_dim__cholesky_solve': False,
'block_dim__contact_sort': False,
'block_dim__energy_vel_kinetic': False,
'block_dim__euler_dense': False,
'block_dim__mul_m_dense': False,
@@ -1406,6 +1415,8 @@ BATCH_DIM = {
'mesh_faceadr': False,
'mesh_graph': False,
'mesh_graphadr': False,
'mesh_normal': False,
'mesh_normaladr': False,
'mesh_polyadr': False,
'mesh_polymap': False,
'mesh_polymapadr': False,
@@ -1415,6 +1426,7 @@ BATCH_DIM = {
'mesh_polyvert': False,
'mesh_polyvertadr': False,
'mesh_polyvertnum': False,
'mesh_quat': False,
'mesh_vert': False,
'mesh_vertadr': False,
'mesh_vertnum': False,
@@ -1449,6 +1461,7 @@ BATCH_DIM = {
'nq': False,
'nsensor': False,
'nsensordata': False,
'nsensortaxel': False,
'nsite': False,
'ntendon': False,
'nu': False,
@@ -1508,11 +1521,13 @@ BATCH_DIM = {
'rangefinder_sensor_adr': False,
'sensor_acc_adr': False,
'sensor_adr': False,
'sensor_contact_adr': False,
'sensor_cutoff': False,
'sensor_datatype': False,
'sensor_dim': False,
'sensor_e_kinetic': False,
'sensor_e_potential': False,
'sensor_intprm': False,
'sensor_limitfrc_adr': False,
'sensor_limitpos_adr': False,
'sensor_limitvel_adr': False,
@@ -1536,6 +1551,8 @@ BATCH_DIM = {
'site_type': False,
'stat__meaninertia': False,
'subtree_mass': True,
'taxel_sensorid': False,
'taxel_vertadr': False,
'ten_wrapadr_site': False,
'ten_wrapnum_site': False,
'tendon_actfrclimited': False,