Import google-deepmind/mujoco_warp from GitHub.
PiperOrigin-RevId: 807766737 Change-Id: I1c22361ffa61f27fd8e0fd534b96488531761916
This commit is contained in:
committed by
Copybara-Service
parent
52da7586dc
commit
f2fcd05b1b
@@ -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,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
@@ -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
@@ -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),
|
||||
|
||||
+352
-228
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,
|
||||
|
||||
+23
-299
@@ -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
|
||||
|
||||
+935
-2172
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+475
-176
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
+2
@@ -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.
|
||||
+296
@@ -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
|
||||
+12011
File diff suppressed because it is too large
Load Diff
BIN
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
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+470
-342
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -36,7 +36,7 @@ dependencies = [
|
||||
|
||||
[project.optional-dependencies]
|
||||
warp = [
|
||||
"warp-lang==1.8.1",
|
||||
"warp-lang==1.9.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
|
||||
Reference in New Issue
Block a user