Import google-deepmind/mujoco_warp from GitHub.

PiperOrigin-RevId: 807766737
Change-Id: I1c22361ffa61f27fd8e0fd534b96488531761916
This commit is contained in:
Taylor Howell
2025-09-16 11:07:51 -07:00
committed by Copybara-Service
parent 52da7586dc
commit f2fcd05b1b
45 changed files with 18028 additions and 3940 deletions
+3 -3
View File
@@ -16,6 +16,6 @@ jax-cuda12-pjrt==0.5.3; python_version >= '3.10' \
jax-cuda12-pjrt==0.4.30; python_version == '3.9' \
--hash=sha256:895d0198ad99638fcaf976c47592e2a543eef79ea15fabd24a402d055390c328 \
--hash=sha256:c36fb1e0c236563bf3a87e70f4d1ab28a31d7cf5d722c9ede30c4172116e8bcb
warp-lang==1.8.1 \
--hash=sha256:cfc59e1070ad71531b5d83186de48162507277af344a102fa33d5df9cdb942f7 \
--hash=sha256:1db9ca92c46902b76bb99565c544347d1a32e9fb875ce902f1cafb94978d1ac3
warp-lang==1.9.0 \
--hash=sha256:23165d3291eeecc5ac47b9a3de0b93d34b4c921414e5a761ab96d68861efa3cb \
--hash=sha256:7ea8057e5d6fdb9b0885d725b51954c5fc6dac62d6ac007828148b1cb13ee615
+15
View File
@@ -15,7 +15,9 @@
"""Tests for io functions."""
import os
import tempfile
from unittest import mock
from absl.testing import absltest
from absl.testing import parameterized
import jax
@@ -31,6 +33,7 @@ from mujoco.mjx._src.types import JacobianType
# pylint: enable=g-importing-member
import mujoco.mjx.warp as mjxw
from mujoco.mjx.warp import types as mjxw_types
from mujoco.mjx.warp import warp as wp # pylint: disable=g-importing-member
import numpy as np
@@ -125,6 +128,12 @@ def _get_name_from_path(path: jax.tree_util.KeyPath) -> str:
class ModelIOTest(parameterized.TestCase):
"""IO tests for mjx.Model."""
def setUp(self):
super().setUp()
if mjxw.WARP_INSTALLED:
self.tempdir = tempfile.TemporaryDirectory()
wp.config.kernel_cache_dir = self.tempdir.name
@parameterized.product(
xml=(_MULTIPLE_CONVEX_OBJECTS, _MULTIPLE_CONSTRAINTS),
impl=('jax', 'c', 'warp'),
@@ -326,6 +335,12 @@ class ModelIOTest(parameterized.TestCase):
class DataIOTest(parameterized.TestCase):
"""IO tests for mjx.Data."""
def setUp(self):
super().setUp()
if mjxw.WARP_INSTALLED:
self.tempdir = tempfile.TemporaryDirectory()
wp.config.kernel_cache_dir = self.tempdir.name
@parameterized.parameters('jax', 'c')
def test_make_data(self, impl: str):
"""Test that make_data returns the correct shapes."""
+2 -2
View File
@@ -21,6 +21,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import Model as Model
from mujoco.mjx.third_party.mujoco_warp._src.types import Data as Data
# isort: on
from ._src import test_util as test_util # used by viewer and testspeed, not meant for public consumption
from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import collision as collision
from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import nxn_broadphase as nxn_broadphase
from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import sap_broadphase as sap_broadphase
@@ -43,6 +44,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.io import get_data_into as get_data
from mujoco.mjx.third_party.mujoco_warp._src.io import make_data as make_data
from mujoco.mjx.third_party.mujoco_warp._src.io import put_data as put_data
from mujoco.mjx.third_party.mujoco_warp._src.io import put_model as put_model
from mujoco.mjx.third_party.mujoco_warp._src.io import reset_data as reset_data
from mujoco.mjx.third_party.mujoco_warp._src.passive import passive as passive
from mujoco.mjx.third_party.mujoco_warp._src.ray import ray as ray
from mujoco.mjx.third_party.mujoco_warp._src.sensor import energy_pos as energy_pos
@@ -66,8 +68,6 @@ 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
@@ -326,9 +326,6 @@ class BroadphaseTest(parameterized.TestCase):
broadphase_caller(m, d)
self.assertEqual(d.ncollision.numpy()[0], 0)
# TODO(team): test margin
# TODO(team): test DisableBit.FILTERPARENT
if __name__ == "__main__":
wp.init()
+391 -170
View File
@@ -19,12 +19,12 @@ from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd
from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import epa_legacy
from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import gjk_legacy
from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import multicontact_legacy
from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_prism_vertex
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import _geom
from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_filter
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact
from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame
from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index
from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR
from mujoco.mjx.third_party.mujoco_warp._src.types import Data
@@ -36,10 +36,11 @@ from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
# TODO(team): improve compile time to enable backward pass
wp.config.enable_backward = False
wp.set_module_options({"enable_backward": False})
MULTI_CONTACT_COUNT = 8
mat3c = wp.types.matrix(shape=(MULTI_CONTACT_COUNT, 3), dtype=float)
mat63 = wp.types.matrix(shape=(6, 3), dtype=float)
_CONVEX_COLLISION_PAIRS = [
(GeomType.HFIELD.value, GeomType.SPHERE.value),
@@ -78,30 +79,6 @@ def _check_convex_collision_pairs():
assert _check_convex_collision_pairs(), "_CONVEX_COLLISION_PAIRS is in invalid order."
@wp.func
def _max_contacts_height_field(
# Model:
ngeom: int,
geom_type: wp.array(dtype=int),
geompair2hfgeompair: wp.array(dtype=int),
# In:
g1: int,
g2: int,
worldid: int,
# Data out:
ncon_hfield_out: wp.array2d(dtype=int),
):
hfield = int(GeomType.HFIELD.value)
if geom_type[g1] == hfield or (geom_type[g2] == hfield):
geompairid = upper_tri_index(ngeom, g1, g2)
hfgeompairid = geompair2hfgeompair[geompairid]
hfncon = wp.atomic_add(ncon_hfield_out[worldid], hfgeompairid, 1)
if hfncon >= MJ_MAXCONPAIR:
return True
return False
@cache_kernel
def ccd_kernel_builder(
legacy_gjk: bool,
@@ -111,12 +88,176 @@ def ccd_kernel_builder(
epa_iterations: int,
epa_exact_neg_distance: bool,
depth_extension: float,
is_hfield: bool,
):
@wp.func
def eval_ccd_write_contact(
# Model:
opt_ccd_tolerance: wp.array(dtype=float),
geom_type: wp.array(dtype=int),
# Data in:
nconmax_in: int,
epa_vert_in: wp.array2d(dtype=wp.vec3),
epa_vert1_in: wp.array2d(dtype=wp.vec3),
epa_vert2_in: wp.array2d(dtype=wp.vec3),
epa_vert_index1_in: wp.array2d(dtype=int),
epa_vert_index2_in: wp.array2d(dtype=int),
epa_face_in: wp.array2d(dtype=wp.vec3i),
epa_pr_in: wp.array2d(dtype=wp.vec3),
epa_norm2_in: wp.array2d(dtype=float),
epa_index_in: wp.array2d(dtype=int),
epa_map_in: wp.array2d(dtype=int),
epa_horizon_in: wp.array2d(dtype=int),
multiccd_polygon_in: wp.array2d(dtype=wp.vec3),
multiccd_clipped_in: wp.array2d(dtype=wp.vec3),
multiccd_pnormal_in: wp.array2d(dtype=wp.vec3),
multiccd_pdist_in: wp.array2d(dtype=float),
multiccd_idx1_in: wp.array2d(dtype=int),
multiccd_idx2_in: wp.array2d(dtype=int),
multiccd_n1_in: wp.array2d(dtype=wp.vec3),
multiccd_n2_in: wp.array2d(dtype=wp.vec3),
multiccd_endvert_in: wp.array2d(dtype=wp.vec3),
multiccd_face1_in: wp.array2d(dtype=wp.vec3),
multiccd_face2_in: wp.array2d(dtype=wp.vec3),
# In:
geom1: Geom,
geom2: Geom,
geoms: wp.vec2i,
worldid: int,
tid: int,
margin: float,
gap: float,
condim: int,
friction: vec5,
solref: wp.vec2,
solreffriction: wp.vec2,
solimp: vec5,
x1: wp.vec3,
x2: wp.vec3,
count: int,
# Data out:
ncon_out: wp.array(dtype=int),
contact_dist_out: wp.array(dtype=float),
contact_pos_out: wp.array(dtype=wp.vec3),
contact_frame_out: wp.array(dtype=wp.mat33),
contact_includemargin_out: wp.array(dtype=float),
contact_friction_out: wp.array(dtype=vec5),
contact_solref_out: wp.array(dtype=wp.vec2),
contact_solreffriction_out: wp.array(dtype=wp.vec2),
contact_solimp_out: wp.array(dtype=vec5),
contact_dim_out: wp.array(dtype=int),
contact_geom_out: wp.array(dtype=wp.vec2i),
contact_worldid_out: wp.array(dtype=int),
) -> int:
# TODO(kbayes): remove legacy GJK once multicontact can be enabled
if wp.static(legacy_gjk):
simplex, normal = gjk_legacy(
gjk_iterations,
geom1,
geom2,
geomtype1,
geomtype2,
)
depth, normal = epa_legacy(
epa_iterations, geom1, geom2, geomtype1, geomtype2, depth_extension, epa_exact_neg_distance, simplex, normal
)
dist = -depth
if dist >= 0.0 or depth < -depth_extension:
return 0
sphere = int(GeomType.SPHERE.value)
ellipsoid = int(GeomType.ELLIPSOID.value)
g1 = geoms[0]
g2 = geoms[1]
if geom_type[g1] == sphere or geom_type[g1] == ellipsoid or geom_type[g2] == sphere or geom_type[g2] == ellipsoid:
ncontact, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 1, 2, 1.0e-5)
else:
ncontact, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 4, 8, 1.0e-1)
frame = make_frame(normal)
else:
points = mat3c()
dist, ncontact, witness1, witness2 = ccd(
False,
opt_ccd_tolerance[worldid],
0.0,
gjk_iterations,
epa_iterations,
geom1,
geom2,
geomtype1,
geomtype2,
x1,
x2,
epa_vert_in[tid],
epa_vert1_in[tid],
epa_vert2_in[tid],
epa_vert_index1_in[tid],
epa_vert_index2_in[tid],
epa_face_in[tid],
epa_pr_in[tid],
epa_norm2_in[tid],
epa_index_in[tid],
epa_map_in[tid],
epa_horizon_in[tid],
multiccd_polygon_in[tid],
multiccd_clipped_in[tid],
multiccd_pnormal_in[tid],
multiccd_pdist_in[tid],
multiccd_idx1_in[tid],
multiccd_idx2_in[tid],
multiccd_n1_in[tid],
multiccd_n2_in[tid],
multiccd_endvert_in[tid],
multiccd_face1_in[tid],
multiccd_face2_in[tid],
)
if dist >= 0.0:
return 0
for i in range(ncontact):
points[i] = 0.5 * (witness1[i] + witness2[i])
normal = witness1[0] - witness2[0]
frame = make_frame(normal)
for i in range(ncontact):
write_contact(
nconmax_in,
dist,
points[i],
frame,
margin,
gap,
condim,
friction,
solref,
solreffriction,
solimp,
geoms,
worldid,
ncon_out,
contact_dist_out,
contact_pos_out,
contact_frame_out,
contact_includemargin_out,
contact_friction_out,
contact_solref_out,
contact_solreffriction_out,
contact_solimp_out,
contact_dim_out,
contact_geom_out,
contact_worldid_out,
)
if count + (i + 1) >= MJ_MAXCONPAIR:
return i + 1
return ncontact
# runs convex collision on a set of geom pairs to recover contact info
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def ccd_kernel(
# Model:
ngeom: int,
opt_ccd_tolerance: wp.array(dtype=float),
geom_type: wp.array(dtype=int),
geom_condim: wp.array(dtype=int),
geom_dataid: wp.array(dtype=int),
@@ -125,6 +266,8 @@ def ccd_kernel_builder(
geom_solref: wp.array2d(dtype=wp.vec2),
geom_solimp: wp.array2d(dtype=vec5),
geom_size: wp.array2d(dtype=wp.vec3),
geom_aabb: wp.array2d(dtype=wp.vec3),
geom_rbound: wp.array2d(dtype=float),
geom_friction: wp.array2d(dtype=wp.vec3),
geom_margin: wp.array2d(dtype=float),
geom_gap: wp.array2d(dtype=float),
@@ -154,13 +297,11 @@ def ccd_kernel_builder(
pair_margin: wp.array2d(dtype=float),
pair_gap: wp.array2d(dtype=float),
pair_friction: wp.array2d(dtype=vec5),
geompair2hfgeompair: wp.array(dtype=int),
# Data in:
nconmax_in: int,
geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.array2d(dtype=wp.mat33),
collision_pair_in: wp.array(dtype=wp.vec2i),
collision_hftri_index_in: wp.array(dtype=int),
collision_pairid_in: wp.array(dtype=int),
collision_worldid_in: wp.array(dtype=int),
ncollision_in: wp.array(dtype=int),
@@ -175,9 +316,19 @@ def ccd_kernel_builder(
epa_index_in: wp.array2d(dtype=int),
epa_map_in: wp.array2d(dtype=int),
epa_horizon_in: wp.array2d(dtype=int),
multiccd_polygon_in: wp.array2d(dtype=wp.vec3),
multiccd_clipped_in: wp.array2d(dtype=wp.vec3),
multiccd_pnormal_in: wp.array2d(dtype=wp.vec3),
multiccd_pdist_in: wp.array2d(dtype=float),
multiccd_idx1_in: wp.array2d(dtype=int),
multiccd_idx2_in: wp.array2d(dtype=int),
multiccd_n1_in: wp.array2d(dtype=wp.vec3),
multiccd_n2_in: wp.array2d(dtype=wp.vec3),
multiccd_endvert_in: wp.array2d(dtype=wp.vec3),
multiccd_face1_in: wp.array2d(dtype=wp.vec3),
multiccd_face2_in: wp.array2d(dtype=wp.vec3),
# Data out:
ncon_out: wp.array(dtype=int),
ncon_hfield_out: wp.array2d(dtype=int),
contact_dist_out: wp.array(dtype=float),
contact_pos_out: wp.array(dtype=wp.vec3),
contact_frame_out: wp.array(dtype=wp.mat33),
@@ -202,6 +353,15 @@ def ccd_kernel_builder(
return
worldid = collision_worldid_in[tid]
# height field filter
if wp.static(is_hfield):
no_hf_collision, xmin, xmax, ymin, ymax, zmin, zmax = hfield_filter(
geom_dataid, geom_aabb, geom_rbound, geom_margin, hfield_size, geom_xpos_in, geom_xmat_in, worldid, g1, g2
)
if no_hf_collision:
return
_, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
geom_condim,
geom_priority,
@@ -224,24 +384,18 @@ def ccd_kernel_builder(
worldid,
)
hftri_index = collision_hftri_index_in[tid]
geom1 = _geom(
geom_type,
geom_dataid,
geom_size,
hfield_adr,
hfield_nrow,
hfield_ncol,
hfield_size,
hfield_data,
mesh_vertadr,
mesh_vertnum,
geom1_dataid = geom_dataid[g1]
geom1 = geom(
geomtype1,
geom1_dataid,
geom_size[worldid, g1],
mesh_vertadr[geom1_dataid],
mesh_vertnum[geom1_dataid],
mesh_vert,
mesh_graphadr,
mesh_graphadr[geom1_dataid],
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynum[geom1_dataid],
mesh_polyadr[geom1_dataid],
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
@@ -249,29 +403,22 @@ def ccd_kernel_builder(
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in,
geom_xmat_in,
worldid,
g1,
hftri_index,
geom_xpos_in[worldid, g1],
geom_xmat_in[worldid, g1],
)
geom2 = _geom(
geom_type,
geom_dataid,
geom_size,
hfield_adr,
hfield_nrow,
hfield_ncol,
hfield_size,
hfield_data,
mesh_vertadr,
mesh_vertnum,
geom2_dataid = geom_dataid[g2]
geom2 = geom(
geomtype2,
geom2_dataid,
geom_size[worldid, g2],
mesh_vertadr[geom2_dataid],
mesh_vertnum[geom2_dataid],
mesh_vert,
mesh_graphadr,
mesh_graphadr[geom2_dataid],
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynum[geom2_dataid],
mesh_polyadr[geom2_dataid],
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
@@ -279,102 +426,169 @@ def ccd_kernel_builder(
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in,
geom_xmat_in,
worldid,
g2,
hftri_index,
geom_xpos_in[worldid, g2],
geom_xmat_in[worldid, g2],
)
# TODO(kbayes): remove legacy GJK once multicontact can be enabled
if wp.static(legacy_gjk):
# find prism center for height field
if geomtype1 == int(GeomType.HFIELD.value):
x1 = wp.vec3(0.0, 0.0, 0.0)
for i in range(6):
x1 += hfield_prism_vertex(geom1.hfprism, i)
x1 = geom1.pos + geom1.rot @ (x1 / 6.0)
geom1.pos = x1
# see MuJoCo mjc_ConvexHField
if wp.static(is_hfield):
# height field subgrid
nrow = hfield_nrow[g1]
ncol = hfield_ncol[g1]
size = hfield_size[g1]
simplex, normal = gjk_legacy(
gjk_iterations,
geom1,
geom2,
geomtype1,
geomtype2,
)
# subgrid
x_scale = 0.5 * float(ncol - 1) / size[0]
y_scale = 0.5 * float(nrow - 1) / size[1]
cmin = wp.max(0, int(wp.floor((xmin + size[0]) * x_scale)))
cmax = wp.min(ncol - 1, int(wp.ceil((xmax + size[0]) * x_scale)))
rmin = wp.max(0, int(wp.floor((ymin + size[1]) * y_scale)))
rmax = wp.min(nrow - 1, int(wp.ceil((ymax + size[1]) * y_scale)))
depth, normal = epa_legacy(
epa_iterations, geom1, geom2, geomtype1, geomtype2, depth_extension, epa_exact_neg_distance, simplex, normal
)
dist = -depth
dx = (2.0 * size[0]) / float(ncol - 1)
dy = (2.0 * size[1]) / float(nrow - 1)
dr = wp.vec2i(1, 0)
if dist >= 0.0 or depth < -depth_extension:
count = 0
return
sphere = int(GeomType.SPHERE.value)
ellipsoid = int(GeomType.ELLIPSOID.value)
if geom_type[g1] == sphere or geom_type[g1] == ellipsoid or geom_type[g2] == sphere or geom_type[g2] == ellipsoid:
count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 1, 2, 1.0e-5)
else:
count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 4, 8, 1.0e-1)
frame = make_frame(normal)
prism = mat63()
# set zbottom value using base size
prism[0, 2] = -size[3]
prism[1, 2] = -size[3]
prism[2, 2] = -size[3]
adr = hfield_adr[geom1_dataid]
# process all prisms in subgrid
count = int(0)
for r in range(rmin, rmax):
nvert = int(0)
for c in range(cmin, cmax + 1):
# add both triangles from this cell
for i in range(2):
# add vert
x = dx * float(c) - size[0]
y = dy * float(r + dr[i]) - size[1]
z = hfield_data[adr + (r + dr[i]) * ncol + c] * size[2] + margin
prism[0] = prism[1]
prism[1] = prism[2]
prism[3] = prism[4]
prism[4] = prism[5]
prism[2, 0] = x
prism[5, 0] = x
prism[2, 1] = y
prism[5, 1] = y
prism[5, 2] = z
nvert += 1
if nvert <= 2:
continue
# prism height test
if prism[3, 2] < zmin and prism[4, 2] < zmin and prism[5, 2] < zmin:
continue
geom1.hfprism = prism
# prism center
x1 = geom1.pos
if wp.static(not legacy_gjk):
x1_ = wp.vec3(0.0, 0.0, 0.0)
for i in range(6):
x1_ += prism[i]
x1 += geom1.rot @ (x1_ / 6.0)
ncontact = eval_ccd_write_contact(
opt_ccd_tolerance,
geom_type,
nconmax_in,
epa_vert_in,
epa_vert1_in,
epa_vert2_in,
epa_vert_index1_in,
epa_vert_index2_in,
epa_face_in,
epa_pr_in,
epa_norm2_in,
epa_index_in,
epa_map_in,
epa_horizon_in,
multiccd_polygon_in,
multiccd_clipped_in,
multiccd_pnormal_in,
multiccd_pdist_in,
multiccd_idx1_in,
multiccd_idx2_in,
multiccd_n1_in,
multiccd_n2_in,
multiccd_endvert_in,
multiccd_face1_in,
multiccd_face2_in,
geom1,
geom2,
geoms,
worldid,
tid,
margin,
gap,
condim,
friction,
solref,
solreffriction,
solimp,
x1,
geom2.pos,
count,
ncon_out,
contact_dist_out,
contact_pos_out,
contact_frame_out,
contact_includemargin_out,
contact_friction_out,
contact_solref_out,
contact_solreffriction_out,
contact_solimp_out,
contact_dim_out,
contact_geom_out,
contact_worldid_out,
)
count += ncontact
if count >= MJ_MAXCONPAIR:
return
else:
points = mat3c()
x1 = geom1.pos
x2 = geom2.pos
# find prism center for height field
if geomtype1 == int(GeomType.HFIELD.value):
x1_ = wp.vec3(0.0, 0.0, 0.0)
for i in range(6):
x1_ += hfield_prism_vertex(geom1.hfprism, i)
x1 += geom1.rot @ (x1_ / 6.0)
dist, count, witness1, witness2 = ccd(
False,
1e-6,
0.0,
gjk_iterations,
epa_iterations,
eval_ccd_write_contact(
opt_ccd_tolerance,
geom_type,
nconmax_in,
epa_vert_in,
epa_vert1_in,
epa_vert2_in,
epa_vert_index1_in,
epa_vert_index2_in,
epa_face_in,
epa_pr_in,
epa_norm2_in,
epa_index_in,
epa_map_in,
epa_horizon_in,
multiccd_polygon_in,
multiccd_clipped_in,
multiccd_pnormal_in,
multiccd_pdist_in,
multiccd_idx1_in,
multiccd_idx2_in,
multiccd_n1_in,
multiccd_n2_in,
multiccd_endvert_in,
multiccd_face1_in,
multiccd_face2_in,
geom1,
geom2,
geomtype1,
geomtype2,
x1,
x2,
epa_vert_in[tid],
epa_vert1_in[tid],
epa_vert2_in[tid],
epa_vert_index1_in[tid],
epa_vert_index2_in[tid],
epa_face_in[tid],
epa_pr_in[tid],
epa_norm2_in[tid],
epa_index_in[tid],
epa_map_in[tid],
epa_horizon_in[tid],
)
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
if _max_contacts_height_field(ngeom, geom_type, geompair2hfgeompair, g1, g2, worldid, ncon_hfield_out):
return
write_contact(
nconmax_in,
dist,
points[i],
frame,
geoms,
worldid,
tid,
margin,
gap,
condim,
@@ -382,8 +596,9 @@ def ccd_kernel_builder(
solref,
solreffriction,
solimp,
geoms,
worldid,
geom1.pos,
geom2.pos,
0,
ncon_out,
contact_dist_out,
contact_pos_out,
@@ -418,20 +633,16 @@ def convex_narrowphase(m: Model, d: Data):
computations for non-existent pair types.
"""
for geom_pair in _CONVEX_COLLISION_PAIRS:
if m.geom_pair_type_count[upper_trid_index(len(GeomType), geom_pair[0], geom_pair[1])]:
g1 = geom_pair[0]
g2 = geom_pair[1]
if m.geom_pair_type_count[upper_trid_index(len(GeomType), g1, g2)]:
wp.launch(
ccd_kernel_builder(
m.opt.legacy_gjk,
geom_pair[0],
geom_pair[1],
m.opt.gjk_iterations,
m.opt.epa_iterations,
True,
1e9,
m.opt.legacy_gjk, g1, g2, m.opt.gjk_iterations, m.opt.epa_iterations, True, 1e9, g1 == int(GeomType.HFIELD.value)
),
dim=d.nconmax,
inputs=[
m.ngeom,
m.opt.ccd_tolerance,
m.geom_type,
m.geom_condim,
m.geom_dataid,
@@ -440,6 +651,8 @@ def convex_narrowphase(m: Model, d: Data):
m.geom_solref,
m.geom_solimp,
m.geom_size,
m.geom_aabb,
m.geom_rbound,
m.geom_friction,
m.geom_margin,
m.geom_gap,
@@ -469,12 +682,10 @@ def convex_narrowphase(m: Model, d: Data):
m.pair_margin,
m.pair_gap,
m.pair_friction,
m.geompair2hfgeompair,
d.nconmax,
d.geom_xpos,
d.geom_xmat,
d.collision_pair,
d.collision_hftri_index,
d.collision_pairid,
d.collision_worldid,
d.ncollision,
@@ -489,10 +700,20 @@ def convex_narrowphase(m: Model, d: Data):
d.epa_index,
d.epa_map,
d.epa_horizon,
d.multiccd_polygon,
d.multiccd_clipped,
d.multiccd_pnormal,
d.multiccd_pdist,
d.multiccd_idx1,
d.multiccd_idx2,
d.multiccd_n1,
d.multiccd_n2,
d.multiccd_endvert,
d.multiccd_face1,
d.multiccd_face2,
],
outputs=[
d.ncon,
d.ncon_hfield,
d.contact.dist,
d.contact.pos,
d.contact.frame,
@@ -18,7 +18,6 @@ from typing import Any
import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src.collision_convex import convex_narrowphase
from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_midphase
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import primitive_narrowphase
from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf_narrowphase
from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index
@@ -37,29 +36,13 @@ wp.set_module_options({"enable_backward": False})
@wp.kernel
def _zero_collision_arrays(
# Data in:
nworld_in: int,
# In:
hfield_geom_pair_in: int,
def _zero_ncon_ncollision(
# Data out:
ncon_out: wp.array(dtype=int),
ncon_hfield_out: wp.array(dtype=int), # kernel_analyzer: ignore
collision_hftri_index_out: wp.array(dtype=int),
ncollision_out: wp.array(dtype=int),
):
tid = wp.tid()
if tid == 0:
# Zero the single collision counter
ncollision_out[0] = 0
ncon_out[0] = 0
if tid < hfield_geom_pair_in * nworld_in:
ncon_hfield_out[tid] = 0
# Zero collision pair indices
collision_hftri_index_out[tid] = 0
ncollision_out[0] = 0
ncon_out[0] = 0
@wp.func
@@ -316,7 +299,6 @@ def _add_geom_pair(
nxnid: int,
# Data out:
collision_pair_out: wp.array(dtype=wp.vec2i),
collision_hftri_index_out: wp.array(dtype=int),
collision_pairid_out: wp.array(dtype=int),
collision_worldid_out: wp.array(dtype=int),
ncollision_out: wp.array(dtype=int),
@@ -338,12 +320,6 @@ def _add_geom_pair(
collision_pairid_out[pairid] = nxn_pairid[nxnid]
collision_worldid_out[pairid] = worldid
# Writing -1 to collision_hftri_index_out[pairid] signals
# hfield_midphase to generate a collision pair for every
# potentially colliding triangle
if type1 == int(GeomType.HFIELD.value) or type2 == int(GeomType.HFIELD.value):
collision_hftri_index_out[pairid] = -1
@wp.func
def _binary_search(values: wp.array(dtype=Any), value: Any, lower: int, upper: int) -> int:
@@ -419,7 +395,7 @@ def _sap_range(
@cache_kernel
def _sap_broadphase(broadphase_filter):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def kernel(
# Model:
ngeom: int,
@@ -439,7 +415,6 @@ def _sap_broadphase(broadphase_filter):
nsweep_in: int,
# Data out:
collision_pair_out: wp.array(dtype=wp.vec2i),
collision_hftri_index_out: wp.array(dtype=int),
collision_pairid_out: wp.array(dtype=int),
collision_worldid_out: wp.array(dtype=int),
ncollision_out: wp.array(dtype=int),
@@ -485,7 +460,6 @@ def _sap_broadphase(broadphase_filter):
worldid,
idx,
collision_pair_out,
collision_hftri_index_out,
collision_pairid_out,
collision_worldid_out,
ncollision_out,
@@ -615,7 +589,6 @@ def sap_broadphase(m: Model, d: Data):
],
outputs=[
d.collision_pair,
d.collision_hftri_index,
d.collision_pairid,
d.collision_worldid,
d.ncollision,
@@ -625,7 +598,7 @@ def sap_broadphase(m: Model, d: Data):
@cache_kernel
def _nxn_broadphase(broadphase_filter):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def kernel(
# Model:
geom_type: wp.array(dtype=int),
@@ -640,7 +613,6 @@ def _nxn_broadphase(broadphase_filter):
geom_xmat_in: wp.array2d(dtype=wp.mat33),
# Data out:
collision_pair_out: wp.array(dtype=wp.vec2i),
collision_hftri_index_out: wp.array(dtype=int),
collision_pairid_out: wp.array(dtype=int),
collision_worldid_out: wp.array(dtype=int),
ncollision_out: wp.array(dtype=int),
@@ -661,7 +633,6 @@ def _nxn_broadphase(broadphase_filter):
worldid,
elementid,
collision_pair_out,
collision_hftri_index_out,
collision_pairid_out,
collision_worldid_out,
ncollision_out,
@@ -702,7 +673,6 @@ def nxn_broadphase(m: Model, d: Data):
],
outputs=[
d.collision_pair,
d.collision_hftri_index,
d.collision_pairid,
d.collision_worldid,
d.ncollision,
@@ -711,10 +681,6 @@ def nxn_broadphase(m: Model, d: Data):
def _narrowphase(m, d):
# Process heightfield collisions
if m.nhfield > 0:
hfield_midphase(m, d)
# TODO(team): we should reject far-away contacts in the narrowphase instead of constraint
# partitioning because we can move some pressure of the atomics
convex_narrowphase(m, d)
@@ -743,19 +709,8 @@ def collision(m: Model, d: Data):
via `m.opt.disableflags` or if `d.nconmax` is 0.
"""
# zero collision-related arrays
wp.launch(
_zero_collision_arrays,
dim=d.nconmax,
inputs=[
d.nworld,
d.ncon_hfield.shape[1],
d.ncon,
d.ncon_hfield.reshape(-1),
d.collision_hftri_index,
d.ncollision,
],
)
# zero contact and collision counters
wp.launch(_zero_ncon_ncollision, dim=1, outputs=[d.ncon, d.ncollision])
if d.nconmax == 0 or m.opt.disableflags & (DisableBit.CONSTRAINT | DisableBit.CONTACT):
return
@@ -25,6 +25,24 @@ from mujoco.mjx.third_party.mujoco_warp.test_data.collision_sdf.utils import reg
from mujoco.mjx.third_party.mujoco_warp._src import test_util
from mujoco.mjx.third_party.mujoco_warp._src import types
from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import VolumeData
from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sample_volume_sdf
@wp.kernel
def sample_sdf_kernel(
# In:
points: wp.array(dtype=wp.vec3),
volume_data: VolumeData,
# Out:
results_out: wp.array(dtype=float),
):
"""Kernel to sample SDF values at given points using Warp volume."""
tid = wp.tid()
point = points[tid]
sdf_value = sample_volume_sdf(point, volume_data)
results_out[tid] = sdf_value
_TOLERANCE = 5e-5
@@ -599,7 +617,8 @@ class CollisionTest(parameterized.TestCase):
self.assertEqual(m.nxn_geom_pair.numpy().shape[0], 3)
np.testing.assert_equal(m.nxn_pairid.numpy(), np.array([-2, -1, -1]))
def test_contact_pair(self):
@parameterized.parameters(list(types.BroadphaseType))
def test_contact_pair(self, broadphase):
"""Tests contact pair."""
# no pairs
_, _, m, _ = test_util.fixture(
@@ -612,7 +631,8 @@ class CollisionTest(parameterized.TestCase):
</body>
</worldbody>
</mujoco>
"""
""",
broadphase=broadphase,
)
self.assertTrue((m.nxn_pairid.numpy() == -1).all())
@@ -798,8 +818,6 @@ class CollisionTest(parameterized.TestCase):
np.testing.assert_allclose(d.contact.solreffriction.numpy()[1], np.array([2.0, 4.0]))
np.testing.assert_allclose(d.contact.solimp.numpy()[1], np.array([0.1, 0.2, 0.3, 0.4, 0.5]))
# TODO(team): test sap_broadphase
@parameterized.parameters(
(True, True),
(True, False),
File diff suppressed because it is too large Load Diff
@@ -15,7 +15,6 @@
import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_prism_vertex
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom
from mujoco.mjx.third_party.mujoco_warp._src.math import gjk_normalize
from mujoco.mjx.third_party.mujoco_warp._src.math import orthonormal
@@ -26,7 +25,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
# TODO(team): improve compile time to enable backward pass
wp.config.enable_backward = False
wp.set_module_options({"enable_backward": False})
FLOAT_MIN = -1e30
FLOAT_MAX = 1e30
@@ -121,7 +120,7 @@ def _gjk_support_geom(geom: Geom, geomtype: int, dir: wp.vec3):
elif geomtype == int(GeomType.HFIELD.value):
max_dist = float(FLOAT_MIN)
for i in range(6):
vert = hfield_prism_vertex(geom.hfprism, i)
vert = geom.hfprism[i]
dist = wp.dot(vert, local_dir)
if dist > max_dist:
max_dist = dist
@@ -30,7 +30,7 @@ MAX_ITERATIONS = 20
def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int, multiccd=False):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def _gjk_kernel(
# Model:
geom_type: wp.array(dtype=int),
@@ -66,6 +66,17 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int, multicc
face_index: wp.array(dtype=int),
face_map: wp.array(dtype=int),
horizon: wp.array(dtype=int),
polygon: wp.array(dtype=wp.vec3),
clipped: wp.array(dtype=wp.vec3),
pnormal: wp.array(dtype=wp.vec3),
pdist: wp.array(dtype=float),
idx1: wp.array(dtype=int),
idx2: wp.array(dtype=int),
n1: wp.array(dtype=wp.vec3),
n2: wp.array(dtype=wp.vec3),
endvert: wp.array(dtype=wp.vec3),
face1: wp.array(dtype=wp.vec3),
face2: wp.array(dtype=wp.vec3),
# Out:
dist_out: wp.array(dtype=float),
ncon_out: wp.array(dtype=int),
@@ -152,6 +163,17 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int, multicc
face_index,
face_map,
horizon,
polygon,
clipped,
pnormal,
pdist,
idx1,
idx2,
n1,
n2,
endvert,
face1,
face2,
)
dist_out[0] = dist
@@ -170,6 +192,17 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int, multicc
face_index = wp.array(shape=(6 * iterations,), dtype=int)
face_map = wp.array(shape=(6 * iterations,), dtype=int)
horizon = wp.array(shape=(6 * iterations,), dtype=int)
polygon = wp.array(shape=(150,), dtype=wp.vec3)
clipped = wp.array(shape=(150,), dtype=wp.vec3)
pnormal = wp.array(shape=(150,), dtype=wp.vec3)
pdist = wp.array(shape=(150,), dtype=float)
idx1 = wp.array(shape=(150,), dtype=int)
idx2 = wp.array(shape=(150,), dtype=int)
n1 = wp.array(shape=(150,), dtype=wp.vec3)
n2 = wp.array(shape=(150,), dtype=wp.vec3)
endvert = wp.array(shape=(150,), dtype=wp.vec3)
face1 = wp.array(shape=(150,), dtype=wp.vec3)
face2 = wp.array(shape=(150,), dtype=wp.vec3)
dist_out = wp.array(shape=(1,), dtype=float)
ncon_out = wp.array(shape=(1,), dtype=int)
pos_out = wp.array(shape=(2,), dtype=wp.vec3)
@@ -208,6 +241,17 @@ def _geom_dist(m: Model, d: Data, gid1: int, gid2: int, iterations: int, multicc
face_index,
face_map,
horizon,
polygon,
clipped,
pnormal,
pdist,
idx1,
idx2,
n1,
n2,
endvert,
face1,
face2,
],
outputs=[
dist_out,
@@ -18,277 +18,56 @@ from typing import Tuple
import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
@wp.func
def _hfield_subgrid(
# In:
nrow: int,
ncol: int,
size: wp.vec4,
xmax: float,
xmin: float,
ymax: float,
ymin: float,
) -> Tuple[int, int, int, int]:
"""Returns height field subgrid that overlaps with geom AABB.
Args:
nrow: height field number of rows
ncol: height field number of columns
size: height field size
xmax: geom maximum x position
xmin: geom minimum x position
ymax: geom maximum y position
ymin: geom minimum y position
Returns:
grid coordinate bounds
"""
# grid resolution
x_scale = 0.5 * float(ncol - 1) / size[0]
y_scale = 0.5 * float(nrow - 1) / size[1]
# subgrid
cmin = wp.max(0, int(wp.floor((xmin + size[0]) * x_scale)))
cmax = wp.min(ncol - 1, int(wp.ceil((xmax + size[0]) * x_scale)))
rmin = wp.max(0, int(wp.floor((ymin + size[1]) * y_scale)))
rmax = wp.min(nrow - 1, int(wp.ceil((ymax + size[1]) * y_scale)))
return cmin, rmin, cmax, rmax
@wp.func
def hfield_triangle_prism(
def hfield_filter(
# Model:
geom_dataid: wp.array(dtype=int),
hfield_adr: wp.array(dtype=int),
hfield_nrow: wp.array(dtype=int),
hfield_ncol: wp.array(dtype=int),
hfield_size: wp.array(dtype=wp.vec4),
hfield_data: wp.array(dtype=float),
# In:
hfieldid: int,
hftri_index: int,
) -> wp.mat33:
"""Returns triangular prism vertex information in compressed representation.
Args:
geom_dataid: geom data ids
hfield_adr: address for height field
hfield_nrow: height field number of rows
hfield_ncol: height field number of columns
hfield_size: height field sizes
hfield_data: height field data
hfieldid: height field geom id
hftri_index: height field triangle index
Returns:
triangular prism vertex information (compressed)
"""
# https://mujoco.readthedocs.io/en/stable/XMLreference.html#asset-hfield
# get heightfield dimensions
dataid = geom_dataid[hfieldid]
if dataid < 0 or hftri_index < 0:
return wp.mat33(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0)
nrow = hfield_nrow[dataid]
ncol = hfield_ncol[dataid]
size = hfield_size[dataid] # (x, y, z_top, z_bottom)
# calculate which triangle in the grid
row = (hftri_index // 2) // (ncol - 1)
col = (hftri_index // 2) % (ncol - 1)
# calculate vertices in 2D grid
x_scale = 2.0 * size[0] / float(ncol - 1)
y_scale = 2.0 * size[1] / float(nrow - 1)
# grid coordinates (i, j) for triangle corners
i0 = col
j0 = row
i1 = i0 + 1
j1 = j0 + 1
# convert grid coordinates to local space x, y coordinates
x0 = float(i0) * x_scale - size[0]
y0 = float(j0) * y_scale - size[1]
x1 = float(i1) * x_scale - size[0]
y1 = float(j1) * y_scale - size[1]
# get height values at corners from hfield_data
base_addr = hfield_adr[dataid]
z00 = hfield_data[base_addr + j0 * ncol + i0]
z01 = hfield_data[base_addr + j1 * ncol + i0]
z10 = hfield_data[base_addr + j0 * ncol + i1]
z11 = hfield_data[base_addr + j1 * ncol + i1]
# scale heights from range [0, 1] to [0, z_top]
z_top = size[2]
z00 = z00 * z_top
z01 = z01 * z_top
z10 = z10 * z_top
z11 = z11 * z_top
x2 = wp.where(hftri_index % 2, 1.0, 0.0)
y2 = wp.where(hftri_index % 2, z10, z01)
z22 = -size[3]
# compress 6 prism vertices into 3x3 matrix, see hfield_prism_vertex for details
return wp.mat33(x0, y0, z00,
x1, y1, z11,
x2, y2, z22) # fmt: off
@wp.func
def hfield_prism_vertex(prism: wp.mat33, vert_index: int) -> wp.vec3:
"""Extracts vertices from a compressed triangular prism representation.
The compression scheme stores a 6-vertex triangular prism using a 3x3 matrix:
- prism[0] = First vertex (x,y,z) - corner (i,j)
- prism[1] = Second vertex (x,y,z) - corner (i+1,j+1)
- prism[2,0] = Triangle type flag: 0 for even triangle (using corner (i,j+1)),
non-zero for odd triangle (using corner (i+1,j))
- prism[2,1] = Z-coordinate of the third vertex
- prism[2,2] = Z-coordinate used for all bottom vertices (common z)
In this way, we can reconstruct all 6 vertices of the prism by reusing
coordinates from the stored vertices.
Args:
prism: 3x3 compressed representation of a triangular prism
vert_index: index of vertex to extract (0-5)
Returns:
3D coordinates of the requested vertex
"""
if vert_index == 0 or vert_index == 1:
return prism[vert_index] # first two vertices stored directly
if vert_index == 2: # third vertex
if prism[2][0] == 0: # even triangle (i, j+1)
return wp.vec3(prism[0][0], prism[1][1], prism[2][1])
else: # odd triangle (i+1, j)
return wp.vec3(prism[1][0], prism[0][1], prism[2][1])
if vert_index == 3 or vert_index == 4: # bottom vertices below 0 and 1
return wp.vec3(prism[vert_index - 3][0], prism[vert_index - 3][1], prism[2][2])
if vert_index == 5: # bottom vertex below 2
if prism[2][0] == 0: # even triangle
return wp.vec3(prism[0][0], prism[1][1], prism[2][2])
else: # odd triangle
return wp.vec3(prism[1][0], prism[0][1], prism[2][2])
@wp.kernel
def _hfield_midphase(
# Model:
geom_type: wp.array(dtype=int),
geom_dataid: wp.array(dtype=int),
geom_aabb: wp.array2d(dtype=wp.vec3),
geom_rbound: wp.array2d(dtype=float),
geom_margin: wp.array2d(dtype=float),
hfield_nrow: wp.array(dtype=int),
hfield_ncol: wp.array(dtype=int),
hfield_size: wp.array(dtype=wp.vec4),
# Data in:
nconmax_in: int,
geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.array2d(dtype=wp.mat33),
collision_pair_in: wp.array(dtype=wp.vec2i),
collision_hftri_index_in: wp.array(dtype=int),
collision_pairid_in: wp.array(dtype=int),
collision_worldid_in: wp.array(dtype=int),
# Data out:
collision_pair_out: wp.array(dtype=wp.vec2i),
collision_hftri_index_out: wp.array(dtype=int),
collision_pairid_out: wp.array(dtype=int),
collision_worldid_out: wp.array(dtype=int),
ncollision_out: wp.array(dtype=int),
):
"""Midphase collision detection for heightfield triangles with other geoms.
# In:
worldid: int,
g1: int,
g2: int,
) -> Tuple[bool, float, float, float, float, float, float]:
"""Filter for height field collisions.
This kernel processes collision pairs where one geom is a heightfield (identified by
collision_hftri_index_in[pairid] == -1) and expands them into multiple collision pairs,
one for each potentially colliding triangle.
Args:
geom_type: geom type
geom_dataid: geom data id
geom_rbound: geom bounding sphere radius
geom_margin: geom margin
hfield_nrow: height field number of rows
hfield_ncol: height field number of columns
hfield_size: height field size
nconmax_in: maximum number of contacts
geom_xpos_in: geom position
geom_xmat_in: geom orientation
collision_pair_in: collision pair
collision_hftri_index_in: triangle indices, -1 for height field pair
collision_pairid_in: collision pair id from broadphase
collision_worldid_in: collision world id from broadphase
collision_pair_out: collision pair from midphase
collision_hftri_index_out: triangle indices from midphase
collision_pairid_out: collision pair id from midphase
collision_worldid_out: collision world id from midphase
ncollision_out: number of collisions from broadphase and midphase
See MuJoCo mjc_ConvexHField.
"""
pairid = wp.tid()
# only process pairs that are marked for height field collision (-1)
if collision_hftri_index_in[pairid] != -1:
return
# collision pair info
worldid = collision_worldid_in[pairid]
pair_id = collision_pairid_in[pairid]
pair = collision_pair_in[pairid]
g1 = pair[0]
g2 = pair[1]
hfieldid = g1
geomid = g2
# SHOULD NOT OCCUR: if the first geom is not a heightfield, swap
if geom_type[g1] != int(GeomType.HFIELD.value):
hfieldid = g2
geomid = g1
# height field info
hfdataid = geom_dataid[hfieldid]
hfdataid = geom_dataid[g1]
size1 = hfield_size[hfdataid]
pos1 = geom_xpos_in[worldid, hfieldid]
mat1 = geom_xmat_in[worldid, hfieldid]
pos1 = geom_xpos_in[worldid, g1]
mat1 = geom_xmat_in[worldid, g1]
mat1T = wp.transpose(mat1)
# geom info
pos2 = geom_xpos_in[worldid, geomid]
pos2 = geom_xpos_in[worldid, g2]
pos = mat1T @ (pos2 - pos1)
r2 = geom_rbound[worldid, geomid]
r2 = geom_rbound[worldid, g2]
# TODO(team): margin?
margin = wp.max(geom_margin[worldid, hfieldid], geom_margin[worldid, geomid])
margin = wp.max(geom_margin[worldid, g1], geom_margin[worldid, g2])
# box-sphere test: horizontal plane
for i in range(2):
if (size1[i] < pos[i] - r2 - margin) or (-size1[i] > pos[i] + r2 + margin):
return
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
# box-sphere test: vertical direction
if size1[2] < pos[2] - r2 - margin: # up
return
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
if -size1[3] > pos[2] + r2 + margin: # down
return
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
mat2 = geom_xmat_in[worldid, geomid]
mat2 = geom_xmat_in[worldid, g2]
mat = mat1T @ mat2
# aabb for geom in height field frame
@@ -299,8 +78,8 @@ def _hfield_midphase(
ymin = MJ_MAXVAL
zmin = MJ_MAXVAL
center2 = geom_aabb[geomid, 0]
size2 = geom_aabb[geomid, 1]
center2 = geom_aabb[g2, 0]
size2 = geom_aabb[g2, 1]
pos += mat1T @ center2
@@ -341,61 +120,6 @@ def _hfield_midphase(
or (zmin - margin > size1[2])
or (zmax + margin < -size1[3])
):
return
# height field subgrid
nrow = hfield_nrow[hfieldid]
ncol = hfield_ncol[hfieldid]
size = hfield_size[hfieldid]
cmin, rmin, cmax, rmax = _hfield_subgrid(nrow, ncol, size, xmax, xmin, ymax, ymin)
# loop over subgrid triangles
for r in range(rmin, rmax):
for c in range(cmin, cmax):
# add both triangles from this cell
for i in range(2):
if r == rmin and c == cmin and i == 0:
# reuse the initial pair for the 1st triangle
new_pairid = pairid
else:
# create a new pair
new_pairid = wp.atomic_add(ncollision_out, 0, 1)
if new_pairid >= nconmax_in:
return
collision_pair_out[new_pairid] = pair
collision_hftri_index_out[new_pairid] = 2 * (r * (ncol - 1) + c) + i
collision_pairid_out[new_pairid] = pair_id
collision_worldid_out[new_pairid] = worldid
def hfield_midphase(m: Model, d: Data):
"""Midphase collision detection for height field triangles with other geoms.
Processes collision pairs from the broadphase where one geom is a height field and expands
them into multiple collision pairs, one for each potentially colliding triangle. The
function directly writes to the same collision buffers used by _add_geom_pair.
"""
wp.launch(
kernel=_hfield_midphase,
dim=d.nconmax,
inputs=[
m.geom_type,
m.geom_dataid,
m.geom_aabb,
m.geom_rbound,
m.geom_margin,
m.hfield_nrow,
m.hfield_ncol,
m.hfield_size,
d.nconmax,
d.geom_xpos,
d.geom_xmat,
d.collision_pair,
d.collision_hftri_index,
d.collision_pairid,
d.collision_worldid,
],
outputs=[d.collision_pair, d.collision_hftri_index, d.collision_pairid, d.collision_worldid, d.ncollision],
)
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
else:
return False, xmin, xmax, ymin, ymax, zmin, zmax
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+475 -176
View File
@@ -18,17 +18,22 @@ from typing import Tuple
import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src import math
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import _geom
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact
from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame
from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_mesh
from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
from mujoco.mjx.third_party.mujoco_warp._src.types import vec8f
from mujoco.mjx.third_party.mujoco_warp._src.types import vec8i
from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
wp.set_module_options({"enable_backward": False})
@wp.struct
class OptimizationParams:
@@ -44,24 +49,79 @@ class AABB:
max: wp.vec3
@wp.struct
class VolumeData:
center: wp.vec3
half_size: wp.vec3
oct_aabb: wp.array2d(dtype=wp.vec3)
oct_child: wp.array(dtype=vec8i)
oct_coeff: wp.array(dtype=vec8f)
valid: bool = False
@wp.struct
class MeshData:
nmeshface: int
mesh_vertadr: wp.array(dtype=int)
mesh_vert: wp.array(dtype=wp.vec3)
mesh_faceadr: wp.array(dtype=int)
mesh_face: wp.array(dtype=wp.vec3i)
data_id: int
data_id: int
pos: wp.vec3
mat: wp.mat33
pnt: wp.vec3
vec: wp.vec3
valid: bool = False
@wp.func
def get_sdf_params(
# Model:
oct_aabb: wp.array2d(dtype=wp.vec3),
oct_child: wp.array(dtype=vec8i),
oct_coeff: wp.array(dtype=vec8f),
plugin: wp.array(dtype=int),
plugin_attr: wp.array(dtype=wp.vec3f),
# In:
g_type: int,
g_size: wp.vec3,
plugin_id: int,
mesh_id: int,
) -> Tuple[wp.vec3, int, VolumeData, MeshData]:
attributes = g_size
plugin_index = -1
volume_data = VolumeData()
if g_type == int(GeomType.SDF.value) and plugin_id != -1:
attributes = plugin_attr[plugin_id]
plugin_index = plugin[plugin_id]
elif g_type == int(GeomType.SDF.value) and mesh_id != -1:
volume_data.center = oct_aabb[mesh_id, 0]
volume_data.half_size = oct_aabb[mesh_id, 1]
volume_data.oct_aabb = oct_aabb
volume_data.oct_child = oct_child
volume_data.oct_coeff = oct_coeff
volume_data.valid = True
return attributes, plugin_index, volume_data, MeshData()
@wp.func
def transform_aabb(aabb_pos: wp.vec3, aabb_size: wp.vec3, pos: wp.vec3, ori: wp.mat33) -> AABB:
aabb = AABB()
aabb.max = wp.vec3(-1000000000.0, -1000000000.0, -1000000000.0)
aabb.min = wp.vec3(1000000000.0, 1000000000.0, 1000000000.0)
for i in range(8):
vec = wp.vec3(
aabb_size.x * (1.0 if (i & 1) else -1.0),
aabb_size.y * (1.0 if (i & 2) else -1.0),
aabb_size.z * (1.0 if (i & 4) else -1.0),
)
frame_vec = ori * (vec + aabb_pos) + pos
aabb.min = wp.min(aabb.min, frame_vec)
aabb.max = wp.max(aabb.max, frame_vec)
return aabb
@@ -134,13 +194,11 @@ def grad_box(p: wp.vec3, size: wp.vec3) -> wp.vec3:
@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])
b = wp.vec3(a[0] / size[0], a[1] / size[1], a[2] / size[2])
k0 = wp.length(a)
k1 = wp.length(b)
invK0 = 1.0 / k0
invK1 = 1.0 / k1
gk0 = b * invK0
gk1 = wp.vec3(
b[0] * invK1 / (size[0] * size[0]),
@@ -149,7 +207,6 @@ def grad_ellipsoid(p: wp.vec3, size: wp.vec3) -> wp.vec3:
)
df_dk0 = (2.0 * k0 - 1.0) * invK1
df_dk1 = k0 * (k0 - 1.0) * invK1 * invK1
raw_grad = gk0 * df_dk0 - gk1 * df_dk1
return raw_grad / wp.length(raw_grad)
@@ -167,7 +224,127 @@ 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:
def find_oct(
oct_aabb: wp.array2d(dtype=wp.vec3), oct_child: wp.array(dtype=vec8i), p: wp.vec3, grad: bool
) -> Tuple[int, Tuple[vec8f, vec8f, vec8f]]:
stack = int(0)
niter = int(100)
rx = vec8f(0.0)
ry = vec8f(0.0)
rz = vec8f(0.0)
while niter > 0:
niter -= 1
node = stack
if node == -1:
wp.printf("ERROR: Invalid node number\n")
return -1, (rx, ry, rz)
vmin = oct_aabb[node, 0] - oct_aabb[node, 1]
vmax = oct_aabb[node, 0] + oct_aabb[node, 1]
coord = wp.cw_div(p - vmin, vmax - vmin)
# check if the node is a leaf
if (
oct_child[node][0] == -1
and oct_child[node][1] == -1
and oct_child[node][2] == -1
and oct_child[node][3] == -1
and oct_child[node][4] == -1
and oct_child[node][5] == -1
and oct_child[node][6] == -1
and oct_child[node][7] == -1
):
for j in range(8):
if not grad:
rx[j] = (
(coord[0] if j & 1 else 1.0 - coord[0])
* (coord[1] if j & 2 else 1.0 - coord[1])
* (coord[2] if j & 4 else 1.0 - coord[2])
)
else:
rx[j] = (1.0 if j & 1 else -1.0) * (coord[1] if j & 2 else 1.0 - coord[1]) * (coord[2] if j & 4 else 1.0 - coord[2])
ry[j] = (coord[0] if j & 1 else 1.0 - coord[0]) * (1.0 if j & 2 else -1.0) * (coord[2] if j & 4 else 1.0 - coord[2])
rz[j] = (coord[0] if j & 1 else 1.0 - coord[0]) * (coord[1] if j & 2 else 1.0 - coord[1]) * (1.0 if j & 4 else -1.0)
return node, (rx, ry, rz)
# compute which of 8 children to visit next
x = 1 if coord[0] < 0.5 else 0
y = 1 if coord[1] < 0.5 else 0
z = 1 if coord[2] < 0.5 else 0
stack = oct_child[node][4 * z + 2 * y + x]
wp.print("ERROR: Node not found\n")
return -1, (rx, ry, rz)
@wp.func
def box_project(center: wp.vec3, half_size: wp.vec3, xyz: wp.vec3) -> Tuple[float, wp.vec3]:
r = xyz - center
q = wp.vec3(wp.abs(r[0]) - half_size[0], wp.abs(r[1]) - half_size[1], wp.abs(r[2]) - half_size[2])
if q[0] <= 0.0 and q[1] <= 0.0 and q[2] <= 0.0:
return 0.0, xyz
else:
dist_sqr = 0.0
eps = 1e-4
point = wp.vec3(xyz[0], xyz[1], xyz[2])
if q[0] >= 0.0:
dist_sqr += q[0] * q[0]
if r[0] > 0.0:
point = wp.vec3(point[0] - (q[0] + eps), point[1], point[2])
else:
point = wp.vec3(point[0] + (q[0] + eps), point[1], point[2])
if q[1] >= 0.0:
dist_sqr += q[1] * q[1]
if r[1] > 0.0:
point = wp.vec3(point[0], point[1] - (q[1] + eps), point[2])
else:
point = wp.vec3(point[0], point[1] + (q[1] + eps), point[2])
if q[2] >= 0.0:
dist_sqr += q[2] * q[2]
if r[2] > 0.0:
point = wp.vec3(point[0], point[1], point[2] - (q[2] + eps))
else:
point = wp.vec3(point[0], point[1], point[2] + (q[2] + eps))
return wp.sqrt(dist_sqr), point
@wp.func
def sample_volume_sdf(xyz: wp.vec3, volume_data: VolumeData) -> float:
dist0, point = box_project(volume_data.center, volume_data.half_size, xyz)
node, weights = find_oct(volume_data.oct_aabb, volume_data.oct_child, point, grad=False)
return dist0 + wp.dot(weights[0], volume_data.oct_coeff[node])
@wp.func
def sample_volume_grad(xyz: wp.vec3, volume_data: VolumeData) -> wp.vec3:
dist0, point = box_project(volume_data.center, volume_data.half_size, xyz)
if dist0 > 0:
h = 1e-4
dx = wp.vec3(h, 0.0, 0.0)
dy = wp.vec3(0.0, h, 0.0)
dz = wp.vec3(0.0, 0.0, h)
f = sample_volume_sdf(xyz, volume_data)
grad_x = (sample_volume_sdf(xyz + dx, volume_data) - f) / h
grad_y = (sample_volume_sdf(xyz + dy, volume_data) - f) / h
grad_z = (sample_volume_sdf(xyz + dz, volume_data) - f) / h
return wp.vec3(grad_x, grad_y, grad_z)
node, weights = find_oct(volume_data.oct_aabb, volume_data.oct_child, point, grad=True)
grad_x = wp.dot(weights[0], volume_data.oct_coeff[node])
grad_y = wp.dot(weights[1], volume_data.oct_coeff[node])
grad_z = wp.dot(weights[2], volume_data.oct_coeff[node])
return wp.vec3(grad_x, grad_y, grad_z)
@wp.func
def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: VolumeData, mesh_data: MeshData) -> float:
if type == int(GeomType.PLANE.value):
return p[2]
elif type == int(GeomType.SPHERE.value):
@@ -176,14 +353,46 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float:
return box(p, attr)
elif type == int(GeomType.ELLIPSOID.value):
return ellipsoid(p, attr)
elif type == int(GeomType.MESH.value) and mesh_data.valid:
mesh_data.pnt = p
mesh_data.vec = -wp.normalize(p)
dist = ray_mesh(
mesh_data.nmeshface,
mesh_data.mesh_vertadr,
mesh_data.mesh_vert,
mesh_data.mesh_faceadr,
mesh_data.mesh_face,
mesh_data.data_id,
mesh_data.pos,
mesh_data.mat,
mesh_data.pnt,
mesh_data.vec,
)
if dist > wp.norm_l2(p):
return -ray_mesh(
mesh_data.nmeshface,
mesh_data.mesh_vertadr,
mesh_data.mesh_vert,
mesh_data.mesh_faceadr,
mesh_data.mesh_face,
mesh_data.data_id,
mesh_data.pos,
mesh_data.mat,
mesh_data.pnt,
-mesh_data.vec,
)
return dist
elif type == int(GeomType.SDF.value):
return user_sdf(p, attr, sdf_type)
if sdf_type == -1:
return sample_volume_sdf(p, volume_data)
else:
return user_sdf(p, attr, sdf_type)
wp.printf("ERROR: SDF type not implemented\n")
return 0.0
@wp.func
def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3:
def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: VolumeData, mesh_data: MeshData) -> wp.vec3:
if type == int(GeomType.PLANE.value):
grad = wp.vec3(0.0, 0.0, 1.0)
return grad
@@ -193,18 +402,53 @@ def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3:
return grad_box(p, attr)
elif type == int(GeomType.ELLIPSOID.value):
return grad_ellipsoid(p, attr)
elif type == int(GeomType.MESH.value) and mesh_data.valid:
mesh_data.pnt = p
mesh_data.vec = -wp.normalize(p)
dist = ray_mesh(
mesh_data.nmeshface,
mesh_data.mesh_vertadr,
mesh_data.mesh_vert,
mesh_data.mesh_faceadr,
mesh_data.mesh_face,
mesh_data.data_id,
mesh_data.pos,
mesh_data.mat,
mesh_data.pnt,
mesh_data.vec,
)
if dist > wp.norm_l2(p):
return wp.vec3(1.0)
else:
return wp.vec3(-1.0)
elif type == int(GeomType.SDF.value):
return user_sdf_grad(p, attr, sdf_type)
if sdf_type == -1:
return sample_volume_grad(p, volume_data)
else:
return user_sdf_grad(p, attr, sdf_type)
wp.printf("ERROR: SDF grad type not implemented\n")
return wp.vec3(0.0)
@wp.func
def clearance(
type1: int, p1: wp.vec3, p2: wp.vec3, s1: wp.vec3, s2: wp.vec3, sdf_type1: int, sdf_type2: int, sfd_intersection: bool
# In:
type1: int,
p1: wp.vec3,
p2: wp.vec3,
s1: wp.vec3,
s2: wp.vec3,
sdf_type1: int,
sdf_type2: int,
sfd_intersection: bool,
volume_data1: VolumeData,
volume_data2: VolumeData,
mesh_data1: MeshData,
mesh_data2: MeshData,
) -> float:
sdf1 = sdf(type1, p1, s1, sdf_type1)
sdf2 = sdf(int(GeomType.SDF.value), p2, s2, sdf_type2)
sdf1 = sdf(type1, p1, s1, sdf_type1, volume_data1, mesh_data1)
sdf2 = sdf(int(GeomType.SDF.value), p2, s2, sdf_type2, volume_data2, mesh_data2)
if sfd_intersection:
return wp.max(sdf1, sdf2)
else:
@@ -213,13 +457,24 @@ def clearance(
@wp.func
def compute_grad(
type1: int, p1: wp.vec3, p2: wp.vec3, params: OptimizationParams, sdf_type1: int, sdf_type2: int, sfd_intersection: bool
# In:
type1: int,
p1: wp.vec3,
p2: wp.vec3,
params: OptimizationParams,
sdf_type1: int,
sdf_type2: int,
sfd_intersection: bool,
volume_data1: VolumeData,
volume_data2: VolumeData,
mesh_data1: MeshData,
mesh_data2: MeshData,
) -> wp.vec3:
A = sdf(type1, p1, params.attr1, sdf_type1)
B = sdf(int(GeomType.SDF.value), p2, params.attr2, sdf_type2)
grad1 = sdf_grad(type1, p1, params.attr1, sdf_type1)
grad2 = sdf_grad(int(GeomType.SDF.value), p2, params.attr2, sdf_type2)
grad1_transformed = params.rel_mat * grad1
A = sdf(type1, p1, params.attr1, sdf_type1, volume_data1, mesh_data1)
B = sdf(int(GeomType.SDF.value), p2, params.attr2, sdf_type2, volume_data2, mesh_data2)
grad1 = sdf_grad(type1, p1, params.attr1, sdf_type1, volume_data1, mesh_data1)
grad2 = sdf_grad(int(GeomType.SDF.value), p2, params.attr2, sdf_type2, volume_data2, mesh_data2)
grad1_transformed = wp.transpose(params.rel_mat) * grad1
if sfd_intersection:
if A > B:
return grad1_transformed
@@ -239,33 +494,67 @@ def compute_grad(
@wp.func
def gradient_step(
type1: int, x: wp.vec3, params: OptimizationParams, sdf_type1: int, sdf_type2: int, niter: int, sfd_intersection: bool
# In:
type1: int,
x: wp.vec3,
params: OptimizationParams,
sdf_type1: int,
sdf_type2: int,
niter: int,
sfd_intersection: bool,
volume_data1: VolumeData,
volume_data2: VolumeData,
mesh_data1: MeshData,
mesh_data2: MeshData,
) -> Tuple[float, wp.vec3]:
amin = 1e-4
rho = 0.5
c = 0.1
dist = float(1e10)
for _ in range(niter):
for i in range(niter):
alpha = float(2.0)
x2 = wp.vec3(x[0], x[1], x[2])
x1 = params.rel_mat * x2 + params.rel_pos
grad = compute_grad(type1, x1, x2, params, sdf_type1, sdf_type2, sfd_intersection)
dist0 = clearance(type1, x1, x, params.attr1, params.attr2, sdf_type1, sdf_type2, sfd_intersection)
grad = compute_grad(
type1, x1, x2, params, sdf_type1, sdf_type2, sfd_intersection, volume_data1, volume_data2, mesh_data1, mesh_data2
)
dist0 = clearance(
type1,
x1,
x,
params.attr1,
params.attr2,
sdf_type1,
sdf_type2,
sfd_intersection,
volume_data1,
volume_data2,
mesh_data1,
mesh_data2,
)
grad_dot = wp.dot(grad, grad)
if grad_dot < 1e-12:
return dist0, x
wolfe = -c * alpha * grad_dot
while True:
alpha *= rho
wolfe *= rho
x = x2 - grad * alpha
x1 = params.rel_mat * x + params.rel_pos
dist = clearance(type1, x1, x, params.attr1, params.attr2, sdf_type1, sdf_type2, sfd_intersection)
dist = clearance(
type1,
x1,
x,
params.attr1,
params.attr2,
sdf_type1,
sdf_type2,
sfd_intersection,
volume_data1,
volume_data2,
mesh_data1,
mesh_data2,
)
if alpha <= amin or (dist - dist0) <= wolfe:
break
if dist > dist0:
@@ -287,29 +576,26 @@ def gradient_descent(
sdf_type1: int,
sdf_type2: int,
sdf_iterations: int,
volume_data1: VolumeData,
volume_data2: VolumeData,
mesh_data1: MeshData,
mesh_data2: MeshData,
) -> Tuple[float, wp.vec3, wp.vec3]:
params = OptimizationParams()
params.rel_mat = wp.transpose(rot1) * rot2
params.rel_pos = wp.transpose(rot1) * (pos2 - pos1)
params.attr1 = attr1
params.attr2 = attr2
# Collision phase (10 iterations, sfd_intersection=False)
dist, x = gradient_step(type1, x0_initial, params, sdf_type1, sdf_type2, sdf_iterations, False)
# Intersection phase (1 iteration, sfd_intersection=True)
dist, x = gradient_step(type1, x, params, sdf_type1, sdf_type2, 1, True)
# Midsurface calculation
dist, x = gradient_step(
type1, x0_initial, params, sdf_type1, sdf_type2, sdf_iterations, False, volume_data1, volume_data2, mesh_data1, mesh_data2
)
dist, x = gradient_step(type1, x, params, sdf_type1, sdf_type2, 1, True, volume_data1, volume_data2, mesh_data1, mesh_data2)
x_1 = params.rel_mat * x + params.rel_pos
grad1 = sdf_grad(type1, x_1, params.attr1, sdf_type1)
grad1 = sdf_grad(type1, x_1, params.attr1, sdf_type1, volume_data1, mesh_data1)
grad1 = wp.transpose(params.rel_mat) * grad1
grad1 = wp.normalize(grad1)
grad2 = sdf_grad(int(GeomType.SDF.value), x, params.attr2, sdf_type2)
grad2 = sdf_grad(int(GeomType.SDF.value), x, params.attr2, sdf_type2, volume_data2, mesh_data2)
grad2 = wp.normalize(grad2)
n = grad1 - grad2
n = wp.normalize(n)
pos = rot2 * x + pos2
@@ -321,6 +607,7 @@ def gradient_descent(
@wp.kernel
def _sdf_narrowphase(
# Model:
nmeshface: int,
geom_type: wp.array(dtype=int),
geom_condim: wp.array(dtype=int),
geom_dataid: wp.array(dtype=int),
@@ -330,8 +617,6 @@ def _sdf_narrowphase(
geom_solimp: wp.array2d(dtype=vec5),
geom_size: wp.array2d(dtype=wp.vec3),
geom_aabb: wp.array2d(dtype=wp.vec3),
geom_pos: wp.array2d(dtype=wp.vec3),
geom_quat: wp.array2d(dtype=wp.quat),
geom_friction: wp.array2d(dtype=wp.vec3),
geom_margin: wp.array2d(dtype=float),
geom_gap: wp.array2d(dtype=float),
@@ -343,6 +628,8 @@ def _sdf_narrowphase(
mesh_vertadr: wp.array(dtype=int),
mesh_vertnum: wp.array(dtype=int),
mesh_vert: 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_polynum: wp.array(dtype=int),
@@ -354,6 +641,9 @@ def _sdf_narrowphase(
mesh_polymapadr: wp.array(dtype=int),
mesh_polymapnum: wp.array(dtype=int),
mesh_polymap: wp.array(dtype=int),
oct_aabb: wp.array2d(dtype=wp.vec3),
oct_child: wp.array(dtype=vec8i),
oct_coeff: wp.array(dtype=vec8f),
pair_dim: wp.array(dtype=int),
pair_solref: wp.array2d(dtype=wp.vec2),
pair_solreffriction: wp.array2d(dtype=wp.vec2),
@@ -370,7 +660,6 @@ def _sdf_narrowphase(
geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.array2d(dtype=wp.mat33),
collision_pair_in: wp.array(dtype=wp.vec2i),
collision_hftri_index_in: wp.array(dtype=int),
collision_pairid_in: wp.array(dtype=int),
collision_worldid_in: wp.array(dtype=int),
ncollision_in: wp.array(dtype=int),
@@ -391,20 +680,17 @@ def _sdf_narrowphase(
contact_geom_out: wp.array(dtype=wp.vec2i),
contact_worldid_out: wp.array(dtype=int),
):
tid = wp.tid()
if tid >= ncollision_in[0]:
i, contact_tid = wp.tid()
if i >= sdf_initpoints:
return
geoms = collision_pair_in[tid]
if contact_tid >= ncollision_in[0]:
return
geoms = collision_pair_in[contact_tid]
g2 = geoms[1]
type2 = geom_type[g2]
if type2 != int(GeomType.SDF.value):
return
worldid = collision_worldid_in[tid]
worldid = collision_worldid_in[contact_tid]
_, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
geom_condim,
geom_priority,
@@ -423,79 +709,62 @@ def _sdf_narrowphase(
pair_friction,
collision_pair_in,
collision_pairid_in,
tid,
contact_tid,
worldid,
)
g1 = geoms[0]
hftri_index = collision_hftri_index_in[tid]
geom1 = _geom(
geom_type,
geom_dataid,
geom_size,
hfield_adr,
hfield_nrow,
hfield_ncol,
hfield_size,
hfield_data,
mesh_vertadr,
mesh_vertnum,
mesh_vert,
mesh_graphadr,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in,
geom_xmat_in,
worldid,
g1,
hftri_index,
)
geom2 = _geom(
geom_type,
geom_dataid,
geom_size,
hfield_adr,
hfield_nrow,
hfield_ncol,
hfield_size,
hfield_data,
mesh_vertadr,
mesh_vertnum,
mesh_vert,
mesh_graphadr,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in,
geom_xmat_in,
worldid,
g2,
hftri_index,
)
type1 = geom_type[g1]
geom1_dataid = geom_dataid[g1]
geom1 = geom(
type1,
geom1_dataid,
geom_size[worldid, g1],
mesh_vertadr[geom1_dataid],
mesh_vertnum[geom1_dataid],
mesh_vert,
mesh_graphadr[geom1_dataid],
mesh_graph,
mesh_polynum[geom1_dataid],
mesh_polyadr[geom1_dataid],
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in[worldid, g1],
geom_xmat_in[worldid, g1],
)
geom2_dataid = geom_dataid[g2]
geom2 = geom(
type2,
geom2_dataid,
geom_size[worldid, g2],
mesh_vertadr[geom2_dataid],
mesh_vertnum[geom2_dataid],
mesh_vert,
mesh_graphadr[geom2_dataid],
mesh_graph,
mesh_polynum[geom2_dataid],
mesh_polyadr[geom2_dataid],
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in[worldid, g2],
geom_xmat_in[worldid, g2],
)
g1_plugin = geom_plugin_index[g1]
g2_plugin = geom_plugin_index[g2]
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]
identity = wp.identity(3, dtype=float)
@@ -503,7 +772,6 @@ def _sdf_narrowphase(
aabb_pos = geom_aabb[g2, 0]
aabb_size = geom_aabb[g2, 1]
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)
@@ -513,69 +781,98 @@ def _sdf_narrowphase(
pos1 = geom1.pos
rot1 = geom1.rot
if type1 == int(GeomType.SDF.value):
attr1 = plugin_attr[g1_plugin]
g1_plugin_id = plugin[g1_plugin]
else:
attr1 = geom1.size
g1_plugin_id = -1
attr1, g1_plugin_id, volume_data1, mesh_data1 = get_sdf_params(
oct_aabb, oct_child, oct_coeff, plugin, plugin_attr, type1, geom1.size, g1_plugin, geom_dataid[g1]
)
if g2_plugin != -1:
attr2 = plugin_attr[g2_plugin]
g2_plugin_id = plugin[g2_plugin]
else:
attr2 = geom2.size
g2_plugin_id = -1
attr2, g2_plugin_id, volume_data2, mesh_data2 = get_sdf_params(
oct_aabb, oct_child, oct_coeff, plugin, plugin_attr, type2, geom2.size, g2_plugin, geom_dataid[g2]
)
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),
aabb_intersection.min[1] + (aabb_intersection.max[1] - aabb_intersection.min[1]) * halton(i, 3),
aabb_intersection.min[2] + (aabb_intersection.max[2] - aabb_intersection.min[2]) * halton(i, 5),
)
mesh_data1.nmeshface = nmeshface
mesh_data1.mesh_vertadr = mesh_vertadr
mesh_data1.mesh_vert = mesh_vert
mesh_data1.mesh_faceadr = mesh_faceadr
mesh_data1.mesh_face = mesh_face
mesh_data1.data_id = geom_dataid[g1]
mesh_data1.pos = geom1.pos
mesh_data1.mat = geom1.rot
mesh_data1.pnt = wp.vec3(-1.0)
mesh_data1.vec = wp.vec3(0.0)
mesh_data1.valid = True
x = geom1.rot * x_g2 + geom1.pos
x0_initial = wp.transpose(rot2) * (x - pos2)
mesh_data2.nmeshface = nmeshface
mesh_data2.mesh_vertadr = mesh_vertadr
mesh_data2.mesh_vert = mesh_vert
mesh_data2.mesh_faceadr = mesh_faceadr
mesh_data2.mesh_face = mesh_face
mesh_data2.data_id = geom_dataid[g2]
mesh_data2.pos = geom2.pos
mesh_data2.mat = geom2.rot
mesh_data2.pnt = wp.vec3(-1.0)
mesh_data2.vec = wp.vec3(0.0)
mesh_data2.valid = True
dist, pos, n = gradient_descent(
type1, x0_initial, attr1, attr2, pos1, rot1, pos2, rot2, g1_plugin_id, g2_plugin_id, sdf_iterations
)
write_contact(
nconmax_in,
dist,
pos,
make_frame(n),
margin,
gap,
condim,
friction,
solref,
solreffriction,
solimp,
geoms,
worldid,
ncon_out,
contact_dist_out,
contact_pos_out,
contact_frame_out,
contact_includemargin_out,
contact_friction_out,
contact_solref_out,
contact_solreffriction_out,
contact_solimp_out,
contact_dim_out,
contact_geom_out,
contact_worldid_out,
)
x_g2 = wp.vec3(
aabb_intersection.min[0] + (aabb_intersection.max[0] - aabb_intersection.min[0]) * halton(i, 2),
aabb_intersection.min[1] + (aabb_intersection.max[1] - aabb_intersection.min[1]) * halton(i, 3),
aabb_intersection.min[2] + (aabb_intersection.max[2] - aabb_intersection.min[2]) * halton(i, 5),
)
x = geom1.rot * x_g2 + geom1.pos
x0_initial = wp.transpose(rot2) * (x - pos2)
dist, pos, n = gradient_descent(
type1,
x0_initial,
attr1,
attr2,
pos1,
rot1,
pos2,
rot2,
g1_plugin_id,
g2_plugin_id,
sdf_iterations,
volume_data1,
volume_data2,
mesh_data1,
mesh_data2,
)
write_contact(
nconmax_in,
dist,
pos,
make_frame(n),
margin,
gap,
condim,
friction,
solref,
solreffriction,
solimp,
geoms,
worldid,
ncon_out,
contact_dist_out,
contact_pos_out,
contact_frame_out,
contact_includemargin_out,
contact_friction_out,
contact_solref_out,
contact_solreffriction_out,
contact_solimp_out,
contact_dim_out,
contact_geom_out,
contact_worldid_out,
)
@event_scope
def sdf_narrowphase(m: Model, d: Data):
wp.launch(
_sdf_narrowphase,
dim=d.nconmax,
dim=(m.opt.sdf_initpoints, d.nconmax),
inputs=[
m.nmeshface,
m.geom_type,
m.geom_condim,
m.geom_dataid,
@@ -585,8 +882,6 @@ def sdf_narrowphase(m: Model, d: Data):
m.geom_solimp,
m.geom_size,
m.geom_aabb,
m.geom_pos,
m.geom_quat,
m.geom_friction,
m.geom_margin,
m.geom_gap,
@@ -598,6 +893,8 @@ def sdf_narrowphase(m: Model, d: Data):
m.mesh_vertadr,
m.mesh_vertnum,
m.mesh_vert,
m.mesh_faceadr,
m.mesh_face,
m.mesh_graphadr,
m.mesh_graph,
m.mesh_polynum,
@@ -609,6 +906,9 @@ def sdf_narrowphase(m: Model, d: Data):
m.mesh_polymapadr,
m.mesh_polymapnum,
m.mesh_polymap,
m.oct_aabb,
m.oct_child,
m.oct_coeff,
m.pair_dim,
m.pair_solref,
m.pair_solreffriction,
@@ -623,7 +923,6 @@ def sdf_narrowphase(m: Model, d: Data):
d.geom_xpos,
d.geom_xmat,
d.collision_pair,
d.collision_hftri_index,
d.collision_pairid,
d.collision_worldid,
d.ncollision,
+1 -1
View File
@@ -23,7 +23,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import vec5
from mujoco.mjx.third_party.mujoco_warp._src.types import vec11
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
wp.config.enable_backward = False
wp.set_module_options({"enable_backward": False})
@wp.kernel
+2 -2
View File
@@ -82,8 +82,8 @@ def _qderiv_actuator_passive(
qderiv += actuator_moment_in[worldid, actid, dofiid] * actuator_moment_in[worldid, actid, dofjid] * vel
if passive_enabled and dofiid == dofjid:
qderiv -= dof_damping[worldid, dofiid] / float(nu)
if passive_enabled and dofiid == dofjid:
qderiv -= dof_damping[worldid, dofiid]
qderiv *= opt_timestep[worldid]
+3 -4
View File
@@ -42,7 +42,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType
from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
wp.set_module_options({"enable_backward": False})
@@ -358,7 +357,7 @@ def _euler_sparse(m: Model, d: Data):
@cache_kernel
def _tile_euler_dense(tile: TileSet):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def euler_dense(
# Model:
dof_damping: wp.array2d(dtype=float),
@@ -558,7 +557,7 @@ def fwd_position(m: Model, d: Data, factorize: bool = True):
def _actuator_velocity(m: Model, d: Data):
NV = m.nv
@kernel
@nested_kernel(module="unique", enable_backward=False)
def actuator_velocity(
# Data in:
qvel_in: wp.array2d(dtype=float),
@@ -590,7 +589,7 @@ def _actuator_velocity(m: Model, d: Data):
def _tendon_velocity(m: Model, d: Data):
NV = m.nv
@kernel
@nested_kernel(module="unique", enable_backward=False)
def tendon_velocity(
# Data in:
qvel_in: wp.array2d(dtype=float),
@@ -27,6 +27,8 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit
from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
wp.set_module_options({"enable_backward": False})
@wp.kernel
def _qfrc_eulerdamp(
+243 -30
View File
@@ -30,16 +30,20 @@ MJ_CCD_ITERATIONS = 12
# max number of worlds supported
MAX_WORLDS = 2**24
# tolerance override for float32
_TOLERANCE_F32 = 1.0e-6
def _hfield_geom_pair(mjm: mujoco.MjModel) -> Tuple[int, np.array]:
geom1, geom2 = np.triu_indices(mjm.ngeom, k=1)
geom_type_hf = mujoco.mjtGeom.mjGEOM_HFIELD
has_hfield = (mjm.geom_type[geom1] == geom_type_hf) | (mjm.geom_type[geom2] == geom_type_hf)
nhfieldgeompair = np.sum(has_hfield)
geompair2hfgeompair = -1 * np.ones(mjm.ngeom * (mjm.ngeom - 1) // 2, dtype=int)
geompair2hfgeompair[has_hfield] = np.arange(nhfieldgeompair)
return nhfieldgeompair, geompair2hfgeompair
def _max_meshdegree(mjm: mujoco.MjModel) -> int:
if mjm.mesh_polyvertnum.size == 0:
return 4
return max(3, mjm.mesh_polymapnum.max())
def _max_npolygon(mjm: mujoco.MjModel) -> int:
if mjm.mesh_polyvertnum.size == 0:
return 4
return max(4, mjm.mesh_polyvertnum.max())
def put_model(mjm: mujoco.MjModel) -> types.Model:
@@ -118,20 +122,6 @@ 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) | (mjm.sensor_intprm[is_contact_sensor, 1] == 2))).any():
raise NotImplementedError(f"Contact sensor: only mindist and maxforce reduction are 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:
@@ -444,8 +434,11 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
npair=mjm.npair,
opt=types.Option(
timestep=create_nmodel_batched_array(np.array(mjm.opt.timestep), dtype=float, expand_dim=False),
tolerance=create_nmodel_batched_array(np.array(mjm.opt.tolerance), dtype=float, expand_dim=False),
tolerance=create_nmodel_batched_array(
np.array(np.maximum(mjm.opt.tolerance, _TOLERANCE_F32)), dtype=float, expand_dim=False
),
ls_tolerance=create_nmodel_batched_array(np.array(mjm.opt.ls_tolerance), dtype=float, expand_dim=False),
ccd_tolerance=create_nmodel_batched_array(np.array(mjm.opt.ccd_tolerance), dtype=float, expand_dim=False),
gravity=create_nmodel_batched_array(mjm.opt.gravity, dtype=wp.vec3, expand_dim=False),
magnetic=create_nmodel_batched_array(mjm.opt.magnetic, dtype=wp.vec3, expand_dim=False),
wind=create_nmodel_batched_array(mjm.opt.wind, dtype=wp.vec3, expand_dim=False),
@@ -474,6 +467,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
sdf_iterations=mjm.opt.sdf_iterations,
run_collision_detection=True,
legacy_gjk=False,
contact_sensor_maxmatch=64,
),
stat=types.Statistic(
meaninertia=mjm.stat.meaninertia,
@@ -635,6 +629,9 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
mesh_polymapadr=wp.array(mjm.mesh_polymapadr, dtype=int),
mesh_polymapnum=wp.array(mjm.mesh_polymapnum, dtype=int),
mesh_polymap=wp.array(mjm.mesh_polymap, dtype=int),
oct_aabb=wp.array2d(mjm.oct_aabb, dtype=wp.vec3),
oct_child=wp.array(mjm.oct_child, dtype=types.vec8i),
oct_coeff=wp.array(mjm.oct_coeff, dtype=types.vec8f),
nhfield=mjm.nhfield,
nhfielddata=mjm.nhfielddata,
hfield_adr=wp.array(mjm.hfield_adr, dtype=int),
@@ -826,7 +823,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
mat_texrepeat=create_nmodel_batched_array(mjm.mat_texrepeat, dtype=wp.vec2),
mat_rgba=create_nmodel_batched_array(mjm.mat_rgba, dtype=wp.vec4),
actuator_trntype_body_adr=wp.array(np.nonzero(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)[0], dtype=int),
geompair2hfgeompair=wp.array(_hfield_geom_pair(mjm)[1], dtype=int),
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)),
@@ -887,6 +883,9 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
condim = np.concatenate((mjm.geom_condim, mjm.pair_dim))
condim_max = np.max(condim) if len(condim) > 0 else 0
max_npolygon = _max_npolygon(mjm)
max_meshdegree = _max_meshdegree(mjm)
if mujoco.mj_isSparse(mjm):
qM = wp.zeros((nworld, 1, mjm.nM), dtype=float)
qLD = wp.zeros((nworld, 1, mjm.nC), dtype=float)
@@ -907,8 +906,6 @@ 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
ne_weld=wp.zeros(nworld, dtype=int), # warp only
@@ -1060,7 +1057,6 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
),
# collision driver
collision_pair=wp.zeros((nconmax,), dtype=wp.vec2i),
collision_hftri_index=wp.zeros((nconmax,), dtype=int),
collision_pairid=wp.zeros((nconmax,), dtype=int),
collision_worldid=wp.zeros((nconmax,), dtype=int),
ncollision=wp.zeros((1,), dtype=int),
@@ -1076,6 +1072,17 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
epa_index=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int),
epa_map=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int),
epa_horizon=wp.zeros(shape=(nconmax, 6 * MJ_CCD_ITERATIONS), dtype=int),
multiccd_polygon=wp.zeros(shape=(nconmax, 2 * max_npolygon), dtype=wp.vec3),
multiccd_clipped=wp.zeros(shape=(nconmax, 2 * max_npolygon), dtype=wp.vec3),
multiccd_pnormal=wp.zeros(shape=(nconmax, max_npolygon), dtype=wp.vec3),
multiccd_pdist=wp.zeros(shape=(nconmax, max_npolygon), dtype=float),
multiccd_idx1=wp.zeros(shape=(nconmax, max_meshdegree), dtype=int),
multiccd_idx2=wp.zeros(shape=(nconmax, max_meshdegree), dtype=int),
multiccd_n1=wp.zeros(shape=(nconmax, max_meshdegree), dtype=wp.vec3),
multiccd_n2=wp.zeros(shape=(nconmax, max_meshdegree), dtype=wp.vec3),
multiccd_endvert=wp.zeros(shape=(nconmax, max_meshdegree), dtype=wp.vec3),
multiccd_face1=wp.zeros(shape=(nconmax, max_npolygon), dtype=wp.vec3),
multiccd_face2=wp.zeros(shape=(nconmax, max_npolygon), dtype=wp.vec3),
# rne_postconstraint
cacc=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
cfrc_int=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector),
@@ -1158,6 +1165,9 @@ def put_data(
if mjd.nefc > njmax:
raise ValueError(f"njmax overflow (njmax must be >= {mjd.nefc})")
max_npolygon = _max_npolygon(mjm)
max_meshdegree = _max_meshdegree(mjm)
# calculate some fields that cannot be easily computed inline:
if mujoco.mj_isSparse(mjm):
qM = np.expand_dims(mjd.qM, axis=0)
@@ -1276,8 +1286,6 @@ 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),
ne_weld=wp.full(shape=(nworld), value=ne_weld),
@@ -1423,7 +1431,6 @@ def put_data(
sap_segment_index=arr(np.array([i * mjm.ngeom if i < nworld + 1 else 0 for i in range(2 * nworld)]).reshape((nworld, 2))),
# collision driver
collision_pair=wp.empty(nconmax, dtype=wp.vec2i),
collision_hftri_index=wp.empty(nconmax, dtype=int),
collision_pairid=wp.empty(nconmax, dtype=int),
collision_worldid=wp.empty(nconmax, dtype=int),
ncollision=wp.zeros(1, dtype=int),
@@ -1439,6 +1446,17 @@ def put_data(
epa_index=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int),
epa_map=wp.zeros(shape=(nconmax, 6 + 6 * MJ_CCD_ITERATIONS), dtype=int),
epa_horizon=wp.zeros(shape=(nconmax, 6 * MJ_CCD_ITERATIONS), dtype=int),
multiccd_polygon=wp.zeros(shape=(nconmax, 2 * max_npolygon), dtype=wp.vec3),
multiccd_clipped=wp.zeros(shape=(nconmax, 2 * max_npolygon), dtype=wp.vec3),
multiccd_pnormal=wp.zeros(shape=(nconmax, max_npolygon), dtype=wp.vec3),
multiccd_pdist=wp.zeros(shape=(nconmax, max_npolygon), dtype=float),
multiccd_idx1=wp.zeros(shape=(nconmax, max_meshdegree), dtype=int),
multiccd_idx2=wp.zeros(shape=(nconmax, max_meshdegree), dtype=int),
multiccd_n1=wp.zeros(shape=(nconmax, max_meshdegree), dtype=wp.vec3),
multiccd_n2=wp.zeros(shape=(nconmax, max_meshdegree), dtype=wp.vec3),
multiccd_endvert=wp.zeros(shape=(nconmax, max_meshdegree), dtype=wp.vec3),
multiccd_face1=wp.zeros(shape=(nconmax, max_npolygon), dtype=wp.vec3),
multiccd_face2=wp.zeros(shape=(nconmax, max_npolygon), dtype=wp.vec3),
# rne_postconstraint but also smooth
cacc=tile(mjd.cacc, dtype=wp.spatial_vector),
cfrc_int=tile(mjd.cfrc_int, dtype=wp.spatial_vector),
@@ -1631,3 +1649,198 @@ def get_data_into(
# sensors
result.sensordata[:] = d.sensordata.numpy()
@wp.kernel
def _reset_nworld(
# Model:
nq: int,
nv: int,
nu: int,
na: int,
neq: int,
nsensordata: int,
qpos0: wp.array2d(dtype=float),
eq_active0: wp.array(dtype=bool),
# Data in:
nworld_in: int,
# Data out:
solver_niter_out: wp.array(dtype=int),
ncon_out: wp.array(dtype=int),
ne_out: wp.array(dtype=int),
ne_connect_out: wp.array(dtype=int),
ne_weld_out: wp.array(dtype=int),
ne_jnt_out: wp.array(dtype=int),
ne_ten_out: wp.array(dtype=int),
nf_out: wp.array(dtype=int),
nl_out: wp.array(dtype=int),
nefc_out: wp.array(dtype=int),
nsolving_out: wp.array(dtype=int),
time_out: wp.array(dtype=float),
energy_out: wp.array(dtype=wp.vec2),
qpos_out: wp.array2d(dtype=float),
qvel_out: wp.array2d(dtype=float),
act_out: wp.array2d(dtype=float),
qacc_warmstart_out: wp.array2d(dtype=float),
ctrl_out: wp.array2d(dtype=float),
qfrc_applied_out: wp.array2d(dtype=float),
eq_active_out: wp.array2d(dtype=bool),
qacc_out: wp.array2d(dtype=float),
act_dot_out: wp.array2d(dtype=float),
sensordata_out: wp.array2d(dtype=float),
):
worldid = wp.tid()
solver_niter_out[worldid] = 0
if worldid == 0:
ncon_out[0] = 0
ne_out[worldid] = 0
ne_connect_out[worldid] = 0
ne_weld_out[worldid] = 0
ne_jnt_out[worldid] = 0
ne_ten_out[worldid] = 0
nf_out[worldid] = 0
nl_out[worldid] = 0
nefc_out[worldid] = 0
if worldid == 0:
nsolving_out[0] = nworld_in
time_out[worldid] = 0.0
energy_out[worldid] = wp.vec2(0.0, 0.0)
for i in range(nq):
qpos_out[worldid, i] = qpos0[worldid, i]
if i < nv:
qvel_out[worldid, i] = 0.0
qacc_warmstart_out[worldid, i] = 0.0
qfrc_applied_out[worldid, i] = 0.0
qacc_out[worldid, i] = 0.0
for i in range(nu):
ctrl_out[worldid, i] = 0.0
if i < na:
act_out[worldid, i] = 0.0
act_dot_out[worldid, i] = 0.0
for i in range(neq):
eq_active_out[worldid, i] = eq_active0[i]
for i in range(nsensordata):
sensordata_out[worldid, i] = 0.0
@wp.kernel
def _reset_mocap(
# Model:
body_mocapid: wp.array(dtype=int),
body_pos: wp.array2d(dtype=wp.vec3),
body_quat: wp.array2d(dtype=wp.quat),
# Data out:
mocap_pos_out: wp.array2d(dtype=wp.vec3),
mocap_quat_out: wp.array2d(dtype=wp.quat),
):
worldid, bodyid = wp.tid()
mocapid = body_mocapid[bodyid]
if mocapid >= 0:
mocap_pos_out[worldid, mocapid] = body_pos[worldid, bodyid]
mocap_quat_out[worldid, mocapid] = body_quat[worldid, bodyid]
@wp.kernel
def _reset_contact(
# Data in:
ncon_in: wp.array(dtype=int),
# In:
nefcaddress: int,
# Data out:
contact_dist_out: wp.array(dtype=float),
contact_pos_out: wp.array(dtype=wp.vec3),
contact_frame_out: wp.array(dtype=wp.mat33),
contact_includemargin_out: wp.array(dtype=float),
contact_friction_out: wp.array(dtype=types.vec5),
contact_solref_out: wp.array(dtype=wp.vec2),
contact_solreffriction_out: wp.array(dtype=wp.vec2),
contact_solimp_out: wp.array(dtype=types.vec5),
contact_dim_out: wp.array(dtype=int),
contact_geom_out: wp.array(dtype=wp.vec2i),
contact_efc_address_out: wp.array2d(dtype=int),
contact_worldid_out: wp.array(dtype=int),
):
conid = wp.tid()
if conid >= ncon_in[0]:
return
contact_dist_out[conid] = 0.0
contact_pos_out[conid] = wp.vec3(0.0)
contact_frame_out[conid] = wp.mat33(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0)
contact_includemargin_out[conid] = 0.0
contact_friction_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0)
contact_solref_out[conid] = wp.vec2(0.0, 0.0)
contact_solreffriction_out[conid] = wp.vec2(0.0, 0.0)
contact_solimp_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0)
contact_dim_out[conid] = 0
contact_geom_out[conid] = wp.vec2i(0, 0)
for i in range(nefcaddress):
contact_efc_address_out[conid, i] = 0
contact_worldid_out[conid] = 0
def reset_data(m: types.Model, d: types.Data):
"""Clear data, set defaults."""
d.xfrc_applied.zero_()
d.qM.zero_()
# set mocap_pos/quat = body_pos/quat for mocap bodies
wp.launch(
_reset_mocap, dim=(d.nworld, m.nbody), inputs=[m.body_mocapid, m.body_pos, m.body_quat], outputs=[d.mocap_pos, d.mocap_quat]
)
# clear contacts
wp.launch(
_reset_contact,
dim=d.nconmax,
inputs=[d.ncon, d.contact.efc_address.shape[1]],
outputs=[
d.contact.dist,
d.contact.pos,
d.contact.frame,
d.contact.includemargin,
d.contact.friction,
d.contact.solref,
d.contact.solreffriction,
d.contact.solimp,
d.contact.dim,
d.contact.geom,
d.contact.efc_address,
d.contact.worldid,
],
)
wp.launch(
_reset_nworld,
dim=d.nworld,
inputs=[m.nq, m.nv, m.nu, m.na, m.neq, m.nsensordata, m.qpos0, m.eq_active0, d.nworld],
outputs=[
d.solver_niter,
d.ncon,
d.ne,
d.ne_connect,
d.ne_weld,
d.ne_jnt,
d.ne_ten,
d.nf,
d.nl,
d.nefc,
d.nsolving,
d.time,
d.energy,
d.qpos,
d.qvel,
d.act,
d.qacc_warmstart,
d.ctrl,
d.qfrc_applied,
d.eq_active,
d.qacc,
d.act_dot,
d.sensordata,
],
)
@@ -0,0 +1,248 @@
# Copyright 2025 The Newton Developers
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
"""Tests for io functions for JAX compatibility.
NOTE: do not edit this file without consulting with the MuJoCo/MJX team.
"""
import dataclasses
import typing
from typing import Any, Dict, Optional, Union
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
_IO_TEST_MODELS = (
"pendula.xml",
"collision_sdf/tactile.xml",
"flex/floppy.xml",
"actuation/tendon_force_limit.xml",
"hfield/hfield.xml",
)
def _dims_match(test_obj, d1: Any, d2: Any, prefix: str = ""):
"""Checks that two dataclasses have fields with the same leading dims."""
fields1, fields2 = dataclasses.fields(d1), dataclasses.fields(d2)
for f1, f2 in zip(fields1, fields2):
full_name = prefix + f1.name
a1, a2 = getattr(d1, f1.name), getattr(d2, f2.name)
if dataclasses.is_dataclass(a1) or dataclasses.is_dataclass(a2):
_dims_match(test_obj, a1, a2, prefix + f1.name + ".")
continue
if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array):
s1, s2 = a1.shape, a2.shape
test_obj.assertEqual(len(s1), len(s2), f"{full_name} dims mismatch. Got {s1} and {s2}.")
test_obj.assertEqual(s1, s2, f"{full_name} dims mismatch. Got {s1} and {s2}.")
def _get_np_scalar_type(val: Any) -> Optional[Union[bool, int, float]]:
"""Returns the python type from a numpy scalar."""
is_np_scalar = list(isinstance(val, t) for t in (np.integer, np.floating, np.bool_))
if any(is_np_scalar):
return [int, float, bool][is_np_scalar.index(True)]
def _check_type_matches_annotation(test_obj, obj: Any, prefix: str = ""):
"""Checks that dataclass annotations match the runtime types."""
assert dataclasses.is_dataclass(obj), prefix + " must be dataclass."
msg = "Type of {val_type} does not match annotation {type_} for field {prefix}{field_name}"
for field in dataclasses.fields(obj):
field_name = field.name
val = getattr(obj, field_name)
val_type = type(val)
type_ = field.type
if dataclasses.is_dataclass(val):
test_obj.assertTrue(dataclasses.is_dataclass(type_), msg.format(**locals()))
_check_type_matches_annotation(test_obj, val, prefix + field_name + ".")
continue
np_scalar_type = _get_np_scalar_type(val)
if np_scalar_type:
test_obj.assertIsInstance(np_scalar_type(val), type_, msg.format(**locals()))
continue
if isinstance(type_, wp.types.array):
test_obj.assertIsInstance(val, wp.types.array, msg.format(**locals()))
continue
origin_type = typing.get_origin(type_)
if tuple in (val_type, origin_type):
test_obj.assertEqual(val_type, origin_type, msg.format(**locals()))
field_name += ".tuple[]"
type_ = typing.get_args(type_)[0]
items = val
for val in items:
val_type = type(val)
if dataclasses.is_dataclass(val):
_check_type_matches_annotation(test_obj, val, prefix + field_name)
continue
np_scalar_type = _get_np_scalar_type(val)
if np_scalar_type:
test_obj.assertIsInstance(np_scalar_type(val), type_, msg.format(**locals()))
continue
if isinstance(type_, wp.types.array):
test_obj.assertIsInstance(val, wp.types.array, msg.format(**locals()))
continue
test_obj.assertEqual(type(val), type_, msg.format(**locals()))
continue
test_obj.assertEqual(type(val), field.type, msg.format(**locals()))
def _check_annotation_compat(
annotations: Dict[str, Any], prefix: str = "", in_cls: bool = False, in_tuple: bool = False
) -> Dict[str, Any]:
"""Checks that dataclass annotations match criteria for JAX API compat."""
for k, v in annotations.items():
full_key = f"{prefix}{k}"
info = f"Found {v} for annotation {full_key}."
if v in (int, bool, float):
continue
if isinstance(v, wp.types.array):
continue
if v in wp.types.vector_types:
raise AssertionError(f"Vector types are not allowed. {info}")
if typing.get_origin(v) == tuple and (in_cls or in_tuple):
raise AssertionError(f"Nested args in Model/Data must not be tuple. {info}")
if typing.get_origin(v) == tuple:
tuple_args = typing.get_args(v)
if len(tuple_args) != 2 and tuple_args[1] != ...:
raise AssertionError(f"Tuple args must be variadic. {info}")
_check_annotation_compat(
{"[]": tuple_args[0]},
prefix=f"{full_key}.tuple",
in_cls=in_cls,
in_tuple=True,
)
continue
if hasattr(v, "__class__") and in_cls:
raise AssertionError(f"Nested object args in Model/Data are not allowed. {info}")
if hasattr(v, "__class__") and not dataclasses.is_dataclass(v):
raise AssertionError(f"Args that are objects must be dataclass. {info}")
if hasattr(v, "__class__") and not v.__module__.startswith("mujoco_warp"):
raise AssertionError(f"dataclass args must be within the mujoco_warp module. {info}")
if hasattr(v, "__class__"):
_check_annotation_compat(v.__annotations__, prefix=f"{full_key}{v.__name__}.", in_cls=True, in_tuple=in_tuple)
continue
raise AssertionError(f"Model/Data annotation is not allowed. {info}")
def _leading_dims_scale_w_nworld(test_obj, d1: Any, d2: Any, nworld1: int, nworld2: int, prefix: str = ""):
"""Checks that dataclass fields that scale with nworld have leading dim nworld."""
msg = "Arrays that scale with nworld should have leading dim nworld."
fields1, fields2 = dataclasses.fields(d1), dataclasses.fields(d2)
for f1, f2 in zip(fields1, fields2):
full_name = prefix + f1.name
a1, a2 = getattr(d1, f1.name), getattr(d2, f2.name)
if dataclasses.is_dataclass(a1) or dataclasses.is_dataclass(a2):
_leading_dims_scale_w_nworld(test_obj, a1, a2, nworld1, nworld2, prefix + f1.name + ".")
continue
if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array):
s1, s2 = a1.shape[0], a2.shape[0]
if s1 == s2:
continue
test_obj.assertEqual(s2, nworld2, full_name + f" has leading dim {s2} with nworld={nworld2}. {msg}")
test_obj.assertEqual(s1, nworld1, full_name + f" has leading dim {s1} with nworld={nworld1}. {msg}")
class IOTest(parameterized.TestCase):
def test_put_model_nworld_array(self):
"""Tests that put_model arrays with nworld leading dim have `_is_batched`."""
mjm, *_ = test_util.fixture("pendula.xml")
m1 = mjwarp.put_model(mjm)
self.assertTrue(hasattr(m1.geom_pos, "_is_batched"))
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], 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"))
self.assertGreater(m1.body_parentid.shape[0], 0)
self.assertGreater(m1.body_parentid.strides[0], 0)
self.assertLen(m1.body_parentid.strides, m1.body_parentid.ndim)
@parameterized.parameters(*_IO_TEST_MODELS)
def test_put_data_nworld_array(self, xml):
"""Tests that put_data arrays that scale with nworld have leading dim nworld."""
mjm, mjd, _, _ = test_util.fixture(xml)
d1 = mjwarp.put_data(mjm, mjd, nworld=1, nconmax=3_000, njmax=3_000)
dn = mjwarp.put_data(mjm, mjd, nworld=133, nconmax=3_000, njmax=3_000)
_leading_dims_scale_w_nworld(self, d1, dn, 1, 133)
def test_public_api_jax_compat(self):
"""Tests that annotations meet a set of criteria for JAX compat."""
_check_annotation_compat(mjwarp.Model.__annotations__, "Model.")
_check_annotation_compat(mjwarp.Data.__annotations__, "Data.")
@parameterized.parameters(*_IO_TEST_MODELS)
def test_types_match_annotations(self, xml):
"""Tests that the types of dataclass fields match the annotations."""
mjm, _, m, d = test_util.fixture(xml)
_check_type_matches_annotation(self, m, "Model.")
_check_type_matches_annotation(self, d, "Data.")
d = mjwarp.make_data(mjm, nworld=2)
_check_type_matches_annotation(self, d, "Data.")
@parameterized.parameters(*_IO_TEST_MODELS)
def test_make_put_data_dims_match(self, xml):
"""Tests that make_data and put_data have matching dimensions."""
mjm, mjd, _, _ = test_util.fixture(xml)
dm2 = mjwarp.make_data(mjm, nworld=2, nconmax=3_000, njmax=4_200)
dm3 = mjwarp.make_data(mjm, nworld=3, nconmax=3_000, njmax=4_200)
dp2 = mjwarp.put_data(mjm, mjd, nworld=2, nconmax=3_000, njmax=4_200)
dp3 = mjwarp.put_data(mjm, mjd, nworld=3, nconmax=3_000, njmax=4_200)
_dims_match(self, dm2, dp2)
_dims_match(self, dm3, dp3)
if __name__ == "__main__":
wp.init()
absltest.main()
+69 -230
View File
@@ -15,10 +15,6 @@
"""Tests for io functions."""
import dataclasses
import typing
from typing import Any, Dict, Optional, Union
import mujoco
import numpy as np
import warp as wp
@@ -28,8 +24,15 @@ 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 _assert_eq(a, b, name):
tol = 5e-4
err_msg = f"mismatch: {name}"
np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol)
# NOTE: modify io_jax_test _IO_TEST_MODELS if changed here.
_IO_TEST_MODELS = (
"pendula.xml",
"collision_sdf/tactile.xml",
@@ -39,150 +42,6 @@ _IO_TEST_MODELS = (
)
def _dims_match(test_obj, d1: Any, d2: Any, prefix: str = ""):
"""Checks that two dataclasses have fields with the same leading dims."""
fields1, fields2 = dataclasses.fields(d1), dataclasses.fields(d2)
for f1, f2 in zip(fields1, fields2):
full_name = prefix + f1.name
a1, a2 = getattr(d1, f1.name), getattr(d2, f2.name)
if dataclasses.is_dataclass(a1) or dataclasses.is_dataclass(a2):
_dims_match(test_obj, a1, a2, prefix + f1.name + ".")
continue
if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array):
s1, s2 = a1.shape, a2.shape
test_obj.assertEqual(len(s1), len(s2), f"{full_name} dims mismatch. Got {s1} and {s2}.")
test_obj.assertEqual(s1, s2, f"{full_name} dims mismatch. Got {s1} and {s2}.")
def _get_np_scalar_type(val: Any) -> Optional[Union[bool, int, float]]:
"""Returns the python type from a numpy scalar."""
is_np_scalar = list(isinstance(val, t) for t in (np.integer, np.floating, np.bool_))
if any(is_np_scalar):
return [int, float, bool][is_np_scalar.index(True)]
def _check_type_matches_annotation(test_obj, obj: Any, prefix: str = ""):
"""Checks that dataclass annotations match the runtime types."""
assert dataclasses.is_dataclass(obj), prefix + " must be dataclass."
msg = "Type of {val_type} does not match annotation {type_} for field {prefix}{field_name}"
for field in dataclasses.fields(obj):
field_name = field.name
val = getattr(obj, field_name)
val_type = type(val)
type_ = field.type
if dataclasses.is_dataclass(val):
test_obj.assertTrue(dataclasses.is_dataclass(type_), msg.format(**locals()))
_check_type_matches_annotation(test_obj, val, prefix + field_name + ".")
continue
np_scalar_type = _get_np_scalar_type(val)
if np_scalar_type:
test_obj.assertIsInstance(np_scalar_type(val), type_, msg.format(**locals()))
continue
if isinstance(type_, wp.types.array):
test_obj.assertIsInstance(val, wp.types.array, msg.format(**locals()))
continue
origin_type = typing.get_origin(type_)
if tuple in (val_type, origin_type):
test_obj.assertEqual(val_type, origin_type, msg.format(**locals()))
field_name += ".tuple[]"
type_ = typing.get_args(type_)[0]
items = val
for val in items:
val_type = type(val)
if dataclasses.is_dataclass(val):
_check_type_matches_annotation(test_obj, val, prefix + field_name)
continue
np_scalar_type = _get_np_scalar_type(val)
if np_scalar_type:
test_obj.assertIsInstance(np_scalar_type(val), type_, msg.format(**locals()))
continue
if isinstance(type_, wp.types.array):
test_obj.assertIsInstance(val, wp.types.array, msg.format(**locals()))
continue
test_obj.assertEqual(type(val), type_, msg.format(**locals()))
continue
test_obj.assertEqual(type(val), field.type, msg.format(**locals()))
def _check_annotation_compat(
annotations: Dict[str, Any], prefix: str = "", in_cls: bool = False, in_tuple: bool = False
) -> Dict[str, Any]:
"""Checks that dataclass annotations match criteria for JAX API compat."""
for k, v in annotations.items():
full_key = f"{prefix}{k}"
info = f"Found {v} for annotation {full_key}."
if v in (int, bool, float):
continue
if isinstance(v, wp.types.array):
continue
if v in wp.types.vector_types:
raise AssertionError(f"Vector types are not allowed. {info}")
if typing.get_origin(v) == tuple and (in_cls or in_tuple):
raise AssertionError(f"Nested args in Model/Data must not be tuple. {info}")
if typing.get_origin(v) == tuple:
tuple_args = typing.get_args(v)
if len(tuple_args) != 2 and tuple_args[1] != ...:
raise AssertionError(f"Tuple args must be variadic. {info}")
_check_annotation_compat(
{"[]": tuple_args[0]},
prefix=f"{full_key}.tuple",
in_cls=in_cls,
in_tuple=True,
)
continue
if hasattr(v, "__class__") and in_cls:
raise AssertionError(f"Nested object args in Model/Data are not allowed. {info}")
if hasattr(v, "__class__") and not dataclasses.is_dataclass(v):
raise AssertionError(f"Args that are objects must be dataclass. {info}")
if hasattr(v, "__class__") and not v.__module__.startswith("mujoco_warp"):
raise AssertionError(f"dataclass args must be within the mujoco_warp module. {info}")
if hasattr(v, "__class__"):
_check_annotation_compat(v.__annotations__, prefix=f"{full_key}{v.__name__}.", in_cls=True, in_tuple=in_tuple)
continue
raise AssertionError(f"Model/Data annotation is not allowed. {info}")
def _leading_dims_scale_w_nworld(test_obj, d1: Any, d2: Any, nworld1: int, nworld2: int, prefix: str = ""):
"""Checks that dataclass fields that scale with nworld have leading dim nworld."""
msg = "Arrays that scale with nworld should have leading dim nworld."
fields1, fields2 = dataclasses.fields(d1), dataclasses.fields(d2)
for f1, f2 in zip(fields1, fields2):
full_name = prefix + f1.name
a1, a2 = getattr(d1, f1.name), getattr(d2, f2.name)
if dataclasses.is_dataclass(a1) or dataclasses.is_dataclass(a2):
_leading_dims_scale_w_nworld(test_obj, a1, a2, nworld1, nworld2, prefix + f1.name + ".")
continue
if isinstance(f1.type, wp.types.array) or isinstance(f2.type, wp.types.array):
s1, s2 = a1.shape[0], a2.shape[0]
if s1 == s2:
continue
test_obj.assertEqual(s2, nworld2, full_name + f" has leading dim {s2} with nworld={nworld2}. {msg}")
test_obj.assertEqual(s1, nworld1, full_name + f" has leading dim {s1} with nworld={nworld1}. {msg}")
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."""
@@ -197,8 +56,6 @@ class IOTest(parameterized.TestCase):
if isinstance(val, wp.array):
self.assertEqual(val.shape, getattr(d, attr).shape, f"{attr} shape mismatch")
# TODO(team): sensors
def test_get_data_into_m(self):
mjm = mujoco.MjModel.from_xml_string("""
<mujoco>
@@ -299,95 +156,77 @@ class IOTest(parameterized.TestCase):
"""
)
def test_put_model_nworld_array(self):
"""Tests that put_model arrays with nworld leading dim have `_is_batched`."""
mjm, *_ = test_util.fixture("pendula.xml")
m1 = mjwarp.put_model(mjm)
self.assertTrue(hasattr(m1.geom_pos, "_is_batched"))
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], 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"))
self.assertGreater(m1.body_parentid.shape[0], 0)
self.assertGreater(m1.body_parentid.strides[0], 0)
self.assertLen(m1.body_parentid.strides, m1.body_parentid.ndim)
@parameterized.parameters(*_IO_TEST_MODELS)
def test_put_data_nworld_array(self, xml):
"""Tests that put_data arrays that scale with nworld have leading dim nworld."""
mjm, mjd, _, _ = test_util.fixture(xml)
d1 = mjwarp.put_data(mjm, mjd, nworld=1, nconmax=3_000, njmax=3_000)
dn = mjwarp.put_data(mjm, mjd, nworld=133, nconmax=3_000, njmax=3_000)
_leading_dims_scale_w_nworld(self, d1, dn, 1, 133)
def test_reset_data(self, xml):
reset_datafield = [
"ncon",
"ne",
"nf",
"nl",
"nefc",
"time",
"energy",
"qpos",
"qvel",
"act",
"ctrl",
"eq_active",
"qfrc_applied",
"xfrc_applied",
"qacc",
"qacc_warmstart",
"act_dot",
"sensordata",
"mocap_pos",
"mocap_quat",
"qM",
]
@parameterized.parameters(*_IO_TEST_MODELS)
def test_make_data_nworld_array(self, xml):
"""Tests that make_data arrays that scale with nworld have leading dim nworld."""
mjm, *_ = test_util.fixture(xml)
d1 = mjwarp.make_data(mjm, nworld=1, nconmax=3_000, njmax=3_000)
dn = mjwarp.make_data(mjm, nworld=133, nconmax=3_000, njmax=3_000)
_leading_dims_scale_w_nworld(self, d1, dn, 1, 133)
nworld = 1
mjm, mjd, m, d = test_util.fixture(xml, nworld=nworld)
nconmax = d.nconmax
def test_public_api_jax_compat(self):
"""Tests that annotations meet a set of criteria for JAX compat."""
_check_annotation_compat(mjwarp.Model.__annotations__, "Model.")
_check_annotation_compat(mjwarp.Data.__annotations__, "Data.")
# data fields
for arr in reset_datafield:
attr = getattr(d, arr)
if attr.dtype == float:
attr.fill_(wp.nan)
else:
attr.fill_(-1)
@parameterized.parameters(*_IO_TEST_MODELS)
def test_types_match_annotations(self, xml):
"""Tests that the types of dataclass fields match the annotations."""
mjm, _, m, d = test_util.fixture(xml)
for arr in d.contact.__dataclass_fields__:
attr = getattr(d.contact, arr)
if attr.dtype == float:
attr.fill_(wp.nan)
else:
attr.fill_(-1)
_check_type_matches_annotation(self, m, "Model.")
_check_type_matches_annotation(self, d, "Data.")
mujoco.mj_resetData(mjm, mjd)
d = mjwarp.make_data(mjm, nworld=2)
_check_type_matches_annotation(self, d, "Data.")
# set ncon in order to zero all contact memory
wp.copy(d.ncon, wp.array([nconmax], dtype=int))
mjwarp.reset_data(m, d)
@parameterized.parameters(*_IO_TEST_MODELS)
def test_make_put_data_dims_match(self, xml):
"""Tests that make_data and put_data have matching dimensions."""
mjm, mjd, _, _ = test_util.fixture(xml)
dm2 = mjwarp.make_data(mjm, nworld=2, nconmax=3_000, njmax=4_200)
dm3 = mjwarp.make_data(mjm, nworld=3, nconmax=3_000, njmax=4_200)
for arr in reset_datafield:
d_arr = getattr(d, arr).numpy()
for i in range(d_arr.shape[0]):
di_arr = d_arr[i]
if arr == "qM":
di_arr = di_arr.reshape(-1)[: mjd.qM.size]
_assert_eq(di_arr, getattr(mjd, arr), arr)
dp2 = mjwarp.put_data(mjm, mjd, nworld=2, nconmax=3_000, njmax=4_200)
dp3 = mjwarp.put_data(mjm, mjd, nworld=3, nconmax=3_000, njmax=4_200)
for arr in d.contact.__dataclass_fields__:
_assert_eq(getattr(d.contact, arr).numpy(), 0.0, arr)
_dims_match(self, dm2, dp2)
_dims_match(self, dm3, dp3)
def test_sdf(self):
"""Tests that an SDF can be loaded."""
mjm, mjd, m, d = test_util.fixture(fname="collision_sdf/cow.xml", qpos0=True)
@parameterized.parameters(
'<contact geom1="plane"/>',
'<contact geom2="plane"/>',
'<contact site="site"/>',
'<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)
self.assertIsInstance(m.oct_aabb, wp.array)
self.assertEqual(m.oct_aabb.dtype, wp.vec3)
self.assertEqual(len(m.oct_aabb.shape), 2)
if m.oct_aabb.size > 0:
self.assertEqual(m.oct_aabb.shape[1], 2)
if __name__ == "__main__":
@@ -24,6 +24,8 @@ 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.warp_util import event_scope
wp.set_module_options({"enable_backward": False})
@wp.kernel
def _spring_damper_dof_passive(
+4 -2
View File
@@ -24,6 +24,8 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
from mujoco.mjx.third_party.mujoco_warp._src.types import vec6
wp.set_module_options({"enable_backward": False})
@wp.func
def _ray_map(pos: wp.vec3, mat: wp.mat33, pnt: wp.vec3, vec: wp.vec3) -> Tuple[wp.vec3, wp.vec3]:
@@ -534,7 +536,7 @@ def _ray_hfield(
@wp.func
def _ray_mesh(
def ray_mesh(
# Model:
nmeshface: int,
mesh_vertadr: wp.array(dtype=int),
@@ -674,7 +676,7 @@ def _ray_geom_mesh(
type = geom_type[geomid]
if type == int(GeomType.MESH.value):
return _ray_mesh(
return ray_mesh(
nmeshface,
mesh_vertadr,
mesh_vert,
+312 -95
View File
@@ -21,8 +21,8 @@ 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 get_sdf_params
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
@@ -37,10 +37,15 @@ 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.types import vec8f
from mujoco.mjx.third_party.mujoco_warp._src.types import vec8i
from mujoco.mjx.third_party.mujoco_warp._src.util_misc import inside_geom
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope
from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel
wp.set_module_options({"enable_backward": False})
@wp.func
def _write_scalar(
@@ -1540,6 +1545,7 @@ def _sensor_acc(
dataspec = sensor_intprm[sensorid, 0]
dim = sensor_dim[sensorid]
objtype = sensor_objtype[sensorid]
reduce = sensor_intprm[sensorid, 1]
# found, force, torque, dist, pos, normal, tangent
# TODO(thowell): precompute slot size
@@ -1580,21 +1586,16 @@ def _sensor_acc(
adr = sensor_adr[sensorid]
contactsensorid = sensor_adr_to_contact_adr[sensorid]
nmatch = sensor_contact_nmatch_in[worldid, contactsensorid]
for i in range(wp.min(nmatch, num)):
# sorted contact id
cid = sensor_contact_matchid_in[worldid, contactsensorid, i]
# contact direction
dir = sensor_contact_direction_in[worldid, contactsensorid, i]
if reduce == 3: # netforce
# compute point: force-weighted centroid of contact position
net_pos = wp.vec3(0.0)
total_force_magnitude = float(0.0)
adr_slot = adr + i * size
for i in range(nmatch):
cid = sensor_contact_matchid_in[worldid, contactsensorid, i]
if found:
out[adr_slot] = float(nmatch)
adr_slot += 1
if force or torque:
contact_forcetorque = support.contact_force_fn(
opt_cone,
njmax_in,
@@ -1608,42 +1609,148 @@ def _sensor_acc(
cid,
False,
)
weight = wp.norm_l2(wp.spatial_top(contact_forcetorque))
net_pos += weight * contact_pos_in[cid]
total_force_magnitude += weight
net_pos /= wp.max(total_force_magnitude, MJ_MINVAL)
# TODO(team): iterate over matches once
# compute total wrench about point, in the global frame
net_force = wp.vec3(0.0)
net_torque = wp.vec3(0.0)
for i in range(nmatch):
cid = sensor_contact_matchid_in[worldid, contactsensorid, i]
dir = sensor_contact_direction_in[worldid, contactsensorid, i]
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,
)
contact_forcetorque *= dir
force_local = wp.spatial_top(contact_forcetorque)
torque_local = wp.spatial_bottom(contact_forcetorque)
frame = contact_frame_in[cid]
frameT = wp.transpose(frame)
force_global = frameT @ force_local
torque_global = frameT @ torque_local
# add to total force, torque
net_force += force_global
net_torque += torque_global
# add induced moment: torque += (pos - point) x force
net_torque += wp.cross(contact_pos_in[cid] - net_pos, force_global)
adr_slot = adr
if found:
out[adr_slot] = float(nmatch)
adr_slot += 1
if force:
out[adr_slot + 0] = contact_forcetorque[0]
out[adr_slot + 1] = contact_forcetorque[1]
out[adr_slot + 2] = dir * contact_forcetorque[2]
out[adr_slot + 0] = net_force[0]
out[adr_slot + 1] = net_force[1]
out[adr_slot + 2] = net_force[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]
out[adr_slot + 0] = net_torque[0]
out[adr_slot + 1] = net_torque[1]
out[adr_slot + 2] = net_torque[2]
adr_slot += 3
if dist:
out[adr_slot] = contact_dist_in[cid]
out[adr_slot] = 0.0
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]
out[adr_slot + 0] = net_pos[0]
out[adr_slot + 1] = net_pos[1]
out[adr_slot + 2] = net_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]
out[adr_slot + 0] = 1.0
out[adr_slot + 1] = 0.0
out[adr_slot + 2] = 0.0
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
out[adr_slot + 0] = 0.0
out[adr_slot + 1] = 1.0
out[adr_slot + 2] = 0.0
else:
for i in range(wp.min(nmatch, num)):
# sorted contact id
cid = sensor_contact_matchid_in[worldid, contactsensorid, i]
# zero remaining slots
for i in range(nmatch, num):
for j in range(size):
out[adr + i * size + j] = 0.0
# 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]
# 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(
@@ -1787,12 +1894,17 @@ def _sensor_tactile(
# Model:
body_rootid: wp.array(dtype=int),
body_weldid: wp.array(dtype=int),
geom_type: wp.array(dtype=int),
geom_bodyid: wp.array(dtype=int),
geom_size: wp.array2d(dtype=wp.vec3),
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),
oct_aabb: wp.array2d(dtype=wp.vec3),
oct_child: wp.array(dtype=vec8i),
oct_coeff: wp.array(dtype=vec8f),
sensor_objid: wp.array(dtype=int),
sensor_refid: wp.array(dtype=int),
sensor_dim: wp.array(dtype=int),
@@ -1852,9 +1964,15 @@ def _sensor_tactile(
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)
contact_type = geom_type[geom]
plugin_attributes, plugin_index, volume_data, mesh_data = get_sdf_params(
oct_aabb, oct_child, oct_coeff, plugin, plugin_attr, contact_type, geom_size[worldid, geom], plugin_id, mesh_id
)
depth = wp.min(sdf(contact_type, lpos, plugin_attributes, plugin_index, volume_data, mesh_data), 0.0)
if depth >= 0.0:
return
@@ -1887,18 +2005,47 @@ def _sensor_tactile(
wp.atomic_add(sensordata_out[worldid], sensor_adr[sensor_id] + 2 * dim + vertid, forceT[2])
@wp.func
def _check_match(body_parentid: wp.array(dtype=int), body: int, geom: int, objtype: int, objid: int) -> bool:
"""Check if a contact body/geom matches a sensor spec (objtype, objid)."""
if objtype == int(ObjType.UNKNOWN.value):
return True
if objtype == int(ObjType.SITE.value):
return True # already passed site filter test
if objtype == int(ObjType.GEOM.value):
return objid == geom
if objtype == int(ObjType.BODY.value):
return objid == body
if objtype == int(ObjType.XBODY.value):
# traverse up the tree from body, return true if we land on id
while body > objid:
body = body_parentid[body]
return body == objid
return False
@wp.kernel
def _contact_match(
# Model:
opt_cone: int,
opt_contact_sensor_maxmatch: int,
body_parentid: wp.array(dtype=int),
geom_bodyid: wp.array(dtype=int),
site_type: wp.array(dtype=int),
site_size: wp.array(dtype=wp.vec3),
sensor_objtype: wp.array(dtype=int),
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_contact_adr: wp.array(dtype=int),
# Data in:
njmax_in: int,
ncon_in: wp.array(dtype=int),
site_xpos_in: wp.array2d(dtype=wp.vec3),
site_xmat_in: wp.array2d(dtype=wp.mat33),
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),
@@ -1919,76 +2066,129 @@ def _contact_match(
return
# sensor information
objtype = sensor_objtype[sensorid]
objid = sensor_objid[sensorid]
reftype = sensor_reftype[sensorid]
refid = sensor_refid[sensorid]
reduce = sensor_intprm[sensorid, 1]
# contact information
geom = contact_geom_in[contactid]
worldid = contact_worldid_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]
# site filter
if objtype == int(ObjType.SITE.value):
if not inside_geom(
site_xpos_in[worldid, objid], site_xmat_in[worldid, objid], site_size[objid], site_type[objid], contact_pos_in[contactid]
):
return
contactmatchid = wp.atomic_add(sensor_contact_nmatch_out[worldid], contactsensorid, 1)
sensor_contact_matchid_out[worldid, contactsensorid, contactmatchid] = contactid
# unknown-unknown match
if objtype == int(ObjType.UNKNOWN.value) and reftype == int(ObjType.UNKNOWN.value):
dir = 1.0
else:
# contact information
geom = contact_geom_in[contactid]
geom1 = geom[0]
geom2 = geom[1]
body1 = geom_bodyid[geom1]
body2 = geom_bodyid[geom2]
if reduce == 1: # mindist
sensor_contact_criteria_out[worldid, contactsensorid, contactmatchid] = contact_dist_in[contactid]
elif reduce == 2: # maxforce
contact_force = 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,
contactid,
False,
)
force_magnitude = (
contact_force[0] * contact_force[0] + contact_force[1] * contact_force[1] + contact_force[2] * contact_force[2]
)
sensor_contact_criteria_out[worldid, contactsensorid, contactmatchid] = -force_magnitude
# TODO(thowell): netforce
# check match of sensor objects with contact objects
match11 = _check_match(body_parentid, body1, geom1, objtype, objid)
match12 = _check_match(body_parentid, body2, geom2, objtype, objid)
match21 = _check_match(body_parentid, body1, geom1, reftype, refid)
match22 = _check_match(body_parentid, body2, geom2, reftype, refid)
# contact direction
if geom1geom0:
sensor_contact_direction_out[worldid, contactsensorid, contactmatchid] = -1.0
else:
sensor_contact_direction_out[worldid, contactsensorid, contactmatchid] = 1.0
# if a sensor object is specified, it must be involved in the contact
if not match11 and not match12:
return
if not match21 and not match22:
return
# determine direction
dir = 1.0
if objtype != int(ObjType.UNKNOWN.value) and reftype != int(ObjType.UNKNOWN.value):
# both obj1 and obj2 specified: direction depends on order
order_regular = match11 and match22
order_reverse = match12 and match21
if not order_regular and not order_reverse:
return
if order_reverse and not order_regular:
dir = -1.0
elif objtype != int(ObjType.UNKNOWN.value):
if not match11:
dir = -1.0
elif reftype != int(ObjType.UNKNOWN.value):
if not match22:
dir = -1.0
contactmatchid = wp.atomic_add(sensor_contact_nmatch_out[worldid], contactsensorid, 1)
if contactmatchid >= opt_contact_sensor_maxmatch:
# TODO(team): alternative to wp.printf for reporting overflow?
wp.printf("contact match overflow: please increase Option.contact_sensor_maxmatch to %u\n", contactmatchid)
return
# TODO(thowell): alternative matching
sensor_contact_matchid_out[worldid, contactsensorid, contactmatchid] = contactid
if reduce == 1: # mindist
sensor_contact_criteria_out[worldid, contactsensorid, contactmatchid] = contact_dist_in[contactid]
elif reduce == 2: # maxforce
contact_force = 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,
contactid,
False,
)
force_magnitude = (
contact_force[0] * contact_force[0] + contact_force[1] * contact_force[1] + contact_force[2] * contact_force[2]
)
sensor_contact_criteria_out[worldid, contactsensorid, contactmatchid] = -force_magnitude
# contact direction
sensor_contact_direction_out[worldid, contactsensorid, contactmatchid] = dir
@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()
def _contact_sort(maxmatch: int):
@nested_kernel(module="unique", enable_backward=False)
def contact_sort(
# Model:
sensor_intprm: wp.array2d(dtype=int),
sensor_contact_adr: wp.array(dtype=int),
# 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]
worldid, contactsensorid = wp.tid()
sensorid = sensor_contact_adr[contactsensorid]
# skip sort
if nmatch <= 1:
return
reduce = sensor_intprm[sensorid, 1]
if reduce == 0 or reduce == 3: # none or netforce
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)
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=maxmatch)
matchid_tile = wp.tile_load(sensor_contact_matchid_in[worldid, contactsensorid], shape=maxmatch)
wp.tile_sort(criteria_tile, matchid_tile)
wp.tile_store(sensor_contact_matchid_out[worldid, contactsensorid], matchid_tile)
return contact_sort
@event_scope
@@ -2031,12 +2231,17 @@ def sensor_acc(m: Model, d: Data):
inputs=[
m.body_rootid,
m.body_weldid,
m.geom_type,
m.geom_bodyid,
m.geom_size,
m.mesh_vertadr,
m.mesh_vert,
m.mesh_normaladr,
m.mesh_normal,
m.mesh_quat,
m.oct_aabb,
m.oct_child,
m.oct_coeff,
m.sensor_objid,
m.sensor_refid,
m.sensor_dim,
@@ -2070,13 +2275,23 @@ def sensor_acc(m: Model, d: Data):
dim=(m.sensor_contact_adr.size, d.nconmax),
inputs=[
m.opt.cone,
m.opt.contact_sensor_maxmatch,
m.body_parentid,
m.geom_bodyid,
m.site_type,
m.site_size,
m.sensor_objtype,
m.sensor_objid,
m.sensor_reftype,
m.sensor_refid,
m.sensor_intprm,
m.sensor_contact_adr,
d.njmax,
d.ncon,
d.site_xpos,
d.site_xmat,
d.contact.dist,
d.contact.pos,
d.contact.frame,
d.contact.friction,
d.contact.dim,
@@ -2095,9 +2310,11 @@ def sensor_acc(m: Model, d: Data):
# sorting
wp.launch_tiled(
_contact_sort,
_contact_sort(m.opt.contact_sensor_maxmatch),
dim=(d.nworld, m.sensor_contact_adr.size),
inputs=[
m.sensor_intprm,
m.sensor_contact_adr,
d.sensor_contact_nmatch,
d.sensor_contact_matchid,
d.sensor_contact_criteria,
@@ -2370,7 +2587,7 @@ def energy_pos(m: Model, d: Data):
_energy_pos_gravity, dim=(d.nworld, m.nbody - 1), inputs=[m.opt.gravity, m.body_mass, d.xipos], outputs=[d.energy]
)
if not m.opt.disableflags & DisableBit.SPRING:
if not m.opt.disableflags & (DisableBit.SPRING | DisableBit.DAMPER):
# add joint-level springs
wp.launch(
_energy_pos_passive_joint,
@@ -2403,7 +2620,7 @@ def energy_pos(m: Model, d: Data):
@cache_kernel
def _energy_vel_kinetic(nv: int):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def energy_vel_kinetic(
# Data in:
qvel_in: wp.array2d(dtype=float),
+153 -10
View File
@@ -409,7 +409,6 @@ 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 .11 .12" euler=".02 .05 .1"',
)
@@ -422,33 +421,40 @@ class SensorTest(parameterized.TestCase):
datas = ["found", "force dist normal", "torque pos tangent", "found force torque dist pos normal tangent"]
for geoms in [
'geom2="plane"',
'geom1="plane" geom2="geom"',
'geom1="geom" geom2="plane"',
'geom1="sphere" geom2="plane"',
'geom1="geom" geom2="sphere"',
'body1="plane"',
'body1="plane" body2="geom"',
'body1="geom" body2="plane"',
]:
for num in [1, 3, 5]:
for reduce in ["mindist", "maxforce"]:
for num in [1, 5]:
for reduce in [None, "mindist", "maxforce"]:
for data in datas:
contact_sensor += f'<contact {geoms} num="{num}" reduce="{reduce}" data="{data}"/>'
contact_sensor += f'<contact {geoms} num="{num}"'
if reduce is not None:
contact_sensor += f' reduce="{reduce}"'
contact_sensor += f' data="{data}"/>'
_MJCF = f"""
<mujoco>
<compiler angle="degree"/>
<option cone="pyramidal"/>
<worldbody>
<geom name="plane" type="plane" size="10 10 .001"/>
<body>
<body name="plane">
<geom name="plane" type="plane" size="10 10 .001"/>
</body>
<body name="geom">
<geom name="geom" {geom}/>
<joint type="slide" axis="0 0 1"/>
</body>
<body>
<body name="sphere">
<geom name="sphere" type="sphere" size=".1"/>
<joint type="slide" axis="0 0 1"/>
</body>
</worldbody>
<keyframe>
<key qpos=".09 1"/>
<key qpos=".08 1"/>
</keyframe>
<sensor>
{contact_sensor}
@@ -461,6 +467,143 @@ class SensorTest(parameterized.TestCase):
d.sensordata.zero_()
mjwarp.forward(m, d)
print(f"ncon: {d.ncon.numpy()}")
sensordata = d.sensordata.numpy()[0]
_assert_eq(sensordata, mjd.sensordata, "sensordata")
self.assertTrue(sensordata.any()) # check that sensordata is not empty
def test_contact_sensor_subtree(self):
"""Test contact sensor with subtree matching semantics."""
_MJCF = f"""
<mujoco>
<worldbody>
<geom type="plane" size="2 2 .01"/>
<body name="thigh" pos="-1 0 .1">
<joint type="slide"/>
<geom type="capsule" size=".1" fromto="0 0 0 .5 0 0"/>
<body name="shin" pos=".7 0 0">
<joint axis="0 1 0"/>
<geom type="capsule" size=".1" fromto="0 0 0 0 0 .5"/>
<body name="foot" pos="0 0 .7">
<joint axis="0 1 0"/>
<geom type="capsule" size=".1" fromto="0 0 0 -.7 0 0"/>
</body>
</body>
</body>
</worldbody>
<sensor>
<contact name="all" reduce="mindist"/>
<contact name="world" subtree1="world" reduce="mindist"/>
<contact name="thigh" subtree1="thigh" reduce="mindist"/>
<contact name="shin" subtree1="shin" reduce="mindist"/>
<contact name="foot" subtree1="foot" reduce="mindist"/>
<contact name="foot_w" subtree1="foot" body2="world" reduce="mindist"/>
<contact name="foot_w2" subtree1="foot" subtree2="world" reduce="mindist"/>
</sensor>
<keyframe>
<key qpos="-6.96651e-05 -0.0478055 -0.746498"/>
</keyframe>
</mujoco>
"""
_, _, m, d = test_util.fixture(xml=_MJCF, nconmax=12, njmax=48, keyframe=0)
d.sensordata.zero_()
mjwarp.forward(m, d)
_assert_eq(d.sensordata.numpy()[0], np.array([4, 4, 4, 2, 1, 0, 1]), "found")
def test_contact_sensor_subtree_found_zero(self):
"""Test contact sensor with subtree matching semantics generates found=0."""
_MJCF = f"""
<mujoco>
<worldbody>
<geom type="plane" size="2 2 .01"/>
<body name="base">
<joint type="slide" axis="1 0 0"/>
<joint type="slide" axis="0 0 1"/>
<geom type="capsule" size=".1" fromto="0 0 0 1 0 0"/>
<body name="body0">
<joint axis="0 1 0"/>
<geom type="capsule" size=".1" fromto="0 0 0 0 0 -1"/>
</body>
<body name="body1" pos="1 0 0">
<joint axis="0 1 0"/>
<geom type="capsule" size=".1" fromto="0 0 0 0 0 -1"/>
</body>
</body>
</worldbody>
<sensor>
<contact subtree1="base" subtree2="base"/>
</sensor>
<keyframe>
<key qpos="0 1 0 0"/>
</keyframe>
</mujoco>
"""
_, _, m, d = test_util.fixture(xml=_MJCF, nconmax=2, njmax=8, keyframe=0)
d.sensordata.fill_(wp.inf)
mjwarp.forward(m, d)
_assert_eq(d.ncon.numpy()[0], 2, "ncon")
_assert_eq(d.sensordata.numpy()[0], 0, "found")
@parameterized.product(site_geom=["sphere", "capsule", "ellipsoid", "cylinder", "box"], key_pos=["0 0 10", "0 0 .09"])
def test_contact_sensor_site(self, site_geom, key_pos):
_, mjd, m, d = test_util.fixture(
xml=f"""
<mujoco>
<worldbody>
<geom type="plane" size="10 10 .001"/>
<site name="site" type="{site_geom}" size=".1 .2 .3"/>
<body>
<geom type="sphere" size=".1"/>
<freejoint/>
</body>
</worldbody>
<sensor>
<contact site="site" reduce="mindist"/>
</sensor>
<keyframe>
<key qpos="{key_pos} 1 0 0 0"/>
</keyframe>
</mujoco>
""",
keyframe=0,
)
d.sensordata.zero_()
mjwarp.forward(m, d)
_assert_eq(d.sensordata.numpy()[0], mjd.sensordata, "sensordata")
def test_contact_sensor_netforce(self):
"""Test contact sensor with netforce reduction."""
_, mjd, m, d = test_util.fixture(
xml="""
<mujoco>
<worldbody>
<geom name="plane" type="plane" size="10 10 .001"/>
<body>
<geom name="box" type="box" size=".1 .1 .1"/>
<freejoint/>
</body>
</worldbody>
<sensor>
<contact geom1="plane" geom2="box" data="found force torque dist pos normal tangent" reduce="netforce" num="2"/>
<contact geom1="box" geom2="plane" data="force torque" reduce="netforce"/>
</sensor>
<keyframe>
<key qpos="0 0 .09 1 0 0 0"/>
</keyframe>
</mujoco>
""",
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
+3 -3
View File
@@ -921,7 +921,7 @@ def _factor_i_sparse(m: Model, d: Data, M: wp.array3d(dtype=float), L: wp.array3
def _tile_cholesky_factorize(tile: TileSet):
"""Returns a kernel for dense Cholesky factorization of a tile."""
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def cholesky_factorize(
# Data In:
qM_in: wp.array3d(dtype=float),
@@ -2362,7 +2362,7 @@ def _solve_LD_sparse(
def _tile_cholesky_solve(tile: TileSet):
"""Returns a kernel for dense Cholesky backsubstitution of a tile."""
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def cholesky_solve(
# In:
L: wp.array3d(dtype=float),
@@ -2441,7 +2441,7 @@ def solve_m(m: Model, d: Data, x: wp.array2d(dtype=float), y: wp.array2d(dtype=f
def _tile_cholesky_factorize_solve(tile: TileSet):
"""Returns a kernel for dense Cholesky factorization and backsubstitution of a tile."""
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def cholesky_factorize_solve(
# In:
M: wp.array3d(dtype=float),
+4 -4
View File
@@ -702,7 +702,7 @@ def linesearch_zero_jv(
@cache_kernel
def linesearch_jv_fused(nv: int, dofs_per_thread: int):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def kernel(
# Data in:
nefc_in: wp.array(dtype=int),
@@ -1211,7 +1211,7 @@ def update_constraint_init_qfrc_constraint(
@cache_kernel
def update_constraint_gauss_cost(nv: int, dofs_per_thread: int):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def kernel(
# Data in:
qacc_in: wp.array2d(dtype=float),
@@ -1603,7 +1603,7 @@ def update_gradient_JTCJ(
@cache_kernel
def update_gradient_cholesky(tile_size: int):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def kernel(
# Data in:
efc_grad_in: wp.array2d(dtype=float),
@@ -1629,7 +1629,7 @@ def update_gradient_cholesky(tile_size: int):
@cache_kernel
def update_gradient_cholesky_blocked(tile_size: int):
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def kernel(
# Data in:
efc_grad_in: wp.array3d(dtype=float),
+1 -1
View File
@@ -86,7 +86,7 @@ def mul_m_sparse_ij(
def mul_m_dense(tile: TileSet):
"""Returns a matmul kernel for some tile size"""
@nested_kernel
@nested_kernel(module="unique", enable_backward=False)
def kernel(
# Data In:
qM_in: wp.array3d(dtype=float),
+66 -6
View File
@@ -28,6 +28,7 @@ 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 BroadphaseType
from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType
from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit
@@ -62,6 +63,7 @@ def fixture(
ls_iterations: Optional[int] = None,
ls_parallel: Optional[bool] = None,
sparse: Optional[bool] = None,
broadphase: Optional[BroadphaseType] = None,
disableflags: Optional[int] = None,
enableflags: Optional[int] = None,
applied: bool = False,
@@ -153,11 +155,51 @@ def fixture(
m = io.put_model(mjm)
if ls_parallel is not None:
m.opt.ls_parallel = ls_parallel
if broadphase is not None:
m.opt.broadphase = broadphase
d = io.put_data(mjm, mjd, nworld=nworld, nconmax=nconmax, njmax=njmax)
return mjm, mjd, m, d
def find_keys(model: mujoco.MjModel, keyname_prefix: str) -> list[int]:
"""Finds keyframes that start with keyname_prefix."""
keys = []
for keyid in range(model.nkey):
name = mujoco.mj_id2name(model, mujoco.mjtObj.mjOBJ_KEY, keyid)
if name.startswith(keyname_prefix):
keys.append(keyid)
return keys
def make_trajectory(model: mujoco.MjModel, keys: list[int]) -> np.ndarray:
"""Make a ctrl trajectory with linear interpolation."""
ctrls = []
prev_ctrl_key = np.zeros(model.nu, dtype=np.float64)
prev_time, time = 0.0, 0.0
for key in keys:
ctrl_key, ctrl_time = model.key_ctrl[key], model.key_time[key]
if not ctrls and ctrl_time != 0.0:
raise ValueError("first keyframe must have time 0.0")
elif ctrls and ctrl_time <= prev_time:
raise ValueError("keyframes must be in time order")
while time < ctrl_time:
frac = (time - prev_time) / (ctrl_time - prev_time)
ctrls.append(prev_ctrl_key * (1 - frac) + ctrl_key * frac)
time += model.opt.timestep
ctrls.append(ctrl_key)
time += model.opt.timestep
prev_ctrl_key = ctrl_key
prev_time = time
return np.array(ctrls)
def _sum(stack1, stack2):
ret = {}
for k in stack1:
@@ -174,6 +216,7 @@ def ctrl_noise(
actuator_ctrllimited: wp.array(dtype=bool),
actuator_ctrlrange: wp.array2d(dtype=wp.vec2),
# In:
ctrl_center: wp.array1d(dtype=float),
step: int,
ctrlnoise: float,
# Data out:
@@ -184,7 +227,9 @@ def ctrl_noise(
center = 0.0
radius = 1.0
ctrlrange = actuator_ctrlrange[0, actid]
if actuator_ctrllimited[actid]:
if ctrl_center.shape[0] > 0:
center = ctrl_center[actid]
elif actuator_ctrllimited[actid]:
center = (ctrlrange[1] + ctrlrange[0]) / 2.0
radius = (ctrlrange[1] - ctrlrange[0]) / 2.0
radius *= ctrlnoise
@@ -197,6 +242,7 @@ def benchmark(
m: Model,
d: Data,
nstep: int,
ctrls: Optional[np.ndarray] = None,
event_trace: bool = False,
measure_alloc: bool = False,
measure_solver_niter: bool = False,
@@ -208,6 +254,8 @@ def benchmark(
m (Model): The model containing kinematic and dynamic information (device).
d (Data): The data object containing the current state and output information (device).
nstep (int): Number of timesteps.
ctrls (list, optional): control sequence to apply during benchmarking.
Default is None.
event_trace (bool, optional): If True, time routines decorated with @event_scope.
Default is False.
measure_alloc (bool, optional): If True, record number of contacts and constraints.
@@ -226,6 +274,7 @@ def benchmark(
trace = {}
ncon, nefc, solver_niter = [], [], []
center = wp.array([], dtype=wp.float32)
with warp_util.EventTracer(enabled=event_trace) as tracer:
# capture the whole function as a CUDA graph
@@ -240,13 +289,14 @@ def benchmark(
time_vec = np.zeros(nstep)
for i in range(nstep):
with wp.ScopedStream(wp.get_stream()):
if ctrls is not None:
center = wp.array(ctrls[i], dtype=wp.float32)
wp.launch(
ctrl_noise,
dim=(d.nworld, m.nu),
inputs=[
m.actuator_ctrllimited, m.actuator_ctrlrange, i, 0.01
],
outputs=[d.ctrl]) # fmt: skip
inputs=[m.actuator_ctrllimited, m.actuator_ctrlrange, center, i, 0.01],
outputs=[d.ctrl],
)
wp.synchronize()
run_beg = time.perf_counter()
@@ -312,12 +362,22 @@ class BenchmarkSuite:
rounds = 1
sample_time = 0
repeat = 1
replay = ""
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)
ctrls = None
if self.replay:
keys = find_keys(mjm, self.replay)
if not keys:
raise ValueError(f"Key prefix not find: {self.replay}")
ctrls = make_trajectory(mjm, keys)
mujoco.mj_resetDataKeyframe(mjm, mjd, keys[0])
if mjm.nkey > 0:
mujoco.mj_resetDataKeyframe(mjm, mjd, 0)
@@ -333,7 +393,7 @@ class BenchmarkSuite:
d = io.put_data(mjm, mjd, self.batch_size, self.nconmax, self.njmax)
free_after = wp.get_device().free_memory
jit_duration, _, trace, _, _, solver_niter, _ = benchmark(forward.step, m, d, 1000, True, False, True)
jit_duration, _, trace, _, _, solver_niter, _ = benchmark(forward.step, m, d, 1000, ctrls, True, False, True)
metrics = {
"jit_duration": jit_duration,
"solver_niter_mean": np.mean(solver_niter),
+43 -9
View File
@@ -522,6 +522,14 @@ class vec6f(wp.types.vector(length=6, dtype=float)):
pass
class vec8f(wp.types.vector(length=8, dtype=float)):
pass
class vec8i(wp.types.vector(length=8, dtype=int)):
pass
class vec10f(wp.types.vector(length=10, dtype=float)):
pass
@@ -545,6 +553,7 @@ class Option:
impratio: ratio of friction-to-normal contact impedance
tolerance: main solver tolerance
ls_tolerance: CG/Newton linesearch tolerance
ccd_tolerance: convex collision solver tolerance
gravity: gravitational acceleration
magnetic: global magnetic flux
integrator: integration mode (IntegratorType)
@@ -572,12 +581,15 @@ class Option:
contacts during the physics step (as opposed to DisableBit.CONTACT which explicitly
zeros out the contacts at each step)
legacy_gjk: run legacy gjk algorithm
contact_sensor_maxmatch: max number of contacts considered by contact sensor matching criteria
contacts matched after this value is exceded will be ignored
"""
timestep: wp.array(dtype=float)
impratio: wp.array(dtype=float)
tolerance: wp.array(dtype=float)
ls_tolerance: wp.array(dtype=float)
ccd_tolerance: wp.array(dtype=float)
gravity: wp.array(dtype=wp.vec3)
magnetic: wp.array(dtype=wp.vec3)
integrator: int
@@ -603,6 +615,7 @@ class Option:
sdf_iterations: int
run_collision_detection: bool # warp only
legacy_gjk: bool
contact_sensor_maxmatch: int # warp only
@dataclasses.dataclass
@@ -885,6 +898,9 @@ 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,)
oct_aabb: octree axis-aligned bounding boxes (noct, 6)
oct_child: octree children (noct, 8)
oct_coeff: octree interpolation coefficients (noct, 8)
eq_type: constraint type (EqType) (neq,)
eq_obj1id: id of object 1 (neq,)
eq_obj2id: id of object 2 (neq,)
@@ -1010,8 +1026,6 @@ class Model:
mat_rgba: rgba (nworld, nmat, 4)
actuator_trntype_body_adr: addresses for actuators (<=nu,)
with body transmission
geompair2hfgeompair: geom pair to geom pair with (ngeom * (ngeom - 1) // 2,)
height field mapping
block_dim: BlockDim
geom_pair_type_count: count of max number of each potential collision
has_sdf_geom: whether the model contains SDF geoms
@@ -1208,6 +1222,9 @@ class Model:
mesh_polymapadr: wp.array(dtype=int)
mesh_polymapnum: wp.array(dtype=int)
mesh_polymap: wp.array(dtype=int)
oct_aabb: wp.array2d(dtype=wp.vec3)
oct_child: wp.array(dtype=vec8i)
oct_coeff: wp.array(dtype=vec8f)
eq_type: wp.array(dtype=int)
eq_obj1id: wp.array(dtype=int)
eq_obj2id: wp.array(dtype=int)
@@ -1326,7 +1343,6 @@ class Model:
mat_texrepeat: wp.array2d(dtype=wp.vec2)
mat_rgba: wp.array2d(dtype=wp.vec4)
actuator_trntype_body_adr: wp.array(dtype=int) # warp only
geompair2hfgeompair: wp.array(dtype=int) # warp only
block_dim: BlockDim # warp only
geom_pair_type_count: tuple[int, ...] # warp only
has_sdf_geom: bool # warp only
@@ -1377,8 +1393,6 @@ class Data:
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 (nworld,)
ne_connect: number of equality connect constraints (nworld,)
ne_weld: number of equality weld constraints (nworld,)
@@ -1475,7 +1489,6 @@ class Data:
sap_segment_index: broadphase context (requires nworld + 1) (nworld, 2)
dyn_geom_aabb: dynamic geometry axis-aligned bounding boxes (nworld, ngeom, 2)
collision_pair: collision pairs from broadphase (nconmax,)
collision_hftri_index: collision index for hfield pairs (nconmax,)
collision_worldid: collision world ids from broadphase (nconmax,)
ncollision: collision count from broadphase
epa_vert: vertices in EPA polytope in Minkowski space (nconmax, 5 + CCDiter)
@@ -1489,6 +1502,17 @@ class Data:
epa_index: index of face in polytope map (nconmax, 6 + 6 * CCDiter)
epa_map: status of faces in polytope (nconmax, 6 + 6 * CCDiter)
epa_horizon: index pair (i j) of edges on horizon (nconmax, 3 * 2 * CCDiter)
multiccd_polygon: clipped contact surface (nconmax, 2 * max_npolygon)
multiccd_clipped: clipped contact surface (intermediate) (nconmax, 2 * max_npolygon)
multiccd_pnormal: plane normal of clipping polygon (nconmax, max_npolygon)
multiccd_pdist: plane distance of clipping polygon (nconmax, max_npolygon)
multiccd_idx1: list of normal index candidates for Geom 1 (nconmax, max_meshdegree)
multiccd_idx2: list of normal index candidates for Geom 2 (nconmax, max_meshdegree)
multiccd_n1: list of normal candidates for Geom 1 (nconmax, max_meshdegree)
multiccd_n2: list of normal candidates for Geom 1 (nconmax, max_meshdegree)
multiccd_endvert: list of edge vertices candidates (nconmax, max_meshdegree)
multiccd_face1: contact face (nconmax, max_npolygon)
multiccd_face2: contact face (nconmax, max_npolygon)
cacc: com-based acceleration (nworld, nbody, 6)
cfrc_int: com-based interaction force with parent (nworld, nbody, 6)
cfrc_ext: com-based external force on body (nworld, nbody, 6)
@@ -1524,8 +1548,6 @@ 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
ne_weld: wp.array(dtype=int) # warp only
@@ -1627,7 +1649,6 @@ class Data:
# collision driver
collision_pair: wp.array(dtype=wp.vec2i)
collision_hftri_index: wp.array(dtype=int)
collision_pairid: wp.array(dtype=int)
collision_worldid: wp.array(dtype=int)
ncollision: wp.array(dtype=int)
@@ -1645,6 +1666,19 @@ class Data:
epa_map: wp.array2d(dtype=int)
epa_horizon: wp.array2d(dtype=int)
# narrowphase collision (multicontact)
multiccd_polygon: wp.array2d(dtype=wp.vec3)
multiccd_clipped: wp.array2d(dtype=wp.vec3)
multiccd_pnormal: wp.array2d(dtype=wp.vec3)
multiccd_pdist: wp.array2d(dtype=float)
multiccd_idx1: wp.array2d(dtype=int)
multiccd_idx2: wp.array2d(dtype=int)
multiccd_n1: wp.array2d(dtype=wp.vec3)
multiccd_n2: wp.array2d(dtype=wp.vec3)
multiccd_endvert: wp.array2d(dtype=wp.vec3)
multiccd_face1: wp.array2d(dtype=wp.vec3)
multiccd_face2: wp.array2d(dtype=wp.vec3)
# rne_postconstraint
cacc: wp.array2d(dtype=wp.spatial_vector)
cfrc_int: wp.array2d(dtype=wp.spatial_vector)
@@ -21,6 +21,7 @@ import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src import math
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import WrapType
from mujoco.mjx.third_party.mujoco_warp._src.types import vec10
@@ -601,3 +602,36 @@ def muscle_dynamics(control: float, activation: float, prm: vec10) -> float:
# filter output
return dctrl / wp.max(MJ_MINVAL, tau)
@wp.func
def inside_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, geomtype: int, point: wp.vec3) -> bool:
"""Return True if point is inside primitive geom, False otherwise."""
# vector from geom to point
vec = point - pos
# quick return for spheres, frame rotation not required
if geomtype == int(GeomType.SPHERE.value):
return wp.dot(vec, vec) < size[0] * size[0]
# rotate into local frame
plocal = wp.transpose(mat) @ vec
# handle other geom types
if geomtype == int(GeomType.CAPSULE.value):
z = plocal[2]
z_clamped = wp.clamp(z, -size[1], size[1])
z_dif = z - z_clamped
z_dist_sq = z_dif * z_dif
return plocal[0] * plocal[0] + plocal[1] * plocal[1] + z_dist_sq < size[0] * size[0]
elif geomtype == int(GeomType.ELLIPSOID.value):
plocalsize = wp.cw_div(plocal, size)
return wp.dot(plocalsize, plocalsize) < 1.0
elif geomtype == int(GeomType.CYLINDER.value):
return (wp.abs(plocal[2]) < size[1]) and (plocal[0] * plocal[0] + plocal[1] * plocal[1] < size[0] * size[0])
elif geomtype == int(GeomType.BOX.value):
return wp.abs(plocal[0]) < size[0] and wp.abs(plocal[1]) < size[1] and wp.abs(plocal[2]) < size[2]
elif geomtype == int(GeomType.PLANE.value):
return plocal[2] < 0.0
return False
@@ -0,0 +1,2 @@
The spot assets were taken from https://www.cs.cmu.edu/~kmcrane/Projects/ModelRepository/ and are
released under the CC0 1.0 Universal (CC0 1.0) Public Domain Dedication license.
@@ -0,0 +1,296 @@
v 0.740000 0.740000 -1.000000
v 0.740000 1.000000 -0.740000
v 1.000000 0.740000 -0.740000
v 0.740000 0.807269 -0.991147
v 0.740000 0.870021 -0.965154
v 0.740000 0.923848 -0.923848
v 0.807269 0.740000 -0.991147
v 0.809387 0.809958 -0.980606
v 0.808254 0.870420 -0.954318
v 0.806858 0.917583 -0.917748
v 0.870021 0.740000 -0.965154
v 0.870490 0.808646 -0.954150
v 0.863847 0.863908 -0.932116
v 0.858399 0.903670 -0.903688
v 0.890111 0.890111 -0.890111
v 0.807269 0.991147 -0.740000
v 0.870021 0.965154 -0.740000
v 0.923848 0.923848 -0.740000
v 0.740000 0.991147 -0.807269
v 0.809958 0.980606 -0.809387
v 0.870420 0.954318 -0.808254
v 0.917583 0.917748 -0.806858
v 0.740000 0.965154 -0.870021
v 0.808646 0.954150 -0.870490
v 0.863908 0.932116 -0.863847
v 0.903670 0.903688 -0.858399
v 0.991147 0.740000 -0.807269
v 0.965154 0.740000 -0.870021
v 0.923848 0.740000 -0.923848
v 0.991147 0.807269 -0.740000
v 0.980606 0.809387 -0.809958
v 0.954318 0.808254 -0.870420
v 0.917748 0.806858 -0.917583
v 0.965154 0.870021 -0.740000
v 0.954150 0.870490 -0.808646
v 0.932116 0.863847 -0.863908
v 0.903688 0.858399 -0.903670
v 0.740000 -1.000000 -0.740000
v 0.740000 -0.740000 -1.000000
v 1.000000 -0.740000 -0.740000
v 0.740000 -0.991147 -0.807269
v 0.740000 -0.965154 -0.870021
v 0.740000 -0.923848 -0.923848
v 0.807269 -0.991147 -0.740000
v 0.809387 -0.980606 -0.809958
v 0.808254 -0.954318 -0.870420
v 0.806858 -0.917748 -0.917583
v 0.870021 -0.965154 -0.740000
v 0.870490 -0.954150 -0.808646
v 0.863847 -0.932116 -0.863908
v 0.858399 -0.903688 -0.903670
v 0.890111 -0.890111 -0.890111
v 0.807269 -0.740000 -0.991147
v 0.870021 -0.740000 -0.965154
v 0.923848 -0.740000 -0.923848
v 0.740000 -0.807269 -0.991147
v 0.809958 -0.809387 -0.980606
v 0.870420 -0.808254 -0.954318
v 0.917583 -0.806858 -0.917748
v 0.740000 -0.870021 -0.965154
v 0.808646 -0.870490 -0.954150
v 0.863908 -0.863847 -0.932116
v 0.903670 -0.858399 -0.903688
v 0.991147 -0.807269 -0.740000
v 0.965154 -0.870021 -0.740000
v 0.923848 -0.923848 -0.740000
v 0.991147 -0.740000 -0.807269
v 0.980606 -0.809958 -0.809387
v 0.954318 -0.870420 -0.808254
v 0.917748 -0.917583 -0.806858
v 0.965154 -0.740000 -0.870021
v 0.954150 -0.808646 -0.870490
v 0.932116 -0.863908 -0.863847
v 0.903688 -0.903670 -0.858399
v 1.000000 0.740000 0.740000
v 0.740000 1.000000 0.740000
v 0.740000 0.740000 1.000000
v 0.991147 0.807269 0.740000
v 0.965154 0.870021 0.740000
v 0.923848 0.923848 0.740000
v 0.991147 0.740000 0.807269
v 0.980606 0.809958 0.809387
v 0.954318 0.870420 0.808254
v 0.917748 0.917583 0.806858
v 0.965154 0.740000 0.870021
v 0.954150 0.808646 0.870490
v 0.932116 0.863908 0.863847
v 0.903688 0.903670 0.858399
v 0.890111 0.890111 0.890111
v 0.740000 0.991147 0.807269
v 0.740000 0.965154 0.870021
v 0.740000 0.923848 0.923848
v 0.807269 0.991147 0.740000
v 0.809387 0.980606 0.809958
v 0.808254 0.954318 0.870420
v 0.806858 0.917748 0.917583
v 0.870021 0.965154 0.740000
v 0.870490 0.954150 0.808646
v 0.863847 0.932116 0.863908
v 0.858399 0.903688 0.903670
v 0.807269 0.740000 0.991147
v 0.870021 0.740000 0.965154
v 0.923848 0.740000 0.923848
v 0.740000 0.807269 0.991147
v 0.809958 0.809387 0.980606
v 0.870420 0.808254 0.954318
v 0.917583 0.806858 0.917748
v 0.740000 0.870021 0.965154
v 0.808646 0.870490 0.954150
v 0.863908 0.863847 0.932116
v 0.903670 0.858399 0.903688
v 1.000000 -0.740000 0.740000
v 0.740000 -0.740000 1.000000
v 0.740000 -1.000000 0.740000
v 0.991147 -0.740000 0.807269
v 0.965154 -0.740000 0.870021
v 0.923848 -0.740000 0.923848
v 0.991147 -0.807269 0.740000
v 0.980606 -0.809387 0.809958
v 0.954318 -0.808254 0.870420
v 0.917748 -0.806858 0.917583
v 0.965154 -0.870021 0.740000
v 0.954150 -0.870490 0.808646
v 0.932116 -0.863847 0.863908
v 0.903688 -0.858399 0.903670
v 0.890111 -0.890111 0.890111
v 0.740000 -0.807269 0.991147
v 0.740000 -0.870021 0.965154
v 0.740000 -0.923848 0.923848
v 0.807269 -0.740000 0.991147
v 0.809387 -0.809958 0.980606
v 0.808254 -0.870420 0.954318
v 0.806858 -0.917583 0.917748
v 0.870021 -0.740000 0.965154
v 0.870490 -0.808646 0.954150
v 0.863847 -0.863908 0.932116
v 0.858399 -0.903670 0.903688
v 0.807269 -0.991147 0.740000
v 0.870021 -0.965154 0.740000
v 0.923848 -0.923848 0.740000
v 0.740000 -0.991147 0.807269
v 0.809958 -0.980606 0.809387
v 0.870420 -0.954318 0.808254
v 0.917583 -0.917748 0.806858
v 0.740000 -0.965154 0.870021
v 0.808646 -0.954150 0.870490
v 0.863908 -0.932116 0.863847
v 0.903670 -0.903688 0.858399
v -0.740000 0.740000 -1.000000
v -1.000000 0.740000 -0.740000
v -0.740000 1.000000 -0.740000
v -0.807269 0.740000 -0.991147
v -0.870021 0.740000 -0.965154
v -0.923848 0.740000 -0.923848
v -0.740000 0.807269 -0.991147
v -0.809958 0.809387 -0.980606
v -0.870420 0.808254 -0.954318
v -0.917583 0.806858 -0.917748
v -0.740000 0.870021 -0.965154
v -0.808646 0.870490 -0.954150
v -0.863908 0.863847 -0.932116
v -0.903670 0.858399 -0.903688
v -0.890111 0.890111 -0.890111
v -0.991147 0.807269 -0.740000
v -0.965154 0.870021 -0.740000
v -0.923848 0.923848 -0.740000
v -0.991147 0.740000 -0.807269
v -0.980606 0.809958 -0.809387
v -0.954318 0.870420 -0.808254
v -0.917748 0.917583 -0.806858
v -0.965154 0.740000 -0.870021
v -0.954150 0.808646 -0.870490
v -0.932116 0.863908 -0.863847
v -0.903688 0.903670 -0.858399
v -0.740000 0.991147 -0.807269
v -0.740000 0.965154 -0.870021
v -0.740000 0.923848 -0.923848
v -0.807269 0.991147 -0.740000
v -0.809387 0.980606 -0.809958
v -0.808254 0.954318 -0.870420
v -0.806858 0.917748 -0.917583
v -0.870021 0.965154 -0.740000
v -0.870490 0.954150 -0.808646
v -0.863847 0.932116 -0.863908
v -0.858399 0.903688 -0.903670
v -1.000000 -0.740000 -0.740000
v -0.740000 -0.740000 -1.000000
v -0.740000 -1.000000 -0.740000
v -0.991147 -0.740000 -0.807269
v -0.965154 -0.740000 -0.870021
v -0.923848 -0.740000 -0.923848
v -0.991147 -0.807269 -0.740000
v -0.980606 -0.809387 -0.809958
v -0.954318 -0.808254 -0.870420
v -0.917748 -0.806858 -0.917583
v -0.965154 -0.870021 -0.740000
v -0.954150 -0.870490 -0.808646
v -0.932116 -0.863847 -0.863908
v -0.903688 -0.858399 -0.903670
v -0.890111 -0.890111 -0.890111
v -0.740000 -0.807269 -0.991147
v -0.740000 -0.870021 -0.965154
v -0.740000 -0.923848 -0.923848
v -0.807269 -0.740000 -0.991147
v -0.809387 -0.809958 -0.980606
v -0.808254 -0.870420 -0.954318
v -0.806858 -0.917583 -0.917748
v -0.870021 -0.740000 -0.965154
v -0.870490 -0.808646 -0.954150
v -0.863847 -0.863908 -0.932116
v -0.858399 -0.903670 -0.903688
v -0.807269 -0.991147 -0.740000
v -0.870021 -0.965154 -0.740000
v -0.923848 -0.923848 -0.740000
v -0.740000 -0.991147 -0.807269
v -0.809958 -0.980606 -0.809387
v -0.870420 -0.954318 -0.808254
v -0.917583 -0.917748 -0.806858
v -0.740000 -0.965154 -0.870021
v -0.808646 -0.954150 -0.870490
v -0.863908 -0.932116 -0.863847
v -0.903670 -0.903688 -0.858399
v -1.000000 0.740000 0.740000
v -0.740000 0.740000 1.000000
v -0.740000 1.000000 0.740000
v -0.991147 0.740000 0.807269
v -0.965154 0.740000 0.870021
v -0.923848 0.740000 0.923848
v -0.991147 0.807269 0.740000
v -0.980606 0.809387 0.809958
v -0.954318 0.808254 0.870420
v -0.917748 0.806858 0.917583
v -0.965154 0.870021 0.740000
v -0.954150 0.870490 0.808646
v -0.932116 0.863847 0.863908
v -0.903688 0.858399 0.903670
v -0.890111 0.890111 0.890111
v -0.740000 0.807269 0.991147
v -0.740000 0.870021 0.965154
v -0.740000 0.923848 0.923848
v -0.807269 0.740000 0.991147
v -0.809387 0.809958 0.980606
v -0.808254 0.870420 0.954318
v -0.806858 0.917583 0.917748
v -0.870021 0.740000 0.965154
v -0.870490 0.808646 0.954150
v -0.863847 0.863908 0.932116
v -0.858399 0.903670 0.903688
v -0.807269 0.991147 0.740000
v -0.870021 0.965154 0.740000
v -0.923848 0.923848 0.740000
v -0.740000 0.991147 0.807269
v -0.809958 0.980606 0.809387
v -0.870420 0.954318 0.808254
v -0.917583 0.917748 0.806858
v -0.740000 0.965154 0.870021
v -0.808646 0.954150 0.870490
v -0.863908 0.932116 0.863847
v -0.903670 0.903688 0.858399
v -0.740000 -1.000000 0.740000
v -0.740000 -0.740000 1.000000
v -1.000000 -0.740000 0.740000
v -0.740000 -0.991147 0.807269
v -0.740000 -0.965154 0.870021
v -0.740000 -0.923848 0.923848
v -0.807269 -0.991147 0.740000
v -0.809387 -0.980606 0.809958
v -0.808254 -0.954318 0.870420
v -0.806858 -0.917748 0.917583
v -0.870021 -0.965154 0.740000
v -0.870490 -0.954150 0.808646
v -0.863847 -0.932116 0.863908
v -0.858399 -0.903688 0.903670
v -0.890111 -0.890111 0.890111
v -0.807269 -0.740000 0.991147
v -0.870021 -0.740000 0.965154
v -0.923848 -0.740000 0.923848
v -0.740000 -0.807269 0.991147
v -0.809958 -0.809387 0.980606
v -0.870420 -0.808254 0.954318
v -0.917583 -0.806858 0.917748
v -0.740000 -0.870021 0.965154
v -0.808646 -0.870490 0.954150
v -0.863908 -0.863847 0.932116
v -0.903670 -0.858399 0.903688
v -0.991147 -0.807269 0.740000
v -0.965154 -0.870021 0.740000
v -0.923848 -0.923848 0.740000
v -0.991147 -0.740000 0.807269
v -0.980606 -0.809958 0.809387
v -0.954318 -0.870420 0.808254
v -0.917748 -0.917583 0.806858
v -0.965154 -0.740000 0.870021
v -0.954150 -0.808646 0.870490
v -0.932116 -0.863908 0.863847
v -0.903688 -0.903670 0.858399
File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 77 KiB

@@ -0,0 +1,53 @@
<mujoco>
<compiler texturedir="asset"/>
<extension>
<plugin plugin="mujoco.sdf.torus">
<instance name="torus">
<config key="radius1" value="0.15"/>
<config key="radius2" value="0.05"/>
</instance>
</plugin>
</extension>
<asset>
<texture name="texspot" type="2d" file="spot.png"/>
<material name="matspot" texture="texspot"/>
<mesh name="spot" file="asset/spot.obj"/>
<mesh name="torus">
<plugin instance="torus"/>
</mesh>
</asset>
<option sdf_iterations="10" sdf_initpoints="40"/>
<visual>
<map force="1000"/>
</visual>
<default>
<geom solref="0.01 1" solimp=".95 .99 .0001" friction="0.5"/>
</default>
<statistic meansize="0.2"/>
<include file="scene.xml"/>
<worldbody>
<body pos="0.1 .25 5.7">
<freejoint/>
<geom type="sdf" mesh="torus" rgba=".2 .8 .2 1">
<plugin instance="torus"/>
</geom>
</body>
<body euler="90 0 0" pos="0 0 .7">
<geom type="sdf" name="cow1" mesh="spot" material="matspot"/>
</body>
<body pos="0.05 .25 2.2">
<freejoint/>
<geom type="sdf" name="cow2" mesh="spot" material="matspot"/>
</body>
<light name="left" pos="0 0 1"/>
<light name="right" pos="1 0 1"/>
</worldbody>
</mujoco>
@@ -0,0 +1,34 @@
<mujoco>
<extension>
<plugin plugin="mujoco.sdf.torus">
<instance name="torus">
<config key="radius1" value="0.35"/>
<config key="radius2" value="0.15"/>
</instance>
</plugin>
</extension>
<option gravity="0 0 -9.81" sdf_iterations="3" sdf_initpoints="10"/>
<asset>
<mesh name="torus">
<plugin instance="torus"/>
</mesh>
<mesh name="die" file="asset/die.obj" scale="1 1 1"/>
</asset>
<include file="scene.xml"/>
<worldbody>
<body pos="0 .05 2.5" euler="90 0 0">
<freejoint/>
<geom type="sdf" mesh="torus" rgba=".8 .17 .15 1" group="1">
<plugin instance="torus"/>
</geom>
</body>
<body>
<geom type="mesh" mesh="die" euler="90 0 0" rgba="0 0 1 .2"/>
</body>
<light pos="1 0 7" dir="0 0 -1" castshadow="false"/>
<light pos="-1 0 7" dir="0 0 -1" castshadow="false"/>
</worldbody>
</mujoco>
@@ -0,0 +1,34 @@
import math
import warp as wp
@wp.func
def torus(p: wp.vec3, attr: wp.vec3) -> wp.float32:
major_radius = attr[0]
minor_radius = attr[1]
q = math.sqrt(p[0] * p[0] + p[1] * p[1]) - major_radius
sdf = math.sqrt(q * q + p[2] * p[2]) - minor_radius
return sdf
@wp.func
def torus_sdf_grad(p: wp.vec3, attr: wp.vec3) -> wp.vec3:
grad = wp.vec3()
major_radius = attr[0]
minor_val = attr[1]
len_xy = math.sqrt(p[0] * p[0] + p[1] * p[1])
q = len_xy - major_radius
len_xy = wp.max(len_xy, 1e-8)
grad_q_x = p[0] / len_xy
grad_q_y = p[1] / len_xy
len_qz = math.sqrt(q * q + p[2] * p[2])
denom = wp.max(len_qz, wp.max(minor_val, 1e-8))
grad[0] = q * grad_q_x / denom
grad[1] = q * grad_q_y / denom
grad[2] = p[2] / denom
return grad
@@ -0,0 +1,8 @@
<mujoco>
<asset>
<mesh file="torus.obj"/>
</asset>
<worldbody>
<geom type="mesh" mesh="torus"/>
</worldbody>
</mujoco>
@@ -26,6 +26,8 @@ from .gear import gear
from .gear import gear_sdf_grad
from .nut import nut
from .nut import nut_sdf_grad
from .torus import torus
from .torus import torus_sdf_grad
class SDFType(enum.Enum):
@@ -33,6 +35,7 @@ class SDFType(enum.Enum):
NUT = "NUT"
BOLT = "BOLT"
TORUS = "TORUS"
GEAR = "GEAR"
@@ -41,16 +44,19 @@ def register_sdf_plugins(mjwarp) -> Dict[str, int]:
<extension>
<plugin plugin="mujoco.sdf.nut"><instance name="n"/></plugin>
<plugin plugin="mujoco.sdf.bolt"><instance name="b"/></plugin>
<plugin plugin="mujoco.sdf.torus"><instance name="t"/></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="tm"><plugin instance="t"/></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="tg" mesh="tm"><plugin instance="t"/></geom></body>
<body><geom type="sdf" name="gg" mesh="gm"><plugin instance="g"/></geom></body>
</worldbody>
</mujoco>"""
@@ -68,6 +74,9 @@ def register_sdf_plugins(mjwarp) -> 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 == "tg":
sdf_types[SDFType.TORUS.value] = int(m.plugin[i])
elif name == "gg":
sdf_types[SDFType.GEAR.value] = int(m.plugin[i])
@@ -78,6 +87,8 @@ def register_sdf_plugins(mjwarp) -> 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.TORUS.value]):
result = torus(p, attr)
elif sdf_type == wp.static(sdf_types[SDFType.GEAR.value]):
result = gear(p, attr)
return result
@@ -88,6 +99,8 @@ def register_sdf_plugins(mjwarp) -> 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.TORUS.value]):
return torus_sdf_grad(p, attr)
elif sdf_type == wp.static(sdf_types[SDFType.GEAR.value]):
return gear_sdf_grad(p, attr)
return wp.vec3()
+24 -14
View File
@@ -22,9 +22,9 @@ Example:
"""
import ast
import copy
import enum
import logging
import pickle
import sys
import time
from typing import Sequence, Union
@@ -52,7 +52,7 @@ _NJMAX = flags.DEFINE_integer("njmax", None, "Maximum number of constraints per
_OVERRIDE = flags.DEFINE_multi_string("override", [], "Model overrides (notation: foo.bar = baz)", short_name="o")
_KEYFRAME = flags.DEFINE_integer("keyframe", 0, "keyframe to initialize simulation.")
_DEVICE = flags.DEFINE_string("device", None, "override the default Warp device")
_REPLAY = flags.DEFINE_string("replay", None, "keyframe sequence to replay, keyframe name must prefix match")
_VIEWER_GLOBAL_STATE = {"running": True, "step_once": False}
@@ -120,12 +120,13 @@ def _override(model: Union[mjw.Model, mujoco.MjModel]):
def _compile_step(m, d):
mjw.step(m, d)
# double warmup to work around issues with compilation during graph capture:
mjw.step(m, d)
print("Compiling physics step...", end="", flush=True)
start = time.time()
# capture the whole step function as a CUDA graph
with wp.ScopedCapture() as capture:
mjw.step(m, d)
elapsed = time.time() - start
print(f"done ({elapsed:0.2g}s).")
return capture.graph
@@ -138,7 +139,15 @@ def _main(argv: Sequence[str]) -> None:
mjm = _load_model(epath.Path(argv[1]))
mjd = mujoco.MjData(mjm)
if mjm.nkey > 0 and _KEYFRAME.value > -1:
ctrls = None
ctrlid = 0
if _REPLAY.value:
keys = mjw.test_util.find_keys(mjm, _REPLAY.value)
if not keys:
raise app.UsageError(f"Key prefix not find: {_REPLAY.value}")
ctrls = mjw.test_util.make_trajectory(mjm, keys)
mujoco.mj_resetDataKeyframe(mjm, mjd, keys[0])
elif mjm.nkey > 0 and _KEYFRAME.value > -1:
mujoco.mj_resetDataKeyframe(mjm, mjd, _KEYFRAME.value)
mujoco.mj_forward(mjm, mjd)
@@ -160,7 +169,6 @@ def _main(argv: Sequence[str]) -> None:
with wp.ScopedDevice(_DEVICE.value):
m = mjw.put_model(mjm)
_override(m)
mjm_hash = pickle.dumps(mjm)
broadphase, filter = mjw.BroadphaseType(m.opt.broadphase).name, mjw.BroadphaseFilter(m.opt.broadphase_filter).name
solver, cone = mjw.SolverType(m.opt.solver).name, mjw.ConeType(m.opt.cone).name
integrator = mjw.IntegratorType(m.opt.integrator).name
@@ -173,17 +181,19 @@ def _main(argv: Sequence[str]) -> None:
)
d = mjw.put_data(mjm, mjd, nconmax=_NCONMAX.value, njmax=_NJMAX.value)
print(f"Data\n nworld: {d.nworld} nconmax: {d.nconmax} njmax: {d.njmax}\n")
print("Compiling physics step...", end="")
start = time.time()
graph = _compile_step(m, d)
elapsed = time.time() - start
print(f"done ({elapsed:0.2}s).")
print(f"MuJoCo Warp simulating with dt = {m.opt.timestep.numpy()[0]:.3f}...")
with mujoco.viewer.launch_passive(mjm, mjd, key_callback=key_callback) as viewer:
opt = copy.copy(mjm.opt)
while True:
start = time.time()
if ctrls is not None and ctrlid < len(ctrls):
mjd.ctrl[:] = ctrls[ctrlid]
ctrlid += 1
if _ENGINE.value == EngineOptions.C:
mujoco.mj_step(mjm, mjd)
else: # mjwarp
@@ -194,9 +204,9 @@ def _main(argv: Sequence[str]) -> None:
wp.copy(d.qvel, wp.array([mjd.qvel.astype(np.float32)]))
wp.copy(d.time, wp.array([mjd.time], dtype=wp.float32))
hash = pickle.dumps(mjm)
if hash != mjm_hash:
mjm_hash = hash
# if the user changed an option in the MuJoCo Simulate UI, go ahead and recompile the step
if mjm.opt != opt:
opt = copy.copy(mjm.opt)
m = mjw.put_model(mjm)
graph = _compile_step(m, d)
+124 -60
View File
@@ -42,6 +42,7 @@ _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
@@ -55,22 +56,21 @@ def _collision_shim(
geom_margin: wp.array2d(dtype=float),
geom_pair_type_count: tuple[int, ...],
geom_plugin_index: wp.array(dtype=int),
geom_pos: wp.array2d(dtype=wp.vec3),
geom_priority: wp.array(dtype=int),
geom_quat: wp.array2d(dtype=wp.quat),
geom_rbound: wp.array2d(dtype=float),
geom_size: wp.array2d(dtype=wp.vec3),
geom_solimp: wp.array2d(dtype=mjwp_types.vec5),
geom_solmix: wp.array2d(dtype=float),
geom_solref: wp.array2d(dtype=wp.vec2),
geom_type: wp.array(dtype=int),
geompair2hfgeompair: wp.array(dtype=int),
has_sdf_geom: bool,
hfield_adr: wp.array(dtype=int),
hfield_data: wp.array(dtype=float),
hfield_ncol: wp.array(dtype=int),
hfield_nrow: wp.array(dtype=int),
hfield_size: wp.array(dtype=wp.vec4),
mesh_face: wp.array(dtype=wp.vec3i),
mesh_faceadr: wp.array(dtype=int),
mesh_graph: wp.array(dtype=int),
mesh_graphadr: wp.array(dtype=int),
mesh_polyadr: wp.array(dtype=int),
@@ -86,10 +86,13 @@ def _collision_shim(
mesh_vertadr: wp.array(dtype=int),
mesh_vertnum: wp.array(dtype=int),
ngeom: int,
nhfield: int,
nmeshface: int,
nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i),
nxn_pairid: wp.array(dtype=int),
nxn_pairid_filtered: wp.array(dtype=int),
oct_aabb: wp.array2d(dtype=wp.vec3),
oct_child: wp.array(dtype=mjwp_types.vec8i),
oct_coeff: wp.array(dtype=mjwp_types.vec8f),
pair_dim: wp.array(dtype=int),
pair_friction: wp.array2d(dtype=mjwp_types.vec5),
pair_gap: wp.array2d(dtype=float),
@@ -101,6 +104,7 @@ def _collision_shim(
plugin_attr: wp.array(dtype=wp.vec3f),
opt__broadphase: int,
opt__broadphase_filter: int,
opt__ccd_tolerance: wp.array(dtype=float),
opt__disableflags: int,
opt__epa_iterations: int,
opt__gjk_iterations: int,
@@ -110,7 +114,6 @@ def _collision_shim(
opt__sdf_iterations: int,
# Data
nconmax: int,
collision_hftri_index: wp.array(dtype=int),
collision_pair: wp.array(dtype=wp.vec2i),
collision_pairid: wp.array(dtype=int),
collision_worldid: wp.array(dtype=int),
@@ -127,9 +130,19 @@ def _collision_shim(
epa_vert_index2: wp.array2d(dtype=int),
geom_xmat: wp.array2d(dtype=wp.mat33),
geom_xpos: wp.array2d(dtype=wp.vec3),
multiccd_clipped: wp.array2d(dtype=wp.vec3),
multiccd_endvert: wp.array2d(dtype=wp.vec3),
multiccd_face1: wp.array2d(dtype=wp.vec3),
multiccd_face2: wp.array2d(dtype=wp.vec3),
multiccd_idx1: wp.array2d(dtype=int),
multiccd_idx2: wp.array2d(dtype=int),
multiccd_n1: wp.array2d(dtype=wp.vec3),
multiccd_n2: wp.array2d(dtype=wp.vec3),
multiccd_pdist: wp.array2d(dtype=float),
multiccd_pnormal: wp.array2d(dtype=wp.vec3),
multiccd_polygon: wp.array2d(dtype=wp.vec3),
ncollision: wp.array(dtype=int),
ncon: wp.array(dtype=int),
ncon_hfield: wp.array2d(dtype=int),
sap_cumulative_sum: wp.array2d(dtype=int),
sap_projection_lower: wp.array3d(dtype=float),
sap_projection_upper: wp.array2d(dtype=float),
@@ -161,22 +174,21 @@ def _collision_shim(
_m.geom_margin = geom_margin
_m.geom_pair_type_count = geom_pair_type_count
_m.geom_plugin_index = geom_plugin_index
_m.geom_pos = geom_pos
_m.geom_priority = geom_priority
_m.geom_quat = geom_quat
_m.geom_rbound = geom_rbound
_m.geom_size = geom_size
_m.geom_solimp = geom_solimp
_m.geom_solmix = geom_solmix
_m.geom_solref = geom_solref
_m.geom_type = geom_type
_m.geompair2hfgeompair = geompair2hfgeompair
_m.has_sdf_geom = has_sdf_geom
_m.hfield_adr = hfield_adr
_m.hfield_data = hfield_data
_m.hfield_ncol = hfield_ncol
_m.hfield_nrow = hfield_nrow
_m.hfield_size = hfield_size
_m.mesh_face = mesh_face
_m.mesh_faceadr = mesh_faceadr
_m.mesh_graph = mesh_graph
_m.mesh_graphadr = mesh_graphadr
_m.mesh_polyadr = mesh_polyadr
@@ -192,12 +204,16 @@ def _collision_shim(
_m.mesh_vertadr = mesh_vertadr
_m.mesh_vertnum = mesh_vertnum
_m.ngeom = ngeom
_m.nhfield = nhfield
_m.nmeshface = nmeshface
_m.nxn_geom_pair_filtered = nxn_geom_pair_filtered
_m.nxn_pairid = nxn_pairid
_m.nxn_pairid_filtered = nxn_pairid_filtered
_m.oct_aabb = oct_aabb
_m.oct_child = oct_child
_m.oct_coeff = oct_coeff
_m.opt.broadphase = opt__broadphase
_m.opt.broadphase_filter = opt__broadphase_filter
_m.opt.ccd_tolerance = opt__ccd_tolerance
_m.opt.disableflags = opt__disableflags
_m.opt.epa_iterations = opt__epa_iterations
_m.opt.gjk_iterations = opt__gjk_iterations
@@ -214,7 +230,6 @@ def _collision_shim(
_m.pair_solreffriction = pair_solreffriction
_m.plugin = plugin
_m.plugin_attr = plugin_attr
_d.collision_hftri_index = collision_hftri_index
_d.collision_pair = collision_pair
_d.collision_pairid = collision_pairid
_d.collision_worldid = collision_worldid
@@ -242,9 +257,19 @@ def _collision_shim(
_d.epa_vert_index2 = epa_vert_index2
_d.geom_xmat = geom_xmat
_d.geom_xpos = geom_xpos
_d.multiccd_clipped = multiccd_clipped
_d.multiccd_endvert = multiccd_endvert
_d.multiccd_face1 = multiccd_face1
_d.multiccd_face2 = multiccd_face2
_d.multiccd_idx1 = multiccd_idx1
_d.multiccd_idx2 = multiccd_idx2
_d.multiccd_n1 = multiccd_n1
_d.multiccd_n2 = multiccd_n2
_d.multiccd_pdist = multiccd_pdist
_d.multiccd_pnormal = multiccd_pnormal
_d.multiccd_polygon = multiccd_polygon
_d.ncollision = ncollision
_d.ncon = ncon
_d.ncon_hfield = ncon_hfield
_d.nconmax = nconmax
_d.sap_cumulative_sum = sap_cumulative_sum
_d.sap_projection_lower = sap_projection_lower
@@ -258,7 +283,6 @@ def _collision_shim(
def _collision_jax_impl(m: types.Model, d: types.Data):
output_dims = {
'collision_hftri_index': d._impl.collision_hftri_index.shape,
'collision_pair': d._impl.collision_pair.shape,
'collision_pairid': d._impl.collision_pairid.shape,
'collision_worldid': d._impl.collision_worldid.shape,
@@ -275,9 +299,19 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
'epa_vert_index2': d._impl.epa_vert_index2.shape,
'geom_xmat': d.geom_xmat.shape,
'geom_xpos': d.geom_xpos.shape,
'multiccd_clipped': d._impl.multiccd_clipped.shape,
'multiccd_endvert': d._impl.multiccd_endvert.shape,
'multiccd_face1': d._impl.multiccd_face1.shape,
'multiccd_face2': d._impl.multiccd_face2.shape,
'multiccd_idx1': d._impl.multiccd_idx1.shape,
'multiccd_idx2': d._impl.multiccd_idx2.shape,
'multiccd_n1': d._impl.multiccd_n1.shape,
'multiccd_n2': d._impl.multiccd_n2.shape,
'multiccd_pdist': d._impl.multiccd_pdist.shape,
'multiccd_pnormal': d._impl.multiccd_pnormal.shape,
'multiccd_polygon': d._impl.multiccd_polygon.shape,
'ncollision': d._impl.ncollision.shape,
'ncon': d._impl.ncon.shape,
'ncon_hfield': d._impl.ncon_hfield.shape,
'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape,
'sap_projection_lower': d._impl.sap_projection_lower.shape,
'sap_projection_upper': d._impl.sap_projection_upper.shape,
@@ -298,11 +332,10 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
}
jf = ffi.jax_callable_variadic_tuple(
_collision_shim,
num_outputs=37,
num_outputs=46,
output_dims=output_dims,
vmap_method=None,
in_out_argnames={
'collision_hftri_index',
'collision_pair',
'collision_pairid',
'collision_worldid',
@@ -319,9 +352,19 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
'epa_vert_index2',
'geom_xmat',
'geom_xpos',
'multiccd_clipped',
'multiccd_endvert',
'multiccd_face1',
'multiccd_face2',
'multiccd_idx1',
'multiccd_idx2',
'multiccd_n1',
'multiccd_n2',
'multiccd_pdist',
'multiccd_pnormal',
'multiccd_polygon',
'ncollision',
'ncon',
'ncon_hfield',
'sap_cumulative_sum',
'sap_projection_lower',
'sap_projection_upper',
@@ -352,22 +395,21 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m.geom_margin,
m._impl.geom_pair_type_count,
m._impl.geom_plugin_index,
m.geom_pos,
m.geom_priority,
m.geom_quat,
m.geom_rbound,
m.geom_size,
m.geom_solimp,
m.geom_solmix,
m.geom_solref,
m.geom_type,
m._impl.geompair2hfgeompair,
m._impl.has_sdf_geom,
m.hfield_adr,
m.hfield_data,
m.hfield_ncol,
m.hfield_nrow,
m.hfield_size,
m.mesh_face,
m.mesh_faceadr,
m.mesh_graph,
m.mesh_graphadr,
m._impl.mesh_polyadr,
@@ -383,10 +425,13 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m.mesh_vertadr,
m.mesh_vertnum,
m.ngeom,
m.nhfield,
m.nmeshface,
m._impl.nxn_geom_pair_filtered,
m._impl.nxn_pairid,
m._impl.nxn_pairid_filtered,
m._impl.oct_aabb,
m._impl.oct_child,
m._impl.oct_coeff,
m.pair_dim,
m.pair_friction,
m.pair_gap,
@@ -398,6 +443,7 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m._impl.plugin_attr,
m.opt._impl.broadphase,
m.opt._impl.broadphase_filter,
m.opt._impl.ccd_tolerance,
m.opt.disableflags,
m.opt._impl.epa_iterations,
m.opt._impl.gjk_iterations,
@@ -406,7 +452,6 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m.opt._impl.sdf_initpoints,
m.opt._impl.sdf_iterations,
d._impl.nconmax,
d._impl.collision_hftri_index,
d._impl.collision_pair,
d._impl.collision_pairid,
d._impl.collision_worldid,
@@ -423,9 +468,19 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
d._impl.epa_vert_index2,
d.geom_xmat,
d.geom_xpos,
d._impl.multiccd_clipped,
d._impl.multiccd_endvert,
d._impl.multiccd_face1,
d._impl.multiccd_face2,
d._impl.multiccd_idx1,
d._impl.multiccd_idx2,
d._impl.multiccd_n1,
d._impl.multiccd_n2,
d._impl.multiccd_pdist,
d._impl.multiccd_pnormal,
d._impl.multiccd_polygon,
d._impl.ncollision,
d._impl.ncon,
d._impl.ncon_hfield,
d._impl.sap_cumulative_sum,
d._impl.sap_projection_lower,
d._impl.sap_projection_upper,
@@ -445,43 +500,52 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
d._impl.contact__worldid,
)
d = d.tree_replace({
'_impl.collision_hftri_index': out[0],
'_impl.collision_pair': out[1],
'_impl.collision_pairid': out[2],
'_impl.collision_worldid': out[3],
'_impl.epa_face': out[4],
'_impl.epa_horizon': out[5],
'_impl.epa_index': out[6],
'_impl.epa_map': out[7],
'_impl.epa_norm2': out[8],
'_impl.epa_pr': out[9],
'_impl.epa_vert': out[10],
'_impl.epa_vert1': out[11],
'_impl.epa_vert2': out[12],
'_impl.epa_vert_index1': out[13],
'_impl.epa_vert_index2': out[14],
'geom_xmat': out[15],
'geom_xpos': out[16],
'_impl.ncollision': out[17],
'_impl.ncon': out[18],
'_impl.ncon_hfield': out[19],
'_impl.sap_cumulative_sum': out[20],
'_impl.sap_projection_lower': out[21],
'_impl.sap_projection_upper': out[22],
'_impl.sap_range': out[23],
'_impl.sap_segment_index': out[24],
'_impl.sap_sort_index': out[25],
'_impl.contact__dim': out[26],
'_impl.contact__dist': out[27],
'_impl.contact__frame': out[28],
'_impl.contact__friction': out[29],
'_impl.contact__geom': out[30],
'_impl.contact__includemargin': out[31],
'_impl.contact__pos': out[32],
'_impl.contact__solimp': out[33],
'_impl.contact__solref': out[34],
'_impl.contact__solreffriction': out[35],
'_impl.contact__worldid': out[36],
'_impl.collision_pair': out[0],
'_impl.collision_pairid': out[1],
'_impl.collision_worldid': out[2],
'_impl.epa_face': out[3],
'_impl.epa_horizon': out[4],
'_impl.epa_index': out[5],
'_impl.epa_map': out[6],
'_impl.epa_norm2': out[7],
'_impl.epa_pr': out[8],
'_impl.epa_vert': out[9],
'_impl.epa_vert1': out[10],
'_impl.epa_vert2': out[11],
'_impl.epa_vert_index1': out[12],
'_impl.epa_vert_index2': out[13],
'geom_xmat': out[14],
'geom_xpos': out[15],
'_impl.multiccd_clipped': out[16],
'_impl.multiccd_endvert': out[17],
'_impl.multiccd_face1': out[18],
'_impl.multiccd_face2': out[19],
'_impl.multiccd_idx1': out[20],
'_impl.multiccd_idx2': out[21],
'_impl.multiccd_n1': out[22],
'_impl.multiccd_n2': out[23],
'_impl.multiccd_pdist': out[24],
'_impl.multiccd_pnormal': out[25],
'_impl.multiccd_polygon': out[26],
'_impl.ncollision': out[27],
'_impl.ncon': out[28],
'_impl.sap_cumulative_sum': out[29],
'_impl.sap_projection_lower': out[30],
'_impl.sap_projection_upper': out[31],
'_impl.sap_range': out[32],
'_impl.sap_segment_index': out[33],
'_impl.sap_sort_index': out[34],
'_impl.contact__dim': out[35],
'_impl.contact__dist': out[36],
'_impl.contact__frame': out[37],
'_impl.contact__friction': out[38],
'_impl.contact__geom': out[39],
'_impl.contact__includemargin': out[40],
'_impl.contact__pos': out[41],
'_impl.contact__solimp': out[42],
'_impl.contact__solref': out[43],
'_impl.contact__solreffriction': out[44],
'_impl.contact__worldid': out[45],
})
return d
File diff suppressed because it is too large Load Diff
+63 -13
View File
@@ -88,6 +88,8 @@ class OptionWarp(PyTreeNode):
"""Derived fields from Option."""
broadphase: int
broadphase_filter: int
ccd_tolerance: jax.Array
contact_sensor_maxmatch: int
epa_iterations: int
gjk_iterations: int
graph_conditional: bool
@@ -134,7 +136,6 @@ class ModelWarp(PyTreeNode):
flexedge_length0: np.ndarray
geom_pair_type_count: Tuple[int, ...]
geom_plugin_index: np.ndarray
geompair2hfgeompair: np.ndarray
has_sdf_geom: bool
jnt_limited_ball_adr: np.ndarray
jnt_limited_slide_hinge_adr: np.ndarray
@@ -167,6 +168,9 @@ class ModelWarp(PyTreeNode):
nxn_geom_pair_filtered: np.ndarray
nxn_pairid: np.ndarray
nxn_pairid_filtered: np.ndarray
oct_aabb: np.ndarray
oct_child: np.ndarray
oct_coeff: np.ndarray
plugin: np.ndarray
plugin_attr: np.ndarray
qLD_updates: Tuple[np.ndarray, ...]
@@ -223,7 +227,6 @@ class DataWarp(PyTreeNode):
cfrc_ext: jax.Array
cfrc_int: jax.Array
cinert: jax.Array
collision_hftri_index: jax.Array
collision_pair: jax.Array
collision_pairid: jax.Array
collision_worldid: jax.Array
@@ -295,10 +298,19 @@ class DataWarp(PyTreeNode):
inverse_mul_m_skip: jax.Array
light_xdir: jax.Array
light_xpos: jax.Array
multiccd_clipped: jax.Array
multiccd_endvert: jax.Array
multiccd_face1: jax.Array
multiccd_face2: jax.Array
multiccd_idx1: jax.Array
multiccd_idx2: jax.Array
multiccd_n1: jax.Array
multiccd_n2: jax.Array
multiccd_pdist: jax.Array
multiccd_pnormal: jax.Array
multiccd_polygon: jax.Array
ncollision: jax.Array
ncon: jax.Array
ncon_hfield: jax.Array
ncon_world: jax.Array
nconmax: int
ne: jax.Array
ne_connect: jax.Array
@@ -359,7 +371,6 @@ class DataWarp(PyTreeNode):
wrap_xpos: jax.Array
shape = property(lambda self: self.cacc.shape)
DATA_NON_VMAP = {
'collision_hftri_index',
'collision_pair',
'collision_pairid',
'collision_worldid',
@@ -387,6 +398,17 @@ DATA_NON_VMAP = {
'epa_vert_index1',
'epa_vert_index2',
'geom_skip',
'multiccd_clipped',
'multiccd_endvert',
'multiccd_face1',
'multiccd_face2',
'multiccd_idx1',
'multiccd_idx2',
'multiccd_n1',
'multiccd_n2',
'multiccd_pdist',
'multiccd_pnormal',
'multiccd_polygon',
'ncollision',
'ncon',
'nconmax',
@@ -440,7 +462,6 @@ _NDIM = {
'cfrc_ext': 3,
'cfrc_int': 3,
'cinert': 3,
'collision_hftri_index': 1,
'collision_pair': 2,
'collision_pairid': 1,
'collision_worldid': 1,
@@ -519,10 +540,19 @@ _NDIM = {
'light_xpos': 3,
'mocap_pos': 3,
'mocap_quat': 3,
'multiccd_clipped': 3,
'multiccd_endvert': 3,
'multiccd_face1': 3,
'multiccd_face2': 3,
'multiccd_idx1': 2,
'multiccd_idx2': 2,
'multiccd_n1': 3,
'multiccd_n2': 3,
'multiccd_pdist': 2,
'multiccd_pnormal': 3,
'multiccd_polygon': 3,
'ncollision': 1,
'ncon': 1,
'ncon_hfield': 2,
'ncon_world': 1,
'nconmax': 0,
'ne': 1,
'ne_connect': 1,
@@ -750,7 +780,6 @@ _NDIM = {
'geom_solmix': 2,
'geom_solref': 3,
'geom_type': 1,
'geompair2hfgeompair': 1,
'has_sdf_geom': 0,
'hfield_adr': 1,
'hfield_data': 1,
@@ -850,9 +879,14 @@ _NDIM = {
'nxn_geom_pair_filtered': 2,
'nxn_pairid': 1,
'nxn_pairid_filtered': 1,
'oct_aabb': 3,
'oct_child': 2,
'oct_coeff': 2,
'opt__broadphase': 0,
'opt__broadphase_filter': 0,
'opt__ccd_tolerance': 1,
'opt__cone': 0,
'opt__contact_sensor_maxmatch': 0,
'opt__density': 1,
'opt__disableflags': 0,
'opt__enableflags': 0,
@@ -971,7 +1005,9 @@ _NDIM = {
'Option': {
'broadphase': 0,
'broadphase_filter': 0,
'ccd_tolerance': 1,
'cone': 0,
'contact_sensor_maxmatch': 0,
'density': 1,
'disableflags': 0,
'enableflags': 0,
@@ -1021,7 +1057,6 @@ _BATCH_DIM = {
'cfrc_ext': True,
'cfrc_int': True,
'cinert': True,
'collision_hftri_index': False,
'collision_pair': False,
'collision_pairid': False,
'collision_worldid': False,
@@ -1100,10 +1135,19 @@ _BATCH_DIM = {
'light_xpos': True,
'mocap_pos': True,
'mocap_quat': True,
'multiccd_clipped': False,
'multiccd_endvert': False,
'multiccd_face1': False,
'multiccd_face2': False,
'multiccd_idx1': False,
'multiccd_idx2': False,
'multiccd_n1': False,
'multiccd_n2': False,
'multiccd_pdist': False,
'multiccd_pnormal': False,
'multiccd_polygon': False,
'ncollision': False,
'ncon': False,
'ncon_hfield': True,
'ncon_world': True,
'nconmax': False,
'ne': True,
'ne_connect': True,
@@ -1331,7 +1375,6 @@ _BATCH_DIM = {
'geom_solmix': True,
'geom_solref': True,
'geom_type': False,
'geompair2hfgeompair': False,
'has_sdf_geom': False,
'hfield_adr': False,
'hfield_data': False,
@@ -1431,9 +1474,14 @@ _BATCH_DIM = {
'nxn_geom_pair_filtered': False,
'nxn_pairid': False,
'nxn_pairid_filtered': False,
'oct_aabb': False,
'oct_child': False,
'oct_coeff': False,
'opt__broadphase': False,
'opt__broadphase_filter': False,
'opt__ccd_tolerance': True,
'opt__cone': False,
'opt__contact_sensor_maxmatch': False,
'opt__density': True,
'opt__disableflags': False,
'opt__enableflags': False,
@@ -1552,7 +1600,9 @@ _BATCH_DIM = {
'Option': {
'broadphase': False,
'broadphase_filter': False,
'ccd_tolerance': True,
'cone': False,
'contact_sensor_maxmatch': False,
'density': True,
'disableflags': False,
'enableflags': False,
+1 -1
View File
@@ -36,7 +36,7 @@ dependencies = [
[project.optional-dependencies]
warp = [
"warp-lang==1.8.1",
"warp-lang==1.9.0",
]
[project.scripts]