diff --git a/mjx/cuda_requirements.txt b/mjx/cuda_requirements.txt index 9a5463ea..0851eb61 100644 --- a/mjx/cuda_requirements.txt +++ b/mjx/cuda_requirements.txt @@ -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 diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 8c0954b3..fb7eafe2 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -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.""" diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py index deb06509..aff524f6 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py index 57a7b439..f1555cc5 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/broadphase_test.py @@ -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() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py index 2601e3ae..7d0580ff 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py @@ -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, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py index 199c0b62..2973c99b 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py index 7df19a84..2003c8a5 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver_test.py @@ -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): - """ + """, + 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), diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py index 969e638e..a4345e62 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py @@ -13,15 +13,16 @@ # limitations under the License. # ============================================================================== +from typing import Tuple + 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.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 @@ -31,13 +32,6 @@ MJ_MINVAL2 = MJ_MINVAL * MJ_MINVAL FACE_TOL = 0.99999872 EDGE_TOL = 0.00159999931 -MAX_POLYVERT = 15 -polyverts = wp.types.matrix(shape=(MAX_POLYVERT, 3), dtype=float) -polyclip = wp.types.matrix(shape=(2 * MAX_POLYVERT, 3), dtype=float) -polyvec = wp.types.vector(MAX_POLYVERT, dtype=float) -polyindices = wp.types.vector(MAX_POLYVERT, dtype=int) - - mat43 = wp.types.matrix(shape=(4, 3), dtype=float) mat63 = wp.types.matrix(shape=(6, 3), dtype=float) @@ -95,14 +89,14 @@ class SupportPoint: @wp.func -def _discrete_geoms(g1: int, g2: int): +def _discrete_geoms(g1: int, g2: int) -> bool: return (g1 == int(GeomType.MESH.value) or g1 == int(GeomType.BOX.value) or g1 == int(GeomType.HFIELD.value)) and ( g2 == int(GeomType.MESH.value) or g2 == int(GeomType.BOX.value) or g2 == int(GeomType.HFIELD.value) ) @wp.func -def _support_margin(geom: Geom, geomtype: int, dir: wp.vec3): +def _support_margin(geom: Geom, geomtype: int, dir: wp.vec3) -> SupportPoint: sp = SupportPoint() sp.cached_index = -1 sp.vertex_index = -1 @@ -118,7 +112,7 @@ def _support_margin(geom: Geom, geomtype: int, dir: wp.vec3): @wp.func -def _support(geom: Geom, geomtype: int, dir: wp.vec3): +def _support(geom: Geom, geomtype: int, dir: wp.vec3) -> SupportPoint: sp = SupportPoint() sp.cached_index = -1 sp.vertex_index = -1 @@ -204,7 +198,7 @@ def _support(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 @@ -215,7 +209,7 @@ def _support(geom: Geom, geomtype: int, dir: wp.vec3): @wp.func -def _attach_face(pt: Polytope, idx: int, v1: int, v2: int, v3: int): +def _attach_face(pt: Polytope, idx: int, v1: int, v2: int, v3: int) -> float: # out of memory, returning 0 will force EPA to return early without contact if pt.nface == pt.face.shape[0]: return 0.0 @@ -235,7 +229,9 @@ def _attach_face(pt: Polytope, idx: int, v1: int, v2: int, v3: int): @wp.func -def _epa_support(pt: Polytope, idx: int, geom1: Geom, geom2: Geom, geom1_type: int, geom2_type: int, dir: wp.vec3): +def _epa_support( + pt: Polytope, idx: int, geom1: Geom, geom2: Geom, geom1_type: int, geom2_type: int, dir: wp.vec3 +) -> Tuple[int, int]: sp = _support(geom1, geom1_type, dir) pt.vert1[idx] = sp.point pt.vert_index1[idx] = sp.vertex_index @@ -252,7 +248,7 @@ def _epa_support(pt: Polytope, idx: int, geom1: Geom, geom2: Geom, geom1_type: i @wp.func -def _linear_combine(n: int, coefs: wp.vec4, mat: mat43): +def _linear_combine(n: int, coefs: wp.vec4, mat: mat43) -> wp.vec3: v = wp.vec3(0.0) if n == 1: v = coefs[0] * mat[0] @@ -266,12 +262,12 @@ def _linear_combine(n: int, coefs: wp.vec4, mat: mat43): @wp.func -def _almost_equal(v1: wp.vec3, v2: wp.vec3): +def _almost_equal(v1: wp.vec3, v2: wp.vec3) -> bool: return wp.abs(v1[0] - v2[0]) < MJ_MINVAL and wp.abs(v1[1] - v2[1]) < MJ_MINVAL and wp.abs(v1[2] - v2[2]) < MJ_MINVAL @wp.func -def _subdistance(n: int, simplex: mat43): +def _subdistance(n: int, simplex: mat43) -> wp.vec4: if n == 4: return _S3D(simplex[0], simplex[1], simplex[2], simplex[3]) if n == 3: @@ -284,12 +280,12 @@ def _subdistance(n: int, simplex: mat43): @wp.func -def _det3(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3): +def _det3(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3) -> float: return wp.dot(v1, wp.cross(v2, v3)) @wp.func -def _same_sign(a: float, b: float): +def _same_sign(a: float, b: float) -> int: if a > 0 and b > 0: return 1 if a < 0 and b < 0: @@ -298,14 +294,14 @@ def _same_sign(a: float, b: float): @wp.func -def _project_origin_line(v1: wp.vec3, v2: wp.vec3): +def _project_origin_line(v1: wp.vec3, v2: wp.vec3) -> wp.vec3: diff = v2 - v1 scl = -(wp.dot(v2, diff) / wp.dot(diff, diff)) return v2 + scl * diff @wp.func -def _project_origin_plane(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3): +def _project_origin_plane(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3) -> Tuple[wp.vec3, int]: z = wp.vec3(0.0) diff21 = v2 - v1 diff31 = v3 - v1 @@ -340,7 +336,7 @@ def _project_origin_plane(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3): @wp.func -def _S3D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3, s4: wp.vec3): +def _S3D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3, s4: wp.vec3) -> wp.vec4: # [[ s1_x, s2_x, s3_x, s4_x ], # [ s1_y, s2_y, s3_y, s4_y ], # [ s1_z, s2_z, s3_z, s4_z ], @@ -416,7 +412,7 @@ def _S3D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3, s4: wp.vec3): @wp.func -def _S2D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3): +def _S2D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3) -> wp.vec3: # project origin onto affine hull of the simplex p_o, ret = _project_origin_plane(s1, s2, s3) if ret: @@ -561,7 +557,7 @@ def _S2D(s1: wp.vec3, s2: wp.vec3, s3: wp.vec3): @wp.func -def _S1D(s1: wp.vec3, s2: wp.vec3): +def _S1D(s1: wp.vec3, s2: wp.vec3) -> wp.vec2: # find projection of origin onto the 1-simplex: p_o = _project_origin_line(s1, s2) @@ -584,7 +580,7 @@ def _S1D(s1: wp.vec3, s2: wp.vec3): @wp.func -def _gjk( +def gjk( # In: tolerance: float, gjk_iterations: int, @@ -596,7 +592,7 @@ def _gjk( geomtype2: int, cutoff: float, use_margin: bool, -): +) -> GJKResult: """Find distance within a tolerance between two geoms.""" is_discrete = _discrete_geoms(geomtype1, geomtype2) cutoff2 = cutoff * cutoff @@ -710,7 +706,7 @@ def _gjk( @wp.func -def _same_side(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3): +def _same_side(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3) -> bool: n = wp.cross(p1 - p0, p2 - p0) dot1 = wp.dot(n, p3 - p0) dot2 = wp.dot(n, -p0) @@ -718,12 +714,12 @@ def _same_side(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3): @wp.func -def _test_tetra(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3): +def _test_tetra(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3) -> bool: return _same_side(p0, p1, p2, p3) and _same_side(p1, p2, p3, p0) and _same_side(p2, p3, p0, p1) and _same_side(p3, p0, p1, p2) @wp.func -def _tri_affine_coord(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3): +def _tri_affine_coord(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3) -> wp.vec3: # compute minors as in S2D M_14 = v2[1] * v3[2] - v2[2] * v3[1] - v1[1] * v3[2] + v1[2] * v3[1] + v1[1] * v2[2] - v1[2] * v2[1] M_24 = v2[0] * v3[2] - v2[2] * v3[0] - v1[0] * v3[2] + v1[2] * v3[0] + v1[0] * v2[2] - v1[2] * v2[0] @@ -766,7 +762,7 @@ def _tri_affine_coord(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3): @wp.func -def _tri_point_intersect(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3): +def _tri_point_intersect(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3) -> bool: coordinates = _tri_affine_coord(v1, v2, v3, p) l1 = coordinates[0] l2 = coordinates[1] @@ -783,7 +779,7 @@ def _tri_point_intersect(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3): @wp.func -def _replace_simplex3(pt: Polytope, v1: int, v2: int, v3: int): +def _replace_simplex3(pt: Polytope, v1: int, v2: int, v3: int) -> GJKResult: result = GJKResult() # reset GJK simplex @@ -822,7 +818,7 @@ def _replace_simplex3(pt: Polytope, v1: int, v2: int, v3: int): @wp.func -def _rotmat(axis: wp.vec3): +def _rotmat(axis: wp.vec3) -> wp.mat33: n = wp.norm_l2(axis) u1 = axis[0] / n u2 = axis[1] / n @@ -844,7 +840,7 @@ def _rotmat(axis: wp.vec3): @wp.func -def _ray_triangle(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, v4: wp.vec3, v5: wp.vec3): +def _ray_triangle(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, v4: wp.vec3, v5: wp.vec3) -> int: vol1 = _det3(v3 - v1, v4 - v1, v2 - v1) vol2 = _det3(v4 - v1, v5 - v1, v2 - v1) vol3 = _det3(v5 - v1, v3 - v1, v2 - v1) @@ -857,7 +853,7 @@ def _ray_triangle(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, v4: wp.vec3, v5: wp.vec @wp.func -def _add_edge(pt: Polytope, e1: int, e2: int): +def _add_edge(pt: Polytope, e1: int, e2: int) -> int: n = pt.nhorizon if n < 0: @@ -881,7 +877,7 @@ def _add_edge(pt: Polytope, e1: int, e2: int): @wp.func -def _delete_face(pt: Polytope, face_id: int): +def _delete_face(pt: Polytope, face_id: int) -> int: index = pt.face_index[face_id] # delete from map if index >= 0: @@ -895,7 +891,7 @@ def _delete_face(pt: Polytope, face_id: int): @wp.func -def _epa_witness(pt: Polytope, face_idx: int): +def _epa_witness(pt: Polytope, face_idx: int) -> Tuple[wp.vec3, wp.vec3]: # compute affine coordinates for witness points on plane defined by face v1 = pt.vert[pt.face[face_idx][0]] v2 = pt.vert[pt.face[face_idx][1]] @@ -941,7 +937,7 @@ def _polytope2( geom2: Geom, geomtype1: int, geomtype2: int, -): +) -> Tuple[Polytope, GJKResult]: """Create polytope for EPA given a 1-simplex from GJK""" diff = simplex[1] - simplex[0] @@ -1040,7 +1036,7 @@ def _polytope3( geom2: Geom, geomtype1: int, geomtype2: int, -): +) -> Polytope: """Create polytope for EPA given a 2-simplex from GJK""" # get normals in both directions n = wp.cross(simplex[1] - simplex[0], simplex[2] - simplex[0]) @@ -1140,7 +1136,7 @@ def _polytope4( geom2: Geom, geomtype1: int, geomtype2: int, -): +) -> Tuple[Polytope, GJKResult]: """Create polytope for EPA given a 3-simplex from GJK""" pt.vert[0] = simplex[0] pt.vert[1] = simplex[1] @@ -1202,14 +1198,16 @@ def _polytope4( @wp.func -def _epa(tolerance: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom2: Geom, geomtype1: int, geomtype2: int): +def _epa( + tolerance: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom2: Geom, geomtype1: int, geomtype2: int +) -> Tuple[float, wp.vec3, wp.vec3, int]: """Recover penetration data from two geoms in contact given an initial polytope.""" is_discrete = _discrete_geoms(geomtype1, geomtype2) upper = FLOAT_MAX upper2 = FLOAT_MAX idx = int(-1) pidx = int(-1) - epsilon = wp.where(is_discrete, 1e-15, tolerance * tolerance) + epsilon = wp.where(is_discrete, 1e-15, tolerance) for k in range(epa_iterations): pidx = int(idx) @@ -1311,19 +1309,19 @@ def _epa(tolerance: float, epa_iterations: int, pt: Polytope, geom1: Geom, geom2 @wp.func -def _area4(a: wp.vec3, b: wp.vec3, c: wp.vec3, d: wp.vec3): +def _area4(a: wp.vec3, b: wp.vec3, c: wp.vec3, d: wp.vec3) -> float: """Computes area of a quadrilateral embedded in 3D space.""" return 0.5 * wp.norm_l2(wp.cross(a - d, d - b) + wp.cross(b - c, c - a)) @wp.func -def _next(n: int, i: int): +def _next(n: int, i: int) -> int: """Returns (i + 1) mod n for 0 <= i <= n - 1.""" return wp.where(i == n - 1, 0, i + 1) @wp.func -def _polygon_quad(polygon: polyclip, npolygon: int): +def _polygon_quad(polygon: wp.array(dtype=wp.vec3), npolygon: int) -> wp.vec4i: """Returns the indices of a quadrilateral of maximum area in a convex polygon.""" b = _next(npolygon, 0) c = _next(npolygon, b) @@ -1363,7 +1361,9 @@ def _polygon_quad(polygon: polyclip, npolygon: int): # return number (1, 2 or 3) of dimensions of a simplex; reorder vertices if necessary @wp.func -def _feature_dim(face: wp.vec3i, vert_index: wp.array(dtype=int), vert: wp.array(dtype=wp.vec3)): +def _feature_dim( + face: wp.vec3i, vert_index: wp.array(dtype=int), vert: wp.array(dtype=wp.vec3) +) -> Tuple[int, wp.vec3i, wp.mat33]: v1i = vert_index[face[0]] v2i = vert_index[face[1]] v3i = vert_index[face[2]] @@ -1387,7 +1387,9 @@ def _feature_dim(face: wp.vec3i, vert_index: wp.array(dtype=int), vert: wp.array # find two normals that are facing each other within a tolerance, return 1 if found @wp.func -def _aligned_faces(vert1: polyverts, len1: int, vert2: polyverts, len2: int): +def _aligned_faces( + vert1: wp.array(dtype=wp.vec3), len1: int, vert2: wp.array(dtype=wp.vec3), len2: int +) -> Tuple[int, wp.vec2i]: res = wp.vec2i() for i in range(len1): for j in range(len2): @@ -1401,7 +1403,9 @@ def _aligned_faces(vert1: polyverts, len1: int, vert2: polyverts, len2: int): # find two normals that are perpendicular to each other within a tolerance # return 1 if found @wp.func -def _aligned_face_edge(edge: polyverts, nedge: int, face: polyverts, nface: int): +def _aligned_face_edge( + edge: wp.array(dtype=wp.vec3), nedge: int, face: wp.array(dtype=wp.vec3), nface: int +) -> Tuple[int, wp.vec2i]: res = wp.vec2i() for i in range(nface): for j in range(nedge): @@ -1414,7 +1418,9 @@ def _aligned_face_edge(edge: polyverts, nedge: int, face: polyverts, nface: int) # find up to n <= 2 common integers of two arrays, return n @wp.func -def _intersect1(a1: wp.array(dtype=int), a2: wp.array(dtype=int), start1: int, start2: int, len1: int, len2: int): +def _intersect1( + a1: wp.array(dtype=int), a2: wp.array(dtype=int), start1: int, start2: int, len1: int, len2: int +) -> Tuple[int, wp.vec2i]: count = int(0) res = wp.vec2i() for i in range(start1, start1 + len1): @@ -1428,7 +1434,7 @@ def _intersect1(a1: wp.array(dtype=int), a2: wp.array(dtype=int), start1: int, s @wp.func -def _intersect2(a1: wp.vec2i, a2: wp.array(dtype=int), start2: int, len1: int, len2: int): +def _intersect2(a1: wp.vec2i, a2: wp.array(dtype=int), start2: int, len1: int, len2: int) -> Tuple[int, wp.vec2i]: count = int(0) res = wp.vec2i() for i in range(len1): @@ -1454,10 +1460,10 @@ def _mesh_normals( polymapadr: wp.array(dtype=int), polymapnum: wp.array(dtype=int), polymap: wp.array(dtype=int), -): - normals = polyverts() - indices = polyindices() - + # Out: + normals_out: wp.array(dtype=wp.vec3), + indices_out: wp.array(dtype=int), +) -> int: v1 = feature_index[0] v2 = feature_index[1] v3 = feature_index[2] @@ -1474,15 +1480,15 @@ def _mesh_normals( faceset = wp.vec2i() n, edgeset = _intersect1(polymap, polymap, v1_adr, v2_adr, v1_num, v2_num) if n == 0: - return 0, normals, indices + return 0 n, faceset = _intersect2(edgeset, polymap, v3_adr, n, v3_num) if n == 0: - return 0, normals, indices + return 0 # three vertices on mesh define a unique face - normals[0] = mat @ polynormal[polyadr + faceset[0]] - indices[0] = faceset[0] - return 1, normals, indices + normals_out[0] = mat @ polynormal[polyadr + faceset[0]] + indices_out[0] = faceset[0] + return 1 if feature_dim == 2: v1_adr = polymapadr[vertadr + v1] @@ -1494,22 +1500,21 @@ def _mesh_normals( # up to two faces as two vertices define an edge n, edgeset = _intersect1(polymap, polymap, v1_adr, v2_adr, v1_num, v2_num) if n == 0: - return 0, normals, indices + return 0 for i in range(n): - normals[i] = mat @ polynormal[polyadr + edgeset[i]] - indices[i] = edgeset[i] - return n, normals, indices + normals_out[i] = mat @ polynormal[polyadr + edgeset[i]] + indices_out[i] = edgeset[i] + return n if feature_dim == 1: v1_adr = polymapadr[vertadr + v1] v1_num = polymapnum[vertadr + v1] - v1_num = wp.where(v1_num <= MAX_POLYVERT, v1_num, MAX_POLYVERT) for i in range(v1_num): index = polymap[v1_adr + i] - normals[i] = mat @ polynormal[polyadr + index] - indices[i] = index - return v1_num, normals, indices - return 0, normals, indices + normals_out[i] = mat @ polynormal[polyadr + index] + indices_out[i] = index + return v1_num + return 0 # compute normal directional vectors along possible edges given by up to two vertices @@ -1531,20 +1536,19 @@ def _mesh_edge_normals( v1: wp.vec3, v2: wp.vec3, v1i: int, -): - normals = polyverts() - endverts = polyverts() - + # Out: + normals_out: wp.array(dtype=wp.vec3), + endverts_out: wp.array(dtype=wp.vec3), +) -> int: # only one edge if dim == 2: - endverts[0] = v2 - normals[0] = wp.normalize(v2 - v1) - return 1, normals, endverts + endverts_out[0] = v2 + normals_out[0] = wp.normalize(v2 - v1) + return 1 if dim == 1: v1_adr = polymapadr[vertadr + v1i] v1_num = polymapnum[vertadr + v1i] - v1_num = wp.where(v1_num <= MAX_POLYVERT, v1_num, MAX_POLYVERT) # loop through all faces with vertex v1 for i in range(v1_num): @@ -1555,18 +1559,22 @@ def _mesh_edge_normals( for j in range(nvert): if polyvert[adr + j] == v1i: k = wp.where(j == 0, nvert - 1, j - 1) - endverts[i] = mat @ vert[vertadr + polyvert[adr + k]] + pos - normals[i] = wp.normalize(endverts[i] - v1) - return v1_num, normals, endverts - return 0, normals, endverts + endverts_out[i] = mat @ vert[vertadr + polyvert[adr + k]] + pos + normals_out[i] = wp.normalize(endverts_out[i] - v1) + return v1_num + return 0 # try recovering box normal from collision normal @wp.func -def _box_normals2(mat: wp.mat33, n: wp.vec3): - normals = polyverts() - indices = polyindices() - +def _box_normals2( + # In: + mat: wp.mat33, + n: wp.vec3, + # Out: + normal_out: wp.array(dtype=wp.vec3), + index_out: wp.array(dtype=int), +) -> int: # list of box face normals face_normals = mat63(1.0, 0.0, 0.0, -1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, -1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, -1.0) @@ -1583,19 +1591,25 @@ def _box_normals2(mat: wp.mat33, n: wp.vec3): # determine if there is a side close to the normal for i in range(6): if wp.dot(local_n, face_normals[i]) > FACE_TOL: - normals[0] = mat @ face_normals[i] - indices[0] = i - return 1, normals, indices + normal_out[0] = mat @ face_normals[i] + index_out[0] = i + return 1 - return 0, normals, indices + return 0 # compute possible face normals of a box given up to 3 vertices @wp.func -def _box_normals(feature_dim: int, feature_index: wp.vec3i, mat: wp.mat33, dir: wp.vec3): - normals = polyverts() - indices = polyindices() - +def _box_normals( + # In: + feature_dim: int, + feature_index: wp.vec3i, + mat: wp.mat33, + dir: wp.vec3, + # Out: + normal_out: wp.array(dtype=wp.vec3), + index_out: wp.array(dtype=int), +) -> int: v1 = feature_index[0] v2 = feature_index[1] v3 = feature_index[2] @@ -1605,67 +1619,76 @@ def _box_normals(feature_dim: int, feature_index: wp.vec3i, mat: wp.mat33, dir: x = float((v1 & 1) and (v2 & 1) and (v3 & 1)) - float(not (v1 & 1) and not (v2 & 1) and not (v3 & 1)) y = float((v1 & 2) and (v2 & 2) and (v3 & 2)) - float(not (v1 & 2) and not (v2 & 2) and not (v3 & 2)) z = float((v1 & 4) and (v2 & 4) and (v3 & 4)) - float(not (v1 & 4) and not (v2 & 4) and not (v3 & 4)) - normals[0] = mat @ wp.vec3(x, y, z) + normal_out[0] = mat @ wp.vec3(x, y, z) sgn = x + y + z if x != 0.0: - indices[c] = 0 + index_out[c] = 0 c += 1 if y != 0.0: - indices[c] = 2 + index_out[c] = 2 c += 1 if z != 0.0: - indices[c] = 4 + index_out[c] = 4 c += 1 if sgn == -1.0: - indices[0] = indices[0] + 1 + index_out[0] = index_out[0] + 1 if c == 1: - return 1, normals, indices - return _box_normals2(mat, dir) + return 1 + return _box_normals2(mat, dir, normal_out, index_out) if feature_dim == 2: c = 0 x = float((v1 & 1) and (v2 & 1)) - float(not (v1 & 1) and not (v2 & 1)) y = float((v1 & 2) and (v2 & 2)) - float(not (v1 & 2) and not (v2 & 2)) z = float((v1 & 4) and (v2 & 4)) - float(not (v1 & 4) and not (v2 & 4)) if x != 0.0: - normals[c] = mat @ wp.vec3(float(x), 0.0, 0.0) - indices[c] = wp.where(x > 0.0, 0, 1) + normal_out[c] = mat @ wp.vec3(float(x), 0.0, 0.0) + index_out[c] = wp.where(x > 0.0, 0, 1) c += 1 if y != 0.0: - normals[c] = mat @ wp.vec3(0.0, y, 0.0) - indices[c] = wp.where(y > 0.0, 2, 3) + normal_out[c] = mat @ wp.vec3(0.0, y, 0.0) + index_out[c] = wp.where(y > 0.0, 2, 3) c += 1 if z != 0.0: - normals[c] = mat @ wp.vec3(0.0, 0.0, z) - indices[c] = wp.where(z > 0.0, 4, 5) + normal_out[c] = mat @ wp.vec3(0.0, 0.0, z) + index_out[c] = wp.where(z > 0.0, 4, 5) c += 1 if c == 2: - return 2, normals, indices - return _box_normals2(mat, dir) + return 2 + return _box_normals2(mat, dir, normal_out, index_out) if feature_dim == 1: x = wp.where(v1 & 1, 1.0, -1.0) y = wp.where(v1 & 2, 1.0, -1.0) z = wp.where(v1 & 4, 1.0, -1.0) - normals[0] = mat @ wp.vec3(x, 0.0, 0.0) - normals[1] = mat @ wp.vec3(0.0, y, 0.0) - normals[2] = mat @ wp.vec3(0.0, 0.0, z) - indices[0] = wp.where(x > 0.0, 0, 1) - indices[1] = wp.where(y > 0.0, 2, 3) - indices[2] = wp.where(z > 0.0, 4, 5) - return 3, normals, indices - return 0, normals, indices + normal_out[0] = mat @ wp.vec3(x, 0.0, 0.0) + normal_out[1] = mat @ wp.vec3(0.0, y, 0.0) + normal_out[2] = mat @ wp.vec3(0.0, 0.0, z) + index_out[0] = wp.where(x > 0.0, 0, 1) + index_out[1] = wp.where(y > 0.0, 2, 3) + index_out[2] = wp.where(z > 0.0, 4, 5) + return 3 + return 0 # compute possible edge normals for box for edge collisions @wp.func -def _box_edge_normals(dim: int, mat: wp.mat33, pos: wp.vec3, size: wp.vec3, v1: wp.vec3, v2: wp.vec3, v1i: int): - normals = polyverts() - endverts = polyverts() - +def _box_edge_normals( + # In: + dim: int, + mat: wp.mat33, + pos: wp.vec3, + size: wp.vec3, + v1: wp.vec3, + v2: wp.vec3, + v1i: int, + # Out: + normal_out: wp.array(dtype=wp.vec3), + endvert_out: wp.array(dtype=wp.vec3), +) -> int: if dim == 2: - endverts[0] = v2 - normals[0] = wp.normalize(v2 - v1) - return 1, normals, endverts + endvert_out[0] = v2 + normal_out[0] = wp.normalize(v2 - v1) + return 1 # return 3 adjacent vertices if dim == 1: @@ -1673,61 +1696,59 @@ def _box_edge_normals(dim: int, mat: wp.mat33, pos: wp.vec3, size: wp.vec3, v1: y = wp.where(v1i & 2, size[1], -size[1]) z = wp.where(v1i & 4, size[2], -size[2]) - endverts[0] = mat @ wp.vec3(-x, y, z) + pos - normals[0] = wp.normalize(endverts[0] - v1) + endvert_out[0] = mat @ wp.vec3(-x, y, z) + pos + normal_out[0] = wp.normalize(endvert_out[0] - v1) - endverts[1] = mat @ wp.vec3(x, -y, z) + pos - normals[1] = wp.normalize(endverts[1] - v1) + endvert_out[1] = mat @ wp.vec3(x, -y, z) + pos + normal_out[1] = wp.normalize(endvert_out[1] - v1) - endverts[2] = mat @ wp.vec3(x, y, -z) + pos - normals[2] = wp.normalize(endverts[2] - v1) - return 3, normals, endverts - return 0, normals, endverts + endvert_out[2] = mat @ wp.vec3(x, y, -z) + pos + normal_out[2] = wp.normalize(endvert_out[2] - v1) + return 3 + return 0 # recover face of a box from its index @wp.func -def _box_face(mat: wp.mat33, pos: wp.vec3, size: wp.vec3, idx: int): - res = polyverts() - +def _box_face(mat: wp.mat33, pos: wp.vec3, size: wp.vec3, idx: int, face_out: wp.array(dtype=wp.vec3)) -> int: # compute global coordinates of the box face and face normal if idx == 0: # right - res[0] = mat @ wp.vec(size[0], size[1], size[2]) + pos - res[1] = mat @ wp.vec(size[0], size[1], -size[2]) + pos - res[2] = mat @ wp.vec(size[0], -size[1], -size[2]) + pos - res[3] = mat @ wp.vec(size[0], -size[1], size[2]) + pos - return 4, res + face_out[0] = mat @ wp.vec(size[0], size[1], size[2]) + pos + face_out[1] = mat @ wp.vec(size[0], size[1], -size[2]) + pos + face_out[2] = mat @ wp.vec(size[0], -size[1], -size[2]) + pos + face_out[3] = mat @ wp.vec(size[0], -size[1], size[2]) + pos + return 4 if idx == 1: # left - res[0] = mat @ wp.vec(-size[0], size[1], -size[2]) + pos - res[1] = mat @ wp.vec(-size[0], size[1], size[2]) + pos - res[2] = mat @ wp.vec(-size[0], -size[1], size[2]) + pos - res[3] = mat @ wp.vec(-size[0], -size[1], -size[2]) + pos - return 4, res + face_out[0] = mat @ wp.vec(-size[0], size[1], -size[2]) + pos + face_out[1] = mat @ wp.vec(-size[0], size[1], size[2]) + pos + face_out[2] = mat @ wp.vec(-size[0], -size[1], size[2]) + pos + face_out[3] = mat @ wp.vec(-size[0], -size[1], -size[2]) + pos + return 4 if idx == 2: # top - res[0] = mat @ wp.vec(-size[0], size[1], -size[2]) + pos - res[1] = mat @ wp.vec(size[0], size[1], -size[2]) + pos - res[2] = mat @ wp.vec(size[0], size[1], size[2]) + pos - res[3] = mat @ wp.vec(-size[0], size[1], size[2]) + pos - return 4, res + face_out[0] = mat @ wp.vec(-size[0], size[1], -size[2]) + pos + face_out[1] = mat @ wp.vec(size[0], size[1], -size[2]) + pos + face_out[2] = mat @ wp.vec(size[0], size[1], size[2]) + pos + face_out[3] = mat @ wp.vec(-size[0], size[1], size[2]) + pos + return 4 if idx == 3: # bottom - res[0] = mat @ wp.vec(-size[0], -size[1], size[2]) + pos - res[1] = mat @ wp.vec(size[0], -size[1], size[2]) + pos - res[2] = mat @ wp.vec(size[0], -size[1], -size[2]) + pos - res[3] = mat @ wp.vec(-size[0], -size[1], -size[2]) + pos - return 4, res + face_out[0] = mat @ wp.vec(-size[0], -size[1], size[2]) + pos + face_out[1] = mat @ wp.vec(size[0], -size[1], size[2]) + pos + face_out[2] = mat @ wp.vec(size[0], -size[1], -size[2]) + pos + face_out[3] = mat @ wp.vec(-size[0], -size[1], -size[2]) + pos + return 4 if idx == 4: # front - res[0] = mat @ wp.vec(-size[0], size[1], size[2]) + pos - res[1] = mat @ wp.vec(size[0], size[1], size[2]) + pos - res[2] = mat @ wp.vec(size[0], -size[1], size[2]) + pos - res[3] = mat @ wp.vec(-size[0], -size[1], size[2]) + pos - return 4, res + face_out[0] = mat @ wp.vec(-size[0], size[1], size[2]) + pos + face_out[1] = mat @ wp.vec(size[0], size[1], size[2]) + pos + face_out[2] = mat @ wp.vec(size[0], -size[1], size[2]) + pos + face_out[3] = mat @ wp.vec(-size[0], -size[1], size[2]) + pos + return 4 if idx == 5: # back - res[0] = mat @ wp.vec(size[0], size[1], -size[2]) + pos - res[1] = mat @ wp.vec(-size[0], size[1], -size[2]) + pos - res[2] = mat @ wp.vec(-size[0], -size[1], -size[2]) + pos - res[3] = mat @ wp.vec(size[0], -size[1], -size[2]) + pos - return 4, res - return 0, res + face_out[0] = mat @ wp.vec(size[0], size[1], -size[2]) + pos + face_out[1] = mat @ wp.vec(-size[0], size[1], -size[2]) + pos + face_out[2] = mat @ wp.vec(-size[0], -size[1], -size[2]) + pos + face_out[3] = mat @ wp.vec(size[0], -size[1], -size[2]) + pos + return 4 + return 0 # recover mesh polygon from its index, return number of edges @@ -1743,34 +1764,33 @@ def _mesh_face( polyvertnum: wp.array(dtype=int), polyvert: wp.array(dtype=int), idx: int, -): - res = polyverts() - + # Out: + face_out: wp.array(dtype=wp.vec3), +) -> int: adr = polyvertadr[polyadr + idx] j = int(0) nvert = polyvertnum[polyadr + idx] - nvert = wp.where(nvert <= MAX_POLYVERT, nvert, MAX_POLYVERT) for i in range(nvert - 1, -1, -1): v = vert[vertadr + polyvert[adr + i]] - res[j] = mat @ v + pos + face_out[j] = mat @ v + pos j += 1 - return nvert, res + return nvert @wp.func -def _plane_normal(v1: wp.vec3, v2: wp.vec3, n: wp.vec3): +def _plane_normal(v1: wp.vec3, v2: wp.vec3, n: wp.vec3) -> Tuple[float, wp.vec3]: v3 = v1 + n res = wp.cross(v2 - v1, v3 - v1) return wp.dot(res, v1), res @wp.func -def _halfspace(a: wp.vec3, n: wp.vec3, p: wp.vec3): +def _halfspace(a: wp.vec3, n: wp.vec3, p: wp.vec3) -> bool: return wp.dot(p - a, n) > -1e-10 @wp.func -def _plane_intersect(pn: wp.vec3, pd: float, a: wp.vec3, b: wp.vec3): +def _plane_intersect(pn: wp.vec3, pd: float, a: wp.vec3, b: wp.vec3) -> Tuple[float, wp.vec3]: res = wp.vec3() ab = b - a temp = wp.dot(pn, ab) @@ -1786,7 +1806,20 @@ def _plane_intersect(pn: wp.vec3, pd: float, a: wp.vec3, b: wp.vec3): # clip a polygon against another polygon @wp.func -def _polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, n: wp.vec3, dir: wp.vec3): +def _polygon_clip( + # In: + plane_normal: wp.array(dtype=wp.vec3), + plane_dist: wp.array(dtype=float), + face1: wp.array(dtype=wp.vec3), + nface1: int, + face2: wp.array(dtype=wp.vec3), + nface2: int, + n: wp.vec3, + dir: wp.vec3, + # Out: + polygon_out: wp.array(dtype=wp.vec3), + clipped_out: wp.array(dtype=wp.vec3), +) -> Tuple[int, mat3c, mat3c]: witness1 = mat3c() witness2 = mat3c() @@ -1795,8 +1828,8 @@ def _polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, return 0, witness1, witness2 # compute plane normal and distance to plane for each vertex - pn = polyverts() - pd = polyvec() + pn = plane_normal + pd = plane_dist for i in range(nface1 - 1): pdi, pni = _plane_normal(face1[i], face1[i + 1], n) pd[i] = pdi @@ -1806,20 +1839,18 @@ def _polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, pn[nface1 - 1] = pni # reserve 2 * max_sides as max sides for a clipped polygon - polygon = polyclip() - clipped = polyclip() npolygon = nface2 nclipped = int(0) for i in range(nface2): - polygon[i] = face2[i] + polygon_out[i] = face2[i] # clip the polygon by one edge e at a time for e in range(nface1): for i in range(npolygon): # get edge PQ of the polygon - P = polygon[i] - Q = wp.where(i < npolygon - 1, polygon[i + 1], polygon[0]) + P = polygon_out[i] + Q = wp.where(i < npolygon - 1, polygon_out[i + 1], polygon_out[0]) # determine if P and Q are in the halfspace of the clipping edge inside1 = _halfspace(face1[e], pn[e], P) @@ -1831,25 +1862,25 @@ def _polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, # edge PQ is inside the clipping edge, add Q if inside1 and inside2: - clipped[nclipped] = Q + clipped_out[nclipped] = Q nclipped += 1 continue # add new vertex to clipped polygon where PQ intersects the clipping edge t, res = _plane_intersect(pn[e], pd[e], P, Q) if t >= 0.0 and t <= 1.0: - clipped[nclipped] = res + clipped_out[nclipped] = res nclipped += 1 # add Q as PQ is now back inside the clipping edge if inside2: - clipped[nclipped] = Q + clipped_out[nclipped] = Q nclipped += 1 # swap clipped and polygon - tmp = polygon - polygon = clipped - clipped = tmp + tmp = polygon_out + polygon_out = clipped_out + clipped_out = tmp npolygon = nclipped nclipped = 0 @@ -1857,33 +1888,57 @@ def _polygon_clip(face1: polyverts, nface1: int, face2: polyverts, nface2: int, return 0, witness1, witness2 if npolygon > 4: - quad = _polygon_quad(polygon, npolygon) + quad = _polygon_quad(polygon_out, npolygon) for i in range(4): - witness2[i] = polygon[quad[i]] + witness2[i] = polygon_out[quad[i]] witness1[i] = witness2[i] - dir return 4, witness1, witness2 # no pruning needed for i in range(npolygon): - witness2[i] = polygon[i] + witness2[i] = polygon_out[i] witness1[i] = witness2[i] - dir return npolygon, witness1, witness2 +@wp.func +def _set_edge( + vert1: wp.array(dtype=wp.vec3), vert2: wp.array(dtype=wp.vec3), start: int, end: int, face_out: wp.array(dtype=wp.vec3) +) -> int: + face_out[0] = vert1[start] + face_out[1] = vert2[end] + return 2 + + # recover multiple contacts from EPA polytope @wp.func def _multicontact( - pt: Polytope, face: wp.vec3i, x1: wp.vec3, x2: wp.vec3, geom1: Geom, geom2: Geom, geomtype1: int, geomtype2: int -): + # In: + polygon: wp.array(dtype=wp.vec3), + clipped: wp.array(dtype=wp.vec3), + plane_normal: wp.array(dtype=wp.vec3), + plane_dist: 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), + pt: Polytope, + face: wp.vec3i, + x1: wp.vec3, + x2: wp.vec3, + geom1: Geom, + geom2: Geom, + geomtype1: int, + geomtype2: int, +) -> Tuple[int, mat3c, mat3c]: witness1 = mat3c() witness2 = mat3c() witness1[0] = x1 witness2[0] = x2 - face1 = polyverts() - face2 = polyverts() - endverts = polyverts() - if geomtype1 == int(GeomType.MESH.value): vert = geom1.vert polynormal = geom1.mesh_polynormal @@ -1912,16 +1967,36 @@ def _multicontact( # get all possible face normals for each geom if geomtype1 == int(GeomType.BOX.value): - nnorms1, n1, idx1 = _box_normals(nface1, feature_index1, geom1.rot, dir_neg) + nnorms1 = _box_normals(nface1, feature_index1, geom1.rot, dir_neg, n1, idx1) elif geomtype1 == int(GeomType.MESH.value): - nnorms1, n1, idx1 = _mesh_normals( - nface1, feature_index1, geom1.rot, geom1.vertadr, geom1.mesh_polyadr, polynormal, polymapadr, polymapnum, polymap + nnorms1 = _mesh_normals( + nface1, + feature_index1, + geom1.rot, + geom1.vertadr, + geom1.mesh_polyadr, + polynormal, + polymapadr, + polymapnum, + polymap, + n1, + idx1, ) if geomtype2 == int(GeomType.BOX.value): - nnorms2, n2, idx2 = _box_normals(nface2, feature_index2, geom2.rot, dir) + nnorms2 = _box_normals(nface2, feature_index2, geom2.rot, dir, n2, idx2) elif geomtype2 == int(GeomType.MESH.value): - nnorms2, n2, idx2 = _mesh_normals( - nface2, feature_index2, geom2.rot, geom2.vertadr, geom2.mesh_polyadr, polynormal, polymapadr, polymapnum, polymap + nnorms2 = _mesh_normals( + nface2, + feature_index2, + geom2.rot, + geom2.vertadr, + geom2.mesh_polyadr, + polynormal, + polymapadr, + polymapnum, + polymap, + n2, + idx2, ) # determine if any two face normals match @@ -1933,11 +2008,11 @@ def _multicontact( if nface1 < 3 and nface1 <= nface2: nnorms1 = 0 if geomtype1 == int(GeomType.BOX.value): - nnorms1, n1, endverts = _box_edge_normals( - nface1, geom1.rot, geom1.pos, geom1.size, feature_vertex1[0], feature_vertex1[1], feature_index1[0] + nnorms1 = _box_edge_normals( + nface1, geom1.rot, geom1.pos, geom1.size, feature_vertex1[0], feature_vertex1[1], feature_index1[0], n1, endvert ) elif geomtype1 == int(GeomType.MESH.value): - nnorms1, n1, endverts = _mesh_edge_normals( + nnorms1 = _mesh_edge_normals( nface1, geom1.rot, geom1.pos, @@ -1953,6 +2028,8 @@ def _multicontact( feature_vertex1[0], feature_vertex1[1], feature_index1[0], + n1, + endvert, ) nres, res = _aligned_face_edge(n1, nnorms1, n2, nnorms2) if not nres: @@ -1963,11 +2040,11 @@ def _multicontact( elif nface2 < 3: nnorms2 = 0 if geomtype2 == int(GeomType.BOX.value): - nnorms2, n2, endverts = _box_edge_normals( - nface2, geom2.rot, geom2.pos, geom2.size, feature_vertex2[0], feature_vertex2[1], feature_index2[0] + nnorms2 = _box_edge_normals( + nface2, geom2.rot, geom2.pos, geom2.size, feature_vertex2[0], feature_vertex2[1], feature_index2[0], n2, endvert ) elif geomtype2 == int(GeomType.MESH.value): - nnorms2, n2, endverts = _mesh_edge_normals( + nnorms2 = _mesh_edge_normals( nface2, geom2.rot, geom2.pos, @@ -1983,6 +2060,8 @@ def _multicontact( feature_vertex2[0], feature_vertex2[1], feature_index2[0], + n2, + endvert, ) nres, res = _aligned_face_edge(n2, nnorms2, n1, nnorms1) if not nres: @@ -1997,29 +2076,43 @@ def _multicontact( # recover geom1 matching edge or face if is_edge_contact_geom1: - face1[0] = pt.vert1[face[0]] - face1[1] = endverts[i] - nface1 = 2 + nface1 = _set_edge(pt.vert1, endvert, face[0], i, face1) else: ind = wp.where(is_edge_contact_geom2, idx1[j], idx1[i]) if geomtype1 == int(GeomType.BOX.value): - nface1, face1 = _box_face(geom1.rot, geom1.pos, geom1.size, ind) + nface1 = _box_face(geom1.rot, geom1.pos, geom1.size, ind, face1) elif geomtype1 == int(GeomType.MESH.value): - nface1, face1 = _mesh_face( - geom1.rot, geom1.pos, geom1.vertadr, geom1.mesh_polyadr, vert, polyvertadr, polyvertnum, polyvert, ind + nface1 = _mesh_face( + geom1.rot, + geom1.pos, + geom1.vertadr, + geom1.mesh_polyadr, + vert, + polyvertadr, + polyvertnum, + polyvert, + ind, + face1, ) # recover geom2 matching edge or face if is_edge_contact_geom2: - face2[0] = pt.vert2[face[0]] - face2[1] = endverts[i] - nface2 = 2 + nface2 = _set_edge(pt.vert2, endvert, face[0], i, face2) else: if geomtype2 == int(GeomType.BOX.value): - nface2, face2 = _box_face(geom2.rot, geom2.pos, geom2.size, idx2[j]) + nface2 = _box_face(geom2.rot, geom2.pos, geom2.size, idx2[j], face2) elif geomtype2 == int(GeomType.MESH.value): - nface2, face2 = _mesh_face( - geom2.rot, geom2.pos, geom2.vertadr, geom2.mesh_polyadr, vert, polyvertadr, polyvertnum, polyvert, idx2[j] + nface2 = _mesh_face( + geom2.rot, + geom2.pos, + geom2.vertadr, + geom2.mesh_polyadr, + vert, + polyvertadr, + polyvertnum, + polyvert, + idx2[j], + face2, ) # TODO(kbayes): this approximates the contact direction, by scaling the face normal by the @@ -2030,20 +2123,20 @@ def _multicontact( # face1 is an edge; clip face1 against face2 if is_edge_contact_geom1: approx_dir = wp.norm_l2(dir) * n2[j] - return _polygon_clip(face2, nface2, face1, nface1, n2[j], approx_dir) + return _polygon_clip(plane_normal, plane_dist, face2, nface2, face1, nface1, n2[j], approx_dir, polygon, clipped) # face2 is an edge; clip face2 against face1 if is_edge_contact_geom2: approx_dir = -wp.norm_l2(dir) * n1[j] - return _polygon_clip(face1, nface1, face2, nface2, n1[j], approx_dir) + return _polygon_clip(plane_normal, plane_dist, face1, nface1, face2, nface2, n1[j], approx_dir, polygon, clipped) # face-face collision approx_dir = wp.norm_l2(dir) * n2[j] - return _polygon_clip(face1, nface1, face2, nface2, n1[i], approx_dir) + return _polygon_clip(plane_normal, plane_dist, face1, nface1, face2, nface2, n1[i], approx_dir, polygon, clipped) @wp.func -def _inflate(dist: float, x1: wp.vec3, x2: wp.vec3, margin1: float, margin2: float): +def _inflate(dist: float, x1: wp.vec3, x2: wp.vec3, margin1: float, margin2: float) -> Tuple[float, wp.vec3, wp.vec3]: n = wp.normalize(x2 - x1) if margin1 > 0.0: x1 += margin1 * n @@ -2079,7 +2172,18 @@ def ccd( 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), + plane_normal: wp.array(dtype=wp.vec3), + plane_dist: 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), +) -> Tuple[float, int, mat3c, mat3c]: """General convex collision detection via GJK/EPA.""" witness1 = mat3c() witness2 = mat3c() @@ -2095,7 +2199,7 @@ def ccd( # special handling for sphere and capsule (shrink to point and line respectively) if margin1 + margin2 > 0.0: cutoff += margin1 + margin2 - result = _gjk(tolerance, gjk_iterations, geom1, geom2, x_1, x_2, geomtype1, geomtype2, cutoff, True) + result = gjk(tolerance, gjk_iterations, geom1, geom2, x_1, x_2, geomtype1, geomtype2, cutoff, True) # shallow penetration, inflate contact if result.dist > tolerance: @@ -2111,7 +2215,7 @@ def ccd( # deep penetration, reset initial conditions and rerun GJK + EPA cutoff -= margin1 + margin2 - result = _gjk(tolerance, gjk_iterations, geom1, geom2, x_1, x_2, geomtype1, geomtype2, cutoff, False) + result = gjk(tolerance, gjk_iterations, geom1, geom2, x_1, x_2, geomtype1, geomtype2, cutoff, False) # no penetration depth to recover if result.dist > tolerance or result.dim < 2: @@ -2210,7 +2314,27 @@ def ccd( and (geomtype1 == int(GeomType.BOX.value) or (geomtype1 == int(GeomType.MESH.value) and geom1.mesh_polyadr > -1)) and (geomtype2 == int(GeomType.BOX.value) or (geomtype2 == int(GeomType.MESH.value) and geom2.mesh_polyadr > -1)) ): - num, w1, w2 = _multicontact(pt, pt.face[idx], x1, x2, geom1, geom2, geomtype1, geomtype2) + num, w1, w2 = _multicontact( + polygon, + clipped, + plane_normal, + plane_dist, + idx1, + idx2, + n1, + n2, + endvert, + face1, + face2, + pt, + pt.face[idx], + x1, + x2, + geom1, + geom2, + geomtype1, + geomtype2, + ) if num > 0: return dist, num, w1, w2 witness1[0] = x1 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py index 0b93ab15..e4b2d28d 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py index 64dd7945..5a3c82de 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_test.py @@ -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, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py index 0458fd6b..b2eed4ab 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py index 9289ccf1..340ca24c 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py @@ -13,13 +13,23 @@ # limitations under the License. # ============================================================================== +from typing import Tuple + import warp as wp -from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_triangle_prism -from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_point -from mujoco.mjx.third_party.mujoco_warp._src.math import closest_segment_to_segment_points +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import box_box +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import capsule_box +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import capsule_capsule +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import plane_box +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import plane_capsule +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import plane_cylinder +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import plane_ellipsoid +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import plane_sphere +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import sphere_box +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import sphere_capsule +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import sphere_cylinder +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import sphere_sphere from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame -from mujoco.mjx.third_party.mujoco_warp._src.math import normalize_with_norm from mujoco.mjx.third_party.mujoco_warp._src.math import safe_div from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU @@ -29,20 +39,18 @@ 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.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}) - -class vec8f(wp.types.vector(length=8, dtype=wp.float32)): - pass +_HUGE_VAL = 1e6 class mat43f(wp.types.matrix(shape=(4, 3), dtype=wp.float32)): pass -class mat83f(wp.types.matrix(shape=(8, 3), dtype=wp.float32)): - pass +mat63 = wp.types.matrix(shape=(6, 3), dtype=float) @wp.struct @@ -51,7 +59,7 @@ class Geom: rot: wp.mat33 normal: wp.vec3 size: wp.vec3 - hfprism: wp.mat33 + hfprism: mat63 vertadr: int vertnum: int vert: wp.array(dtype=wp.vec3) @@ -70,23 +78,19 @@ class Geom: @wp.func -def _geom( +def geom( + # kernel_analyzer: off # Model: - geom_type: wp.array(dtype=int), - geom_dataid: wp.array(dtype=int), - geom_size: wp.array2d(dtype=wp.vec3), - 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), - mesh_vertadr: wp.array(dtype=int), - mesh_vertnum: wp.array(dtype=int), + geom_type: int, + geom_dataid: int, + geom_size: wp.vec3, + mesh_vertadr: int, + mesh_vertnum: int, mesh_vert: wp.array(dtype=wp.vec3), - mesh_graphadr: wp.array(dtype=int), + mesh_graphadr: int, mesh_graph: wp.array(dtype=int), - mesh_polynum: wp.array(dtype=int), - mesh_polyadr: wp.array(dtype=int), + mesh_polynum: int, + mesh_polyadr: int, mesh_polynormal: wp.array(dtype=wp.vec3), mesh_polyvertadr: wp.array(dtype=int), mesh_polyvertnum: wp.array(dtype=int), @@ -95,28 +99,23 @@ def _geom( mesh_polymapnum: wp.array(dtype=int), mesh_polymap: wp.array(dtype=int), # Data in: - geom_xpos_in: wp.array2d(dtype=wp.vec3), - geom_xmat_in: wp.array2d(dtype=wp.mat33), - # In: - worldid: int, - gid: int, - hftri_index: int, + geom_xpos_in: wp.vec3, + geom_xmat_in: wp.mat33, + # kernel_analyzer: on ) -> Geom: geom = Geom() - geom.pos = geom_xpos_in[worldid, gid] - rot = geom_xmat_in[worldid, gid] - geom.rot = rot - geom.size = geom_size[worldid, gid] - geom.normal = wp.vec3(rot[0, 2], rot[1, 2], rot[2, 2]) # plane - dataid = geom_dataid[gid] + geom.pos = geom_xpos_in + geom.rot = geom_xmat_in + geom.size = geom_size + geom.normal = wp.vec3(geom_xmat_in[0, 2], geom_xmat_in[1, 2], geom_xmat_in[2, 2]) # plane # If geom is MESH, get mesh verts - if dataid >= 0 and geom_type[gid] == int(GeomType.MESH.value): - geom.vertadr = mesh_vertadr[dataid] - geom.vertnum = mesh_vertnum[dataid] - geom.graphadr = mesh_graphadr[dataid] - geom.mesh_polynum = mesh_polynum[dataid] - geom.mesh_polyadr = mesh_polyadr[dataid] + if geom_dataid >= 0 and geom_type == int(GeomType.MESH.value): + geom.vertadr = mesh_vertadr + geom.vertnum = mesh_vertnum + geom.graphadr = mesh_graphadr + geom.mesh_polynum = mesh_polynum + geom.mesh_polyadr = mesh_polyadr else: geom.vertadr = -1 geom.vertnum = -1 @@ -124,7 +123,7 @@ def _geom( geom.mesh_polynum = -1 geom.mesh_polyadr = -1 - if geom_type[gid] == int(GeomType.MESH.value): + if geom_type == int(GeomType.MESH.value): geom.vert = mesh_vert geom.graph = mesh_graph geom.mesh_polynormal = mesh_polynormal @@ -135,773 +134,33 @@ def _geom( geom.mesh_polymapnum = mesh_polymapnum geom.mesh_polymap = mesh_polymap - # If geom is HFIELD triangle, compute triangle prism verts - if geom_type[gid] == int(GeomType.HFIELD.value): - geom.hfprism = hfield_triangle_prism( - geom_dataid, hfield_adr, hfield_nrow, hfield_ncol, hfield_size, hfield_data, gid, hftri_index - ) - geom.index = -1 return geom @wp.func -def write_contact( - # Data in: - nconmax_in: int, - # In: - dist_in: float, - pos_in: wp.vec3, - frame_in: wp.mat33, - margin_in: float, - gap_in: float, - condim_in: int, - friction_in: vec5, - solref_in: wp.vec2f, - solreffriction_in: wp.vec2f, - solimp_in: vec5, - geoms_in: wp.vec2i, - worldid_in: 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), -): - active = (dist_in - margin_in) < 0 - if active: - cid = wp.atomic_add(ncon_out, 0, 1) - if cid < nconmax_in: - contact_dist_out[cid] = dist_in - contact_pos_out[cid] = pos_in - contact_frame_out[cid] = frame_in - contact_geom_out[cid] = geoms_in - contact_worldid_out[cid] = worldid_in - includemargin = margin_in - gap_in - contact_includemargin_out[cid] = includemargin - contact_dim_out[cid] = condim_in - contact_friction_out[cid] = friction_in - contact_solref_out[cid] = solref_in - contact_solreffriction_out[cid] = solreffriction_in - contact_solimp_out[cid] = solimp_in +def plane_convex(plane_normal: wp.vec3, plane_pos: wp.vec3, convex: Geom) -> Tuple[wp.vec4, mat43f, wp.vec3]: + """Core contact geometry calculation for plane-convex collision. + Args: + plane_normal: Normal vector of the plane + plane_pos: Position point on the plane + convex: Convex geometry object containing position, rotation, and mesh data -@wp.func -def _plane_sphere(plane_normal: wp.vec3, plane_pos: wp.vec3, sphere_pos: wp.vec3, sphere_radius: float): - dist = wp.dot(sphere_pos - plane_pos, plane_normal) - sphere_radius - pos = sphere_pos - plane_normal * (sphere_radius + 0.5 * dist) - return dist, pos + Returns: + Tuple containing: + contact_dist: Vector of contact distances (wp.inf for unpopulated contacts) + contact_pos: Matrix of contact positions (one per row) + contact_normal: Matrix of contact normal vectors (one per row) + """ - -@wp.func -def plane_sphere( - # Data in: - nconmax_in: int, - # In: - plane: Geom, - sphere: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - dist, pos = _plane_sphere(plane.normal, plane.pos, sphere.pos, sphere.size[0]) - - write_contact( - nconmax_in, - dist, - pos, - make_frame(plane.normal), - 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, - ) - - -@wp.func -def _sphere_sphere( - # Data in: - nconmax_in: int, - # In: - pos1: wp.vec3, - radius1: float, - pos2: wp.vec3, - radius2: float, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - dir = pos2 - pos1 - dist = wp.length(dir) - if dist == 0.0: - n = wp.vec3(1.0, 0.0, 0.0) - else: - n = dir / dist - dist = dist - (radius1 + radius2) - pos = pos1 + n * (radius1 + 0.5 * dist) - - 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, - ) - - -@wp.func -def _sphere_sphere_ext( - # Data in: - nconmax_in: int, - # In: - pos1: wp.vec3, - radius1: float, - pos2: wp.vec3, - radius2: float, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - mat1: wp.mat33, - mat2: wp.mat33, - # 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), -): - dir = pos2 - pos1 - dist = wp.length(dir) - if dist == 0.0: - # Use cross product of z axes like MuJoCo - axis1 = wp.vec3(mat1[0, 2], mat1[1, 2], mat1[2, 2]) - axis2 = wp.vec3(mat2[0, 2], mat2[1, 2], mat2[2, 2]) - n = wp.cross(axis1, axis2) - n = wp.normalize(n) - else: - n = dir / dist - dist = dist - (radius1 + radius2) - pos = pos1 + n * (radius1 + 0.5 * dist) - - 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, - ) - - -@wp.func -def sphere_sphere( - # Data in: - nconmax_in: int, - # In: - sphere1: Geom, - sphere2: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - _sphere_sphere( - nconmax_in, - sphere1.pos, - sphere1.size[0], - sphere2.pos, - sphere2.size[0], - worldid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - 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, - ) - - -@wp.func -def sphere_capsule( - # Data in: - nconmax_in: int, - # In: - sphere: Geom, - cap: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - """Calculates one contact between a sphere and a capsule.""" - axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) - length = cap.size[1] - segment = axis * length - - # Find closest point on capsule centerline to sphere center - pt = closest_segment_point(cap.pos - segment, cap.pos + segment, sphere.pos) - - # Treat as sphere-sphere collision between sphere and closest point - _sphere_sphere( - nconmax_in, - sphere.pos, - sphere.size[0], - pt, - cap.size[0], - worldid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - 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, - ) - - -@wp.func -def capsule_capsule( - # Data in: - nconmax_in: int, - # In: - cap1: Geom, - cap2: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - axis1 = wp.vec3(cap1.rot[0, 2], cap1.rot[1, 2], cap1.rot[2, 2]) - axis2 = wp.vec3(cap2.rot[0, 2], cap2.rot[1, 2], cap2.rot[2, 2]) - length1 = cap1.size[1] - length2 = cap2.size[1] - seg1 = axis1 * length1 - seg2 = axis2 * length2 - - pt1, pt2 = closest_segment_to_segment_points( - cap1.pos - seg1, - cap1.pos + seg1, - cap2.pos - seg2, - cap2.pos + seg2, - ) - - _sphere_sphere( - nconmax_in, - pt1, - cap1.size[0], - pt2, - cap2.size[0], - worldid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - 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, - ) - - -@wp.func -def plane_capsule( - # Data in: - nconmax_in: int, - # In: - plane: Geom, - cap: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - """Calculates two contacts between a capsule and a plane.""" - n = plane.normal - axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) - # align contact frames with capsule axis - b, b_norm = normalize_with_norm(axis - n * wp.dot(n, axis)) - - if b_norm < 0.5: - if -0.5 < n[1] and n[1] < 0.5: - b = wp.vec3(0.0, 1.0, 0.0) - else: - b = wp.vec3(0.0, 0.0, 1.0) - - c = wp.cross(n, b) - frame = wp.mat33(n[0], n[1], n[2], b[0], b[1], b[2], c[0], c[1], c[2]) - segment = axis * cap.size[1] - - dist1, pos1 = _plane_sphere(n, plane.pos, cap.pos + segment, cap.size[0]) - write_contact( - nconmax_in, - dist1, - pos1, - 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, - ) - - dist2, pos2 = _plane_sphere(n, plane.pos, cap.pos - segment, cap.size[0]) - write_contact( - nconmax_in, - dist2, - pos2, - 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, - ) - - -@wp.func -def plane_ellipsoid( - # Data in: - nconmax_in: int, - # In: - plane: Geom, - ellipsoid: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - sphere_support = -wp.normalize(wp.cw_mul(wp.transpose(ellipsoid.rot) @ plane.normal, ellipsoid.size)) - pos = ellipsoid.pos + ellipsoid.rot @ wp.cw_mul(sphere_support, ellipsoid.size) - dist = wp.dot(plane.normal, pos - plane.pos) - pos = pos - plane.normal * dist * 0.5 - - write_contact( - nconmax_in, - dist, - pos, - make_frame(plane.normal), - 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, - ) - - -@wp.func -def plane_box( - # Data in: - nconmax_in: int, - # In: - plane: Geom, - box: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - count = int(0) - corner = wp.vec3() - dist = wp.dot(box.pos - plane.pos, plane.normal) - - # test all corners, pick bottom 4 - for i in range(8): - # get corner in local coordinates - corner.x = wp.where(i & 1, box.size.x, -box.size.x) - corner.y = wp.where(i & 2, box.size.y, -box.size.y) - corner.z = wp.where(i & 4, box.size.z, -box.size.z) - - # get corner in global coordinates relative to box center - corner = box.rot * corner - - # compute distance to plane, skip if too far or pointing up - ldist = wp.dot(plane.normal, corner) - if dist + ldist > margin or ldist > 0: - continue - - cdist = dist + ldist - frame = make_frame(plane.normal) - pos = corner + box.pos + (plane.normal * cdist / -2.0) - write_contact( - nconmax_in, - cdist, - pos, - 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, - ) - count += 1 - if count >= 4: - break - - -_HUGE_VAL = 1e6 - - -@wp.func -def plane_convex( - # Data in: - nconmax_in: int, - # In: - plane: Geom, - convex: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - """Calculates contacts between a plane and a convex object.""" + contact_dist = wp.vec4(wp.inf) + contact_pos = mat43f() + contact_count = int(0) # get points in the convex frame - plane_pos = wp.transpose(convex.rot) @ (plane.pos - convex.pos) - n = wp.transpose(convex.rot) @ plane.normal + plane_pos_local = wp.transpose(convex.rot) @ (plane_pos - convex.pos) + n = wp.transpose(convex.rot) @ plane_normal # Store indices in vec4 indices = wp.vec4i(-1, -1, -1, -1) @@ -911,14 +170,14 @@ def plane_convex( # Find support points max_support = wp.float32(-_HUGE_VAL) for i in range(convex.vertnum): - support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + i], n) max_support = wp.max(support, max_support) threshold = wp.max(0.0, max_support - 1e-3) # Find point a (first support point) a_dist = wp.float32(-_HUGE_VAL) for i in range(convex.vertnum): - support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + i], n) dist = wp.where(support > threshold, support, -_HUGE_VAL) if dist > a_dist: indices[0] = i @@ -928,7 +187,7 @@ def plane_convex( # Find point b (furthest from a) b_dist = wp.float32(-_HUGE_VAL) for i in range(convex.vertnum): - support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + i], n) dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) dist = wp.length_sq(a - convex.vert[convex.vertadr + i]) + dist_mask if dist > b_dist: @@ -940,7 +199,7 @@ def plane_convex( ab = wp.cross(n, a - b) c_dist = wp.float32(-_HUGE_VAL) for i in range(convex.vertnum): - support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + i], n) dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) ap = a - convex.vert[convex.vertadr + i] dist = wp.abs(wp.dot(ap, ab)) + dist_mask @@ -954,7 +213,7 @@ def plane_convex( bc = wp.cross(n, b - c) d_dist = wp.float32(-_HUGE_VAL) for i in range(convex.vertnum): - support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + i], n) dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) ap = a - convex.vert[convex.vertadr + i] bp = b - convex.vert[convex.vertadr + i] @@ -983,7 +242,7 @@ def plane_convex( while convex.graph[edge_localid + i] >= 0: subidx = convex.graph[edge_localid + i] idx = convex.graph[vert_globalid + subidx] - support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + idx], n) if support > max_support: max_support = support imax = int(subidx) @@ -1000,7 +259,7 @@ def plane_convex( while convex.graph[edge_localid + i] >= 0: subidx = convex.graph[edge_localid + i] idx = convex.graph[vert_globalid + subidx] - support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + idx], n) dist = wp.where(support > threshold, support, -_HUGE_VAL) if dist > a_dist: a_dist = dist @@ -1008,9 +267,9 @@ def plane_convex( i += int(1) if imax == prev: break - imax = convex.graph[vert_globalid + imax] - a = convex.vert[convex.vertadr + imax] - indices[0] = imax + imax_global = convex.graph[vert_globalid + imax] + a = convex.vert[convex.vertadr + imax_global] + indices[0] = imax_global # Find point b (furthest from a) b_dist = wp.float32(-_HUGE_VAL) @@ -1020,7 +279,7 @@ def plane_convex( while convex.graph[edge_localid + i] >= 0: subidx = convex.graph[edge_localid + i] idx = convex.graph[vert_globalid + subidx] - support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + idx], n) dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) dist = wp.length_sq(a - convex.vert[convex.vertadr + idx]) + dist_mask if dist > b_dist: @@ -1029,9 +288,9 @@ def plane_convex( i += int(1) if imax == prev: break - imax = convex.graph[vert_globalid + imax] - b = convex.vert[convex.vertadr + imax] - indices[1] = imax + imax_global = convex.graph[vert_globalid + imax] + b = convex.vert[convex.vertadr + imax_global] + indices[1] = imax_global # Find point c (furthest along axis orthogonal to a-b) ab = wp.cross(n, a - b) @@ -1042,9 +301,9 @@ def plane_convex( while convex.graph[edge_localid + i] >= 0: subidx = convex.graph[edge_localid + i] idx = convex.graph[vert_globalid + subidx] - support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + idx], n) dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) - ap = a - convex.vert[convex.vertadr + i] + ap = a - convex.vert[convex.vertadr + idx] dist = wp.abs(wp.dot(ap, ab)) + dist_mask if dist > c_dist: c_dist = dist @@ -1052,9 +311,9 @@ def plane_convex( i += int(1) if imax == prev: break - imax = convex.graph[vert_globalid + imax] - c = convex.vert[convex.vertadr + imax] - indices[2] = imax + imax_global = convex.graph[vert_globalid + imax] + c = convex.vert[convex.vertadr + imax_global] + indices[2] = imax_global # Find point d (furthest from other triangle edges) ac = wp.cross(n, a - c) @@ -1066,7 +325,7 @@ def plane_convex( while convex.graph[edge_localid + i] >= 0: subidx = convex.graph[edge_localid + i] idx = convex.graph[vert_globalid + subidx] - support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + idx], n) dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL) ap = a - convex.vert[convex.vertadr + idx] bp = b - convex.vert[convex.vertadr + idx] @@ -1078,11 +337,10 @@ def plane_convex( i += int(1) if imax == prev: break - imax = convex.graph[vert_globalid + imax] - indices[3] = imax + imax_global = convex.graph[vert_globalid + imax] + indices[3] = imax_global - # Write contacts - frame = make_frame(plane.normal) + # Collect contacts from unique indices for i in range(3, -1, -1): idx = indices[i] count = int(0) @@ -1094,54 +352,34 @@ def plane_convex( if count == 1: pos = convex.vert[convex.vertadr + idx] pos = convex.pos + convex.rot @ pos - support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n) + support = wp.dot(plane_pos_local - convex.vert[convex.vertadr + idx], n) dist = -support - pos = pos - 0.5 * dist * plane.normal - write_contact( - nconmax_in, - dist, - pos, - 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, - ) + pos = pos - 0.5 * dist * plane_normal + + contact_dist[contact_count] = dist + contact_pos[contact_count] = pos + contact_count = contact_count + 1 + + return contact_dist, contact_pos, plane_normal @wp.func -def sphere_cylinder( +def write_contact( # Data in: nconmax_in: int, # In: - sphere: Geom, - cylinder: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, + dist_in: float, + pos_in: wp.vec3, + frame_in: wp.mat33, + margin_in: float, + gap_in: float, + condim_in: int, + friction_in: vec5, + solref_in: wp.vec2, + solreffriction_in: wp.vec2, + solimp_in: vec5, + geoms_in: wp.vec2i, + worldid_in: int, # Data out: ncon_out: wp.array(dtype=int), contact_dist_out: wp.array(dtype=float), @@ -1156,349 +394,21 @@ def sphere_cylinder( contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), ): - axis = wp.vec3( - cylinder.rot[0, 2], - cylinder.rot[1, 2], - cylinder.rot[2, 2], - ) - - vec = sphere.pos - cylinder.pos - x = wp.dot(vec, axis) - - a_proj = axis * x - p_proj = vec - a_proj - p_proj_sqr = wp.dot(p_proj, p_proj) - - collide_side = wp.abs(x) < cylinder.size[1] - collide_cap = p_proj_sqr < (cylinder.size[0] * cylinder.size[0]) - - if collide_side and collide_cap: - dist_cap = cylinder.size[1] - wp.abs(x) - dist_radius = cylinder.size[0] - wp.sqrt(p_proj_sqr) - - if dist_cap < dist_radius: - collide_side = False - else: - collide_cap = False - - # Side collision - if collide_side: - pos_target = cylinder.pos + a_proj - - _sphere_sphere_ext( - nconmax_in, - sphere.pos, - sphere.size[0], - pos_target, - cylinder.size[0], - worldid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - sphere.rot, - cylinder.rot, - 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, - ) - return - - # Cap collision - if collide_cap: - if x > 0.0: - # top cap - pos_cap = cylinder.pos + axis * cylinder.size[1] - plane_normal = axis - else: - # bottom cap - pos_cap = cylinder.pos - axis * cylinder.size[1] - plane_normal = -axis - - dist, pos_contact = _plane_sphere(plane_normal, pos_cap, sphere.pos, sphere.size[0]) - plane_normal = -plane_normal # Flip normal after position calculation - - write_contact( - nconmax_in, - dist, - pos_contact, - make_frame(plane_normal), - 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, - ) - - return - - # Corner collision - inv_len = safe_div(1.0, wp.sqrt(p_proj_sqr)) - p_proj = p_proj * (cylinder.size[0] * inv_len) - - cap_offset = axis * (wp.sign(x) * cylinder.size[1]) - pos_corner = cylinder.pos + cap_offset + p_proj - - _sphere_sphere_ext( - nconmax_in, - sphere.pos, - sphere.size[0], - pos_corner, - 0.0, - worldid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - sphere.rot, - cylinder.rot, - 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, - ) - - -@wp.func -def plane_cylinder( - # Data in: - nconmax_in: int, - # In: - plane: Geom, - cylinder: Geom, - worldid: int, - margin: float, - gap: float, - condim: int, - friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, - solimp: vec5, - geoms: wp.vec2i, - # 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), -): - """Calculates contacts between a cylinder and a plane.""" - # Extract plane normal and cylinder axis - n = plane.normal - axis = wp.vec3(cylinder.rot[0, 2], cylinder.rot[1, 2], cylinder.rot[2, 2]) - - # Project, make sure axis points toward plane - prjaxis = wp.dot(n, axis) - if prjaxis > 0: - axis = -axis - prjaxis = -prjaxis - - # Compute normal distance from plane to cylinder center - dist0 = wp.dot(cylinder.pos - plane.pos, n) - - # Remove component of -normal along cylinder axis - vec = axis * prjaxis - n - len_sqr = wp.dot(vec, vec) - - # If vector is nondegenerate, normalize and scale by radius - # Otherwise use cylinder's x-axis scaled by radius - vec = wp.where( - len_sqr >= 1e-12, - vec * safe_div(cylinder.size[0], wp.sqrt(len_sqr)), - wp.vec3(cylinder.rot[0, 0], cylinder.rot[1, 0], cylinder.rot[2, 0]) * cylinder.size[0], - ) - - # Project scaled vector on normal - prjvec = wp.dot(vec, n) - - # Scale cylinder axis by half-length - axis = axis * cylinder.size[1] - prjaxis = prjaxis * cylinder.size[1] - - frame = make_frame(n) - - # First contact point (end cap closer to plane) - dist1 = dist0 + prjaxis + prjvec - if dist1 <= margin: - pos1 = cylinder.pos + vec + axis - n * (dist1 * 0.5) - write_contact( - nconmax_in, - dist1, - pos1, - 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, - ) - else: - # If nearest point is above margin, no contacts - return - - # Second contact point (end cap farther from plane) - dist2 = dist0 - prjaxis + prjvec - if dist2 <= margin: - pos2 = cylinder.pos + vec - axis - n * (dist2 * 0.5) - write_contact( - nconmax_in, - dist2, - pos2, - make_frame(plane.normal), - 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, - ) - - # Try triangle contact points on side closer to plane - prjvec1 = -prjvec * 0.5 - dist3 = dist0 + prjaxis + prjvec1 - if dist3 <= margin: - # Compute sideways vector scaled by radius*sqrt(3)/2 - vec1 = wp.cross(vec, axis) - vec1 = wp.normalize(vec1) * (cylinder.size[0] * wp.sqrt(3.0) * 0.5) - - # Add contact point A - adjust to closest side - pos3 = cylinder.pos + vec1 + axis - vec * 0.5 - n * (dist3 * 0.5) - write_contact( - nconmax_in, - dist3, - pos3, - 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, - ) - - # Add contact point B - adjust to closest side - pos4 = cylinder.pos - vec1 + axis - vec * 0.5 - n * (dist3 * 0.5) - write_contact( - nconmax_in, - dist3, - pos4, - 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 dist_in - margin_in < 0.0: + cid = wp.atomic_add(ncon_out, 0, 1) + if cid < nconmax_in: + contact_dist_out[cid] = dist_in + contact_pos_out[cid] = pos_in + contact_frame_out[cid] = frame_in + contact_geom_out[cid] = geoms_in + contact_worldid_out[cid] = worldid_in + includemargin = margin_in - gap_in + contact_includemargin_out[cid] = includemargin + contact_dim_out[cid] = condim_in + contact_friction_out[cid] = friction_in + contact_solref_out[cid] = solref_in + contact_solreffriction_out[cid] = solreffriction_in + contact_solimp_out[cid] = solimp_in @wp.func @@ -1592,22 +502,19 @@ def contact_params( @wp.func -def _sphere_box( +def plane_sphere_wrapper( # Data in: nconmax_in: int, # In: - sphere_pos: wp.vec3, - sphere_size: float, - box_pos: wp.vec3, - box_rot: wp.mat33, - box_size: wp.vec3, + plane: Geom, + sphere: Geom, worldid: int, margin: float, gap: float, condim: int, friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, + solref: wp.vec2, + solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, # Data out: @@ -1624,68 +531,668 @@ def _sphere_box( contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), ): - center = wp.transpose(box_rot) @ (sphere_pos - box_pos) + """Calculates contact between a sphere and a plane.""" + dist, pos = plane_sphere(plane.normal, plane.pos, sphere.pos, sphere.size[0]) - clamped = wp.max(-box_size, wp.min(box_size, center)) - clamped_dir, dist = normalize_with_norm(clamped - center) - - if dist - sphere_size > margin: - return - - # sphere center inside box - if dist <= MJ_MINVAL: - closest = 2.0 * (box_size[0] + box_size[1] + box_size[2]) - k = wp.int32(0) - for i in range(6): - face_dist = wp.abs(wp.where(i % 2, 1.0, -1.0) * box_size[i // 2] - center[i // 2]) - if closest > face_dist: - closest = face_dist - k = i - - nearest = wp.vec3(0.0) - nearest[k // 2] = wp.where(k % 2, -1.0, 1.0) - pos = center + nearest * (sphere_size - closest) / 2.0 - contact_normal = box_rot @ nearest - contact_dist = -closest - sphere_size - - else: - deepest = center + clamped_dir * sphere_size - pos = 0.5 * (clamped + deepest) - contact_normal = box_rot @ clamped_dir - contact_dist = dist - sphere_size - - contact_pos = box_pos + box_rot @ pos - write_contact( - nconmax_in, - contact_dist, - contact_pos, - make_frame(contact_normal), - 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 dist - margin < 0: + write_contact( + nconmax_in, + dist, + pos, + make_frame(plane.normal), + 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, + ) @wp.func -def sphere_box( +def sphere_sphere_wrapper( + # Data in: + nconmax_in: int, + # In: + sphere1: Geom, + sphere2: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates contact between two spheres.""" + dist, pos, normal = sphere_sphere(sphere1.pos, sphere1.size[0], sphere2.pos, sphere2.size[0]) + + if dist - margin < 0: + write_contact( + nconmax_in, + dist, + pos, + make_frame(normal), + 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, + ) + + +@wp.func +def sphere_capsule_wrapper( + # Data in: + nconmax_in: int, + # In: + sphere: Geom, + cap: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates one contact between a sphere and a capsule.""" + # capsule axis + axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) + + dist, pos, normal = sphere_capsule(sphere.pos, sphere.size[0], cap.pos, axis, cap.size[0], cap.size[1]) + + if dist - margin < 0: + write_contact( + nconmax_in, + dist, + pos, + make_frame(normal), + 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, + ) + + +@wp.func +def capsule_capsule_wrapper( + # Data in: + nconmax_in: int, + # In: + cap1: Geom, + cap2: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates contacts between two capsules.""" + # capsule axes + cap1_axis = wp.vec3(cap1.rot[0, 2], cap1.rot[1, 2], cap1.rot[2, 2]) + cap2_axis = wp.vec3(cap2.rot[0, 2], cap2.rot[1, 2], cap2.rot[2, 2]) + + dist, pos, normal = capsule_capsule( + cap1.pos, + cap1_axis, + cap1.size[0], # radius1 + cap1.size[1], # half_length1 + cap2.pos, + cap2_axis, + cap2.size[0], # radius2 + cap2.size[1], # half_length2 + ) + + if dist - margin < 0: + write_contact( + nconmax_in, + dist, + pos, + make_frame(normal), + 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, + ) + + +@wp.func +def plane_capsule_wrapper( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + cap: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates contacts between a capsule and a plane.""" + # capsule axis + capsule_axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) + + dist, pos, frame = plane_capsule( + plane.normal, + plane.pos, + cap.pos, + capsule_axis, + cap.size[0], # radius + cap.size[1], # half_length + ) + + for i in range(2): + disti = dist[i] + if disti - margin < 0.0: + write_contact( + nconmax_in, + disti, + pos[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, + ) + + +@wp.func +def plane_ellipsoid_wrapper( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + ellipsoid: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates contacts between an ellipsoid and a plane.""" + dist, pos, normal = plane_ellipsoid(plane.normal, plane.pos, ellipsoid.pos, ellipsoid.rot, ellipsoid.size) + + if dist - margin < 0: + write_contact( + nconmax_in, + dist, + pos, + make_frame(normal), + 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, + ) + + +@wp.func +def plane_box_wrapper( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + box: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates contacts between a box and a plane.""" + dist, pos, normal = plane_box(plane.normal, plane.pos, box.pos, box.rot, box.size) + frame = make_frame(normal) + + for i in range(4): + disti = dist[i] + if disti - margin < 0.0: + write_contact( + nconmax_in, + disti, + pos[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, + ) + + +_HUGE_VAL = 1e6 + + +@wp.func +def plane_convex_wrapper( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + convex: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates contacts between a plane and a convex object.""" + dist, pos, normal = plane_convex(plane.normal, plane.pos, convex) + + frame = make_frame(normal) + for i in range(4): + disti = dist[i] + if disti - margin < 0.0: + write_contact( + nconmax_in, + disti, + pos[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, + ) + + +@wp.func +def sphere_cylinder_wrapper( + # Data in: + nconmax_in: int, + # In: + sphere: Geom, + cylinder: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates contacts between a sphere and a cylinder.""" + # cylinder axis + cylinder_axis = wp.vec3(cylinder.rot[0, 2], cylinder.rot[1, 2], cylinder.rot[2, 2]) + + dist, pos, normal = sphere_cylinder( + sphere.pos, + sphere.size[0], # sphere radius + cylinder.pos, + cylinder_axis, + cylinder.size[0], # cylinder radius + cylinder.size[1], # cylinder half_height + ) + + if dist - margin < 0.0: + write_contact( + nconmax_in, + dist, + pos, + make_frame(normal), + 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, + ) + + +@wp.func +def plane_cylinder_wrapper( + # Data in: + nconmax_in: int, + # In: + plane: Geom, + cylinder: Geom, + worldid: int, + margin: float, + gap: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solreffriction: wp.vec2, + solimp: vec5, + geoms: wp.vec2i, + # 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), +): + """Calculates contacts between a cylinder and a plane.""" + # cylinder axis + cylinder_axis = wp.vec3(cylinder.rot[0, 2], cylinder.rot[1, 2], cylinder.rot[2, 2]) + + dist, pos, normal = plane_cylinder( + plane.normal, + plane.pos, + cylinder.pos, + cylinder_axis, + cylinder.size[0], # radius + cylinder.size[1], # half_height + ) + + frame = make_frame(normal) + for i in range(4): + disti = dist[i] + if disti - margin < 0.0: + write_contact( + nconmax_in, + disti, + pos[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, + ) + + +@wp.func +def sphere_box_wrapper( # Data in: nconmax_in: int, # In: @@ -1696,8 +1203,8 @@ def sphere_box( gap: float, condim: int, friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, + solref: wp.vec2, + solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, # Data out: @@ -1714,39 +1221,40 @@ def sphere_box( contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), ): - _sphere_box( - nconmax_in, - sphere.pos, - sphere.size[0], - box.pos, - box.rot, - box.size, - worldid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - 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, - ) + dist, pos, normal = sphere_box(sphere.pos, sphere.size[0], box.pos, box.rot, box.size) + + if dist - margin < 0.0: + write_contact( + nconmax_in, + dist, + pos, + make_frame(normal), + 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, + ) @wp.func -def capsule_box( +def capsule_box_wrapper( # Data in: nconmax_in: int, # In: @@ -1757,8 +1265,8 @@ def capsule_box( gap: float, condim: int, friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, + solref: wp.vec2, + solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, # Data out: @@ -1776,396 +1284,55 @@ def capsule_box( contact_worldid_out: wp.array(dtype=int), ): """Calculates contacts between a capsule and a box.""" - # Based on the mjc implementation - boxmatT = wp.transpose(box.rot) - pos = boxmatT @ (cap.pos - box.pos) - axis = boxmatT @ wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) - halfaxis = axis * cap.size[1] # halfaxis is the capsule direction - axisdir = wp.int32(halfaxis[0] > 0.0) + 2 * wp.int32(halfaxis[1] > 0.0) + 4 * wp.int32(halfaxis[2] > 0.0) + # Extract capsule axis + axis = wp.vec3(cap.rot[0, 2], cap.rot[1, 2], cap.rot[2, 2]) - bestdistmax = margin + 2.0 * (cap.size[0] + cap.size[1] + box.size[0] + box.size[1] + box.size[2]) - - # keep track of closest point - bestdist = wp.float32(bestdistmax) - bestsegmentpos = wp.float32(-12) - - # cltype: encoded collision configuration - # cltype / 3 == 0 : lower corner is closest to the capsule - # == 2 : upper corner is closest to the capsule - # == 1 : middle of the edge is closest to the capsule - # cltype % 3 == 0 : lower corner is closest to the box - # == 2 : upper corner is closest to the box - # == 1 : middle of the capsule is closest to the box - cltype = wp.int32(-4) - - # clface: index of the closest face of the box to the capsule - # -1: no face is closest (edge or corner is closest) - # 0, 1, 2: index of the axis perpendicular to the closest face - clface = wp.int32(-12) - - # first: consider cases where a face of the box is closest - for i in range(-1, 2, 2): - axisTip = pos + wp.float32(i) * halfaxis - boxPoint = wp.vec3(axisTip) - - n_out = wp.int32(0) - ax_out = wp.int32(-1) - - for j in range(3): - if boxPoint[j] < -box.size[j]: - n_out += 1 - ax_out = j - boxPoint[j] = -box.size[j] - elif boxPoint[j] > box.size[j]: - n_out += 1 - ax_out = j - boxPoint[j] = box.size[j] - - if n_out > 1: - continue - - dist = wp.length_sq(boxPoint - axisTip) - - if dist < bestdist: - bestdist = dist - bestsegmentpos = wp.float32(i) - cltype = -2 + i - clface = ax_out - - # second: consider cases where an edge of the box is closest - clcorner = wp.int32(-123) # which corner is the closest - cledge = wp.int32(-123) # which axis - bestboxpos = wp.float32(0.0) - - for i in range(8): - for j in range(3): - if i & (1 << j) != 0: - continue - - c2 = wp.int32(-123) - - # box_pt is the starting point (corner) on the box - box_pt = wp.cw_mul( - wp.vec3( - wp.where(i & 1, 1.0, -1.0), - wp.where(i & 2, 1.0, -1.0), - wp.where(i & 4, 1.0, -1.0), - ), - box.size, - ) - box_pt[j] = 0.0 - - # find closest point between capsule and the edge - dif = box_pt - pos - - u = -box.size[j] * dif[j] - v = wp.dot(halfaxis, dif) - ma = box.size[j] * box.size[j] - mb = -box.size[j] * halfaxis[j] - mc = cap.size[1] * cap.size[1] - det = ma * mc - mb * mb - if wp.abs(det) < MJ_MINVAL: - continue - - idet = 1.0 / det - # sX : X=1 means middle of segment. X=0 or 2 one or the other end - - x1 = wp.float32((mc * u - mb * v) * idet) - x2 = wp.float32((ma * v - mb * u) * idet) - - s1 = wp.int32(1) - s2 = wp.int32(1) - - if x1 > 1: - x1 = 1.0 - s1 = 2 - x2 = safe_div(v - mb, mc) - elif x1 < -1: - x1 = -1.0 - s1 = 0 - x2 = safe_div(v + mb, mc) - - x2_over = x2 > 1.0 - if x2_over or x2 < -1.0: - if x2_over: - x2 = 1.0 - s2 = 2 - x1 = safe_div(u - mb, ma) - else: - x2 = -1.0 - s2 = 0 - x1 = safe_div(u + mb, ma) - - if x1 > 1: - x1 = 1.0 - s1 = 2 - elif x1 < -1: - x1 = -1.0 - s1 = 0 - - dif -= halfaxis * x2 - dif[j] += box.size[j] * x1 - - # encode relative positions of the closest points - ct = s1 * 3 + s2 - - dif_sq = wp.length_sq(dif) - if dif_sq < bestdist - MJ_MINVAL: - bestdist = dif_sq - bestsegmentpos = x2 - bestboxpos = x1 - # ct<6 means closest point on box is at lower end or middle of edge - c2 = ct // 6 - - clcorner = i + (1 << j) * c2 # index of closest box corner - cledge = j # axis index of closest box edge - cltype = ct # encoded collision configuration - - best = wp.float32(0.0) - - p = wp.vec2(pos.x, pos.y) - dd = wp.vec2(halfaxis.x, halfaxis.y) - s = wp.vec2(box.size.x, box.size.y) - secondpos = wp.float32(-4.0) - - uu = dd.x * s.y - vv = dd.y * s.x - w_neg = dd.x * p.y - dd.y * p.x < 0 - - best = wp.float32(-1.0) - - ee1 = uu - vv - ee2 = uu + vv - - if wp.abs(ee1) > best: - best = wp.abs(ee1) - c1 = wp.where((ee1 < 0) == w_neg, 0, 3) - - if wp.abs(ee2) > best: - best = wp.abs(ee2) - c1 = wp.where((ee2 > 0) == w_neg, 1, 2) - - if cltype == -4: # invalid type - return - - if cltype >= 0 and cltype // 3 != 1: # closest to a corner of the box - c1 = axisdir ^ clcorner - # Calculate relative orientation between capsule and corner - # There are two possible configurations: - # 1. Capsule axis points toward/away from corner - # 2. Capsule axis aligns with a face or edge - if c1 != 0 and c1 != 7: # create second contact point - if c1 == 1 or c1 == 2 or c1 == 4: - mul = 1 - else: - mul = -1 - c1 = 7 - c1 - - # "de" and "dp" distance from first closest point on the capsule to both ends of it - # mul is a direction along the capsule's axis - - if c1 == 1: - ax = 0 - ax1 = 1 - ax2 = 2 - elif c1 == 2: - ax = 1 - ax1 = 2 - ax2 = 0 - elif c1 == 4: - ax = 2 - ax1 = 0 - ax2 = 1 - - if axis[ax] * axis[ax] > 0.5: # second point along the edge of the box - m = 2.0 * safe_div(box.size[ax], wp.abs(halfaxis[ax])) - secondpos = min(1.0 - wp.float32(mul) * bestsegmentpos, m) - else: # second point along a face of the box - # check for overshoot again - m = 2.0 * min( - safe_div(box.size[ax1], wp.abs(halfaxis[ax1])), - safe_div(box.size[ax2], wp.abs(halfaxis[ax2])), - ) - secondpos = -min(1.0 + wp.float32(mul) * bestsegmentpos, m) - secondpos *= wp.float32(mul) - - elif cltype >= 0 and cltype // 3 == 1: # we are on box's edge - # Calculate relative orientation between capsule and edge - # Two possible configurations: - # - T configuration: c1 = 2^n (no additional contacts) - # - X configuration: c1 != 2^n (potential additional contacts) - c1 = axisdir ^ clcorner - c1 &= 7 - (1 << cledge) # mask out edge axis to determine configuration - - if c1 == 1 or c1 == 2 or c1 == 4: # create second contact point - if cledge == 0: - ax1 = 1 - ax2 = 2 - if cledge == 1: - ax1 = 2 - ax2 = 0 - if cledge == 2: - ax1 = 0 - ax2 = 1 - ax = cledge - - # find which face the capsule has a lower angle, and switch the axis - if wp.abs(axis[ax1]) > wp.abs(axis[ax2]): - ax1 = ax2 - ax2 = 3 - ax - ax1 - - # mul determines direction along capsule axis for second contact point - if c1 & (1 << ax2): - mul = 1 - secondpos = 1.0 - bestsegmentpos - else: - mul = -1 - secondpos = 1.0 + bestsegmentpos - - # now find out whether we point towards the opposite side or towards one of the sides - # and also find the farthest point along the capsule that is above the box - - e1 = 2.0 * safe_div(box.size[ax2], wp.abs(halfaxis[ax2])) - secondpos = min(e1, secondpos) - - if ((axisdir & (1 << ax)) != 0) == ((c1 & (1 << ax2)) != 0): - e2 = 1.0 - bestboxpos - else: - e2 = 1.0 + bestboxpos - - e1 = box.size[ax] * safe_div(e2, wp.abs(halfaxis[ax])) - - secondpos = min(e1, secondpos) - secondpos *= wp.float32(mul) - - elif cltype < 0: - # similarly we handle the case when one capsule's end is closest to a face of the box - # and find where is the other end pointing to and clamping to the farthest point - # of the capsule that's above the box - # if the closest point is inside the box there's no need for a second point - - if clface != -1: # create second contact point - mul = wp.where(cltype == -3, 1, -1) - secondpos = 2.0 - - tmp1 = pos - halfaxis * wp.float32(mul) - - for i in range(3): - if i != clface: - ha_r = safe_div(wp.float32(mul), halfaxis[i]) - e1 = (box.size[i] - tmp1[i]) * ha_r - if 0 < e1 and e1 < secondpos: - secondpos = e1 - - e1 = (-box.size[i] - tmp1[i]) * ha_r - if 0 < e1 and e1 < secondpos: - secondpos = e1 - - secondpos *= wp.float32(mul) - - # create sphere in original orientation at first contact point - s1_pos_l = pos + halfaxis * bestsegmentpos - s1_pos_g = box.rot @ s1_pos_l + box.pos - - # collide with sphere - _sphere_box( - nconmax_in, - s1_pos_g, - cap.size[0], + # Call the core function to get contact geometry + dist, pos, normal = capsule_box( + cap.pos, + axis, + cap.size[0], # capsule radius + cap.size[1], # capsule half length box.pos, box.rot, box.size, - worldid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - 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 secondpos > -3: # secondpos was modified - s2_pos_l = pos + halfaxis * (secondpos + bestsegmentpos) - s2_pos_g = box.rot @ s2_pos_l + box.pos - _sphere_box( - nconmax_in, - s2_pos_g, - cap.size[0], - box.pos, - box.rot, - box.size, - worldid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - 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, - ) + # Loop over the contacts and write them + for i in range(2): + disti = dist[i] + if disti - margin < 0.0: + write_contact( + nconmax_in, + disti, + pos[i], + make_frame(normal[i]), + 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, + ) @wp.func -def _compute_rotmore(face_idx: int) -> wp.mat33: - rotmore = wp.mat33(0.0) - - if face_idx == 0: - rotmore[0, 2] = -1.0 - rotmore[1, 1] = +1.0 - rotmore[2, 0] = +1.0 - elif face_idx == 1: - rotmore[0, 0] = +1.0 - rotmore[1, 2] = -1.0 - rotmore[2, 1] = +1.0 - elif face_idx == 2: - rotmore[0, 0] = +1.0 - rotmore[1, 1] = +1.0 - rotmore[2, 2] = +1.0 - elif face_idx == 3: - rotmore[0, 2] = +1.0 - rotmore[1, 1] = +1.0 - rotmore[2, 0] = -1.0 - elif face_idx == 4: - rotmore[0, 0] = +1.0 - rotmore[1, 2] = +1.0 - rotmore[2, 1] = -1.0 - elif face_idx == 5: - rotmore[0, 0] = -1.0 - rotmore[1, 1] = +1.0 - rotmore[2, 2] = -1.0 - - return rotmore - - -@wp.func -def box_box( +def box_box_wrapper( # Data in: nconmax_in: int, # In: @@ -2176,8 +1343,8 @@ def box_box( gap: float, condim: int, friction: vec5, - solref: wp.vec2f, - solreffriction: wp.vec2f, + solref: wp.vec2, + solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, # Data out: @@ -2194,450 +1361,64 @@ def box_box( contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), ): - # Compute transforms between box's frames + """Calculates contacts between two boxes.""" + # Call the core function to get contact geometry + dist, pos, normal = box_box( + box1.pos, + box1.rot, + box1.size, + box2.pos, + box2.rot, + box2.size, + ) - pos21 = wp.transpose(box1.rot) @ (box2.pos - box1.pos) - pos12 = wp.transpose(box2.rot) @ (box1.pos - box2.pos) + for i in range(8): + if dist[i] - margin >= 0.0: + continue - rot21 = wp.transpose(box1.rot) @ box2.rot - rot12 = wp.transpose(rot21) - - rot21abs = wp.matrix_from_rows(wp.abs(rot21[0]), wp.abs(rot21[1]), wp.abs(rot21[2])) - rot12abs = wp.transpose(rot21abs) - - plen2 = rot21abs @ box2.size - plen1 = rot12abs @ box1.size - - # Compute axis of maximum separation - s_sum_3 = 3.0 * (box1.size + box2.size) - separation = wp.float32(margin + s_sum_3[0] + s_sum_3[1] + s_sum_3[2]) - axis_code = wp.int32(-1) - - # First test: consider boxes' face normals - for i in range(3): - c1 = -wp.abs(pos21[i]) + box1.size[i] + plen2[i] - - c2 = -wp.abs(pos12[i]) + box2.size[i] + plen1[i] - - if c1 < -margin or c2 < -margin: - return - - if c1 < separation: - separation = c1 - axis_code = i + 3 * wp.int32(pos21[i] < 0) + 0 # Face of box1 - if c2 < separation: - separation = c2 - axis_code = i + 3 * wp.int32(pos12[i] < 0) + 6 # Face of box2 - - clnorm = wp.vec3(0.0) - inv = wp.bool(False) - cle1 = wp.int32(0) - cle2 = wp.int32(0) - - # Second test: consider cross products of boxes' edges - for i in range(3): - for j in range(3): - # Compute cross product of box edges (potential separating axis) - if i == 0: - cross_axis = wp.vec3(0.0, -rot12[j, 2], rot12[j, 1]) - elif i == 1: - cross_axis = wp.vec3(rot12[j, 2], 0.0, -rot12[j, 0]) - else: - cross_axis = wp.vec3(-rot12[j, 1], rot12[j, 0], 0.0) - - cross_length = wp.length(cross_axis) - if cross_length < MJ_MINVAL: - continue - - cross_axis /= cross_length - - box_dist = wp.dot(pos21, cross_axis) - c3 = wp.float32(0.0) - - # Project box half-sizes onto the potential separating axis - for k in range(3): - if k != i: - c3 += box1.size[k] * wp.abs(cross_axis[k]) - if k != j: - c3 += box2.size[k] * rot21abs[i, 3 - k - j] / cross_length - - c3 -= wp.abs(box_dist) - - # Early exit: no collision if separated along this axis - if c3 < -margin: - return - - # Track minimum separation and which edge-edge pair it occurs on - if c3 < separation * (1.0 - 1e-12): - separation = c3 - # Determine which corners/edges are closest - cle1 = 0 - cle2 = 0 - - for k in range(3): - if k != i and (int(cross_axis[k] > 0) ^ int(box_dist < 0)): - cle1 += 1 << k - if k != j: - if int(rot21[i, 3 - k - j] > 0) ^ int(box_dist < 0) ^ int((k - j + 3) % 3 == 1): - cle2 += 1 << k - - axis_code = 12 + i * 3 + j - clnorm = cross_axis - inv = box_dist < 0 - - # No axis with separation < margin found - if axis_code == -1: - return - - points = mat83f() - depth = vec8f() - max_con_pair = 8 - # 8 contacts should suffice for most configurations - - if axis_code < 12: - # Handle face-vertex collision - face_idx = axis_code % 6 - box_idx = axis_code // 6 - rotmore = _compute_rotmore(face_idx) - - r = rotmore @ wp.where(box_idx, rot12, rot21) - p = rotmore @ wp.where(box_idx, pos12, pos21) - ss = wp.abs(rotmore @ wp.where(box_idx, box2.size, box1.size)) - s = wp.where(box_idx, box1.size, box2.size) - rt = wp.transpose(r) - - lx, ly, hz = ss[0], ss[1], ss[2] - p[2] -= hz - - clcorner = wp.int32(0) # corner of non-face box with least axis separation - - for i in range(3): - if r[2, i] < 0: - clcorner += 1 << i - - lp = p - for i in range(wp.static(3)): - lp += rt[i] * s[i] * wp.where(clcorner & 1 << i, 1.0, -1.0) - - m = wp.int32(1) - dirs = wp.int32(0) - - cn1 = wp.vec3(0.0) - cn2 = wp.vec3(0.0) - - for i in range(3): - if wp.abs(r[2, i]) < 0.5: - if not dirs: - cn1 = rt[i] * s[i] * wp.where(clcorner & (1 << i), -2.0, 2.0) - else: - cn2 = rt[i] * s[i] * wp.where(clcorner & (1 << i), -2.0, 2.0) - - dirs += 1 - - k = dirs * dirs - - # Find potential contact points - - n = wp.int32(0) - - for i in range(k): - for q in range(2): - # lines_a and lines_b (lines between corners) computed on the fly - lav = lp + wp.where(i < 2, wp.vec3(0.0), wp.where(i == 2, cn1, cn2)) - lbv = wp.where(i == 0 or i == 3, cn1, cn2) - - if wp.abs(lbv[q]) > MJ_MINVAL: - br = 1.0 / lbv[q] - for j in range(-1, 2, 2): - l = ss[q] * wp.float32(j) - c1 = (l - lav[q]) * br - if c1 < 0 or c1 > 1: - continue - c2 = lav[1 - q] + lbv[1 - q] * c1 - if wp.abs(c2) > ss[1 - q]: - continue - - points[n] = lav + c1 * lbv - n += 1 - - if dirs == 2: - ax = cn1[0] - bx = cn2[0] - ay = cn1[1] - by = cn2[1] - C = safe_div(1.0, ax * by - bx * ay) - - for i in range(4): - llx = wp.where(i // 2, lx, -lx) - lly = wp.where(i % 2, ly, -ly) - - x = llx - lp[0] - y = lly - lp[1] - - u = (x * by - y * bx) * C - v = (y * ax - x * ay) * C - - if u > 0 and v > 0 and u < 1 and v < 1: - points[n] = wp.vec3(llx, lly, lp[2] + u * cn1[2] + v * cn2[2]) - n += 1 - - for i in range(1 << dirs): - tmpv = lp + wp.float32(i & 1) * cn1 + wp.float32((i & 2) != 0) * cn2 - if tmpv[0] > -lx and tmpv[0] < lx and tmpv[1] > -ly and tmpv[1] < ly: - points[n] = tmpv - n += 1 - - m = n - n = wp.int32(0) - - for i in range(m): - if points[i][2] > margin: - continue - if i != n: - points[n] = points[i] - - points[n, 2] *= 0.5 - depth[n] = points[n, 2] - n += 1 - - # Set up contact frame - rw = wp.where(box_idx, box2.rot, box1.rot) @ wp.transpose(rotmore) - pw = wp.where(box_idx, box2.pos, box1.pos) - normal = wp.where(box_idx, -1.0, 1.0) * wp.transpose(rw)[2] - - else: - # Handle edge-edge collision - edge1 = (axis_code - 12) // 3 - edge2 = (axis_code - 12) % 3 - - # Set up non-contacting edges ax1, ax2 for box2 and pax1, pax2 for box 1 - ax1 = wp.int(1 - (edge2 & 1)) - ax2 = wp.int(2 - (edge2 & 2)) - - pax1 = wp.int(1 - (edge1 & 1)) - pax2 = wp.int(2 - (edge1 & 2)) - - if rot21abs[edge1, ax1] < rot21abs[edge1, ax2]: - ax1, ax2 = ax2, ax1 - - if rot12abs[edge2, pax1] < rot12abs[edge2, pax2]: - pax1, pax2 = pax2, pax1 - - rotmore = _compute_rotmore(wp.where(cle1 & (1 << pax2), pax2, pax2 + 3)) - - # Transform coordinates for edge-edge contact calculation - p = rotmore @ pos21 - rnorm = rotmore @ clnorm - r = rotmore @ rot21 - rt = wp.transpose(r) - s = wp.abs(wp.transpose(rotmore) @ box1.size) - - lx, ly, hz = s[0], s[1], s[2] - p[2] -= hz - - # Calculate closest box2 face - - points[0] = ( - p - + rt[ax1] * box2.size[ax1] * wp.where(cle2 & (1 << ax1), 1.0, -1.0) - + rt[ax2] * box2.size[ax2] * wp.where(cle2 & (1 << ax2), 1.0, -1.0) + write_contact( + nconmax_in, + dist[i], + pos[i], + make_frame(normal[i]), + 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, ) - points[1] = points[0] - rt[edge2] * box2.size[edge2] - points[0] += rt[edge2] * box2.size[edge2] - - points[2] = ( - p - + rt[ax1] * box2.size[ax1] * wp.where(cle2 & (1 << ax1), -1.0, 1.0) - + rt[ax2] * box2.size[ax2] * wp.where(cle2 & (1 << ax2), 1.0, -1.0) - ) - - points[3] = points[2] - rt[edge2] * box2.size[edge2] - points[2] += rt[edge2] * box2.size[edge2] - - n = 4 - - # Set up coordinate axes for contact face of box2 - axi_lp = points[0] - axi_cn1 = points[1] - points[0] - axi_cn2 = points[2] - points[0] - - # Check if contact normal is valid - if wp.abs(rnorm[2]) < MJ_MINVAL: - return # Shouldn't happen - - # Calculate inverse normal for projection - innorm = wp.where(inv, -1.0, 1.0) / rnorm[2] - - pu = mat43f() - - # Project points onto contact plane - for i in range(4): - pu[i] = points[i] - c_scl = points[i, 2] * wp.where(inv, -1.0, 1.0) * innorm - points[i] -= rnorm * c_scl - - pts_lp = points[0] - pts_cn1 = points[1] - points[0] - pts_cn2 = points[2] - points[0] - - n = wp.int32(0) - - for i in range(4): - for q in range(2): - la = pts_lp[q] + wp.where(i < 2, 0.0, wp.where(i == 2, pts_cn1[q], pts_cn2[q])) - lb = wp.where(i == 0 or i == 3, pts_cn1[q], pts_cn2[q]) - lc = pts_lp[1 - q] + wp.where(i < 2, 0.0, wp.where(i == 2, pts_cn1[1 - q], pts_cn2[1 - q])) - ld = wp.where(i == 0 or i == 3, pts_cn1[1 - q], pts_cn2[1 - q]) - - # linesu_a and linesu_b (lines between corners) computed on the fly - lua = axi_lp + wp.where(i < 2, wp.vec3(0.0), wp.where(i == 2, axi_cn1, axi_cn2)) - lub = wp.where(i == 0 or i == 3, axi_cn1, axi_cn2) - - if wp.abs(lb) > MJ_MINVAL: - br = 1.0 / lb - for j in range(-1, 2, 2): - if n == max_con_pair: - break - l = s[q] * wp.float32(j) - c1 = (l - la) * br - if c1 < 0 or c1 > 1: - continue - c2 = lc + ld * c1 - if wp.abs(c2) > s[1 - q]: - continue - if (lua[2] + lub[2] * c1) * innorm > margin: - continue - - points[n] = lua * 0.5 + c1 * lub * 0.5 - points[n, q] += 0.5 * l - points[n, 1 - q] += 0.5 * c2 - depth[n] = points[n, 2] * innorm * 2.0 - n += 1 - - nl = n - - ax = pts_cn1[0] - bx = pts_cn2[0] - ay = pts_cn1[1] - by = pts_cn2[1] - C = safe_div(1.0, ax * by - bx * ay) - - for i in range(4): - if n == max_con_pair: - break - llx = wp.where(i // 2, lx, -lx) - lly = wp.where(i % 2, ly, -ly) - - x = llx - pts_lp[0] - y = lly - pts_lp[1] - - u = (x * by - y * bx) * C - v = (y * ax - x * ay) * C - - if nl == 0: - if (u < 0 or u > 1) and (v < 0 or v > 1): - continue - elif u < 0 or v < 0 or u > 1 or v > 1: - continue - - u = wp.clamp(u, 0.0, 1.0) - v = wp.clamp(v, 0.0, 1.0) - w = 1.0 - u - v - vtmp = pu[0] * w + pu[1] * u + pu[2] * v - - points[n] = wp.vec3(llx, lly, 0.0) - - vtmp2 = points[n] - vtmp - tc1 = wp.length_sq(vtmp2) - if vtmp[2] > 0 and tc1 > margin * margin: - continue - - points[n] = 0.5 * (points[n] + vtmp) - - depth[n] = wp.sqrt(tc1) * wp.where(vtmp[2] < 0, -1.0, 1.0) - n += 1 - - nf = n - - for i in range(4): - if n >= max_con_pair: - break - x = pu[i, 0] - y = pu[i, 1] - if nl == 0 and nf != 0: - if (x < -lx or x > lx) and (y < -ly or y > ly): - continue - elif x < -lx or x > lx or y < -ly or y > ly: - continue - - c1 = wp.float32(0) - - for j in range(2): - if pu[i, j] < -s[j]: - c1 += (pu[i, j] + s[j]) * (pu[i, j] + s[j]) - elif pu[i, j] > s[j]: - c1 += (pu[i, j] - s[j]) * (pu[i, j] - s[j]) - - c1 += pu[i, 2] * innorm * pu[i, 2] * innorm - - if pu[i, 2] > 0 and c1 > margin * margin: - continue - - tmp_p = wp.vec3(pu[i, 0], pu[i, 1], 0.0) - - for j in range(2): - if pu[i, j] < -s[j]: - tmp_p[j] = -s[j] * 0.5 - elif pu[i, j] > s[j]: - tmp_p[j] = +s[j] * 0.5 - - tmp_p += pu[i] - points[n] = tmp_p * 0.5 - - depth[n] = wp.sqrt(c1) * wp.where(pu[i, 2] < 0, -1.0, 1.0) - n += 1 - - # Set up contact data for all points - rw = box1.rot @ wp.transpose(rotmore) - pw = box1.pos - normal = wp.where(inv, -1.0, 1.0) * rw @ rnorm - - frame = make_frame(normal) - coff = wp.atomic_add(ncon_out, 0, n) - - for i in range(min(nconmax_in - coff, n)): - points[i, 2] += hz - pos = rw @ points[i] + pw - - cid = coff + i - - contact_dist_out[cid] = depth[i] - contact_pos_out[cid] = pos - contact_frame_out[cid] = frame - contact_geom_out[cid] = geoms - contact_worldid_out[cid] = worldid - contact_includemargin_out[cid] = margin - gap - contact_dim_out[cid] = condim - contact_friction_out[cid] = friction - contact_solref_out[cid] = solref - contact_solreffriction_out[cid] = solreffriction - contact_solimp_out[cid] = solimp _PRIMITIVE_COLLISIONS = { - (GeomType.PLANE.value, GeomType.SPHERE.value): plane_sphere, - (GeomType.PLANE.value, GeomType.CAPSULE.value): plane_capsule, - (GeomType.PLANE.value, GeomType.ELLIPSOID.value): plane_ellipsoid, - (GeomType.PLANE.value, GeomType.CYLINDER.value): plane_cylinder, - (GeomType.PLANE.value, GeomType.BOX.value): plane_box, - (GeomType.PLANE.value, GeomType.MESH.value): plane_convex, - (GeomType.SPHERE.value, GeomType.SPHERE.value): sphere_sphere, - (GeomType.SPHERE.value, GeomType.CAPSULE.value): sphere_capsule, - (GeomType.SPHERE.value, GeomType.CYLINDER.value): sphere_cylinder, - (GeomType.SPHERE.value, GeomType.BOX.value): sphere_box, - (GeomType.CAPSULE.value, GeomType.CAPSULE.value): capsule_capsule, - (GeomType.CAPSULE.value, GeomType.BOX.value): capsule_box, - (GeomType.BOX.value, GeomType.BOX.value): box_box, + (GeomType.PLANE.value, GeomType.SPHERE.value): plane_sphere_wrapper, + (GeomType.PLANE.value, GeomType.CAPSULE.value): plane_capsule_wrapper, + (GeomType.PLANE.value, GeomType.ELLIPSOID.value): plane_ellipsoid_wrapper, + (GeomType.PLANE.value, GeomType.CYLINDER.value): plane_cylinder_wrapper, + (GeomType.PLANE.value, GeomType.BOX.value): plane_box_wrapper, + (GeomType.PLANE.value, GeomType.MESH.value): plane_convex_wrapper, + (GeomType.SPHERE.value, GeomType.SPHERE.value): sphere_sphere_wrapper, + (GeomType.SPHERE.value, GeomType.CAPSULE.value): sphere_capsule_wrapper, + (GeomType.SPHERE.value, GeomType.CYLINDER.value): sphere_cylinder_wrapper, + (GeomType.SPHERE.value, GeomType.BOX.value): sphere_box_wrapper, + (GeomType.CAPSULE.value, GeomType.CAPSULE.value): capsule_capsule_wrapper, + (GeomType.CAPSULE.value, GeomType.BOX.value): capsule_box_wrapper, + (GeomType.BOX.value, GeomType.BOX.value): box_box_wrapper, } @@ -2665,7 +1446,7 @@ def _primitive_narrowphase_builder(m: Model): _primitive_collisions_types.append(types) _primitive_collisions_func.append(func) - @wp.kernel + @nested_kernel(module="unique", enable_backward=False) def _primitive_narrowphase( # Model: geom_type: wp.array(dtype=int), @@ -2710,7 +1491,6 @@ def _primitive_narrowphase_builder(m: Model): 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), @@ -2764,24 +1544,18 @@ def _primitive_narrowphase_builder(m: Model): 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( + type1, + 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, @@ -2789,29 +1563,22 @@ def _primitive_narrowphase_builder(m: Model): 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( + type2, + 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, @@ -2819,11 +1586,8 @@ def _primitive_narrowphase_builder(m: Model): 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], ) for i in range(wp.static(len(_primitive_collisions_func))): @@ -2924,7 +1688,6 @@ def primitive_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, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py new file mode 100644 index 00000000..f20df572 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py @@ -0,0 +1,1433 @@ +# 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. +# ============================================================================== + +from typing import Any, Tuple + +import warp as wp + +MJ_MINVAL = 1e-15 + + +wp.set_module_options({"enable_backward": False}) + + +@wp.func +def safe_div(x: Any, y: Any) -> Any: + return x / wp.where(y != 0.0, y, MJ_MINVAL) + + +@wp.func +def normalize_with_norm(x: Any): + norm = wp.length(x) + if norm == 0.0: + return x, 0.0 + return x / norm, norm + + +@wp.func +def closest_segment_point(a: wp.vec3, b: wp.vec3, pt: wp.vec3) -> wp.vec3: + """Returns the closest point on the a-b line segment to a point pt.""" + ab = b - a + t = wp.dot(pt - a, ab) / (wp.dot(ab, ab) + 1e-6) + return a + wp.clamp(t, 0.0, 1.0) * ab + + +@wp.func +def closest_segment_point_and_dist(a: wp.vec3, b: wp.vec3, pt: wp.vec3) -> Tuple[wp.vec3, float]: + """Returns closest point on the line segment and the distance squared.""" + closest = closest_segment_point(a, b, pt) + dist = wp.dot((pt - closest), (pt - closest)) + return closest, dist + + +@wp.func +def closest_segment_to_segment_points(a0: wp.vec3, a1: wp.vec3, b0: wp.vec3, b1: wp.vec3) -> Tuple[wp.vec3, wp.vec3]: + """Returns closest points between two line segments.""" + + dir_a, len_a = normalize_with_norm(a1 - a0) + dir_b, len_b = normalize_with_norm(b1 - b0) + + half_len_a = len_a * 0.5 + half_len_b = len_b * 0.5 + a_mid = a0 + dir_a * half_len_a + b_mid = b0 + dir_b * half_len_b + + trans = a_mid - b_mid + + dira_dot_dirb = wp.dot(dir_a, dir_b) + dira_dot_trans = wp.dot(dir_a, trans) + dirb_dot_trans = wp.dot(dir_b, trans) + denom = 1.0 - dira_dot_dirb * dira_dot_dirb + + orig_t_a = (-dira_dot_trans + dira_dot_dirb * dirb_dot_trans) / (denom + 1e-6) + orig_t_b = dirb_dot_trans + orig_t_a * dira_dot_dirb + t_a = wp.clamp(orig_t_a, -half_len_a, half_len_a) + t_b = wp.clamp(orig_t_b, -half_len_b, half_len_b) + + best_a = a_mid + dir_a * t_a + best_b = b_mid + dir_b * t_b + + new_a, d1 = closest_segment_point_and_dist(a0, a1, best_b) + new_b, d2 = closest_segment_point_and_dist(b0, b1, best_a) + if d1 < d2: + return new_a, best_b + return best_a, new_b + + +class vec8f(wp.types.vector(length=8, dtype=wp.float32)): + pass + + +class mat23f(wp.types.matrix(shape=(2, 3), dtype=wp.float32)): + pass + + +class mat43f(wp.types.matrix(shape=(4, 3), dtype=wp.float32)): + pass + + +class mat83f(wp.types.matrix(shape=(8, 3), dtype=wp.float32)): + pass + + +# core +@wp.func +def plane_sphere(plane_normal: wp.vec3, plane_pos: wp.vec3, sphere_pos: wp.vec3, sphere_radius: float) -> Tuple[float, wp.vec3]: + # TODO(team): docstring + dist = wp.dot(sphere_pos - plane_pos, plane_normal) - sphere_radius + pos = sphere_pos - plane_normal * (sphere_radius + 0.5 * dist) + return dist, pos + + +@wp.func +def sphere_sphere( + # In: + pos1: wp.vec3, + radius1: float, + pos2: wp.vec3, + radius2: float, +) -> Tuple[float, wp.vec3, wp.vec3]: + """Sphere-sphere collision calculation. + + Args: + pos1: Center position of the first sphere + radius1: Radius of the first sphere + pos2: Center position of the second sphere + radius2: Radius of the second sphere + + Returns: + Tuple containing: + dist: Distance between sphere surfaces (negative if overlapping) + pos: Contact position + n: Contact normal vector + """ + dir = pos2 - pos1 + dist = wp.length(dir) + if dist == 0.0: + n = wp.vec3(1.0, 0.0, 0.0) + else: + n = dir / dist + dist = dist - (radius1 + radius2) + pos = pos1 + n * (radius1 + 0.5 * dist) + return dist, pos, n + + +@wp.func +def sphere_capsule( + # In: + sphere_pos: wp.vec3, + sphere_radius: float, + capsule_pos: wp.vec3, + capsule_axis: wp.vec3, + capsule_radius: float, + capsule_half_length: float, +) -> Tuple[float, wp.vec3, wp.vec3]: + """Core contact geometry calculation for sphere-capsule collision. + + Args: + sphere_pos: Center position of the sphere + sphere_radius: Radius of the sphere + capsule_pos: Center position of the capsule + capsule_axis: Axis direction of the capsule + capsule_radius: Radius of the capsule + capsule_half_length: Half length of the capsule + + Returns: + Tuple containing: + contact_dist: Vector of contact distances + contact_pos: Matrix of contact positions (one per row) + contact_normals: Matrix of contact normal vectors (one per row) + """ + + # Calculate capsule segment + segment = capsule_axis * capsule_half_length + + # Find closest point on capsule centerline to sphere center + pt = closest_segment_point(capsule_pos - segment, capsule_pos + segment, sphere_pos) + + # Use sphere-sphere collision between sphere and closest point + return sphere_sphere(sphere_pos, sphere_radius, pt, capsule_radius) + + +@wp.func +def capsule_capsule( + # In: + cap1_pos: wp.vec3, + cap1_axis: wp.vec3, + cap1_radius: float, + cap1_half_length: float, + cap2_pos: wp.vec3, + cap2_axis: wp.vec3, + cap2_radius: float, + cap2_half_length: float, +) -> Tuple[float, wp.vec3, wp.vec3]: + """Core contact geometry calculation for capsule-capsule collision. + + Args: + cap1_pos: Center position of the first capsule + cap1_axis: Axis direction of the first capsule + cap1_radius: Radius of the first capsule + cap1_half_length: Half length of the first capsule + cap2_pos: Center position of the second capsule + cap2_axis: Axis direction of the second capsule + cap2_radius: Radius of the second capsule + cap2_half_length: Half length of the second capsule + + Returns: + Tuple containing: + contact_dist: Vector of contact distances + contact_pos: Matrix of contact positions (one per row) + contact_normals: Matrix of contact normal vectors (one per row) + """ + + # TODO(team): parallel axes case + + # Calculate capsule segments + seg1 = cap1_axis * cap1_half_length + seg2 = cap2_axis * cap2_half_length + + # Find closest points between capsule centerlines + pt1, pt2 = closest_segment_to_segment_points( + cap1_pos - seg1, + cap1_pos + seg1, + cap2_pos - seg2, + cap2_pos + seg2, + ) + + # Use sphere-sphere collision between closest points + return sphere_sphere(pt1, cap1_radius, pt2, cap2_radius) + + +@wp.func +def plane_capsule( + # In: + plane_normal: wp.vec3, + plane_pos: wp.vec3, + capsule_pos: wp.vec3, + capsule_axis: wp.vec3, + capsule_radius: float, + capsule_half_length: float, +) -> Tuple[wp.vec2, mat23f, wp.mat33]: + """Core contact geometry calculation for plane-capsule collision. + + Args: + plane_normal: Normal vector of the plane + plane_pos: Position point on the plane + capsule_pos: Center position of the capsule + capsule_axis: Axis direction of the capsule + capsule_radius: Radius of the capsule + capsule_half_length: Half length of the capsule + + Returns: + Tuple containing: + contact_dist: Vector of contact distances + contact_pos: Matrix of contact positions (one per row) + contact_frame: Contact frame for both contacts + """ + + n = plane_normal + axis = capsule_axis + + # align contact frames with capsule axis + b, b_norm = normalize_with_norm(axis - n * wp.dot(n, axis)) + + if b_norm < 0.5: + if -0.5 < n[1] and n[1] < 0.5: + b = wp.vec3(0.0, 1.0, 0.0) + else: + b = wp.vec3(0.0, 0.0, 1.0) + + c = wp.cross(n, b) + frame = wp.mat33(n[0], n[1], n[2], b[0], b[1], b[2], c[0], c[1], c[2]) + segment = axis * capsule_half_length + + # First contact (positive end of capsule) + dist1, pos1 = plane_sphere(n, plane_pos, capsule_pos + segment, capsule_radius) + + # Second contact (negative end of capsule) + dist2, pos2 = plane_sphere(n, plane_pos, capsule_pos - segment, capsule_radius) + + dist = wp.vec2(dist1, dist2) + pos = mat23f(pos1[0], pos1[1], pos1[2], pos2[0], pos2[1], pos2[2]) + + return dist, pos, frame + + +@wp.func +def plane_ellipsoid( + # In: + plane_normal: wp.vec3, + plane_pos: wp.vec3, + ellipsoid_pos: wp.vec3, + ellipsoid_rot: wp.mat33, + ellipsoid_size: wp.vec3, +) -> Tuple[float, wp.vec3, wp.vec3]: + """Core contact geometry calculation for plane-ellipsoid collision. + + Args: + plane_normal: Normal vector of the plane + plane_pos: Position point on the plane + ellipsoid_pos: Center position of the ellipsoid + ellipsoid_rot: Rotation matrix of the ellipsoid + ellipsoid_size: Size (radii) of the ellipsoid along each axis + + Returns: + Tuple containing: + contact_dist: Vector of contact distances + contact_pos: Matrix of contact positions (one per row) + contact_normals: Matrix of contact normal vectors (one per row) + """ + sphere_support = -wp.normalize(wp.cw_mul(wp.transpose(ellipsoid_rot) @ plane_normal, ellipsoid_size)) + pos = ellipsoid_pos + ellipsoid_rot @ wp.cw_mul(sphere_support, ellipsoid_size) + dist = wp.dot(plane_normal, pos - plane_pos) + pos = pos - plane_normal * dist * 0.5 + + return dist, pos, plane_normal + + +@wp.func +def plane_box( + # In: + plane_normal: wp.vec3, + plane_pos: wp.vec3, + box_pos: wp.vec3, + box_rot: wp.mat33, + box_size: wp.vec3, +) -> Tuple[wp.vec4, mat43f, wp.vec3]: + """Core contact geometry calculation for plane-box collision. + + Args: + plane_normal: Normal vector of the plane + plane_pos: Position point on the plane + box_pos: Center position of the box + box_rot: Rotation matrix of the box + box_size: Half-extents of the box along each axis + + Returns: + Tuple containing: + contact_dist: Vector of contact distances (wp.inf for unpopulated contacts) + contact_pos: Matrix of contact positions (one per row) + contact_normal: contact normal vector + """ + + corner = wp.vec3() + center_dist = wp.dot(box_pos - plane_pos, plane_normal) + + dist = wp.vec4(wp.inf) + pos = mat43f() + + # test all corners, pick bottom 4 + ncontact = int(0) + for i in range(8): + # get corner in local coordinates + corner.x = wp.where(i & 1, box_size.x, -box_size.x) + corner.y = wp.where(i & 2, box_size.y, -box_size.y) + corner.z = wp.where(i & 4, box_size.z, -box_size.z) + + # get corner in global coordinates relative to box center + corner = box_rot * corner + + # compute distance to plane, skip if too far or pointing up + ldist = wp.dot(plane_normal, corner) + if center_dist + ldist > 0 or ldist > 0: + continue + + cdist = center_dist + ldist + + dist[ncontact] = cdist + pos[ncontact] = corner + box_pos - 0.5 * plane_normal * cdist + ncontact += 1 + + if ncontact >= 4: + break + + return dist, pos, plane_normal + + +@wp.func +def sphere_cylinder( + # In: + sphere_pos: wp.vec3, + sphere_radius: float, + cylinder_pos: wp.vec3, + cylinder_axis: wp.vec3, + cylinder_radius: float, + cylinder_half_height: float, +) -> Tuple[float, wp.vec3, wp.vec3]: + """Core contact geometry calculation for sphere-cylinder collision. + + Args: + sphere_pos: Center position of the sphere + sphere_radius: Radius of the sphere + cylinder_pos: Center position of the cylinder + cylinder_axis: Axis direction of the cylinder + cylinder_radius: Radius of the cylinder + cylinder_half_height: Half height of the cylinder + + Returns: + Tuple containing: + contact_dist: Vector of contact distances + contact_pos: Matrix of contact positions (one per row) + contact_normals: Matrix of contact normal vectors (one per row) + """ + vec = sphere_pos - cylinder_pos + x = wp.dot(vec, cylinder_axis) + + a_proj = cylinder_axis * x + p_proj = vec - a_proj + p_proj_sqr = wp.dot(p_proj, p_proj) + + collide_side = wp.abs(x) < cylinder_half_height + collide_cap = p_proj_sqr < (cylinder_radius * cylinder_radius) + + if collide_side and collide_cap: + dist_cap = cylinder_half_height - wp.abs(x) + dist_radius = cylinder_radius - wp.sqrt(p_proj_sqr) + + if dist_cap < dist_radius: + collide_side = False + else: + collide_cap = False + + # side collision + if collide_side: + pos_target = cylinder_pos + a_proj + return sphere_sphere(sphere_pos, sphere_radius, pos_target, cylinder_radius) + # cap collision + elif collide_cap: + if x > 0.0: + # top cap + pos_cap = cylinder_pos + cylinder_axis * cylinder_half_height + plane_normal = cylinder_axis + else: + # bottom cap + pos_cap = cylinder_pos - cylinder_axis * cylinder_half_height + plane_normal = -cylinder_axis + + dist, pos = plane_sphere(plane_normal, pos_cap, sphere_pos, sphere_radius) + return dist, pos, -plane_normal # flip normal after position calculation + # corner collision + else: + inv_len = safe_div(1.0, wp.sqrt(p_proj_sqr)) + p_proj = p_proj * (cylinder_radius * inv_len) + + cap_offset = cylinder_axis * (wp.sign(x) * cylinder_half_height) + pos_corner = cylinder_pos + cap_offset + p_proj + + return sphere_sphere(sphere_pos, sphere_radius, pos_corner, 0.0) + + +@wp.func +def plane_cylinder( + # In: + plane_normal: wp.vec3, + plane_pos: wp.vec3, + cylinder_center: wp.vec3, + cylinder_axis: wp.vec3, + cylinder_radius: float, + cylinder_half_height: float, +) -> Tuple[wp.vec4, mat43f, wp.vec3]: + """Core contact geometry calculation for plane-cylinder collision. + + Args: + plane_normal: Normal vector of the plane + plane_pos: Position point on the plane + cylinder_center: Center position of the cylinder + cylinder_axis: Axis direction of the cylinder + cylinder_radius: Radius of the cylinder + cylinder_half_height: Half height of the cylinder + + Returns: + Tuple containing: + contact_dist: Vector of contact distances + contact_pos: Matrix of contact positions (one per row) + contact_normals: Matrix of contact normal vectors (one per row) + """ + + # Initialize output matrices + contact_dist = wp.vec4(wp.inf) + contact_pos = mat43f() + contact_count = 0 + + n = plane_normal + axis = cylinder_axis + + # Project, make sure axis points toward plane + prjaxis = wp.dot(n, axis) + if prjaxis > 0: + axis = -axis + prjaxis = -prjaxis + + # Compute normal distance from plane to cylinder center + dist0 = wp.dot(cylinder_center - plane_pos, n) + + # Remove component of -normal along cylinder axis + vec = axis * prjaxis - n + len_sqr = wp.dot(vec, vec) + + # If vector is nondegenerate, normalize and scale by radius + # Otherwise use cylinder's x-axis scaled by radius + vec = wp.where( + len_sqr >= 1e-12, + vec * safe_div(cylinder_radius, wp.sqrt(len_sqr)), + wp.vec3(1.0, 0.0, 0.0) * cylinder_radius, # Default x-axis when degenerate + ) + + # Project scaled vector on normal + prjvec = wp.dot(vec, n) + + # Scale cylinder axis by half-length + axis = axis * cylinder_half_height + prjaxis = prjaxis * cylinder_half_height + + # First contact point (end cap closer to plane) + dist1 = dist0 + prjaxis + prjvec + pos1 = cylinder_center + vec + axis - n * (dist1 * 0.5) + contact_dist[contact_count] = dist1 + contact_pos[contact_count] = pos1 + contact_count = contact_count + 1 + + # Second contact point (end cap farther from plane) + dist2 = dist0 - prjaxis + prjvec + pos2 = cylinder_center + vec - axis - n * (dist2 * 0.5) + contact_dist[contact_count] = dist2 + contact_pos[contact_count] = pos2 + contact_count = contact_count + 1 + + # Try triangle contact points on side closer to plane + prjvec1 = -prjvec * 0.5 + dist3 = dist0 + prjaxis + prjvec1 + # Compute sideways vector scaled by radius*sqrt(3)/2 + vec1 = wp.cross(vec, axis) + vec1 = wp.normalize(vec1) * (cylinder_radius * wp.sqrt(3.0) * 0.5) + + # Add contact point A - adjust to closest side + pos3 = cylinder_center + vec1 + axis - vec * 0.5 - n * (dist3 * 0.5) + contact_dist[contact_count] = dist3 + contact_pos[contact_count] = pos3 + contact_count = contact_count + 1 + + # Add contact point B - adjust to closest side + pos4 = cylinder_center - vec1 + axis - vec * 0.5 - n * (dist3 * 0.5) + contact_dist[contact_count] = dist3 + contact_pos[contact_count] = pos4 + contact_count = contact_count + 1 + + return contact_dist, contact_pos, n + + +@wp.func +def _compute_rotmore(face_idx: int) -> wp.mat33: + rotmore = wp.mat33(0.0) + + if face_idx == 0: + rotmore[0, 2] = -1.0 + rotmore[1, 1] = +1.0 + rotmore[2, 0] = +1.0 + elif face_idx == 1: + rotmore[0, 0] = +1.0 + rotmore[1, 2] = -1.0 + rotmore[2, 1] = +1.0 + elif face_idx == 2: + rotmore[0, 0] = +1.0 + rotmore[1, 1] = +1.0 + rotmore[2, 2] = +1.0 + elif face_idx == 3: + rotmore[0, 2] = +1.0 + rotmore[1, 1] = +1.0 + rotmore[2, 0] = -1.0 + elif face_idx == 4: + rotmore[0, 0] = +1.0 + rotmore[1, 2] = +1.0 + rotmore[2, 1] = -1.0 + elif face_idx == 5: + rotmore[0, 0] = -1.0 + rotmore[1, 1] = +1.0 + rotmore[2, 2] = -1.0 + + return rotmore + + +@wp.func +def box_box( + # In: + box1_pos: wp.vec3, + box1_rot: wp.mat33, + box1_size: wp.vec3, + box2_pos: wp.vec3, + box2_rot: wp.mat33, + box2_size: wp.vec3, +) -> Tuple[vec8f, mat83f, mat83f]: + """Core contact geometry calculation for box-box collision. + + Args: + box1_pos: Center position of the first box + box1_rot: Rotation matrix of the first box + box1_size: Half-extents of the first box along each axis + box2_pos: Center position of the second box + box2_rot: Rotation matrix of the second box + box2_size: Half-extents of the second box along each axis + + Returns: + Tuple containing: + contact_dist: Vector of contact distances (wp.inf for unpopulated contacts) + contact_pos: Matrix of contact positions (one per row) + contact_normals: Matrix of contact normal vectors (one per row) + """ + + # Initialize output matrices + contact_dist = vec8f() + for i in range(8): + contact_dist[i] = wp.inf + contact_pos = mat83f() + contact_normals = mat83f() + contact_count = 0 + + # Compute transforms between box's frames + pos21 = wp.transpose(box1_rot) @ (box2_pos - box1_pos) + pos12 = wp.transpose(box2_rot) @ (box1_pos - box2_pos) + + rot21 = wp.transpose(box1_rot) @ box2_rot + rot12 = wp.transpose(rot21) + + rot21abs = wp.matrix_from_rows(wp.abs(rot21[0]), wp.abs(rot21[1]), wp.abs(rot21[2])) + rot12abs = wp.transpose(rot21abs) + + plen2 = rot21abs @ box2_size + plen1 = rot12abs @ box1_size + + # Compute axis of maximum separation + s_sum_3 = 3.0 * (box1_size + box2_size) + separation = wp.float32(s_sum_3[0] + s_sum_3[1] + s_sum_3[2]) + axis_code = wp.int32(-1) + + # First test: consider boxes' face normals + for i in range(3): + c1 = -wp.abs(pos21[i]) + box1_size[i] + plen2[i] + + c2 = -wp.abs(pos12[i]) + box2_size[i] + plen1[i] + + if c1 < 0.0 or c2 < 0.0: + return contact_dist, contact_pos, contact_normals + + if c1 < separation: + separation = c1 + axis_code = i + 3 * wp.int32(pos21[i] < 0) + 0 # Face of box1 + if c2 < separation: + separation = c2 + axis_code = i + 3 * wp.int32(pos12[i] < 0) + 6 # Face of box2 + + clnorm = wp.vec3(0.0) + inv = wp.bool(False) + cle1 = wp.int32(0) + cle2 = wp.int32(0) + + # Second test: consider cross products of boxes' edges + for i in range(3): + for j in range(3): + # Compute cross product of box edges (potential separating axis) + if i == 0: + cross_axis = wp.vec3(0.0, -rot12[j, 2], rot12[j, 1]) + elif i == 1: + cross_axis = wp.vec3(rot12[j, 2], 0.0, -rot12[j, 0]) + else: + cross_axis = wp.vec3(-rot12[j, 1], rot12[j, 0], 0.0) + + cross_length = wp.length(cross_axis) + if cross_length < MJ_MINVAL: + continue + + cross_axis /= cross_length + + box_dist = wp.dot(pos21, cross_axis) + c3 = wp.float32(0.0) + + # Project box half-sizes onto the potential separating axis + for k in range(3): + if k != i: + c3 += box1_size[k] * wp.abs(cross_axis[k]) + if k != j: + c3 += box2_size[k] * rot21abs[i, 3 - k - j] / cross_length + + c3 -= wp.abs(box_dist) + + # Early exit: no collision if separated along this axis + if c3 < 0.0: + return contact_dist, contact_pos, contact_normals + + # Track minimum separation and which edge-edge pair it occurs on + if c3 < separation * (1.0 - 1e-12): + separation = c3 + # Determine which corners/edges are closest + cle1 = 0 + cle2 = 0 + + for k in range(3): + if k != i and (int(cross_axis[k] > 0) ^ int(box_dist < 0)): + cle1 += 1 << k + if k != j: + if int(rot21[i, 3 - k - j] > 0) ^ int(box_dist < 0) ^ int((k - j + 3) % 3 == 1): + cle2 += 1 << k + + axis_code = 12 + i * 3 + j + clnorm = cross_axis + inv = box_dist < 0 + + # No axis with separation < margin found + if axis_code == -1: + return contact_dist, contact_pos, contact_normals + + points = mat83f() + depth = vec8f() + max_con_pair = 8 + # 8 contacts should suffice for most configurations + + if axis_code < 12: + # Handle face-vertex collision + face_idx = axis_code % 6 + box_idx = axis_code // 6 + rotmore = _compute_rotmore(face_idx) + + r = rotmore @ wp.where(box_idx, rot12, rot21) + p = rotmore @ wp.where(box_idx, pos12, pos21) + ss = wp.abs(rotmore @ wp.where(box_idx, box2_size, box1_size)) + s = wp.where(box_idx, box1_size, box2_size) + rt = wp.transpose(r) + + lx, ly, hz = ss[0], ss[1], ss[2] + p[2] -= hz + + clcorner = wp.int32(0) # corner of non-face box with least axis separation + + for i in range(3): + if r[2, i] < 0: + clcorner += 1 << i + + lp = p + for i in range(wp.static(3)): + lp += rt[i] * s[i] * wp.where(clcorner & 1 << i, 1.0, -1.0) + + m = wp.int32(1) + dirs = wp.int32(0) + + cn1 = wp.vec3(0.0) + cn2 = wp.vec3(0.0) + + for i in range(3): + if wp.abs(r[2, i]) < 0.5: + if not dirs: + cn1 = rt[i] * s[i] * wp.where(clcorner & (1 << i), -2.0, 2.0) + else: + cn2 = rt[i] * s[i] * wp.where(clcorner & (1 << i), -2.0, 2.0) + + dirs += 1 + + k = dirs * dirs + + # Find potential contact points + + n = wp.int32(0) + + for i in range(k): + for q in range(2): + # lines_a and lines_b (lines between corners) computed on the fly + lav = lp + wp.where(i < 2, wp.vec3(0.0), wp.where(i == 2, cn1, cn2)) + lbv = wp.where(i == 0 or i == 3, cn1, cn2) + + if wp.abs(lbv[q]) > MJ_MINVAL: + br = 1.0 / lbv[q] + for j in range(-1, 2, 2): + l = ss[q] * wp.float32(j) + c1 = (l - lav[q]) * br + if c1 < 0 or c1 > 1: + continue + c2 = lav[1 - q] + lbv[1 - q] * c1 + if wp.abs(c2) > ss[1 - q]: + continue + + points[n] = lav + c1 * lbv + n += 1 + + if dirs == 2: + ax = cn1[0] + bx = cn2[0] + ay = cn1[1] + by = cn2[1] + C = safe_div(1.0, ax * by - bx * ay) + + for i in range(4): + llx = wp.where(i // 2, lx, -lx) + lly = wp.where(i % 2, ly, -ly) + + x = llx - lp[0] + y = lly - lp[1] + + u = (x * by - y * bx) * C + v = (y * ax - x * ay) * C + + if u > 0 and v > 0 and u < 1 and v < 1: + points[n] = wp.vec3(llx, lly, lp[2] + u * cn1[2] + v * cn2[2]) + n += 1 + + for i in range(1 << dirs): + tmpv = lp + wp.float32(i & 1) * cn1 + wp.float32((i & 2) != 0) * cn2 + if tmpv[0] > -lx and tmpv[0] < lx and tmpv[1] > -ly and tmpv[1] < ly: + points[n] = tmpv + n += 1 + + m = n + n = wp.int32(0) + + for i in range(m): + if points[i][2] > 0.0: + continue + if i != n: + points[n] = points[i] + + points[n, 2] *= 0.5 + depth[n] = points[n, 2] + n += 1 + + # Set up contact frame + rw = wp.where(box_idx, box2_rot, box1_rot) @ wp.transpose(rotmore) + pw = wp.where(box_idx, box2_pos, box1_pos) + normal = wp.where(box_idx, -1.0, 1.0) * wp.transpose(rw)[2] + + else: + # Handle edge-edge collision + edge1 = (axis_code - 12) // 3 + edge2 = (axis_code - 12) % 3 + + # Set up non-contacting edges ax1, ax2 for box2 and pax1, pax2 for box 1 + ax1 = wp.int(1 - (edge2 & 1)) + ax2 = wp.int(2 - (edge2 & 2)) + + pax1 = wp.int(1 - (edge1 & 1)) + pax2 = wp.int(2 - (edge1 & 2)) + + if rot21abs[edge1, ax1] < rot21abs[edge1, ax2]: + ax1, ax2 = ax2, ax1 + + if rot12abs[edge2, pax1] < rot12abs[edge2, pax2]: + pax1, pax2 = pax2, pax1 + + rotmore = _compute_rotmore(wp.where(cle1 & (1 << pax2), pax2, pax2 + 3)) + + # Transform coordinates for edge-edge contact calculation + p = rotmore @ pos21 + rnorm = rotmore @ clnorm + r = rotmore @ rot21 + rt = wp.transpose(r) + s = wp.abs(wp.transpose(rotmore) @ box1_size) + + lx, ly, hz = s[0], s[1], s[2] + p[2] -= hz + + # Calculate closest box2 face + + points[0] = ( + p + + rt[ax1] * box2_size[ax1] * wp.where(cle2 & (1 << ax1), 1.0, -1.0) + + rt[ax2] * box2_size[ax2] * wp.where(cle2 & (1 << ax2), 1.0, -1.0) + ) + points[1] = points[0] - rt[edge2] * box2_size[edge2] + points[0] += rt[edge2] * box2_size[edge2] + + points[2] = ( + p + + rt[ax1] * box2_size[ax1] * wp.where(cle2 & (1 << ax1), -1.0, 1.0) + + rt[ax2] * box2_size[ax2] * wp.where(cle2 & (1 << ax2), 1.0, -1.0) + ) + + points[3] = points[2] - rt[edge2] * box2_size[edge2] + points[2] += rt[edge2] * box2_size[edge2] + + n = 4 + + # Set up coordinate axes for contact face of box2 + axi_lp = points[0] + axi_cn1 = points[1] - points[0] + axi_cn2 = points[2] - points[0] + + # Check if contact normal is valid + if wp.abs(rnorm[2]) < MJ_MINVAL: + return contact_dist, contact_pos, contact_normals # Shouldn't happen + + # Calculate inverse normal for projection + innorm = wp.where(inv, -1.0, 1.0) / rnorm[2] + + pu = mat43f() + + # Project points onto contact plane + for i in range(4): + pu[i] = points[i] + c_scl = points[i, 2] * wp.where(inv, -1.0, 1.0) * innorm + points[i] -= rnorm * c_scl + + pts_lp = points[0] + pts_cn1 = points[1] - points[0] + pts_cn2 = points[2] - points[0] + + n = wp.int32(0) + + for i in range(4): + for q in range(2): + la = pts_lp[q] + wp.where(i < 2, 0.0, wp.where(i == 2, pts_cn1[q], pts_cn2[q])) + lb = wp.where(i == 0 or i == 3, pts_cn1[q], pts_cn2[q]) + lc = pts_lp[1 - q] + wp.where(i < 2, 0.0, wp.where(i == 2, pts_cn1[1 - q], pts_cn2[1 - q])) + ld = wp.where(i == 0 or i == 3, pts_cn1[1 - q], pts_cn2[1 - q]) + + # linesu_a and linesu_b (lines between corners) computed on the fly + lua = axi_lp + wp.where(i < 2, wp.vec3(0.0), wp.where(i == 2, axi_cn1, axi_cn2)) + lub = wp.where(i == 0 or i == 3, axi_cn1, axi_cn2) + + if wp.abs(lb) > MJ_MINVAL: + br = 1.0 / lb + for j in range(-1, 2, 2): + if n == max_con_pair: + break + l = s[q] * wp.float32(j) + c1 = (l - la) * br + if c1 < 0 or c1 > 1: + continue + c2 = lc + ld * c1 + if wp.abs(c2) > s[1 - q]: + continue + if (lua[2] + lub[2] * c1) * innorm > 0.0: + continue + + points[n] = lua * 0.5 + c1 * lub * 0.5 + points[n, q] += 0.5 * l + points[n, 1 - q] += 0.5 * c2 + depth[n] = points[n, 2] * innorm * 2.0 + n += 1 + + nl = n + + ax = pts_cn1[0] + bx = pts_cn2[0] + ay = pts_cn1[1] + by = pts_cn2[1] + C = safe_div(1.0, ax * by - bx * ay) + + for i in range(4): + if n == max_con_pair: + break + llx = wp.where(i // 2, lx, -lx) + lly = wp.where(i % 2, ly, -ly) + + x = llx - pts_lp[0] + y = lly - pts_lp[1] + + u = (x * by - y * bx) * C + v = (y * ax - x * ay) * C + + if nl == 0: + if (u < 0 or u > 1) and (v < 0 or v > 1): + continue + elif u < 0 or v < 0 or u > 1 or v > 1: + continue + + u = wp.clamp(u, 0.0, 1.0) + v = wp.clamp(v, 0.0, 1.0) + w = 1.0 - u - v + vtmp = pu[0] * w + pu[1] * u + pu[2] * v + + points[n] = wp.vec3(llx, lly, 0.0) + + vtmp2 = points[n] - vtmp + tc1 = wp.length_sq(vtmp2) + if vtmp[2] > 0 and tc1 > 0.0: + continue + + points[n] = 0.5 * (points[n] + vtmp) + + depth[n] = wp.sqrt(tc1) * wp.where(vtmp[2] < 0, -1.0, 1.0) + n += 1 + + nf = n + + for i in range(4): + if n >= max_con_pair: + break + x = pu[i, 0] + y = pu[i, 1] + if nl == 0 and nf != 0: + if (x < -lx or x > lx) and (y < -ly or y > ly): + continue + elif x < -lx or x > lx or y < -ly or y > ly: + continue + + c1 = wp.float32(0) + + for j in range(2): + if pu[i, j] < -s[j]: + c1 += (pu[i, j] + s[j]) * (pu[i, j] + s[j]) + elif pu[i, j] > s[j]: + c1 += (pu[i, j] - s[j]) * (pu[i, j] - s[j]) + + c1 += pu[i, 2] * innorm * pu[i, 2] * innorm + + if pu[i, 2] > 0 and c1 > 0.0: + continue + + tmp_p = wp.vec3(pu[i, 0], pu[i, 1], 0.0) + + for j in range(2): + if pu[i, j] < -s[j]: + tmp_p[j] = -s[j] * 0.5 + elif pu[i, j] > s[j]: + tmp_p[j] = +s[j] * 0.5 + + tmp_p += pu[i] + points[n] = tmp_p * 0.5 + + depth[n] = wp.sqrt(c1) * wp.where(pu[i, 2] < 0, -1.0, 1.0) + n += 1 + + # Set up contact data for all points + rw = box1_rot @ wp.transpose(rotmore) + pw = box1_pos + normal = wp.where(inv, -1.0, 1.0) * rw @ rnorm + + contact_count = n + + # Copy contact data to output matrices + for i in range(contact_count): + points[i, 2] += hz + pos = rw @ points[i] + pw + contact_dist[i] = depth[i] + contact_pos[i] = pos + contact_normals[i] = normal + + return contact_dist, contact_pos, contact_normals + + +@wp.func +def sphere_box( + # In: + sphere_pos: wp.vec3, + sphere_radius: float, + box_pos: wp.vec3, + box_rot: wp.mat33, + box_size: wp.vec3, +) -> Tuple[float, wp.vec3, wp.vec3]: + """Core contact geometry calculation for sphere-box collision. + + Args: + sphere_pos: Center position of the sphere + sphere_radius: Radius of the sphere + box_pos: Center position of the box + box_rot: Rotation matrix of the box + box_size: Half-extents of the box along each axis + + Returns: + Tuple containing: + contact_dist: Vector of contact distances + contact_pos: contact positions + contact_normal: contact normal vectors + """ + + center = wp.transpose(box_rot) @ (sphere_pos - box_pos) + + clamped = wp.max(-box_size, wp.min(box_size, center)) + clamped_dir, dist = normalize_with_norm(clamped - center) + + # sphere center inside box + if dist <= MJ_MINVAL: + closest = 2.0 * (box_size[0] + box_size[1] + box_size[2]) + k = wp.int32(0) + for i in range(6): + face_dist = wp.abs(wp.where(i % 2, 1.0, -1.0) * box_size[i // 2] - center[i // 2]) + if closest > face_dist: + closest = face_dist + k = i + + nearest = wp.vec3(0.0) + nearest[k // 2] = wp.where(k % 2, -1.0, 1.0) + pos = center + nearest * (sphere_radius - closest) / 2.0 + contact_normal = box_rot @ nearest + contact_distance = -closest - sphere_radius + + else: + deepest = center + clamped_dir * sphere_radius + pos = 0.5 * (clamped + deepest) + contact_normal = box_rot @ clamped_dir + contact_distance = dist - sphere_radius + + contact_position = box_pos + box_rot @ pos + + return contact_distance, contact_position, contact_normal + + +@wp.func +def capsule_box( + # In: + capsule_pos: wp.vec3, + capsule_axis: wp.vec3, + capsule_radius: float, + capsule_half_length: float, + box_pos: wp.vec3, + box_rot: wp.mat33, + box_size: wp.vec3, +) -> Tuple[wp.vec2, mat23f, mat23f]: + """Core contact geometry calculation for capsule-box collision. + + Args: + capsule_pos: Center position of the capsule + capsule_axis: Axis direction of the capsule + capsule_radius: Radius of the capsule + capsule_half_length: Half length of the capsule + box_pos: Center position of the box + box_rot: Rotation matrix of the box + box_size: Half-extents of the box along each axis + + Returns: + Tuple containing: + contact_dist: Vector of contact distances (wp.inf for unpopulated contacts) + contact_pos: Matrix of contact positions (one per row) + contact_normals: Matrix of contact normal vectors (one per row) + """ + + # Based on the mjc implementation + boxmatT = wp.transpose(box_rot) + pos = boxmatT @ (capsule_pos - box_pos) + axis = boxmatT @ capsule_axis + halfaxis = axis * capsule_half_length # halfaxis is the capsule direction + axisdir = wp.int32(halfaxis[0] > 0.0) + 2 * wp.int32(halfaxis[1] > 0.0) + 4 * wp.int32(halfaxis[2] > 0.0) + + bestdistmax = 2.0 * (capsule_radius + capsule_half_length + box_size[0] + box_size[1] + box_size[2]) + + # keep track of closest point + bestdist = wp.float32(bestdistmax) + bestsegmentpos = wp.float32(-12) + + # cltype: encoded collision configuration + # cltype / 3 == 0 : lower corner is closest to the capsule + # == 2 : upper corner is closest to the capsule + # == 1 : middle of the edge is closest to the capsule + # cltype % 3 == 0 : lower corner is closest to the box + # == 2 : upper corner is closest to the box + # == 1 : middle of the capsule is closest to the box + cltype = wp.int32(-4) + + # clface: index of the closest face of the box to the capsule + # -1: no face is closest (edge or corner is closest) + # 0, 1, 2: index of the axis perpendicular to the closest face + clface = wp.int32(-12) + + # first: consider cases where a face of the box is closest + for i in range(-1, 2, 2): + axisTip = pos + wp.float32(i) * halfaxis + boxPoint = wp.vec3(axisTip) + + n_out = wp.int32(0) + ax_out = wp.int32(-1) + + for j in range(3): + if boxPoint[j] < -box_size[j]: + n_out += 1 + ax_out = j + boxPoint[j] = -box_size[j] + elif boxPoint[j] > box_size[j]: + n_out += 1 + ax_out = j + boxPoint[j] = box_size[j] + + if n_out > 1: + continue + + dist = wp.length_sq(boxPoint - axisTip) + + if dist < bestdist: + bestdist = dist + bestsegmentpos = wp.float32(i) + cltype = -2 + i + clface = ax_out + + # second: consider cases where an edge of the box is closest + clcorner = wp.int32(-123) # which corner is the closest + cledge = wp.int32(-123) # which axis + bestboxpos = wp.float32(0.0) + + for i in range(8): + for j in range(3): + if i & (1 << j) != 0: + continue + + c2 = wp.int32(-123) + + # box_pt is the starting point (corner) on the box + box_pt = wp.cw_mul( + wp.vec3( + wp.where(i & 1, 1.0, -1.0), + wp.where(i & 2, 1.0, -1.0), + wp.where(i & 4, 1.0, -1.0), + ), + box_size, + ) + box_pt[j] = 0.0 + + # find closest point between capsule and the edge + dif = box_pt - pos + + u = -box_size[j] * dif[j] + v = wp.dot(halfaxis, dif) + ma = box_size[j] * box_size[j] + mb = -box_size[j] * halfaxis[j] + mc = capsule_half_length * capsule_half_length + det = ma * mc - mb * mb + if wp.abs(det) < MJ_MINVAL: + continue + + idet = 1.0 / det + # sX : X=1 means middle of segment. X=0 or 2 one or the other end + + x1 = wp.float32((mc * u - mb * v) * idet) + x2 = wp.float32((ma * v - mb * u) * idet) + + s1 = wp.int32(1) + s2 = wp.int32(1) + + if x1 > 1: + x1 = 1.0 + s1 = 2 + x2 = safe_div(v - mb, mc) + elif x1 < -1: + x1 = -1.0 + s1 = 0 + x2 = safe_div(v + mb, mc) + + x2_over = x2 > 1.0 + if x2_over or x2 < -1.0: + if x2_over: + x2 = 1.0 + s2 = 2 + x1 = safe_div(u - mb, ma) + else: + x2 = -1.0 + s2 = 0 + x1 = safe_div(u + mb, ma) + + if x1 > 1: + x1 = 1.0 + s1 = 2 + elif x1 < -1: + x1 = -1.0 + s1 = 0 + + dif -= halfaxis * x2 + dif[j] += box_size[j] * x1 + + # encode relative positions of the closest points + ct = s1 * 3 + s2 + + dif_sq = wp.length_sq(dif) + if dif_sq < bestdist - MJ_MINVAL: + bestdist = dif_sq + bestsegmentpos = x2 + bestboxpos = x1 + # ct<6 means closest point on box is at lower end or middle of edge + c2 = ct // 6 + + clcorner = i + (1 << j) * c2 # index of closest box corner + cledge = j # axis index of closest box edge + cltype = ct # encoded collision configuration + + best = wp.float32(0.0) + + p = wp.vec2(pos.x, pos.y) + dd = wp.vec2(halfaxis.x, halfaxis.y) + s = wp.vec2(box_size.x, box_size.y) + secondpos = wp.float32(-4.0) + + uu = dd.x * s.y + vv = dd.y * s.x + w_neg = dd.x * p.y - dd.y * p.x < 0 + + best = wp.float32(-1.0) + + ee1 = uu - vv + ee2 = uu + vv + + if wp.abs(ee1) > best: + best = wp.abs(ee1) + c1 = wp.where((ee1 < 0) == w_neg, 0, 3) + + if wp.abs(ee2) > best: + best = wp.abs(ee2) + c1 = wp.where((ee2 > 0) == w_neg, 1, 2) + + if cltype == -4: # invalid type + return wp.vec2(wp.inf), mat23f(), mat23f() + + if cltype >= 0 and cltype // 3 != 1: # closest to a corner of the box + c1 = axisdir ^ clcorner + # Calculate relative orientation between capsule and corner + # There are two possible configurations: + # 1. Capsule axis points toward/away from corner + # 2. Capsule axis aligns with a face or edge + if c1 != 0 and c1 != 7: # create second contact point + if c1 == 1 or c1 == 2 or c1 == 4: + mul = 1 + else: + mul = -1 + c1 = 7 - c1 + + # "de" and "dp" distance from first closest point on the capsule to both ends of it + # mul is a direction along the capsule's axis + + if c1 == 1: + ax = 0 + ax1 = 1 + ax2 = 2 + elif c1 == 2: + ax = 1 + ax1 = 2 + ax2 = 0 + elif c1 == 4: + ax = 2 + ax1 = 0 + ax2 = 1 + + if axis[ax] * axis[ax] > 0.5: # second point along the edge of the box + m = 2.0 * safe_div(box_size[ax], wp.abs(halfaxis[ax])) + secondpos = min(1.0 - wp.float32(mul) * bestsegmentpos, m) + else: # second point along a face of the box + # check for overshoot again + m = 2.0 * min( + safe_div(box_size[ax1], wp.abs(halfaxis[ax1])), + safe_div(box_size[ax2], wp.abs(halfaxis[ax2])), + ) + secondpos = -min(1.0 + wp.float32(mul) * bestsegmentpos, m) + secondpos *= wp.float32(mul) + + elif cltype >= 0 and cltype // 3 == 1: # we are on box's edge + # Calculate relative orientation between capsule and edge + # Two possible configurations: + # - T configuration: c1 = 2^n (no additional contacts) + # - X configuration: c1 != 2^n (potential additional contacts) + c1 = axisdir ^ clcorner + c1 &= 7 - (1 << cledge) # mask out edge axis to determine configuration + + if c1 == 1 or c1 == 2 or c1 == 4: # create second contact point + if cledge == 0: + ax1 = 1 + ax2 = 2 + if cledge == 1: + ax1 = 2 + ax2 = 0 + if cledge == 2: + ax1 = 0 + ax2 = 1 + ax = cledge + + # find which face the capsule has a lower angle, and switch the axis + if wp.abs(axis[ax1]) > wp.abs(axis[ax2]): + ax1 = ax2 + ax2 = 3 - ax - ax1 + + # mul determines direction along capsule axis for second contact point + if c1 & (1 << ax2): + mul = 1 + secondpos = 1.0 - bestsegmentpos + else: + mul = -1 + secondpos = 1.0 + bestsegmentpos + + # now find out whether we point towards the opposite side or towards one of the sides + # and also find the farthest point along the capsule that is above the box + + e1 = 2.0 * safe_div(box_size[ax2], wp.abs(halfaxis[ax2])) + secondpos = min(e1, secondpos) + + if ((axisdir & (1 << ax)) != 0) == ((c1 & (1 << ax2)) != 0): + e2 = 1.0 - bestboxpos + else: + e2 = 1.0 + bestboxpos + + e1 = box_size[ax] * safe_div(e2, wp.abs(halfaxis[ax])) + + secondpos = min(e1, secondpos) + secondpos *= wp.float32(mul) + + elif cltype < 0: + # similarly we handle the case when one capsule's end is closest to a face of the box + # and find where is the other end pointing to and clamping to the farthest point + # of the capsule that's above the box + # if the closest point is inside the box there's no need for a second point + + if clface != -1: # create second contact point + mul = wp.where(cltype == -3, 1, -1) + secondpos = 2.0 + + tmp1 = pos - halfaxis * wp.float32(mul) + + for i in range(3): + if i != clface: + ha_r = safe_div(wp.float32(mul), halfaxis[i]) + e1 = (box_size[i] - tmp1[i]) * ha_r + if 0 < e1 and e1 < secondpos: + secondpos = e1 + + e1 = (-box_size[i] - tmp1[i]) * ha_r + if 0 < e1 and e1 < secondpos: + secondpos = e1 + + secondpos *= wp.float32(mul) + + # create sphere in original orientation at first contact point + s1_pos_l = pos + halfaxis * bestsegmentpos + s1_pos_g = box_rot @ s1_pos_l + box_pos + + # collide with sphere using core function + dist1, pos1, normal1 = sphere_box(s1_pos_g, capsule_radius, box_pos, box_rot, box_size) + + if secondpos > -3: # secondpos was modified + s2_pos_l = pos + halfaxis * (secondpos + bestsegmentpos) + s2_pos_g = box_rot @ s2_pos_l + box_pos + + # collide with sphere using core function + dist2, pos2, normal2 = sphere_box(s2_pos_g, capsule_radius, box_pos, box_rot, box_size) + else: + dist2 = wp.inf + pos2 = wp.vec3() + normal2 = wp.vec3() + + return ( + wp.vec2(dist1, dist2), + mat23f(pos1[0], pos1[1], pos1[2], pos2[0], pos2[1], pos2[2]), + mat23f(normal1[0], normal1[1], normal1[2], normal2[0], normal2[1], normal2[2]), + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py index d00ccdb1..c08e838d 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py @@ -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, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py index 6f99949f..201eab65 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py @@ -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 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py index f4e1efc1..6222d438 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py @@ -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] diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py index b021e612..1a71a2e1 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py @@ -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), diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py index 3856bfc1..8f98bfb3 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py @@ -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( diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py index a6dafbf5..cb264208 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -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, + ], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_jax_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_jax_test.py new file mode 100644 index 00000000..5390761d --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_jax_test.py @@ -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() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py index 1560550e..1d0eebcd 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py @@ -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(""" @@ -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( - '', - '', - '', - '', - '', - ) - def test_contact_sensor(self, contact_sensor): - mjm = mujoco.MjModel.from_xml_string(f""" - - - - - - - - - - - {contact_sensor} - - - """) - - 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__": diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py index d2864033..244d9ea3 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -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( diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py index 787c4cd6..98d7962d 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py @@ -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, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py index 15e8e655..5b13b591 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -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), diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py index 96dcdda0..9c5f5181 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor_test.py @@ -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_sensor += f'' _MJCF = f""" """ @@ -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() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py index 56b02e0f..3bf30602 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py @@ -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) diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index 0b90c028..1f89ad46 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -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 diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index d80eae30..58191252 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -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 _forward_shim( # Model @@ -157,7 +158,6 @@ def _forward_shim( 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), @@ -220,7 +220,6 @@ def _forward_shim( nflexvert: int, ngeom: int, ngravcomp: int, - nhfield: int, njnt: int, nlight: int, nlsp: int, @@ -234,6 +233,9 @@ def _forward_shim( 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), @@ -317,7 +319,9 @@ def _forward_shim( wrap_type: wp.array(dtype=int), opt__broadphase: int, opt__broadphase_filter: int, + opt__ccd_tolerance: wp.array(dtype=float), opt__cone: int, + opt__contact_sensor_maxmatch: int, opt__density: wp.array(dtype=float), opt__disableflags: int, opt__enableflags: int, @@ -362,7 +366,6 @@ def _forward_shim( cfrc_ext: wp.array2d(dtype=wp.spatial_vector), cfrc_int: wp.array2d(dtype=wp.spatial_vector), cinert: wp.array2d(dtype=mjwp_types.vec10), - 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), @@ -393,9 +396,19 @@ def _forward_shim( light_xpos: wp.array2d(dtype=wp.vec3), mocap_pos: wp.array2d(dtype=wp.vec3), mocap_quat: wp.array2d(dtype=wp.quat), + 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), ne: wp.array(dtype=int), ne_connect: wp.array(dtype=int), ne_jnt: wp.array(dtype=int), @@ -627,7 +640,6 @@ def _forward_shim( _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 @@ -690,7 +702,6 @@ def _forward_shim( _m.nflexvert = nflexvert _m.ngeom = ngeom _m.ngravcomp = ngravcomp - _m.nhfield = nhfield _m.njnt = njnt _m.nlight = nlight _m.nlsp = nlsp @@ -704,9 +715,14 @@ def _forward_shim( _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.cone = opt__cone + _m.opt.contact_sensor_maxmatch = opt__contact_sensor_maxmatch _m.opt.density = opt__density _m.opt.disableflags = opt__disableflags _m.opt.enableflags = opt__enableflags @@ -829,7 +845,6 @@ def _forward_shim( _d.cfrc_ext = cfrc_ext _d.cfrc_int = cfrc_int _d.cinert = cinert - _d.collision_hftri_index = collision_hftri_index _d.collision_pair = collision_pair _d.collision_pairid = collision_pairid _d.collision_worldid = collision_worldid @@ -906,9 +921,19 @@ def _forward_shim( _d.light_xpos = light_xpos _d.mocap_pos = mocap_pos _d.mocap_quat = mocap_quat + _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.ne = ne _d.ne_connect = ne_connect @@ -1001,7 +1026,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'cfrc_ext': d._impl.cfrc_ext.shape, 'cfrc_int': d._impl.cfrc_int.shape, 'cinert': d._impl.cinert.shape, - '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, @@ -1032,9 +1056,19 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'light_xpos': d._impl.light_xpos.shape, 'mocap_pos': d.mocap_pos.shape, 'mocap_quat': d.mocap_quat.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, 'ne': d._impl.ne.shape, 'ne_connect': d._impl.ne_connect.shape, 'ne_jnt': d._impl.ne_jnt.shape, @@ -1153,7 +1187,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _forward_shim, - num_outputs=164, + num_outputs=173, output_dims=output_dims, vmap_method=None, in_out_argnames={ @@ -1172,7 +1206,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'cfrc_ext', 'cfrc_int', 'cinert', - 'collision_hftri_index', 'collision_pair', 'collision_pairid', 'collision_worldid', @@ -1203,9 +1236,19 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'light_xpos', 'mocap_pos', 'mocap_quat', + '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', 'ne', 'ne_connect', 'ne_jnt', @@ -1436,7 +1479,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.geom_solmix, m.geom_solref, m.geom_type, - m._impl.geompair2hfgeompair, m._impl.has_sdf_geom, m.hfield_adr, m.hfield_data, @@ -1499,7 +1541,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.nflexvert, m.ngeom, m.ngravcomp, - m.nhfield, m.njnt, m.nlight, m._impl.nlsp, @@ -1513,6 +1554,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 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, @@ -1596,7 +1640,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.wrap_type, m.opt._impl.broadphase, m.opt._impl.broadphase_filter, + m.opt._impl.ccd_tolerance, m.opt.cone, + m.opt._impl.contact_sensor_maxmatch, m.opt.density, m.opt.disableflags, m.opt.enableflags, @@ -1640,7 +1686,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.cfrc_ext, d._impl.cfrc_int, d._impl.cinert, - d._impl.collision_hftri_index, d._impl.collision_pair, d._impl.collision_pairid, d._impl.collision_worldid, @@ -1671,9 +1716,19 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.light_xpos, d.mocap_pos, d.mocap_quat, + 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.ne, d._impl.ne_connect, d._impl.ne_jnt, @@ -1806,155 +1861,164 @@ def _forward_jax_impl(m: types.Model, d: types.Data): '_impl.cfrc_ext': out[12], '_impl.cfrc_int': out[13], '_impl.cinert': out[14], - '_impl.collision_hftri_index': out[15], - '_impl.collision_pair': out[16], - '_impl.collision_pairid': out[17], - '_impl.collision_worldid': out[18], - '_impl.crb': out[19], - 'ctrl': out[20], - 'cvel': out[21], - '_impl.energy': out[22], - '_impl.epa_face': out[23], - '_impl.epa_horizon': out[24], - '_impl.epa_index': out[25], - '_impl.epa_map': out[26], - '_impl.epa_norm2': out[27], - '_impl.epa_pr': out[28], - '_impl.epa_vert': out[29], - '_impl.epa_vert1': out[30], - '_impl.epa_vert2': out[31], - '_impl.epa_vert_index1': out[32], - '_impl.epa_vert_index2': out[33], - 'eq_active': out[34], - '_impl.flexedge_length': out[35], - '_impl.flexedge_velocity': out[36], - '_impl.flexvert_xpos': out[37], - '_impl.fluid_applied': out[38], - '_impl.geom_skip': out[39], - 'geom_xmat': out[40], - 'geom_xpos': out[41], - '_impl.light_xdir': out[42], - '_impl.light_xpos': out[43], - 'mocap_pos': out[44], - 'mocap_quat': out[45], - '_impl.ncollision': out[46], - '_impl.ncon': out[47], - '_impl.ncon_hfield': out[48], - '_impl.ne': out[49], - '_impl.ne_connect': out[50], - '_impl.ne_jnt': out[51], - '_impl.ne_ten': out[52], - '_impl.ne_weld': out[53], - '_impl.nefc': out[54], - '_impl.nf': out[55], - '_impl.nl': out[56], - '_impl.nsolving': out[57], - '_impl.qLD': out[58], - '_impl.qLDiagInv': out[59], - '_impl.qM': out[60], - 'qacc': out[61], - 'qacc_smooth': out[62], - 'qacc_warmstart': out[63], - 'qfrc_actuator': out[64], - 'qfrc_applied': out[65], - 'qfrc_bias': out[66], - 'qfrc_constraint': out[67], - '_impl.qfrc_damper': out[68], - 'qfrc_fluid': out[69], - 'qfrc_gravcomp': out[70], - 'qfrc_passive': out[71], - 'qfrc_smooth': out[72], - '_impl.qfrc_spring': out[73], - 'qpos': out[74], - 'qvel': out[75], - '_impl.sap_cumulative_sum': out[76], - '_impl.sap_projection_lower': out[77], - '_impl.sap_projection_upper': out[78], - '_impl.sap_range': out[79], - '_impl.sap_segment_index': out[80], - '_impl.sap_sort_index': out[81], - '_impl.sensor_contact_criteria': out[82], - '_impl.sensor_contact_direction': out[83], - '_impl.sensor_contact_matchid': out[84], - '_impl.sensor_contact_nmatch': out[85], - '_impl.sensor_rangefinder_dist': out[86], - '_impl.sensor_rangefinder_geomid': out[87], - '_impl.sensor_rangefinder_pnt': out[88], - '_impl.sensor_rangefinder_vec': out[89], - 'sensordata': out[90], - 'site_xmat': out[91], - 'site_xpos': out[92], - '_impl.solver_niter': out[93], - '_impl.subtree_angmom': out[94], - '_impl.subtree_bodyvel': out[95], - 'subtree_com': out[96], - '_impl.subtree_linvel': out[97], - '_impl.ten_J': out[98], - '_impl.ten_Jdot': out[99], - '_impl.ten_actfrc': out[100], - '_impl.ten_bias_coef': out[101], - 'ten_length': out[102], - '_impl.ten_velocity': out[103], - '_impl.ten_wrapadr': out[104], - '_impl.ten_wrapnum': out[105], - 'time': out[106], - '_impl.wrap_geom_xpos': out[107], - '_impl.wrap_obj': out[108], - '_impl.wrap_xpos': out[109], - 'xanchor': out[110], - 'xaxis': out[111], - 'xfrc_applied': out[112], - 'ximat': out[113], - 'xipos': out[114], - 'xmat': out[115], - 'xpos': out[116], - 'xquat': out[117], - '_impl.contact__dim': out[118], - '_impl.contact__dist': out[119], - '_impl.contact__efc_address': out[120], - '_impl.contact__frame': out[121], - '_impl.contact__friction': out[122], - '_impl.contact__geom': out[123], - '_impl.contact__includemargin': out[124], - '_impl.contact__pos': out[125], - '_impl.contact__solimp': out[126], - '_impl.contact__solref': out[127], - '_impl.contact__solreffriction': out[128], - '_impl.contact__worldid': out[129], - '_impl.efc__D': out[130], - '_impl.efc__J': out[131], - '_impl.efc__Jaref': out[132], - '_impl.efc__Ma': out[133], - '_impl.efc__Mgrad': out[134], - '_impl.efc__alpha': out[135], - '_impl.efc__aref': out[136], - '_impl.efc__beta': out[137], - '_impl.efc__cholesky_L_tmp': out[138], - '_impl.efc__cholesky_y_tmp': out[139], - '_impl.efc__cost': out[140], - '_impl.efc__cost_candidate': out[141], - '_impl.efc__done': out[142], - '_impl.efc__force': out[143], - '_impl.efc__frictionloss': out[144], - '_impl.efc__gauss': out[145], - '_impl.efc__grad': out[146], - '_impl.efc__grad_dot': out[147], - '_impl.efc__h': out[148], - '_impl.efc__id': out[149], - '_impl.efc__jv': out[150], - '_impl.efc__margin': out[151], - '_impl.efc__mv': out[152], - '_impl.efc__pos': out[153], - '_impl.efc__prev_Mgrad': out[154], - '_impl.efc__prev_cost': out[155], - '_impl.efc__prev_grad': out[156], - '_impl.efc__quad': out[157], - '_impl.efc__quad_gauss': out[158], - '_impl.efc__search': out[159], - '_impl.efc__search_dot': out[160], - '_impl.efc__state': out[161], - '_impl.efc__type': out[162], - '_impl.efc__vel': out[163], + '_impl.collision_pair': out[15], + '_impl.collision_pairid': out[16], + '_impl.collision_worldid': out[17], + '_impl.crb': out[18], + 'ctrl': out[19], + 'cvel': out[20], + '_impl.energy': out[21], + '_impl.epa_face': out[22], + '_impl.epa_horizon': out[23], + '_impl.epa_index': out[24], + '_impl.epa_map': out[25], + '_impl.epa_norm2': out[26], + '_impl.epa_pr': out[27], + '_impl.epa_vert': out[28], + '_impl.epa_vert1': out[29], + '_impl.epa_vert2': out[30], + '_impl.epa_vert_index1': out[31], + '_impl.epa_vert_index2': out[32], + 'eq_active': out[33], + '_impl.flexedge_length': out[34], + '_impl.flexedge_velocity': out[35], + '_impl.flexvert_xpos': out[36], + '_impl.fluid_applied': out[37], + '_impl.geom_skip': out[38], + 'geom_xmat': out[39], + 'geom_xpos': out[40], + '_impl.light_xdir': out[41], + '_impl.light_xpos': out[42], + 'mocap_pos': out[43], + 'mocap_quat': out[44], + '_impl.multiccd_clipped': out[45], + '_impl.multiccd_endvert': out[46], + '_impl.multiccd_face1': out[47], + '_impl.multiccd_face2': out[48], + '_impl.multiccd_idx1': out[49], + '_impl.multiccd_idx2': out[50], + '_impl.multiccd_n1': out[51], + '_impl.multiccd_n2': out[52], + '_impl.multiccd_pdist': out[53], + '_impl.multiccd_pnormal': out[54], + '_impl.multiccd_polygon': out[55], + '_impl.ncollision': out[56], + '_impl.ncon': out[57], + '_impl.ne': out[58], + '_impl.ne_connect': out[59], + '_impl.ne_jnt': out[60], + '_impl.ne_ten': out[61], + '_impl.ne_weld': out[62], + '_impl.nefc': out[63], + '_impl.nf': out[64], + '_impl.nl': out[65], + '_impl.nsolving': out[66], + '_impl.qLD': out[67], + '_impl.qLDiagInv': out[68], + '_impl.qM': out[69], + 'qacc': out[70], + 'qacc_smooth': out[71], + 'qacc_warmstart': out[72], + 'qfrc_actuator': out[73], + 'qfrc_applied': out[74], + 'qfrc_bias': out[75], + 'qfrc_constraint': out[76], + '_impl.qfrc_damper': out[77], + 'qfrc_fluid': out[78], + 'qfrc_gravcomp': out[79], + 'qfrc_passive': out[80], + 'qfrc_smooth': out[81], + '_impl.qfrc_spring': out[82], + 'qpos': out[83], + 'qvel': out[84], + '_impl.sap_cumulative_sum': out[85], + '_impl.sap_projection_lower': out[86], + '_impl.sap_projection_upper': out[87], + '_impl.sap_range': out[88], + '_impl.sap_segment_index': out[89], + '_impl.sap_sort_index': out[90], + '_impl.sensor_contact_criteria': out[91], + '_impl.sensor_contact_direction': out[92], + '_impl.sensor_contact_matchid': out[93], + '_impl.sensor_contact_nmatch': out[94], + '_impl.sensor_rangefinder_dist': out[95], + '_impl.sensor_rangefinder_geomid': out[96], + '_impl.sensor_rangefinder_pnt': out[97], + '_impl.sensor_rangefinder_vec': out[98], + 'sensordata': out[99], + 'site_xmat': out[100], + 'site_xpos': out[101], + '_impl.solver_niter': out[102], + '_impl.subtree_angmom': out[103], + '_impl.subtree_bodyvel': out[104], + 'subtree_com': out[105], + '_impl.subtree_linvel': out[106], + '_impl.ten_J': out[107], + '_impl.ten_Jdot': out[108], + '_impl.ten_actfrc': out[109], + '_impl.ten_bias_coef': out[110], + 'ten_length': out[111], + '_impl.ten_velocity': out[112], + '_impl.ten_wrapadr': out[113], + '_impl.ten_wrapnum': out[114], + 'time': out[115], + '_impl.wrap_geom_xpos': out[116], + '_impl.wrap_obj': out[117], + '_impl.wrap_xpos': out[118], + 'xanchor': out[119], + 'xaxis': out[120], + 'xfrc_applied': out[121], + 'ximat': out[122], + 'xipos': out[123], + 'xmat': out[124], + 'xpos': out[125], + 'xquat': out[126], + '_impl.contact__dim': out[127], + '_impl.contact__dist': out[128], + '_impl.contact__efc_address': out[129], + '_impl.contact__frame': out[130], + '_impl.contact__friction': out[131], + '_impl.contact__geom': out[132], + '_impl.contact__includemargin': out[133], + '_impl.contact__pos': out[134], + '_impl.contact__solimp': out[135], + '_impl.contact__solref': out[136], + '_impl.contact__solreffriction': out[137], + '_impl.contact__worldid': out[138], + '_impl.efc__D': out[139], + '_impl.efc__J': out[140], + '_impl.efc__Jaref': out[141], + '_impl.efc__Ma': out[142], + '_impl.efc__Mgrad': out[143], + '_impl.efc__alpha': out[144], + '_impl.efc__aref': out[145], + '_impl.efc__beta': out[146], + '_impl.efc__cholesky_L_tmp': out[147], + '_impl.efc__cholesky_y_tmp': out[148], + '_impl.efc__cost': out[149], + '_impl.efc__cost_candidate': out[150], + '_impl.efc__done': out[151], + '_impl.efc__force': out[152], + '_impl.efc__frictionloss': out[153], + '_impl.efc__gauss': out[154], + '_impl.efc__grad': out[155], + '_impl.efc__grad_dot': out[156], + '_impl.efc__h': out[157], + '_impl.efc__id': out[158], + '_impl.efc__jv': out[159], + '_impl.efc__margin': out[160], + '_impl.efc__mv': out[161], + '_impl.efc__pos': out[162], + '_impl.efc__prev_Mgrad': out[163], + '_impl.efc__prev_cost': out[164], + '_impl.efc__prev_grad': out[165], + '_impl.efc__quad': out[166], + '_impl.efc__quad_gauss': out[167], + '_impl.efc__search': out[168], + '_impl.efc__search_dot': out[169], + '_impl.efc__state': out[170], + '_impl.efc__type': out[171], + '_impl.efc__vel': out[172], }) return d @@ -1989,6 +2053,7 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _step_shim( # Model @@ -2105,7 +2170,6 @@ def _step_shim( 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), @@ -2168,7 +2232,6 @@ def _step_shim( nflexvert: int, ngeom: int, ngravcomp: int, - nhfield: int, njnt: int, nlight: int, nlsp: int, @@ -2182,6 +2245,9 @@ def _step_shim( 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), @@ -2265,7 +2331,9 @@ def _step_shim( wrap_type: wp.array(dtype=int), opt__broadphase: int, opt__broadphase_filter: int, + opt__ccd_tolerance: wp.array(dtype=float), opt__cone: int, + opt__contact_sensor_maxmatch: int, opt__density: wp.array(dtype=float), opt__disableflags: int, opt__enableflags: int, @@ -2313,7 +2381,6 @@ def _step_shim( cfrc_ext: wp.array2d(dtype=wp.spatial_vector), cfrc_int: wp.array2d(dtype=wp.spatial_vector), cinert: wp.array2d(dtype=mjwp_types.vec10), - 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), @@ -2345,9 +2412,19 @@ def _step_shim( light_xpos: wp.array2d(dtype=wp.vec3), mocap_pos: wp.array2d(dtype=wp.vec3), mocap_quat: wp.array2d(dtype=wp.quat), + 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), ne: wp.array(dtype=int), ne_connect: wp.array(dtype=int), ne_jnt: wp.array(dtype=int), @@ -2589,7 +2666,6 @@ def _step_shim( _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 @@ -2652,7 +2728,6 @@ def _step_shim( _m.nflexvert = nflexvert _m.ngeom = ngeom _m.ngravcomp = ngravcomp - _m.nhfield = nhfield _m.njnt = njnt _m.nlight = nlight _m.nlsp = nlsp @@ -2666,9 +2741,14 @@ def _step_shim( _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.cone = opt__cone + _m.opt.contact_sensor_maxmatch = opt__contact_sensor_maxmatch _m.opt.density = opt__density _m.opt.disableflags = opt__disableflags _m.opt.enableflags = opt__enableflags @@ -2794,7 +2874,6 @@ def _step_shim( _d.cfrc_ext = cfrc_ext _d.cfrc_int = cfrc_int _d.cinert = cinert - _d.collision_hftri_index = collision_hftri_index _d.collision_pair = collision_pair _d.collision_pairid = collision_pairid _d.collision_worldid = collision_worldid @@ -2872,9 +2951,19 @@ def _step_shim( _d.light_xpos = light_xpos _d.mocap_pos = mocap_pos _d.mocap_quat = mocap_quat + _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.ne = ne _d.ne_connect = ne_connect @@ -2978,7 +3067,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'cfrc_ext': d._impl.cfrc_ext.shape, 'cfrc_int': d._impl.cfrc_int.shape, 'cinert': d._impl.cinert.shape, - '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, @@ -3010,9 +3098,19 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'light_xpos': d._impl.light_xpos.shape, 'mocap_pos': d.mocap_pos.shape, 'mocap_quat': d.mocap_quat.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, 'ne': d._impl.ne.shape, 'ne_connect': d._impl.ne_connect.shape, 'ne_jnt': d._impl.ne_jnt.shape, @@ -3140,7 +3238,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _step_shim, - num_outputs=176, + num_outputs=185, output_dims=output_dims, vmap_method=None, in_out_argnames={ @@ -3161,7 +3259,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'cfrc_ext', 'cfrc_int', 'cinert', - 'collision_hftri_index', 'collision_pair', 'collision_pairid', 'collision_worldid', @@ -3193,9 +3290,19 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'light_xpos', 'mocap_pos', 'mocap_quat', + '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', 'ne', 'ne_connect', 'ne_jnt', @@ -3436,7 +3543,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.geom_solmix, m.geom_solref, m.geom_type, - m._impl.geompair2hfgeompair, m._impl.has_sdf_geom, m.hfield_adr, m.hfield_data, @@ -3499,7 +3605,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.nflexvert, m.ngeom, m.ngravcomp, - m.nhfield, m.njnt, m.nlight, m._impl.nlsp, @@ -3513,6 +3618,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): 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, @@ -3596,7 +3704,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.wrap_type, m.opt._impl.broadphase, m.opt._impl.broadphase_filter, + m.opt._impl.ccd_tolerance, m.opt.cone, + m.opt._impl.contact_sensor_maxmatch, m.opt.density, m.opt.disableflags, m.opt.enableflags, @@ -3643,7 +3753,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.cfrc_ext, d._impl.cfrc_int, d._impl.cinert, - d._impl.collision_hftri_index, d._impl.collision_pair, d._impl.collision_pairid, d._impl.collision_worldid, @@ -3675,9 +3784,19 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.light_xpos, d.mocap_pos, d.mocap_quat, + 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.ne, d._impl.ne_connect, d._impl.ne_jnt, @@ -3821,165 +3940,174 @@ def _step_jax_impl(m: types.Model, d: types.Data): '_impl.cfrc_ext': out[14], '_impl.cfrc_int': out[15], '_impl.cinert': out[16], - '_impl.collision_hftri_index': out[17], - '_impl.collision_pair': out[18], - '_impl.collision_pairid': out[19], - '_impl.collision_worldid': out[20], - '_impl.crb': out[21], - 'ctrl': out[22], - 'cvel': out[23], - '_impl.energy': out[24], - '_impl.epa_face': out[25], - '_impl.epa_horizon': out[26], - '_impl.epa_index': out[27], - '_impl.epa_map': out[28], - '_impl.epa_norm2': out[29], - '_impl.epa_pr': out[30], - '_impl.epa_vert': out[31], - '_impl.epa_vert1': out[32], - '_impl.epa_vert2': out[33], - '_impl.epa_vert_index1': out[34], - '_impl.epa_vert_index2': out[35], - 'eq_active': out[36], - '_impl.flexedge_length': out[37], - '_impl.flexedge_velocity': out[38], - '_impl.flexvert_xpos': out[39], - '_impl.fluid_applied': out[40], - '_impl.geom_skip': out[41], - 'geom_xmat': out[42], - 'geom_xpos': out[43], - '_impl.inverse_mul_m_skip': out[44], - '_impl.light_xdir': out[45], - '_impl.light_xpos': out[46], - 'mocap_pos': out[47], - 'mocap_quat': out[48], - '_impl.ncollision': out[49], - '_impl.ncon': out[50], - '_impl.ncon_hfield': out[51], - '_impl.ne': out[52], - '_impl.ne_connect': out[53], - '_impl.ne_jnt': out[54], - '_impl.ne_ten': out[55], - '_impl.ne_weld': out[56], - '_impl.nefc': out[57], - '_impl.nf': out[58], - '_impl.nl': out[59], - '_impl.nsolving': out[60], - '_impl.qLD': out[61], - '_impl.qLD_integration': out[62], - '_impl.qLDiagInv': out[63], - '_impl.qLDiagInv_integration': out[64], - '_impl.qM': out[65], - '_impl.qM_integration': out[66], - 'qacc': out[67], - '_impl.qacc_integration': out[68], - '_impl.qacc_rk': out[69], - 'qacc_smooth': out[70], - 'qacc_warmstart': out[71], - 'qfrc_actuator': out[72], - 'qfrc_applied': out[73], - 'qfrc_bias': out[74], - 'qfrc_constraint': out[75], - '_impl.qfrc_damper': out[76], - 'qfrc_fluid': out[77], - 'qfrc_gravcomp': out[78], - '_impl.qfrc_integration': out[79], - 'qfrc_passive': out[80], - 'qfrc_smooth': out[81], - '_impl.qfrc_spring': out[82], - 'qpos': out[83], - '_impl.qpos_t0': out[84], - 'qvel': out[85], - '_impl.qvel_rk': out[86], - '_impl.qvel_t0': out[87], - '_impl.sap_cumulative_sum': out[88], - '_impl.sap_projection_lower': out[89], - '_impl.sap_projection_upper': out[90], - '_impl.sap_range': out[91], - '_impl.sap_segment_index': out[92], - '_impl.sap_sort_index': out[93], - '_impl.sensor_contact_criteria': out[94], - '_impl.sensor_contact_direction': out[95], - '_impl.sensor_contact_matchid': out[96], - '_impl.sensor_contact_nmatch': out[97], - '_impl.sensor_rangefinder_dist': out[98], - '_impl.sensor_rangefinder_geomid': out[99], - '_impl.sensor_rangefinder_pnt': out[100], - '_impl.sensor_rangefinder_vec': out[101], - 'sensordata': out[102], - 'site_xmat': out[103], - 'site_xpos': out[104], - '_impl.solver_niter': out[105], - '_impl.subtree_angmom': out[106], - '_impl.subtree_bodyvel': out[107], - 'subtree_com': out[108], - '_impl.subtree_linvel': out[109], - '_impl.ten_J': out[110], - '_impl.ten_Jdot': out[111], - '_impl.ten_actfrc': out[112], - '_impl.ten_bias_coef': out[113], - 'ten_length': out[114], - '_impl.ten_velocity': out[115], - '_impl.ten_wrapadr': out[116], - '_impl.ten_wrapnum': out[117], - 'time': out[118], - '_impl.wrap_geom_xpos': out[119], - '_impl.wrap_obj': out[120], - '_impl.wrap_xpos': out[121], - 'xanchor': out[122], - 'xaxis': out[123], - 'xfrc_applied': out[124], - 'ximat': out[125], - 'xipos': out[126], - 'xmat': out[127], - 'xpos': out[128], - 'xquat': out[129], - '_impl.contact__dim': out[130], - '_impl.contact__dist': out[131], - '_impl.contact__efc_address': out[132], - '_impl.contact__frame': out[133], - '_impl.contact__friction': out[134], - '_impl.contact__geom': out[135], - '_impl.contact__includemargin': out[136], - '_impl.contact__pos': out[137], - '_impl.contact__solimp': out[138], - '_impl.contact__solref': out[139], - '_impl.contact__solreffriction': out[140], - '_impl.contact__worldid': out[141], - '_impl.efc__D': out[142], - '_impl.efc__J': out[143], - '_impl.efc__Jaref': out[144], - '_impl.efc__Ma': out[145], - '_impl.efc__Mgrad': out[146], - '_impl.efc__alpha': out[147], - '_impl.efc__aref': out[148], - '_impl.efc__beta': out[149], - '_impl.efc__cholesky_L_tmp': out[150], - '_impl.efc__cholesky_y_tmp': out[151], - '_impl.efc__cost': out[152], - '_impl.efc__cost_candidate': out[153], - '_impl.efc__done': out[154], - '_impl.efc__force': out[155], - '_impl.efc__frictionloss': out[156], - '_impl.efc__gauss': out[157], - '_impl.efc__grad': out[158], - '_impl.efc__grad_dot': out[159], - '_impl.efc__h': out[160], - '_impl.efc__id': out[161], - '_impl.efc__jv': out[162], - '_impl.efc__margin': out[163], - '_impl.efc__mv': out[164], - '_impl.efc__pos': out[165], - '_impl.efc__prev_Mgrad': out[166], - '_impl.efc__prev_cost': out[167], - '_impl.efc__prev_grad': out[168], - '_impl.efc__quad': out[169], - '_impl.efc__quad_gauss': out[170], - '_impl.efc__search': out[171], - '_impl.efc__search_dot': out[172], - '_impl.efc__state': out[173], - '_impl.efc__type': out[174], - '_impl.efc__vel': out[175], + '_impl.collision_pair': out[17], + '_impl.collision_pairid': out[18], + '_impl.collision_worldid': out[19], + '_impl.crb': out[20], + 'ctrl': out[21], + 'cvel': out[22], + '_impl.energy': out[23], + '_impl.epa_face': out[24], + '_impl.epa_horizon': out[25], + '_impl.epa_index': out[26], + '_impl.epa_map': out[27], + '_impl.epa_norm2': out[28], + '_impl.epa_pr': out[29], + '_impl.epa_vert': out[30], + '_impl.epa_vert1': out[31], + '_impl.epa_vert2': out[32], + '_impl.epa_vert_index1': out[33], + '_impl.epa_vert_index2': out[34], + 'eq_active': out[35], + '_impl.flexedge_length': out[36], + '_impl.flexedge_velocity': out[37], + '_impl.flexvert_xpos': out[38], + '_impl.fluid_applied': out[39], + '_impl.geom_skip': out[40], + 'geom_xmat': out[41], + 'geom_xpos': out[42], + '_impl.inverse_mul_m_skip': out[43], + '_impl.light_xdir': out[44], + '_impl.light_xpos': out[45], + 'mocap_pos': out[46], + 'mocap_quat': out[47], + '_impl.multiccd_clipped': out[48], + '_impl.multiccd_endvert': out[49], + '_impl.multiccd_face1': out[50], + '_impl.multiccd_face2': out[51], + '_impl.multiccd_idx1': out[52], + '_impl.multiccd_idx2': out[53], + '_impl.multiccd_n1': out[54], + '_impl.multiccd_n2': out[55], + '_impl.multiccd_pdist': out[56], + '_impl.multiccd_pnormal': out[57], + '_impl.multiccd_polygon': out[58], + '_impl.ncollision': out[59], + '_impl.ncon': out[60], + '_impl.ne': out[61], + '_impl.ne_connect': out[62], + '_impl.ne_jnt': out[63], + '_impl.ne_ten': out[64], + '_impl.ne_weld': out[65], + '_impl.nefc': out[66], + '_impl.nf': out[67], + '_impl.nl': out[68], + '_impl.nsolving': out[69], + '_impl.qLD': out[70], + '_impl.qLD_integration': out[71], + '_impl.qLDiagInv': out[72], + '_impl.qLDiagInv_integration': out[73], + '_impl.qM': out[74], + '_impl.qM_integration': out[75], + 'qacc': out[76], + '_impl.qacc_integration': out[77], + '_impl.qacc_rk': out[78], + 'qacc_smooth': out[79], + 'qacc_warmstart': out[80], + 'qfrc_actuator': out[81], + 'qfrc_applied': out[82], + 'qfrc_bias': out[83], + 'qfrc_constraint': out[84], + '_impl.qfrc_damper': out[85], + 'qfrc_fluid': out[86], + 'qfrc_gravcomp': out[87], + '_impl.qfrc_integration': out[88], + 'qfrc_passive': out[89], + 'qfrc_smooth': out[90], + '_impl.qfrc_spring': out[91], + 'qpos': out[92], + '_impl.qpos_t0': out[93], + 'qvel': out[94], + '_impl.qvel_rk': out[95], + '_impl.qvel_t0': out[96], + '_impl.sap_cumulative_sum': out[97], + '_impl.sap_projection_lower': out[98], + '_impl.sap_projection_upper': out[99], + '_impl.sap_range': out[100], + '_impl.sap_segment_index': out[101], + '_impl.sap_sort_index': out[102], + '_impl.sensor_contact_criteria': out[103], + '_impl.sensor_contact_direction': out[104], + '_impl.sensor_contact_matchid': out[105], + '_impl.sensor_contact_nmatch': out[106], + '_impl.sensor_rangefinder_dist': out[107], + '_impl.sensor_rangefinder_geomid': out[108], + '_impl.sensor_rangefinder_pnt': out[109], + '_impl.sensor_rangefinder_vec': out[110], + 'sensordata': out[111], + 'site_xmat': out[112], + 'site_xpos': out[113], + '_impl.solver_niter': out[114], + '_impl.subtree_angmom': out[115], + '_impl.subtree_bodyvel': out[116], + 'subtree_com': out[117], + '_impl.subtree_linvel': out[118], + '_impl.ten_J': out[119], + '_impl.ten_Jdot': out[120], + '_impl.ten_actfrc': out[121], + '_impl.ten_bias_coef': out[122], + 'ten_length': out[123], + '_impl.ten_velocity': out[124], + '_impl.ten_wrapadr': out[125], + '_impl.ten_wrapnum': out[126], + 'time': out[127], + '_impl.wrap_geom_xpos': out[128], + '_impl.wrap_obj': out[129], + '_impl.wrap_xpos': out[130], + 'xanchor': out[131], + 'xaxis': out[132], + 'xfrc_applied': out[133], + 'ximat': out[134], + 'xipos': out[135], + 'xmat': out[136], + 'xpos': out[137], + 'xquat': out[138], + '_impl.contact__dim': out[139], + '_impl.contact__dist': out[140], + '_impl.contact__efc_address': out[141], + '_impl.contact__frame': out[142], + '_impl.contact__friction': out[143], + '_impl.contact__geom': out[144], + '_impl.contact__includemargin': out[145], + '_impl.contact__pos': out[146], + '_impl.contact__solimp': out[147], + '_impl.contact__solref': out[148], + '_impl.contact__solreffriction': out[149], + '_impl.contact__worldid': out[150], + '_impl.efc__D': out[151], + '_impl.efc__J': out[152], + '_impl.efc__Jaref': out[153], + '_impl.efc__Ma': out[154], + '_impl.efc__Mgrad': out[155], + '_impl.efc__alpha': out[156], + '_impl.efc__aref': out[157], + '_impl.efc__beta': out[158], + '_impl.efc__cholesky_L_tmp': out[159], + '_impl.efc__cholesky_y_tmp': out[160], + '_impl.efc__cost': out[161], + '_impl.efc__cost_candidate': out[162], + '_impl.efc__done': out[163], + '_impl.efc__force': out[164], + '_impl.efc__frictionloss': out[165], + '_impl.efc__gauss': out[166], + '_impl.efc__grad': out[167], + '_impl.efc__grad_dot': out[168], + '_impl.efc__h': out[169], + '_impl.efc__id': out[170], + '_impl.efc__jv': out[171], + '_impl.efc__margin': out[172], + '_impl.efc__mv': out[173], + '_impl.efc__pos': out[174], + '_impl.efc__prev_Mgrad': out[175], + '_impl.efc__prev_cost': out[176], + '_impl.efc__prev_grad': out[177], + '_impl.efc__quad': out[178], + '_impl.efc__quad_gauss': out[179], + '_impl.efc__search': out[180], + '_impl.efc__search_dot': out[181], + '_impl.efc__state': out[182], + '_impl.efc__type': out[183], + '_impl.efc__vel': out[184], }) return d diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index ea56b2b5..0dc1dd87 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -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, diff --git a/mjx/pyproject.toml b/mjx/pyproject.toml index 454df05e..dd5072f4 100644 --- a/mjx/pyproject.toml +++ b/mjx/pyproject.toml @@ -36,7 +36,7 @@ dependencies = [ [project.optional-dependencies] warp = [ - "warp-lang==1.8.1", + "warp-lang==1.9.0", ] [project.scripts]