diff --git a/mujoco_warp/_src/collision_core.py b/mujoco_warp/_src/collision_core.py index 120bff8c1..b1898ed99 100644 --- a/mujoco_warp/_src/collision_core.py +++ b/mujoco_warp/_src/collision_core.py @@ -504,77 +504,6 @@ def sap_range( range_out[worldid, sortedid] = limit - sortedid -@wp.kernel -def sap_sweep( - # In: - n: int, - sort_index_in: wp.array2d[int], - cumulative_sum_in: wp.array[int], - nsweep_in: int, - aabb_lower_in: wp.array2d[wp.vec3], - aabb_upper_in: wp.array2d[wp.vec3], - max_pairs: int, - # Out: - npairs_out: wp.array[int], - pair_id1_out: wp.array[int], - pair_id2_out: wp.array[int], - pair_worldid_out: wp.array[int], -): - """Generic SAP sweep: output AABB-overlapping pairs. - - This is the GPU equivalent of MuJoCo's mj_SAP function. It takes - axis-aligned bounding boxes and outputs pairs whose AABBs overlap - on all 3 axes. Domain-specific filtering is done by the caller. - """ - worldelemid = wp.tid() - - nworldelem = cumulative_sum_in.shape[0] - nworkpackages = cumulative_sum_in[nworldelem - 1] - - while worldelemid < nworkpackages: - # Binary search to find sortedid (i) and partner sortedid (j) - i = sap_binary_search(cumulative_sum_in, worldelemid, 0, nworldelem) - j = i + worldelemid + 1 - if i > 0: - j -= cumulative_sum_in[i - 1] - - worldid = i // n - i = i % n - j = j % n - - # Get actual element indices from sorted order - elem1 = sort_index_in[worldid, i] - elem2 = sort_index_in[worldid, j] - - # Ensure elem1 < elem2 for consistent ordering - if elem1 > elem2: - tmp = elem1 - elem1 = elem2 - elem2 = tmp - - worldelemid += nsweep_in - - # AABB overlap test on all 3 axes - lower1 = aabb_lower_in[worldid, elem1] - upper1 = aabb_upper_in[worldid, elem1] - lower2 = aabb_lower_in[worldid, elem2] - upper2 = aabb_upper_in[worldid, elem2] - - if lower1[0] > upper2[0] or lower2[0] > upper1[0]: - continue - if lower1[1] > upper2[1] or lower2[1] > upper1[1]: - continue - if lower1[2] > upper2[2] or lower2[2] > upper1[2]: - continue - - # Output this pair - idx = wp.atomic_add(npairs_out, 0, 1) - if idx < max_pairs: - pair_id1_out[idx] = elem1 - pair_id2_out[idx] = elem2 - pair_worldid_out[idx] = worldid - - def create_collision_context(naconmax: int) -> CollisionContext: """Create a CollisionContext with allocated arrays.""" return CollisionContext( diff --git a/mujoco_warp/_src/collision_flex.py b/mujoco_warp/_src/collision_flex.py index 28790171e..7c50a94a1 100644 --- a/mujoco_warp/_src/collision_flex.py +++ b/mujoco_warp/_src/collision_flex.py @@ -19,25 +19,28 @@ from mujoco_warp._src import collision_primitive_core from mujoco_warp._src.collision_core import Geom +from mujoco_warp._src.collision_core import sap_binary_search from mujoco_warp._src.collision_core import sap_range -from mujoco_warp._src.collision_core import sap_sweep # TODO(team): consolidate _flex_sap_project with geom _sap_project from mujoco_warp._src.collision_gjk import ccd from mujoco_warp._src.math import make_frame from mujoco_warp._src.types import MJ_MAX_EPAFACES from mujoco_warp._src.types import MJ_MAX_EPAHORIZON +from mujoco_warp._src.types import MJ_MAXCONPAIR from mujoco_warp._src.types import MJ_MAXVAL from mujoco_warp._src.types import MJ_MINMU from mujoco_warp._src.types import MJ_MINVAL -from mujoco_warp._src.types import ContactType from mujoco_warp._src.types import Data from mujoco_warp._src.types import GeomType from mujoco_warp._src.types import Model from mujoco_warp._src.types import OverflowType from mujoco_warp._src.types import vec5 +from mujoco_warp._src.warp_util import cache_kernel from mujoco_warp._src.warp_util import event_scope wp.set_module_options({"enable_backward": False, "default_grid_stride": False}) +NFLEXELEM_SAP = 32 + @wp.func def _flex_element_aabb_filter( @@ -94,377 +97,229 @@ def _flex_broadphase_bounds( flex_aabb_max_out[worldid, flexid] = max_bound + inflate -@wp.func -def _flex_triangle_geom_broadphase( - # Model: - ngeom: int, - opt_warn_overflow: bool, - geom_type: wp.array[int], - geom_aabb: wp.array3d[wp.vec3], - geom_margin: wp.array2d[float], - # Data in: - geom_xpos_in: wp.array2d[wp.vec3], - geom_xmat_in: wp.array2d[wp.mat33], - naconmax_in: int, - # In: - worldid: int, - element_or_shell_id: int, - flexid: int, - geomid: int, - t1: wp.vec3, - t2: wp.vec3, - t3: wp.vec3, - tri_radius: float, - tri_margin: float, - flex_aabb_min_val: wp.vec3, - flex_aabb_max_val: wp.vec3, - # Data out: - ncollision_out: wp.array[int], - # Data out: - overflow_out: wp.array[int], - # Out: - collision_pair_out: wp.array[wp.vec2i], - collision_worldid_out: wp.array[int], -): - gtype = geom_type[geomid] - if ( - gtype != int(GeomType.SPHERE) - and gtype != int(GeomType.CAPSULE) - and gtype != int(GeomType.BOX) - and gtype != int(GeomType.CYLINDER) - and gtype != int(GeomType.MESH) - and gtype != int(GeomType.ELLIPSOID) +@cache_kernel +def flex_broadphase_builder(warn_overflow: bool): + @wp.kernel(module="unique", enable_backward=False) + def _flex_broadphase( + # Model: + geom_type: wp.array[int], + geom_aabb: wp.array3d[wp.vec3], + geom_margin: wp.array2d[float], + flex_margin: wp.array[float], + flex_vertadr: wp.array[int], + flex_radius: wp.array[float], + # Data in: + geom_xpos_in: wp.array2d[wp.vec3], + geom_xmat_in: wp.array2d[wp.mat33], + flexvert_xpos_in: wp.array2d[wp.vec3], + naconmax_in: int, + flex_aabb_min_in: wp.array2d[wp.vec3], + flex_aabb_max_in: wp.array2d[wp.vec3], + # In: + triadr: wp.array[int], + tridataadr: wp.array[int], + tri: wp.array[int], + pairs_filtered: wp.array[wp.vec2i], + triflexid: wp.array[int], + # Data out: + ncollision_out: wp.array[int], + # Data out: + overflow_out: wp.array[int], + # Out: + collision_pair_out: wp.array[wp.vec2i], + collision_worldid_out: wp.array[int], ): - return + worldid, pairid = wp.tid() - geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] - margin = geom_margin_val + tri_margin - - aabb_id = worldid % geom_aabb.shape[0] - geom_center_local = geom_aabb[aabb_id, geomid, 0] - geom_half_size_local = geom_aabb[aabb_id, geomid, 1] - - geom_pos = geom_xpos_in[worldid, geomid] - geom_rot = geom_xmat_in[worldid, geomid] - - # Stage 1 Filter: Coarse flex AABB vs Geom world AABB check - # Transform center to global frame - geom_center_global = geom_rot @ geom_center_local + geom_pos - - # Project local half-size onto world axes using absolute rotation matrix entries - geom_half_size_global = wp.vec3( - wp.abs(geom_rot[0, 0]) * geom_half_size_local[0] - + wp.abs(geom_rot[0, 1]) * geom_half_size_local[1] - + wp.abs(geom_rot[0, 2]) * geom_half_size_local[2], - wp.abs(geom_rot[1, 0]) * geom_half_size_local[0] - + wp.abs(geom_rot[1, 1]) * geom_half_size_local[1] - + wp.abs(geom_rot[1, 2]) * geom_half_size_local[2], - wp.abs(geom_rot[2, 0]) * geom_half_size_local[0] - + wp.abs(geom_rot[2, 1]) * geom_half_size_local[1] - + wp.abs(geom_rot[2, 2]) * geom_half_size_local[2], - ) + pair = pairs_filtered[pairid] + tri_id = pair[0] + geomid = pair[1] - inflate = wp.vec3(margin, margin, margin) - geom_box_min = geom_center_global - geom_half_size_global - inflate - geom_box_max = geom_center_global + geom_half_size_global + inflate + flexid = triflexid[tri_id] - if _flex_element_aabb_filter(geom_box_min, geom_box_max, flex_aabb_min_val, flex_aabb_max_val): - return + vert_adr = flex_vertadr[flexid] + tri_radius = flex_radius[flexid] + tri_margin = flex_margin[flexid] - # Stage 2 Filter: Element AABB vs Geom world AABB check - tri_min = wp.min(t1, wp.min(t2, t3)) - wp.vec3(tri_radius, tri_radius, tri_radius) - tri_max = wp.max(t1, wp.max(t2, t3)) + wp.vec3(tri_radius, tri_radius, tri_radius) + tri_data_idx = tridataadr[flexid] + (tri_id - triadr[flexid]) * 3 + v0_local = tri[tri_data_idx] + v1_local = tri[tri_data_idx + 1] + v2_local = tri[tri_data_idx + 2] - if _flex_element_aabb_filter(geom_box_min, geom_box_max, tri_min, tri_max): - return + t1 = flexvert_xpos_in[worldid, vert_adr + v0_local] + t2 = flexvert_xpos_in[worldid, vert_adr + v1_local] + t3 = flexvert_xpos_in[worldid, vert_adr + v2_local] - # Stage 3 Filter: Project Geom onto triangle plane normal - normal = wp.normalize(wp.cross(t2 - t1, t3 - t1)) - signed_dist = wp.dot(geom_pos - t1, normal) + gtype = geom_type[geomid] + if ( + gtype != int(GeomType.SPHERE) + and gtype != int(GeomType.CAPSULE) + and gtype != int(GeomType.BOX) + and gtype != int(GeomType.CYLINDER) + and gtype != int(GeomType.MESH) + and gtype != int(GeomType.ELLIPSOID) + ): + return - r_extent = float(0.0) - if gtype == int(GeomType.SPHERE): - r_extent = geom_half_size_local[0] - elif gtype == int(GeomType.CAPSULE): - r_extent = geom_half_size_local[0] + geom_half_size_local[1] - elif gtype == int(GeomType.CYLINDER): - r_extent = wp.sqrt(geom_half_size_local[0] * geom_half_size_local[0] + geom_half_size_local[1] * geom_half_size_local[1]) - elif gtype == int(GeomType.BOX): - r_extent = wp.length(geom_half_size_local) - elif gtype == int(GeomType.MESH): - r_extent = wp.length(geom_half_size_local) - elif gtype == int(GeomType.ELLIPSOID): - r_extent = wp.length(geom_half_size_local) + geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] + margin = geom_margin_val + tri_margin - if wp.abs(signed_dist) > r_extent + margin + tri_radius: - return + aabb_id = worldid % geom_aabb.shape[0] + geom_center_local = geom_aabb[aabb_id, geomid, 0] + geom_half_size_local = geom_aabb[aabb_id, geomid, 1] - # Overlap found! Save candidate. - idx = wp.atomic_add(ncollision_out, 0, 1) - if idx >= naconmax_in: - if opt_warn_overflow: - wp.printf("Collision buffer overflow in flex broadphase - please increase naconmax to %u\n", idx + 1) - wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) - return - collision_pair_out[idx] = wp.vec2i(element_or_shell_id, geomid) - collision_worldid_out[idx] = worldid + geom_pos = geom_xpos_in[worldid, geomid] + geom_rot = geom_xmat_in[worldid, geomid] + # Stage 1 Filter: Coarse flex AABB vs Geom world AABB check + geom_center_global = geom_rot @ geom_center_local + geom_pos + + # Project local half-size onto world axes using absolute rotation matrix entries + geom_half_size_global = wp.vec3( + wp.abs(geom_rot[0, 0]) * geom_half_size_local[0] + + wp.abs(geom_rot[0, 1]) * geom_half_size_local[1] + + wp.abs(geom_rot[0, 2]) * geom_half_size_local[2], + wp.abs(geom_rot[1, 0]) * geom_half_size_local[0] + + wp.abs(geom_rot[1, 1]) * geom_half_size_local[1] + + wp.abs(geom_rot[1, 2]) * geom_half_size_local[2], + wp.abs(geom_rot[2, 0]) * geom_half_size_local[0] + + wp.abs(geom_rot[2, 1]) * geom_half_size_local[1] + + wp.abs(geom_rot[2, 2]) * geom_half_size_local[2], + ) -@wp.kernel -def _flex_broadphase_unified( - # Model: - ngeom: int, - nflex: int, - opt_warn_overflow: bool, - geom_type: wp.array[int], - geom_size: wp.array2d[wp.vec3], - geom_aabb: wp.array3d[wp.vec3], - geom_rbound: wp.array2d[float], - geom_margin: wp.array2d[float], - flex_margin: wp.array[float], - flex_dim: wp.array[int], - flex_vertadr: wp.array[int], - flex_radius: wp.array[float], - # Data in: - geom_xpos_in: wp.array2d[wp.vec3], - geom_xmat_in: wp.array2d[wp.mat33], - flexvert_xpos_in: wp.array2d[wp.vec3], - naconmax_in: int, - flex_aabb_min_in: wp.array2d[wp.vec3], - flex_aabb_max_in: wp.array2d[wp.vec3], - # In: - triadr: wp.array[int], - tridataadr: wp.array[int], - tri: wp.array[int], - pairs_filtered: wp.array[wp.vec2i], - triflexid: wp.array[int], - # Data out: - ncollision_out: wp.array[int], - # Data out: - overflow_out: wp.array[int], - # Out: - collision_pair_out: wp.array[wp.vec2i], - collision_worldid_out: wp.array[int], -): - worldid, pairid = wp.tid() + inflate = wp.vec3(margin, margin, margin) + geom_box_min = geom_center_global - geom_half_size_global - inflate + geom_box_max = geom_center_global + geom_half_size_global + inflate - pair = pairs_filtered[pairid] - tri_id = pair[0] - geomid = pair[1] + flex_aabb_min_val = flex_aabb_min_in[worldid, flexid] + flex_aabb_max_val = flex_aabb_max_in[worldid, flexid] - flexid = triflexid[tri_id] + if _flex_element_aabb_filter(geom_box_min, geom_box_max, flex_aabb_min_val, flex_aabb_max_val): + return - vert_adr = flex_vertadr[flexid] - tri_radius = flex_radius[flexid] - tri_margin = flex_margin[flexid] + # Stage 2 Filter: Element AABB vs Geom world AABB check + tri_min = wp.min(t1, wp.min(t2, t3)) - wp.vec3(tri_radius, tri_radius, tri_radius) + tri_max = wp.max(t1, wp.max(t2, t3)) + wp.vec3(tri_radius, tri_radius, tri_radius) - tri_data_idx = tridataadr[flexid] + (tri_id - triadr[flexid]) * 3 - v0_local = tri[tri_data_idx] - v1_local = tri[tri_data_idx + 1] - v2_local = tri[tri_data_idx + 2] - - t1 = flexvert_xpos_in[worldid, vert_adr + v0_local] - t2 = flexvert_xpos_in[worldid, vert_adr + v1_local] - t3 = flexvert_xpos_in[worldid, vert_adr + v2_local] - - _flex_triangle_geom_broadphase( - ngeom, - opt_warn_overflow, - geom_type, - geom_aabb, - geom_margin, - geom_xpos_in, - geom_xmat_in, - naconmax_in, - worldid, - tri_id, - flexid, - geomid, - t1, - t2, - t3, - tri_radius, - tri_margin, - flex_aabb_min_in[worldid, flexid], - flex_aabb_max_in[worldid, flexid], - # Data out: - ncollision_out, - overflow_out, - collision_pair_out, - collision_worldid_out, - ) + if _flex_element_aabb_filter(geom_box_min, geom_box_max, tri_min, tri_max): + return + # Stage 3 Filter: Project Geom onto triangle plane normal + normal = wp.normalize(wp.cross(t2 - t1, t3 - t1)) + signed_dist = wp.dot(geom_pos - t1, normal) -@wp.kernel -def _flex_broadphase_plane( - # Model: - ngeom: int, - opt_warn_overflow: bool, - geom_type: wp.array[int], - geom_margin: wp.array2d[float], - flex_margin: wp.array[float], - flex_vertadr: wp.array[int], - flex_radius: wp.array[float], - flexvert_geom_pair_filtered: wp.array[wp.vec2i], - flex_vertflexid: wp.array[int], - # Data in: - geom_xpos_in: wp.array2d[wp.vec3], - geom_xmat_in: wp.array2d[wp.mat33], - flexvert_xpos_in: wp.array2d[wp.vec3], - naconmax_in: int, - flex_aabb_min_in: wp.array2d[wp.vec3], - flex_aabb_max_in: wp.array2d[wp.vec3], - # Data out: - ncollision_out: wp.array[int], - # Data out: - overflow_out: wp.array[int], - # Out: - collision_pair_out: wp.array[wp.vec2i], - collision_worldid_out: wp.array[int], -): - worldid, pairid = wp.tid() + r_extent = float(0.0) + if gtype == int(GeomType.SPHERE): + r_extent = geom_half_size_local[0] + elif gtype == int(GeomType.CAPSULE): + r_extent = geom_half_size_local[0] + geom_half_size_local[1] + elif gtype == int(GeomType.CYLINDER): + r_extent = wp.sqrt(geom_half_size_local[0] * geom_half_size_local[0] + geom_half_size_local[1] * geom_half_size_local[1]) + elif gtype == int(GeomType.BOX): + r_extent = wp.length(geom_half_size_local) + elif gtype == int(GeomType.MESH): + r_extent = wp.length(geom_half_size_local) + elif gtype == int(GeomType.ELLIPSOID): + r_extent = wp.length(geom_half_size_local) - pair = flexvert_geom_pair_filtered[pairid] - vertid = pair[0] - geomid = pair[1] + if wp.abs(signed_dist) > r_extent + margin + tri_radius: + return - flexid = flex_vertflexid[vertid] - radius = flex_radius[flexid] - flex_margin_val = flex_margin[flexid] + # Overlap found! Save candidate. + idx = wp.atomic_add(ncollision_out, 0, 1) + if idx >= naconmax_in: + if wp.static(warn_overflow): + wp.printf("Collision buffer overflow in flex broadphase - please increase naconmax to %u\n", idx + 1) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) + return + collision_pair_out[idx] = wp.vec2i(tri_id, geomid) + collision_worldid_out[idx] = worldid - vert = flexvert_xpos_in[worldid, vertid] + return _flex_broadphase + + +@cache_kernel +def flex_broadphase_plane_builder(warn_overflow: bool): + @wp.kernel(module="unique", enable_backward=False) + def _flex_broadphase_plane( + # Model: + geom_type: wp.array[int], + geom_margin: wp.array2d[float], + flex_margin: wp.array[float], + flex_radius: wp.array[float], + flexvert_geom_pair_filtered: wp.array[wp.vec2i], + flex_vertflexid: wp.array[int], + # Data in: + geom_xpos_in: wp.array2d[wp.vec3], + geom_xmat_in: wp.array2d[wp.mat33], + flexvert_xpos_in: wp.array2d[wp.vec3], + naconmax_in: int, + flex_aabb_min_in: wp.array2d[wp.vec3], + flex_aabb_max_in: wp.array2d[wp.vec3], + # Data out: + ncollision_out: wp.array[int], + # Data out: + overflow_out: wp.array[int], + # Out: + collision_pair_out: wp.array[wp.vec2i], + collision_worldid_out: wp.array[int], + ): + worldid, pairid = wp.tid() - flex_aabb_min = flex_aabb_min_in[worldid, flexid] - flex_aabb_max = flex_aabb_max_in[worldid, flexid] + pair = flexvert_geom_pair_filtered[pairid] + vertid = pair[0] + geomid = pair[1] - gtype = geom_type[geomid] - if gtype != int(GeomType.PLANE): - return + flexid = flex_vertflexid[vertid] + radius = flex_radius[flexid] + flex_margin_val = flex_margin[flexid] - margin = geom_margin[worldid % geom_margin.shape[0], geomid] + flex_margin_val - geom_pos = geom_xpos_in[worldid, geomid] - geom_rot = geom_xmat_in[worldid, geomid] - plane_normal = wp.vec3(geom_rot[0, 2], geom_rot[1, 2], geom_rot[2, 2]) - - # Stage 1 filter: Bounding box of flex vs plane - flex_center = 0.5 * (flex_aabb_min + flex_aabb_max) - flex_half_size = 0.5 * (flex_aabb_max - flex_aabb_min) - - proj_half = ( - wp.abs(flex_half_size[0] * plane_normal[0]) - + wp.abs(flex_half_size[1] * plane_normal[1]) - + wp.abs(flex_half_size[2] * plane_normal[2]) - ) + vert = flexvert_xpos_in[worldid, vertid] - diff_center = flex_center - geom_pos - dist_center = wp.dot(diff_center, plane_normal) + flex_aabb_min = flex_aabb_min_in[worldid, flexid] + flex_aabb_max = flex_aabb_max_in[worldid, flexid] - if dist_center - proj_half > margin: - return + gtype = geom_type[geomid] + if gtype != int(GeomType.PLANE): + return - diff = vert - geom_pos - signed_dist = wp.dot(diff, plane_normal) - dist = signed_dist - radius + margin = geom_margin[worldid % geom_margin.shape[0], geomid] + flex_margin_val + geom_pos = geom_xpos_in[worldid, geomid] + geom_rot = geom_xmat_in[worldid, geomid] + plane_normal = wp.vec3(geom_rot[0, 2], geom_rot[1, 2], geom_rot[2, 2]) - if dist < margin: - # Append Candidate to Context - idx = wp.atomic_add(ncollision_out, 0, 1) - if idx >= naconmax_in: - if opt_warn_overflow: - wp.printf("Collision buffer overflow in flex plane broadphase - please increase naconmax to %u\n", idx + 1) - wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) + # Stage 1 filter: Bounding box of flex vs plane + flex_center = 0.5 * (flex_aabb_min + flex_aabb_max) + flex_half_size = 0.5 * (flex_aabb_max - flex_aabb_min) + + proj_half = ( + wp.abs(flex_half_size[0] * plane_normal[0]) + + wp.abs(flex_half_size[1] * plane_normal[1]) + + wp.abs(flex_half_size[2] * plane_normal[2]) + ) + + diff_center = flex_center - geom_pos + dist_center = wp.dot(diff_center, plane_normal) + + if dist_center - proj_half > margin: return - collision_pair_out[idx] = wp.vec2i(vertid, geomid) - collision_worldid_out[idx] = worldid + diff = vert - geom_pos + signed_dist = wp.dot(diff, plane_normal) + dist = signed_dist - radius -@wp.func -def _write_flex_contact( - # Model: - geom_condim: wp.array[int], - geom_priority: wp.array[int], - geom_solmix: wp.array2d[float], - geom_solref: wp.array2d[wp.vec2], - geom_solimp: wp.array2d[vec5], - geom_friction: wp.array2d[wp.vec3], - geom_gap: wp.array2d[float], - flex_condim: wp.array[int], - flex_priority: wp.array[int], - flex_solmix: wp.array[float], - flex_solref: wp.array[wp.vec2], - flex_solimp: wp.array[vec5], - flex_friction: wp.array[wp.vec3], - flex_margin: wp.array[float], - flex_gap: wp.array[float], - # Data in: - naconmax_in: int, - # In: - collisionid: int, - dist: float, - pos: wp.vec3, - normal: wp.vec3, - margin: float, - geomid: int, - flexid: int, - elemid: int, - vertid: int, - worldid: int, - # Data out: - contact_dist_out: wp.array[float], - contact_pos_out: wp.array[wp.vec3], - contact_frame_out: wp.array[wp.mat33], - contact_includemargin_out: wp.array[float], - contact_friction_out: wp.array[vec5], - contact_solref_out: wp.array[wp.vec2], - contact_solreffriction_out: wp.array[wp.vec2], - contact_solimp_out: wp.array[vec5], - contact_dim_out: wp.array[int], - contact_geom_out: wp.array[wp.vec2i], - contact_flex_out: wp.array[wp.vec2i], - contact_elem_out: wp.array[wp.vec2i], - contact_vert_out: wp.array[wp.vec2i], - contact_worldid_out: wp.array[int], - contact_type_out: wp.array[int], - contact_geomcollisionid_out: wp.array[int], - nacon_out: wp.array[int], -): - condim, gap, solref, solimp, friction = _mix_flex_contact_params( - geom_condim[geomid], - geom_priority[geomid], - geom_solmix[worldid % geom_solmix.shape[0], geomid], - geom_solref[worldid % geom_solref.shape[0], geomid], - geom_solimp[worldid % geom_solimp.shape[0], geomid], - geom_friction[worldid % geom_friction.shape[0], geomid], - geom_gap[worldid % geom_gap.shape[0], geomid], - flex_condim[flexid], - flex_priority[flexid], - flex_solmix[flexid], - flex_solref[flexid], - flex_solimp[flexid], - flex_friction[flexid], - flex_gap[flexid], - ) + if dist < margin: + # Append Candidate to Context + idx = wp.atomic_add(ncollision_out, 0, 1) + if idx >= naconmax_in: + if wp.static(warn_overflow): + wp.printf("Collision buffer overflow in flex plane broadphase - please increase naconmax to %u\n", idx + 1) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) + return + collision_pair_out[idx] = wp.vec2i(vertid, geomid) + collision_worldid_out[idx] = worldid - c_idx = wp.atomic_add(nacon_out, 0, 1) - if c_idx < naconmax_in: - frame = make_frame(normal) - - contact_dist_out[c_idx] = dist - contact_pos_out[c_idx] = pos - contact_frame_out[c_idx] = frame - contact_includemargin_out[c_idx] = margin - contact_friction_out[c_idx] = friction - contact_solref_out[c_idx] = solref - contact_solreffriction_out[c_idx] = solref # Same for flex contacts - contact_solimp_out[c_idx] = solimp - contact_dim_out[c_idx] = condim - contact_geom_out[c_idx] = wp.vec2i(geomid, -1) - contact_flex_out[c_idx] = wp.vec2i(-1, flexid) - contact_elem_out[c_idx] = wp.vec2i(-1, elemid) - contact_vert_out[c_idx] = wp.vec2i(-1, vertid) - contact_worldid_out[c_idx] = worldid - contact_type_out[c_idx] = ContactType.CONSTRAINT - contact_geomcollisionid_out[c_idx] = collisionid + return _flex_broadphase_plane # TODO(team): generalize into a shared contact parameter mixing function @@ -555,14 +410,15 @@ def _mix_flex_contact_params( @wp.func -def _write_candidate_contact( +def _write_candidate( # In: max_candidates: int, dist: float, pos: wp.vec3, nrm: wp.vec3, geom: int, - flexid: int, + flexid1: int, + flexid2: int, elemid: int, vertid: int, worldid: int, @@ -598,17 +454,17 @@ def _write_candidate_contact( cand_nrm_out[candid] = nrm if geom >= 0: cand_geom_out[candid] = wp.vec2i(geom, -1) - cand_flex_out[candid] = wp.vec2i(-1, flexid) + cand_flex_out[candid] = wp.vec2i(flexid1, flexid2) cand_elem_out[candid] = wp.vec2i(-1, elemid) cand_vert_out[candid] = wp.vec2i(-1, vertid) elif geom == -2: cand_geom_out[candid] = wp.vec2i(-1, -1) - cand_flex_out[candid] = wp.vec2i(flexid, flexid) + cand_flex_out[candid] = wp.vec2i(flexid1, flexid2) cand_elem_out[candid] = wp.vec2i(elemid, vertid) cand_vert_out[candid] = wp.vec2i(-1, -1) else: cand_geom_out[candid] = wp.vec2i(-1, -1) - cand_flex_out[candid] = wp.vec2i(flexid, flexid) + cand_flex_out[candid] = wp.vec2i(flexid1, flexid2) cand_elem_out[candid] = wp.vec2i(-1, elemid) cand_vert_out[candid] = wp.vec2i(vertid, -1) cand_worldid_out[candid] = worldid @@ -653,12 +509,13 @@ def _collide_geom_triangle_detect( sphere_radius = size_val[0] dist, contact_pos, nrm = collision_primitive_core.sphere_triangle(pos, sphere_radius, t1, t2, t3, tri_radius) if dist < margin: - _write_candidate_contact( + _write_candidate( max_candidates, dist, contact_pos, nrm, geomid, + -1, flexid, elemid, vertex_id, @@ -704,12 +561,13 @@ def _collide_geom_triangle_detect( if dists[0] < margin: p1 = wp.vec3(poss[0, 0], poss[0, 1], poss[0, 2]) n1 = wp.vec3(nrms[0, 0], nrms[0, 1], nrms[0, 2]) - _write_candidate_contact( + _write_candidate( max_candidates, dists[0], p1, n1, geomid, + -1, flexid, elemid, vertex_id, @@ -730,12 +588,13 @@ def _collide_geom_triangle_detect( if dists[1] < margin: p2 = wp.vec3(poss[1, 0], poss[1, 1], poss[1, 2]) n2 = wp.vec3(nrms[1, 0], nrms[1, 1], nrms[1, 2]) - _write_candidate_contact( + _write_candidate( max_candidates, dists[1], p2, n2, geomid, + -1, flexid, elemid, vertex_id, @@ -783,9 +642,6 @@ def _collide_mesh_triangle( geomid: int, flexid: int, elemid: int, - v0_local: int, - v1_local: int, - v2_local: int, worldid: int, did: int, epa_vert: wp.array[wp.vec3], @@ -933,21 +789,19 @@ def _collide_mesh_triangle( min_dist = wp.min(dist_v0, wp.min(dist_v1, dist_v2)) if min_dist < margin: - deepest_vert = v0_local pos = t1 - normal * (tri_radius + 0.5 * dist_v0) if dist_v1 < dist_v0 and dist_v1 < dist_v2: - deepest_vert = v1_local pos = t2 - normal * (tri_radius + 0.5 * dist_v1) elif dist_v2 < dist_v0 and dist_v2 < dist_v1: - deepest_vert = v2_local pos = t3 - normal * (tri_radius + 0.5 * dist_v2) - _write_candidate_contact( + _write_candidate( max_candidates, min_dist, pos, normal, geomid, + -1, flexid, elemid, -1, @@ -972,25 +826,9 @@ def _collide_mesh_triangle( @wp.kernel def _flex_plane_narrowphase( # Model: - ngeom: int, - nflexvert: int, geom_type: wp.array[int], - geom_condim: wp.array[int], - geom_priority: wp.array[int], - geom_solmix: wp.array2d[float], - geom_solref: wp.array2d[wp.vec2], - geom_solimp: wp.array2d[vec5], - geom_friction: wp.array2d[wp.vec3], geom_margin: wp.array2d[float], - geom_gap: wp.array2d[float], - flex_condim: wp.array[int], - flex_priority: wp.array[int], - flex_solmix: wp.array[float], - flex_solref: wp.array[wp.vec2], - flex_solimp: wp.array[vec5], - flex_friction: wp.array[wp.vec3], flex_margin: wp.array[float], - flex_gap: wp.array[float], flex_vertadr: wp.array[int], flex_radius: wp.array[float], flex_vertflexid: wp.array[int], @@ -998,30 +836,25 @@ def _flex_plane_narrowphase( geom_xpos_in: wp.array2d[wp.vec3], geom_xmat_in: wp.array2d[wp.mat33], flexvert_xpos_in: wp.array2d[wp.vec3], - nworld_in: int, - naconmax_in: int, ncollision_in: wp.array[int], # In: collision_pair_in: wp.array[wp.vec2i], collision_worldid_in: wp.array[int], + max_candidates: int, # Data out: - contact_dist_out: wp.array[float], - contact_pos_out: wp.array[wp.vec3], - contact_frame_out: wp.array[wp.mat33], - contact_includemargin_out: wp.array[float], - contact_friction_out: wp.array[vec5], - contact_solref_out: wp.array[wp.vec2], - contact_solreffriction_out: wp.array[wp.vec2], - contact_solimp_out: wp.array[vec5], - contact_dim_out: wp.array[int], - contact_geom_out: wp.array[wp.vec2i], - contact_flex_out: wp.array[wp.vec2i], - contact_elem_out: wp.array[wp.vec2i], - contact_vert_out: wp.array[wp.vec2i], - contact_worldid_out: wp.array[int], - contact_type_out: wp.array[int], - contact_geomcollisionid_out: wp.array[int], - nacon_out: wp.array[int], + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], ): collisionid = wp.tid() if collisionid >= ncollision_in[0] or collisionid >= collision_pair_in.shape[0]: @@ -1057,50 +890,29 @@ def _flex_plane_narrowphase( if dist < margin: contact_pos = vert - plane_normal * (dist * 0.5 + radius) - _write_flex_contact( - geom_condim, - geom_priority, - geom_solmix, - geom_solref, - geom_solimp, - geom_friction, - geom_gap, - flex_condim, - flex_priority, - flex_solmix, - flex_solref, - flex_solimp, - flex_friction, - flex_margin, - flex_gap, - naconmax_in, - collisionid, + _write_candidate( + max_candidates, dist, contact_pos, plane_normal, - margin, geomid, + -1, flexid, -1, local_vertid, worldid, - 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_flex_out, - contact_elem_out, - contact_vert_out, - contact_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, ) @@ -1108,7 +920,6 @@ def _flex_plane_narrowphase( def _flex_geom_vertex_narrowphase_detect( # Model: ngeom: int, - nflexvert: int, geom_type: wp.array[int], geom_contype: wp.array[int], geom_conaffinity: wp.array[int], @@ -1125,7 +936,6 @@ def _flex_geom_vertex_narrowphase_detect( geom_xpos_in: wp.array2d[wp.vec3], geom_xmat_in: wp.array2d[wp.mat33], flexvert_xpos_in: wp.array2d[wp.vec3], - nworld_in: int, # In: max_candidates: int, # Data out: @@ -1204,12 +1014,13 @@ def _flex_geom_vertex_narrowphase_detect( ) if dist < margin: - _write_candidate_contact( + _write_candidate( max_candidates, dist, contact_pos, nrm, geomid, + -1, flexid, -1, local_vertid, @@ -1292,14 +1103,11 @@ def _plane_vertex( @wp.kernel(module="unique", enable_backward=False, grid_stride=False) def _flex_internal_collisions_detect( # Model: - nflex: int, flex_margin: wp.array[float], flex_internal: wp.array[int], flex_dim: wp.array[int], flex_vertadr: wp.array[int], flex_elemdataadr: wp.array[int], - flex_evpairadr: wp.array[int], - flex_evpairnum: wp.array[int], flex_elem: wp.array[int], flex_evpair: wp.array[wp.vec2i], flex_radius: wp.array[float], @@ -1373,13 +1181,14 @@ def _flex_internal_collisions_detect( dist, contact_pos, nrm = _sphere_tetrahedron(sphere_pos, radius, p0, p1, p2, p3, radius) if dist < margin: - _write_candidate_contact( + _write_candidate( max_candidates, dist, contact_pos, nrm, -1, flexid, + flexid, e, v, worldid, @@ -1401,11 +1210,9 @@ def _flex_internal_collisions_detect( @wp.kernel(module="unique", enable_backward=False, grid_stride=False) def _flex_tet_internal_collisions_detect( # Model: - nflex: int, flex_dim: wp.array[int], flex_vertadr: wp.array[int], flex_elemadr: wp.array[int], - flex_elemnum: wp.array[int], flex_elemdataadr: wp.array[int], flex_elem: wp.array[int], flex_radius: wp.array[float], @@ -1454,13 +1261,14 @@ def _flex_tet_internal_collisions_detect( # Test face (0,1,2) vs Vertex 3 ok0, dist0, pos0, nrm0 = _plane_vertex(p3, radius, p0, p1, p2) if ok0: - _write_candidate_contact( + _write_candidate( max_candidates, dist0, pos0, nrm0, -1, flexid, + flexid, local_elemid, v3, worldid, @@ -1481,13 +1289,14 @@ def _flex_tet_internal_collisions_detect( # Test face (0,2,3) vs Vertex 1 ok1, dist1, pos1, nrm1 = _plane_vertex(p1, radius, p0, p2, p3) if ok1: - _write_candidate_contact( + _write_candidate( max_candidates, dist1, pos1, nrm1, -1, flexid, + flexid, local_elemid, v1, worldid, @@ -1508,13 +1317,14 @@ def _flex_tet_internal_collisions_detect( # Test face (0,3,1) vs Vertex 2 ok2, dist2, pos2, nrm2 = _plane_vertex(p2, radius, p0, p3, p1) if ok2: - _write_candidate_contact( + _write_candidate( max_candidates, dist2, pos2, nrm2, -1, flexid, + flexid, local_elemid, v2, worldid, @@ -1535,13 +1345,14 @@ def _flex_tet_internal_collisions_detect( # Test face (1,3,2) vs Vertex 0 ok3, dist3, pos3, nrm3 = _plane_vertex(p0, radius, p1, p3, p2) if ok3: - _write_candidate_contact( + _write_candidate( max_candidates, dist3, pos3, nrm3, -1, flexid, + flexid, local_elemid, v0, worldid, @@ -1743,7 +1554,7 @@ def _flex_sap_project( aabb_max = wp.max(aabb_max, p3) radius = flex_radius[flexid] - rbound = 2.0 * radius + rbound = radius inflate = wp.vec3(rbound, rbound, rbound) aabb_min = aabb_min - inflate aabb_max = aabb_max + inflate @@ -1757,8 +1568,8 @@ def _flex_sap_project( proj_center = wp.dot(direction, center) proj_radius = wp.abs(direction[0]) * halfsize[0] + wp.abs(direction[1]) * halfsize[1] + wp.abs(direction[2]) * halfsize[2] - # If self-collision is disabled for this flex, push to infinity - if flex_selfcollide[flexid] == 0: + # If self-collision is disabled and there are no other flexes, push to infinity + if nflex == 1 and flex_selfcollide[flexid] == 0: projection_lower_out[worldid, elemid] = MJ_MAXVAL projection_upper_out[worldid, elemid] = MJ_MAXVAL else: @@ -1772,84 +1583,200 @@ def _flex_sap_project( segmented_index_out[nworld_in] = nworld_in * nelem -@wp.kernel -def _flex_sap_filter( - # Model: - flex_selfcollide: wp.array[int], - flex_dim: wp.array[int], - flex_vertadr: wp.array[int], - flex_elemadr: wp.array[int], - flex_elemnum: wp.array[int], - flex_elemdataadr: wp.array[int], - flex_vertbodyid: wp.array[int], - flex_elem: wp.array[int], - flex_elemflexid: wp.array[int], - # In: - raw_npairs_in: wp.array[int], - raw_pair_elem1_in: wp.array[int], - raw_pair_elem2_in: wp.array[int], - raw_pair_worldid_in: wp.array[int], - # Out: - npairs_out: wp.array[int], - pair_elem1_out: wp.array[int], - pair_elem2_out: wp.array[int], - pair_worldid_out: wp.array[int], -): - """Filter raw SAP pairs for flex self-collision.""" - tid = wp.tid() - - # Skip if beyond actual pair count - n_raw = raw_npairs_in[0] - if tid >= n_raw: - return - - elem1_global = raw_pair_elem1_in[tid] - elem2_global = raw_pair_elem2_in[tid] - - # Both elements must belong to the same flex - flexid1 = flex_elemflexid[elem1_global] - flexid2 = flex_elemflexid[elem2_global] - if flexid1 != flexid2: - return - - flexid = flexid1 - if flex_selfcollide[flexid] == 0: - return +@cache_kernel +def flex_sap_sweep_builder(warn_overflow: bool): + @wp.kernel(module="unique", enable_backward=False) + def _flex_sap_sweep( + # Model: + flex_contype: wp.array[int], + flex_conaffinity: wp.array[int], + flex_selfcollide: wp.array[int], + flex_dim: wp.array[int], + flex_vertadr: wp.array[int], + flex_elemadr: wp.array[int], + flex_elemdataadr: wp.array[int], + flex_vertbodyid: wp.array[int], + flex_elem: wp.array[int], + flex_elemflexid: wp.array[int], + # In: + n: int, + sort_index_in: wp.array2d[int], + cumulative_sum_in: wp.array[int], + nsweep_in: int, + aabb_lower_in: wp.array2d[wp.vec3], + aabb_upper_in: wp.array2d[wp.vec3], + max_self_pairs: int, + max_ff_pairs: int, + # Data out: + overflow_out: wp.array[int], + # Out: + self_npairs_out: wp.array[int], + self_pair_id1_out: wp.array[int], + self_pair_id2_out: wp.array[int], + self_pair_worldid_out: wp.array[int], + ff_npairs_out: wp.array[int], + ff_pair_id1_out: wp.array[int], + ff_pair_id2_out: wp.array[int], + ff_pair_worldid_out: wp.array[int], + ): + """Generic SAP sweep: output filtered self and flex-flex pairs.""" + worldelemid = wp.tid() + + nworldelem = cumulative_sum_in.shape[0] + nworkpackages = cumulative_sum_in[nworldelem - 1] + + while worldelemid < nworkpackages: + # Binary search to find sortedid (i) and partner sortedid (j) + i = sap_binary_search(cumulative_sum_in, worldelemid, 0, nworldelem) + j = i + worldelemid + 1 + if i > 0: + j -= cumulative_sum_in[i - 1] + + worldid = i // n + i = i % n + j = j % n + + # Get actual element indices from sorted order + elem1 = sort_index_in[worldid, i] + elem2 = sort_index_in[worldid, j] + + # Skip self-collision if disabled + flexid1 = flex_elemflexid[elem1] + flexid2 = flex_elemflexid[elem2] + if flexid1 == flexid2 and flex_selfcollide[flexid1] == 0: + worldelemid += nsweep_in + continue + + # Ensure elem1 < elem2 for consistent ordering + if elem1 > elem2: + tmpelem = elem1 + elem1 = elem2 + elem2 = tmpelem + tmpid = flexid1 + flexid1 = flexid2 + flexid2 = tmpid + + worldelemid += nsweep_in + + # AABB overlap test on all 3 axes + lower1 = aabb_lower_in[worldid, elem1] + upper1 = aabb_upper_in[worldid, elem1] + lower2 = aabb_lower_in[worldid, elem2] + upper2 = aabb_upper_in[worldid, elem2] + + if lower1[0] > upper2[0] or lower2[0] > upper1[0]: + continue + if lower1[1] > upper2[1] or lower2[1] > upper1[1]: + continue + if lower1[2] > upper2[2] or lower2[2] > upper1[2]: + continue + + # Output this pair + e1 = 0 + e2 = 0 + elem_data_idx1 = 0 + elem_data_idx2 = 0 + v1_indices = wp.vec4i(-1, -1, -1, -1) + v2_indices = wp.vec4i(-1, -1, -1, -1) + + if flexid1 == flexid2: + # Self-collision: check topological adjacency + dim = flex_dim[flexid1] + vert_adr = flex_vertadr[flexid1] + elem_adr = flex_elemadr[flexid1] + + e1 = elem1 - elem_adr + e2 = elem2 - elem_adr + + elem_data_idx1 = flex_elemdataadr[flexid1] + e1 * (dim + 1) + elem_data_idx2 = flex_elemdataadr[flexid1] + e2 * (dim + 1) + + v1_indices = _get_element_vertices(flex_elem, dim, elem_data_idx1) + v2_indices = _get_element_vertices(flex_elem, dim, elem_data_idx2) + + if _exclude_self_collision(flex_vertbodyid, v1_indices, dim + 1, v2_indices, dim + 1, vert_adr): + continue - # Exclude elements sharing vertices/bodies - dim = flex_dim[flexid] - vert_adr = flex_vertadr[flexid] - elem_adr = flex_elemadr[flexid] + idx = wp.atomic_add(self_npairs_out, 0, 1) + if idx < max_self_pairs: + self_pair_id1_out[idx] = elem1 + self_pair_id2_out[idx] = elem2 + self_pair_worldid_out[idx] = worldid + else: + if wp.static(warn_overflow): + wp.printf("Self-collision SAP buffer overflow - please increase naconmax to %u\n", idx + 1) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) + else: + # Flex-flex collision: check contype/conaffinity + contype1 = flex_contype[flexid1] + conaffinity1 = flex_conaffinity[flexid1] + contype2 = flex_contype[flexid2] + conaffinity2 = flex_conaffinity[flexid2] + if not ((contype1 & conaffinity2) or (contype2 & conaffinity1)): + continue - e1 = elem1_global - elem_adr - e2 = elem2_global - elem_adr - elem_data_idx1 = flex_elemdataadr[flexid] + e1 * (dim + 1) - elem_data_idx2 = flex_elemdataadr[flexid] + e2 * (dim + 1) - v1_indices = _get_element_vertices(flex_elem, dim, elem_data_idx1) - v2_indices = _get_element_vertices(flex_elem, dim, elem_data_idx2) + # Check shared bodies + dim1 = flex_dim[flexid1] + dim2 = flex_dim[flexid2] + vert_adr1 = flex_vertadr[flexid1] + vert_adr2 = flex_vertadr[flexid2] + elem_adr1 = flex_elemadr[flexid1] + elem_adr2 = flex_elemadr[flexid2] + + e1 = elem1 - elem_adr1 + e2 = elem2 - elem_adr2 + + elem_data_idx1 = flex_elemdataadr[flexid1] + e1 * (dim1 + 1) + elem_data_idx2 = flex_elemdataadr[flexid2] + e2 * (dim2 + 1) + + v1_indices = _get_element_vertices(flex_elem, dim1, elem_data_idx1) + v2_indices = _get_element_vertices(flex_elem, dim2, elem_data_idx2) + + shared_body = bool(False) + for ii in range(dim1 + 1): + idx1 = v1_indices[ii] + if idx1 >= 0: + b1 = flex_vertbodyid[vert_adr1 + idx1] + for jj in range(dim2 + 1): + idx2 = v2_indices[jj] + if idx2 >= 0 and b1 >= 0: + b2 = flex_vertbodyid[vert_adr2 + idx2] + if b1 == b2: + shared_body = bool(True) + break + if shared_body: + break + + if shared_body: + continue - if _exclude_self_collision(flex_vertbodyid, v1_indices, dim + 1, v2_indices, dim + 1, vert_adr): - return + idx = wp.atomic_add(ff_npairs_out, 0, 1) + if idx < max_ff_pairs: + ff_pair_id1_out[idx] = elem1 + ff_pair_id2_out[idx] = elem2 + ff_pair_worldid_out[idx] = worldid + else: + if wp.static(warn_overflow): + wp.printf("Flex-flex SAP buffer overflow - please increase naconmax to %u\n", idx + 1) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.BROADPHASE)) - # Output this pair - idx = wp.atomic_add(npairs_out, 0, 1) - if idx < pair_elem1_out.shape[0]: - pair_elem1_out[idx] = elem1_global - pair_elem2_out[idx] = elem2_global - pair_worldid_out[idx] = raw_pair_worldid_in[tid] + return _flex_sap_sweep @wp.kernel -def _flex_selfcollision_narrowphase( +def _flex_narrowphase( # Model: - nflex: int, opt_ccd_tolerance: wp.array[float], opt_warn_overflow: bool, + flex_margin: wp.array[float], + flex_gap: wp.array[float], + flex_activelayers: wp.array[int], flex_dim: wp.array[int], flex_vertadr: wp.array[int], flex_elemadr: wp.array[int], flex_elemdataadr: wp.array[int], flex_elem: wp.array[int], + flex_elemlayer: wp.array[int], flex_radius: wp.array[float], flex_elemflexid: wp.array[int], # Data in: @@ -1863,7 +1790,6 @@ def _flex_selfcollision_narrowphase( pair_elem2_in: wp.array[int], pair_worldid_in: wp.array[int], max_pairs: int, - n_total_elems: int, # Data out: overflow_out: wp.array[int], # Out: @@ -1886,7 +1812,7 @@ def _flex_selfcollision_narrowphase( cand_geomcollisionid_out: wp.array[int], ncand_out: wp.array[int], ): - """Process SAP-identified pairs through narrowphase (GJK/EPA).""" + """Process broadphase-identified flex pairs through narrowphase.""" pairid = wp.tid() # Check bounds @@ -1898,26 +1824,41 @@ def _flex_selfcollision_narrowphase( elem2_global = pair_elem2_in[pairid] worldid = pair_worldid_in[pairid] - flexid = flex_elemflexid[elem1_global] - radius = flex_radius[flexid] - dim = flex_dim[flexid] - vert_adr = flex_vertadr[flexid] - elem_adr = flex_elemadr[flexid] + flexid1 = flex_elemflexid[elem1_global] + flexid2 = flex_elemflexid[elem2_global] - e1 = elem1_global - elem_adr - e2 = elem2_global - elem_adr - elem_data_idx1 = flex_elemdataadr[flexid] + e1 * (dim + 1) - elem_data_idx2 = flex_elemdataadr[flexid] + e2 * (dim + 1) + radius1 = flex_radius[flexid1] + radius2 = flex_radius[flexid2] + dim1 = flex_dim[flexid1] + dim2 = flex_dim[flexid2] + vert_adr1 = flex_vertadr[flexid1] + vert_adr2 = flex_vertadr[flexid2] + elem_adr1 = flex_elemadr[flexid1] + elem_adr2 = flex_elemadr[flexid2] - v1_indices = _get_element_vertices(flex_elem, dim, elem_data_idx1) - v2_indices = _get_element_vertices(flex_elem, dim, elem_data_idx2) + e1 = elem1_global - elem_adr1 + e2 = elem2_global - elem_adr2 + + if not _elem_active(flex_activelayers, flex_dim, flex_elemadr, flex_elemlayer, flexid1, e1): + return + if not _elem_active(flex_activelayers, flex_dim, flex_elemadr, flex_elemlayer, flexid2, e2): + return + elem_data_idx1 = flex_elemdataadr[flexid1] + e1 * (dim1 + 1) + elem_data_idx2 = flex_elemdataadr[flexid2] + e2 * (dim2 + 1) + + v1_indices = _get_element_vertices(flex_elem, dim1, elem_data_idx1) + v2_indices = _get_element_vertices(flex_elem, dim2, elem_data_idx2) - # Workspace for this pair + # Workspace for this pair (8 slots per pair) offset1 = pairid * 8 - for idx in range(dim + 1): - workspace_verts_out[offset1 + idx] = flexvert_xpos_in[worldid, vert_adr + v1_indices[idx]] + for idx in range(dim1 + 1): + workspace_verts_out[offset1 + idx] = flexvert_xpos_in[worldid, vert_adr1 + v1_indices[idx]] - if dim == 1: + mix_margin = 0.0 + if flexid1 != flexid2: + mix_margin = flex_margin[flexid1] + flex_margin[flexid2] + flex_gap[flexid1] + flex_gap[flexid2] + + if dim1 == 1 and dim2 == 1: # Capsule-capsule collision p0 = workspace_verts_out[offset1] p1 = workspace_verts_out[offset1 + 1] @@ -1925,28 +1866,27 @@ def _flex_selfcollision_narrowphase( cap1_axis = wp.normalize(p1 - p0) cap1_half_len = 0.5 * wp.length(p1 - p0) - p2_0 = flexvert_xpos_in[worldid, vert_adr + v2_indices[0]] - p2_1 = flexvert_xpos_in[worldid, vert_adr + v2_indices[1]] + p2_0 = flexvert_xpos_in[worldid, vert_adr2 + v2_indices[0]] + p2_1 = flexvert_xpos_in[worldid, vert_adr2 + v2_indices[1]] cap2_pos = 0.5 * (p2_0 + p2_1) cap2_axis = wp.normalize(p2_1 - p2_0) cap2_half_len = 0.5 * wp.length(p2_1 - p2_0) - margin = 0.0 - contact_dist, contact_pos, contact_normal = collision_primitive_core.capsule_capsule( - cap1_pos, cap1_axis, radius, cap1_half_len, cap2_pos, cap2_axis, radius, cap2_half_len, margin + cap1_pos, cap1_axis, radius1, cap1_half_len, cap2_pos, cap2_axis, radius2, cap2_half_len, mix_margin ) for c in range(2): d_val = contact_dist[c] if d_val < 0.0: - _write_candidate_contact( + _write_candidate( max_candidates, d_val, contact_pos[c], contact_normal[c], -2, - flexid, + flexid1, + flexid2, e1, e2, worldid, @@ -1964,19 +1904,19 @@ def _flex_selfcollision_narrowphase( ncand_out, ) else: - # GJK/EPA for dim >= 2 + # GJK/EPA narrowphase offset2 = pairid * 8 + 4 - for idx in range(dim + 1): - workspace_verts_out[offset2 + idx] = flexvert_xpos_in[worldid, vert_adr + v2_indices[idx]] + for idx in range(dim2 + 1): + workspace_verts_out[offset2 + idx] = flexvert_xpos_in[worldid, vert_adr2 + v2_indices[idx]] geom1 = Geom() geom1.pos = wp.vec3(0.0) geom1.rot = wp.identity(n=3, dtype=float) geom1.size = wp.vec3(0.0) - geom1.margin = 2.0 * radius + geom1.margin = 2.0 * radius1 + mix_margin geom1.vert = workspace_verts_out geom1.vertadr = offset1 - geom1.vertnum = dim + 1 + geom1.vertnum = dim1 + 1 geom1.graphadr = -1 geom1.index = -1 @@ -1984,28 +1924,28 @@ def _flex_selfcollision_narrowphase( geom2.pos = wp.vec3(0.0) geom2.rot = wp.identity(n=3, dtype=float) geom2.size = wp.vec3(0.0) - geom2.margin = 2.0 * radius + geom2.margin = 2.0 * radius2 + mix_margin geom2.vert = workspace_verts_out geom2.vertadr = offset2 - geom2.vertnum = dim + 1 + geom2.vertnum = dim2 + 1 geom2.graphadr = -1 geom2.index = -1 center1 = wp.vec3(0.0) - for idx in range(dim + 1): + for idx in range(dim1 + 1): center1 += workspace_verts_out[offset1 + idx] - center1 = center1 / float(dim + 1) + center1 = center1 / float(dim1 + 1) center2 = wp.vec3(0.0) - for idx in range(dim + 1): + for idx in range(dim2 + 1): center2 += workspace_verts_out[offset2 + idx] - center2 = center2 / float(dim + 1) + center2 = center2 / float(dim2 + 1) tol = opt_ccd_tolerance[0 % opt_ccd_tolerance.shape[0]] dist, ncontact, w1, w2, _ = ccd( tol, - 2.0 * radius, + radius1 + radius2 + mix_margin, gjk_iterations, epa_iterations, geom1, @@ -2025,26 +1965,18 @@ def _flex_selfcollision_narrowphase( overflow_out, ) - phys_dist = dist - if ncontact > 0 and phys_dist < 0.0: - p1_0 = workspace_verts_out[offset1] - p1_1 = workspace_verts_out[offset1 + 1] - p1_2 = workspace_verts_out[offset1 + 2] - p2_0 = workspace_verts_out[offset2] - p2_1 = workspace_verts_out[offset2 + 1] - p2_2 = workspace_verts_out[offset2 + 2] - if not (_inside_triangle(w1, p1_0, p1_1, p1_2, 0.2) and _inside_triangle(w2, p2_0, p2_1, p2_2, 0.2)): - return - + if ncontact > 0 and dist < 0.0: + phys_dist = dist + mix_margin pos = 0.5 * (w1 + w2) nrm = wp.normalize(w1 - w2) - _write_candidate_contact( + _write_candidate( max_candidates, phys_dist, pos, nrm, -2, - flexid, + flexid1, + flexid2, e1, e2, worldid, @@ -2063,13 +1995,29 @@ def _flex_selfcollision_narrowphase( ) +@wp.func +def _elem_active( + # Model: + flex_activelayers: wp.array[int], + flex_dim: wp.array[int], + flex_elemadr: wp.array[int], + flex_elemlayer: wp.array[int], + # In: + flexid: int, + elemid: int, +) -> bool: + if flex_dim[flexid] < 3: + return True + return flex_elemlayer[flex_elemadr[flexid] + elemid] < flex_activelayers[flexid] + + @wp.kernel(module="unique", enable_backward=False, grid_stride=False) def _flex_active_element_collisions_detect( # Model: - nflex: int, opt_ccd_tolerance: wp.array[float], opt_warn_overflow: bool, flex_selfcollide: wp.array[int], + flex_activelayers: wp.array[int], flex_dim: wp.array[int], flex_vertadr: wp.array[int], flex_elemadr: wp.array[int], @@ -2077,6 +2025,7 @@ def _flex_active_element_collisions_detect( flex_elemdataadr: wp.array[int], flex_vertbodyid: wp.array[int], flex_elem: wp.array[int], + flex_elemlayer: wp.array[int], flex_radius: wp.array[float], flex_elemflexid: wp.array[int], # Data in: @@ -2121,6 +2070,9 @@ def _flex_active_element_collisions_detect( elem_num = flex_elemnum[flexid] e1 = elem1_global - elem_adr + if not _elem_active(flex_activelayers, flex_dim, flex_elemadr, flex_elemlayer, flexid, e1): + return + elem_data_idx1 = flex_elemdataadr[flexid] + e1 * (dim + 1) v1_indices = _get_element_vertices(flex_elem, dim, elem_data_idx1) @@ -2132,6 +2084,8 @@ def _flex_active_element_collisions_detect( workspace_verts_out[offset1 + idx] = flexvert_xpos_in[worldid, vert_adr + v1_indices[idx]] for e2 in range(e1 + 1, elem_num): + if not _elem_active(flex_activelayers, flex_dim, flex_elemadr, flex_elemlayer, flexid, e2): + continue elem_data_idx2 = flex_elemdataadr[flexid] + e2 * (dim + 1) v2_indices = _get_element_vertices(flex_elem, dim, elem_data_idx2) @@ -2164,13 +2118,14 @@ def _flex_active_element_collisions_detect( for c in range(2): d_val = contact_dist[c] if d_val < 0.0: - _write_candidate_contact( + _write_candidate( max_candidates, d_val, contact_pos[c], contact_normal[c], -2, flexid, + flexid, e1, e2, worldid, @@ -2248,8 +2203,7 @@ def _flex_active_element_collisions_detect( overflow_out, ) - phys_dist = dist - if ncontact > 0 and phys_dist < 0.0: + if ncontact > 0 and dist < 0.0: p1_0 = workspace_verts_out[offset1] p1_1 = workspace_verts_out[offset1 + 1] p1_2 = workspace_verts_out[offset1 + 2] @@ -2261,13 +2215,14 @@ def _flex_active_element_collisions_detect( pos = 0.5 * (w1 + w2) nrm = wp.normalize(w1 - w2) - _write_candidate_contact( + _write_candidate( max_candidates, - phys_dist, + dist, pos, nrm, -2, flexid, + flexid, e1, e2, worldid, @@ -2286,153 +2241,273 @@ def _flex_active_element_collisions_detect( ) -@wp.kernel -def _flex_narrowphase_unified( - # Model: - ngeom: int, - nflex: int, - opt_ccd_tolerance: wp.array[float], - opt_warn_overflow: bool, - geom_type: wp.array[int], - geom_condim: wp.array[int], - geom_dataid: wp.array2d[int], - geom_priority: wp.array[int], - geom_solmix: wp.array2d[float], - geom_solref: wp.array2d[wp.vec2], - geom_solimp: wp.array2d[vec5], - geom_size: wp.array2d[wp.vec3], - geom_friction: wp.array2d[wp.vec3], - geom_margin: wp.array2d[float], - geom_gap: wp.array2d[float], - flex_condim: wp.array[int], - flex_priority: wp.array[int], - flex_solmix: wp.array[float], - flex_solref: wp.array[wp.vec2], - flex_solimp: wp.array[vec5], - flex_friction: wp.array[wp.vec3], - flex_margin: wp.array[float], - flex_gap: wp.array[float], - flex_dim: wp.array[int], - flex_vertadr: wp.array[int], - flex_radius: wp.array[float], - mesh_vertadr: wp.array[int], - mesh_vertnum: wp.array[int], - mesh_graphadr: wp.array[int], - mesh_vert: wp.array[wp.vec3], - mesh_graph: wp.array[int], - mesh_pos: wp.array[wp.vec3], - mesh_polynormal: wp.array[wp.vec3], - mesh_polyvertadr: wp.array[int], - mesh_polyvert: wp.array[int], - mesh_polymapadr: wp.array[int], - mesh_polymapnum: wp.array[int], - mesh_polymap: wp.array[int], - # Data in: - geom_xpos_in: wp.array2d[wp.vec3], - geom_xmat_in: wp.array2d[wp.mat33], - flexvert_xpos_in: wp.array2d[wp.vec3], - nworld_in: int, - naconmax_in: int, - naccdmax_in: int, - ncollision_in: wp.array[int], - # In: - triadr: wp.array[int], - flex_tridataadr: wp.array[int], - flex_tri: wp.array[int], - flex_triflexid: wp.array[int], - collision_pair_in: wp.array[wp.vec2i], - collision_worldid_in: wp.array[int], - epa_vert: wp.array2d[wp.vec3], - epa_vert_index: wp.array2d[int], - epa_face: wp.array2d[int], - epa_pr: wp.array2d[wp.vec3], - epa_norm2: wp.array2d[float], - epa_horizon: wp.array2d[int], - nccd: wp.array[int], - ccd_iterations: int, - max_candidates: int, - # Data out: - overflow_out: wp.array[int], - # Out: - cand_dist_out: wp.array[float], - cand_pos_out: wp.array[wp.vec3], - cand_nrm_out: wp.array[wp.vec3], - cand_geom_out: wp.array[wp.vec2i], - cand_flex_out: wp.array[wp.vec2i], - cand_elem_out: wp.array[wp.vec2i], - cand_vert_out: wp.array[wp.vec2i], - cand_worldid_out: wp.array[int], - cand_type_out: wp.array[int], - cand_geomcollisionid_out: wp.array[int], - ncand_out: wp.array[int], -): - collisionid = wp.tid() - if collisionid >= ncollision_in[0] or collisionid >= collision_pair_in.shape[0]: - return +@cache_kernel +def flex_narrowphase_unified_builder(warn_overflow: bool): + @wp.kernel(module="unique", enable_backward=False) + def _flex_narrowphase_unified( + # Model: + opt_ccd_tolerance: wp.array[float], + opt_warn_overflow: bool, + geom_type: wp.array[int], + geom_dataid: wp.array2d[int], + geom_size: wp.array2d[wp.vec3], + geom_margin: wp.array2d[float], + flex_margin: wp.array[float], + flex_vertadr: wp.array[int], + flex_radius: wp.array[float], + mesh_vertadr: wp.array[int], + mesh_vertnum: wp.array[int], + mesh_graphadr: wp.array[int], + mesh_vert: wp.array[wp.vec3], + mesh_graph: wp.array[int], + mesh_pos: wp.array[wp.vec3], + mesh_polynormal: wp.array[wp.vec3], + mesh_polyvertadr: wp.array[int], + mesh_polyvert: wp.array[int], + mesh_polymapadr: wp.array[int], + mesh_polymapnum: wp.array[int], + mesh_polymap: wp.array[int], + # Data in: + geom_xpos_in: wp.array2d[wp.vec3], + geom_xmat_in: wp.array2d[wp.mat33], + flexvert_xpos_in: wp.array2d[wp.vec3], + naccdmax_in: int, + ncollision_in: wp.array[int], + # In: + triadr: wp.array[int], + flex_tridataadr: wp.array[int], + flex_tri: wp.array[int], + flex_triflexid: wp.array[int], + collision_pair_in: wp.array[wp.vec2i], + collision_worldid_in: wp.array[int], + epa_vert: wp.array2d[wp.vec3], + epa_vert_index: wp.array2d[int], + epa_face: wp.array2d[int], + epa_pr: wp.array2d[wp.vec3], + epa_norm2: wp.array2d[float], + epa_horizon: wp.array2d[int], + nccd: wp.array[int], + ccd_iterations: int, + max_candidates: int, + # Data out: + overflow_out: wp.array[int], + # Out: + cand_dist_out: wp.array[float], + cand_pos_out: wp.array[wp.vec3], + cand_nrm_out: wp.array[wp.vec3], + cand_geom_out: wp.array[wp.vec2i], + cand_flex_out: wp.array[wp.vec2i], + cand_elem_out: wp.array[wp.vec2i], + cand_vert_out: wp.array[wp.vec2i], + cand_worldid_out: wp.array[int], + cand_type_out: wp.array[int], + cand_geomcollisionid_out: wp.array[int], + ncand_out: wp.array[int], + ): + collisionid = wp.tid() + if collisionid >= ncollision_in[0] or collisionid >= collision_pair_in.shape[0]: + return - pair = collision_pair_in[collisionid] - tri_id = pair[0] - geomid = pair[1] + pair = collision_pair_in[collisionid] + tri_id = pair[0] + geomid = pair[1] - gtype = geom_type[geomid] - if ( - gtype != int(GeomType.SPHERE) - and gtype != int(GeomType.CAPSULE) - and gtype != int(GeomType.BOX) - and gtype != int(GeomType.CYLINDER) - and gtype != int(GeomType.MESH) - and gtype != int(GeomType.ELLIPSOID) - ): - return + gtype = geom_type[geomid] + if ( + gtype != int(GeomType.SPHERE) + and gtype != int(GeomType.CAPSULE) + and gtype != int(GeomType.BOX) + and gtype != int(GeomType.CYLINDER) + and gtype != int(GeomType.MESH) + and gtype != int(GeomType.ELLIPSOID) + ): + return - worldid = collision_worldid_in[collisionid] + worldid = collision_worldid_in[collisionid] - flexid = flex_triflexid[tri_id] + flexid = flex_triflexid[tri_id] - vert_adr = flex_vertadr[flexid] - tri_radius = flex_radius[flexid] - tri_margin = flex_margin[flexid] + vert_adr = flex_vertadr[flexid] + tri_radius = flex_radius[flexid] + tri_margin = flex_margin[flexid] + + local_tri_id = tri_id - triadr[flexid] + tri_data_idx = flex_tridataadr[flexid] + local_tri_id * 3 + v0_local = flex_tri[tri_data_idx] + v1_local = flex_tri[tri_data_idx + 1] + v2_local = flex_tri[tri_data_idx + 2] + + t1 = flexvert_xpos_in[worldid, vert_adr + v0_local] + t2 = flexvert_xpos_in[worldid, vert_adr + v1_local] + t3 = flexvert_xpos_in[worldid, vert_adr + v2_local] + + geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] + margin = geom_margin_val + tri_margin + + geom_pos = geom_xpos_in[worldid, geomid] + geom_rot = geom_xmat_in[worldid, geomid] + geom_size_val = geom_size[worldid % geom_size.shape[0], geomid] + + if gtype == int(GeomType.MESH): + ccdid = wp.atomic_add(nccd, 0, 1) + if ccdid >= naccdmax_in: + if wp.static(warn_overflow): + wp.printf("CCD overflow in flex narrowphase - please increase naccdmax to %u\n", ccdid) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.CCD)) + else: + did = geom_dataid[worldid % geom_dataid.shape[0], geomid] + tolerance = opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]] + _collide_mesh_triangle( + mesh_vertadr, + mesh_vertnum, + mesh_graphadr, + mesh_vert, + mesh_graph, + mesh_pos, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + max_candidates, + geom_pos, + geom_rot, + geom_size_val, + t1, + t2, + t3, + tri_radius, + margin, + geomid, + flexid, + local_tri_id, + worldid, + did, + epa_vert[ccdid], + epa_vert_index[ccdid], + epa_face[ccdid], + epa_pr[ccdid], + epa_norm2[ccdid], + epa_horizon[ccdid], + tolerance, + ccd_iterations, + warn_overflow, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) + elif gtype == int(GeomType.ELLIPSOID): + ccdid = wp.atomic_add(nccd, 0, 1) + if ccdid >= naccdmax_in: + if wp.static(warn_overflow): + wp.printf("CCD overflow in flex narrowphase - please increase naccdmax to %u\n", ccdid) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.CCD)) + else: + tolerance = opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]] + + # Construct Ellipsoid Geom (geom1) + geom1 = Geom() + geom1.pos = geom_pos + geom1.rot = geom_rot + geom1.size = geom_size_val + geom1.margin = 0.0 + geom1.index = -1 + + # Construct Triangle Geom (geom2) + geom2 = Geom() + geom2.pos = wp.vec3(0.0, 0.0, 0.0) + geom2.rot = wp.mat33(t1[0], t1[1], t1[2], t2[0], t2[1], t2[2], t3[0], t3[1], t3[2]) + geom2.margin = 0.0 + geom2.index = -1 + + centroid = (t1 + t2 + t3) * wp.static(1.0 / 3.0) + r_geom = wp.length(geom_size_val) + d1 = wp.length(t1 - centroid) + d2 = wp.length(t2 - centroid) + d3 = wp.length(t3 - centroid) + r_tri = wp.max(d1, wp.max(d2, d3)) + + if wp.length(centroid - geom_pos) <= r_geom + r_tri + margin + tri_radius + 0.04: + dist, ncontact, w1, w2, idx = ccd( + tolerance, + margin + tri_radius, + ccd_iterations, + ccd_iterations, + geom1, + geom2, + int(GeomType.ELLIPSOID), + int(GeomType.TRIANGLE), + geom_pos, + centroid, + epa_vert[ccdid], + epa_vert_index[ccdid], + epa_face[ccdid], + epa_pr[ccdid], + epa_norm2[ccdid], + epa_horizon[ccdid], + warn_overflow, + worldid, + overflow_out, + ) - local_tri_id = tri_id - triadr[flexid] - tri_data_idx = flex_tridataadr[flexid] + local_tri_id * 3 - v0_local = flex_tri[tri_data_idx] - v1_local = flex_tri[tri_data_idx + 1] - v2_local = flex_tri[tri_data_idx + 2] - - t1 = flexvert_xpos_in[worldid, vert_adr + v0_local] - t2 = flexvert_xpos_in[worldid, vert_adr + v1_local] - t3 = flexvert_xpos_in[worldid, vert_adr + v2_local] - - geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] - margin = geom_margin_val + tri_margin - - geom_pos = geom_xpos_in[worldid, geomid] - geom_rot = geom_xmat_in[worldid, geomid] - geom_size_val = geom_size[worldid % geom_size.shape[0], geomid] - - if gtype == int(GeomType.MESH): - ccdid = wp.atomic_add(nccd, 0, 1) - if ccdid >= naccdmax_in: - if opt_warn_overflow: - wp.printf("CCD overflow in flex narrowphase - please increase naccdmax to %u\n", ccdid) - wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.CCD)) + if ncontact > 0 and dist < margin + tri_radius: + if _inside_triangle(w2, t1, t2, t3, 0.2): + if dist < 0.0: + normal = wp.normalize(w1 - w2) + else: + normal = wp.normalize(w2 - w1) + + # Project triangle vertices onto normal to find deepest penetration + dist_v0 = wp.dot(t1 - w1, normal) - tri_radius + dist_v1 = wp.dot(t2 - w1, normal) - tri_radius + dist_v2 = wp.dot(t3 - w1, normal) - tri_radius + + min_dist = wp.min(dist_v0, wp.min(dist_v1, dist_v2)) + if min_dist < margin: + pos = t1 - normal * (tri_radius + 0.5 * dist_v0) + if dist_v1 < dist_v0 and dist_v1 < dist_v2: + pos = t2 - normal * (tri_radius + 0.5 * dist_v1) + elif dist_v2 < dist_v0 and dist_v2 < dist_v1: + pos = t3 - normal * (tri_radius + 0.5 * dist_v2) + + _write_candidate( + max_candidates, + min_dist, + pos, + normal, + geomid, + -1, + flexid, + local_tri_id, + -1, + worldid, + overflow_out, + cand_dist_out, + cand_pos_out, + cand_nrm_out, + cand_geom_out, + cand_flex_out, + cand_elem_out, + cand_vert_out, + cand_worldid_out, + cand_type_out, + cand_geomcollisionid_out, + ncand_out, + ) else: - did = geom_dataid[worldid % geom_dataid.shape[0], geomid] - tolerance = opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]] - _collide_mesh_triangle( - mesh_vertadr, - mesh_vertnum, - mesh_graphadr, - mesh_vert, - mesh_graph, - mesh_pos, - mesh_polynormal, - mesh_polyvertadr, - mesh_polyvert, - mesh_polymapadr, - mesh_polymapnum, - mesh_polymap, + _collide_geom_triangle_detect( max_candidates, + gtype, geom_pos, geom_rot, geom_size_val, @@ -2444,20 +2519,8 @@ def _flex_narrowphase_unified( geomid, flexid, local_tri_id, - v0_local, - v1_local, - v2_local, + -1, worldid, - did, - epa_vert[ccdid], - epa_vert_index[ccdid], - epa_face[ccdid], - epa_pr[ccdid], - epa_norm2[ccdid], - epa_horizon[ccdid], - tolerance, - ccd_iterations, - opt_warn_overflow, overflow_out, cand_dist_out, cand_pos_out, @@ -2471,136 +2534,8 @@ def _flex_narrowphase_unified( cand_geomcollisionid_out, ncand_out, ) - elif gtype == int(GeomType.ELLIPSOID): - ccdid = wp.atomic_add(nccd, 0, 1) - if ccdid >= naccdmax_in: - if opt_warn_overflow: - wp.printf("CCD overflow in flex narrowphase - please increase naccdmax to %u\n", ccdid) - wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.CCD)) - else: - tolerance = opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]] - - # Construct Ellipsoid Geom (geom1) - geom1 = Geom() - geom1.pos = geom_pos - geom1.rot = geom_rot - geom1.size = geom_size_val - geom1.margin = 0.0 - geom1.index = -1 - - # Construct Triangle Geom (geom2) - geom2 = Geom() - geom2.pos = wp.vec3(0.0, 0.0, 0.0) - geom2.rot = wp.mat33(t1[0], t1[1], t1[2], t2[0], t2[1], t2[2], t3[0], t3[1], t3[2]) - geom2.margin = 0.0 - geom2.index = -1 - - centroid = (t1 + t2 + t3) * wp.static(1.0 / 3.0) - r_geom = wp.length(geom_size_val) - d1 = wp.length(t1 - centroid) - d2 = wp.length(t2 - centroid) - d3 = wp.length(t3 - centroid) - r_tri = wp.max(d1, wp.max(d2, d3)) - - if wp.length(centroid - geom_pos) <= r_geom + r_tri + margin + tri_radius + 0.04: - dist, ncontact, w1, w2, idx = ccd( - tolerance, - margin + tri_radius, - ccd_iterations, - ccd_iterations, - geom1, - geom2, - int(GeomType.ELLIPSOID), - int(GeomType.TRIANGLE), - geom_pos, - centroid, - epa_vert[ccdid], - epa_vert_index[ccdid], - epa_face[ccdid], - epa_pr[ccdid], - epa_norm2[ccdid], - epa_horizon[ccdid], - opt_warn_overflow, - worldid, - overflow_out, - ) - if ncontact > 0 and dist < margin + tri_radius: - if _inside_triangle(w2, t1, t2, t3, 0.2): - if dist < 0.0: - normal = wp.normalize(w1 - w2) - else: - normal = wp.normalize(w2 - w1) - - # Project triangle vertices onto normal to find deepest penetration - dist_v0 = wp.dot(t1 - w1, normal) - tri_radius - dist_v1 = wp.dot(t2 - w1, normal) - tri_radius - dist_v2 = wp.dot(t3 - w1, normal) - tri_radius - - min_dist = wp.min(dist_v0, wp.min(dist_v1, dist_v2)) - if min_dist < margin: - deepest_vert = v0_local - pos = t1 - normal * (tri_radius + 0.5 * dist_v0) - if dist_v1 < dist_v0 and dist_v1 < dist_v2: - deepest_vert = v1_local - pos = t2 - normal * (tri_radius + 0.5 * dist_v1) - elif dist_v2 < dist_v0 and dist_v2 < dist_v1: - deepest_vert = v2_local - pos = t3 - normal * (tri_radius + 0.5 * dist_v2) - - _write_candidate_contact( - max_candidates, - min_dist, - pos, - normal, - geomid, - flexid, - local_tri_id, - -1, - worldid, - overflow_out, - cand_dist_out, - cand_pos_out, - cand_nrm_out, - cand_geom_out, - cand_flex_out, - cand_elem_out, - cand_vert_out, - cand_worldid_out, - cand_type_out, - cand_geomcollisionid_out, - ncand_out, - ) - else: - _collide_geom_triangle_detect( - max_candidates, - gtype, - geom_pos, - geom_rot, - geom_size_val, - t1, - t2, - t3, - tri_radius, - margin, - geomid, - flexid, - local_tri_id, - -1, - worldid, - overflow_out, - cand_dist_out, - cand_pos_out, - cand_nrm_out, - cand_geom_out, - cand_flex_out, - cand_elem_out, - cand_vert_out, - cand_worldid_out, - cand_type_out, - cand_geomcollisionid_out, - ncand_out, - ) + return _flex_narrowphase_unified @wp.kernel @@ -2627,7 +2562,6 @@ def _flex_narrowphase_tet_detect( geom_xpos_in: wp.array2d[wp.vec3], geom_xmat_in: wp.array2d[wp.mat33], flexvert_xpos_in: wp.array2d[wp.vec3], - nworld_in: int, # In: max_candidates: int, # Data out: @@ -2768,9 +2702,10 @@ def _compute_filter_key( ncand: wp.array[int], cand_geom: wp.array[wp.vec2i], cand_flex: wp.array[wp.vec2i], + cand_elem: wp.array[wp.vec2i], cand_worldid: wp.array[int], # Out: - key_out: wp.array[int], + key_out: wp.array[wp.int64], val_out: wp.array[int], ): """Compute sort key for candidate grouping. @@ -2781,20 +2716,32 @@ def _compute_filter_key( """ i = wp.tid() if i >= ncand[0]: - key_out[i] = 2147483647 # INT_MAX: sort unused entries to end + key_out[i] = wp.int64(9223372036854775807) # sort unused entries to end val_out[i] = i return worldid = cand_worldid[i] flex_id = cand_flex[i][1] geom_id = cand_geom[i][0] + elem1 = cand_elem[i][0] + elem2 = cand_elem[i][1] + flex1 = cand_flex[i][0] + + if geom_id >= 0: + group = geom_id * nflex + flex_id + else: + f1 = wp.min(flex1, flex_id) + f2 = wp.max(flex1, flex_id) + group = ngeom * nflex + f1 * nflex + f2 + + group_id = worldid * (ngeom * nflex + nflex * nflex) + group + is_self = int(flex1 == flex_id) if geom_id < 0 else 0 - # Map self-collision (geom_id < 0) to sentinel - g = geom_id - if g < 0: - g = ngeom + # High 32 bits: group identifier and self-collision flag; Low 32 bits: element pair + group_key = (wp.int64(group_id) << wp.int64(1)) | wp.int64(is_self) + elem_key = (wp.int64(elem1) << wp.int64(16)) | wp.int64(elem2) - key_out[i] = worldid * (nflex + 1) * (ngeom + 2) + flex_id * (ngeom + 2) + g + key_out[i] = (group_key << wp.int64(32)) | elem_key val_out[i] = i @@ -2803,10 +2750,11 @@ def _filter_flex_candidates_sorted( # In: ncand: wp.array[int], epsilon: float, - sort_key: wp.array[int], + sort_key: wp.array[wp.int64], sort_val: wp.array[int], cand_dist: wp.array[float], cand_pos: wp.array[wp.vec3], + cand_elem: wp.array[wp.vec2i], # Out: cand_active_out: wp.array[int], ): @@ -2817,21 +2765,26 @@ def _filter_flex_candidates_sorted( complexity from O(n^2) to O(n * k) where k is the average group size. """ si = wp.tid() - if si >= ncand[0]: + ncand_limit = wp.min(ncand[0], sort_key.shape[0]) + if si >= ncand_limit: return i = sort_val[si] my_key = sort_key[si] + my_group = my_key >> wp.int64(32) pos_i = cand_pos[i] dist_i = cand_dist[i] eps2 = epsilon * epsilon + elem1 = cand_elem[i][0] + elem2 = cand_elem[i][1] + keep = int(1) # Compare with same-key neighbors (backward) j = si - 1 while j >= 0: - if sort_key[j] != my_key: + if (sort_key[j] >> wp.int64(32)) != my_group: break oj = sort_val[j] diff = pos_i - cand_pos[oj] @@ -2839,14 +2792,27 @@ def _filter_flex_candidates_sorted( dist_j = cand_dist[oj] if dist_j < dist_i: keep = 0 - elif dist_j == dist_i and oj < i: - keep = 0 + elif dist_j == dist_i: + elem1_j = cand_elem[oj][0] + elem2_j = cand_elem[oj][1] + if elem1_j < elem1: + keep = 0 + elif elem1_j == elem1 and elem2_j < elem2: + keep = 0 + elif elem1_j == elem1 and elem2_j == elem2: + pos_j = cand_pos[oj] + if pos_j[0] < pos_i[0]: + keep = 0 + elif pos_j[0] == pos_i[0] and pos_j[1] < pos_i[1]: + keep = 0 + elif pos_j[0] == pos_i[0] and pos_j[1] == pos_i[1] and pos_j[2] < pos_i[2]: + keep = 0 j -= 1 # Compare with same-key neighbors (forward) j = si + 1 - while j < ncand[0]: - if sort_key[j] != my_key: + while j < ncand_limit: + if (sort_key[j] >> wp.int64(32)) != my_group: break oj = sort_val[j] diff = pos_i - cand_pos[oj] @@ -2854,223 +2820,398 @@ def _filter_flex_candidates_sorted( dist_j = cand_dist[oj] if dist_j < dist_i: keep = 0 - elif dist_j == dist_i and oj < i: - keep = 0 + elif dist_j == dist_i: + elem1_j = cand_elem[oj][0] + elem2_j = cand_elem[oj][1] + if elem1_j < elem1: + keep = 0 + elif elem1_j == elem1 and elem2_j < elem2: + keep = 0 + elif elem1_j == elem1 and elem2_j == elem2: + pos_j = cand_pos[oj] + if pos_j[0] < pos_i[0]: + keep = 0 + elif pos_j[0] == pos_i[0] and pos_j[1] < pos_i[1]: + keep = 0 + elif pos_j[0] == pos_i[0] and pos_j[1] == pos_i[1] and pos_j[2] < pos_i[2]: + keep = 0 j += 1 cand_active_out[i] = keep +@cache_kernel +def _write_filtered_contacts_builder(warn_overflow: bool): + @wp.kernel(module="unique", enable_backward=False) + def _write_filtered_contacts( + # Model: + geom_type: wp.array[int], + geom_condim: wp.array[int], + geom_priority: wp.array[int], + geom_solmix: wp.array2d[float], + geom_solref: wp.array2d[wp.vec2], + geom_solimp: wp.array2d[vec5], + geom_friction: wp.array2d[wp.vec3], + geom_margin: wp.array2d[float], + geom_gap: wp.array2d[float], + flex_condim: wp.array[int], + flex_priority: wp.array[int], + flex_solmix: wp.array[float], + flex_solref: wp.array[wp.vec2], + flex_solimp: wp.array[vec5], + flex_friction: wp.array[wp.vec3], + flex_margin: wp.array[float], + flex_gap: wp.array[float], + flex_dim: wp.array[int], + # Data in: + naconmax_in: int, + # In: + ncand: wp.array[int], + cand_dist: wp.array[float], + cand_pos: wp.array[wp.vec3], + cand_nrm: wp.array[wp.vec3], + cand_geom: wp.array[wp.vec2i], + cand_flex: wp.array[wp.vec2i], + cand_elem: wp.array[wp.vec2i], + cand_vert: wp.array[wp.vec2i], + cand_worldid: wp.array[int], + cand_type: wp.array[int], + cand_geomcollisionid: wp.array[int], + cand_active: wp.array[int], + # Data out: + contact_dist_out: wp.array[float], + contact_pos_out: wp.array[wp.vec3], + contact_frame_out: wp.array[wp.mat33], + contact_includemargin_out: wp.array[float], + contact_friction_out: wp.array[vec5], + contact_solref_out: wp.array[wp.vec2], + contact_solreffriction_out: wp.array[wp.vec2], + contact_solimp_out: wp.array[vec5], + contact_dim_out: wp.array[int], + contact_geom_out: wp.array[wp.vec2i], + contact_flex_out: wp.array[wp.vec2i], + contact_elem_out: wp.array[wp.vec2i], + contact_vert_out: wp.array[wp.vec2i], + contact_worldid_out: wp.array[int], + contact_type_out: wp.array[int], + contact_geomcollisionid_out: wp.array[int], + nacon_out: wp.array[int], + # Data out: + overflow_out: wp.array[int], + ): + i = wp.tid() + if i >= ncand[0]: + return + + if cand_active[i] == 0: + return + + geomid = cand_geom[i][0] + worldid = cand_worldid[i] + + condim = int(0) + margin = float(0.0) + gap = float(0.0) + solref = wp.vec2(0.0, 0.0) + solimp = vec5(0.0, 0.0, 0.0, 0.0, 0.0) + friction = vec5(0.0, 0.0, 0.0, 0.0, 0.0) + + if geomid >= 0: + flexid = cand_flex[i][1] + geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] + tri_margin = flex_margin[flexid] + margin = geom_margin_val + tri_margin + + condim, gap, solref, solimp, friction = _mix_flex_contact_params( + geom_condim[geomid], + geom_priority[geomid], + geom_solmix[worldid % geom_solmix.shape[0], geomid], + geom_solref[worldid % geom_solref.shape[0], geomid], + geom_solimp[worldid % geom_solimp.shape[0], geomid], + geom_friction[worldid % geom_friction.shape[0], geomid], + geom_gap[worldid % geom_gap.shape[0], geomid], + flex_condim[flexid], + flex_priority[flexid], + flex_solmix[flexid], + flex_solref[flexid], + flex_solimp[flexid], + flex_friction[flexid], + flex_gap[flexid], + ) + else: + flex1 = cand_flex[i][0] + flex2 = cand_flex[i][1] + is_self = flex1 == flex2 + if is_self: + margin = float(0.0) + gap1 = float(0.0) + gap2 = float(0.0) + else: + margin = flex_margin[flex1] + flex_margin[flex2] + gap1 = flex_gap[flex1] + gap2 = flex_gap[flex2] + mixed_condim, gap, solref, solimp, friction = _mix_flex_contact_params( + flex_condim[flex1], + flex_priority[flex1], + flex_solmix[flex1], + flex_solref[flex1], + flex_solimp[flex1], + flex_friction[flex1], + gap1, + flex_condim[flex2], + flex_priority[flex2], + flex_solmix[flex2], + flex_solref[flex2], + flex_solimp[flex2], + flex_friction[flex2], + gap2, + ) + + if cand_vert[i][0] >= 0 and flex_dim[flex1] == 3: + condim = 1 + else: + condim = mixed_condim + + if cand_dist[i] >= margin: + return + + id_ = wp.atomic_add(nacon_out, 0, 1) + if id_ >= naconmax_in: + if wp.static(warn_overflow): + wp.printf( + "flex contact overflow - please increase naconmax to %u\n", + id_ + 1, + ) + wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.NARROWPHASE)) + return + + contact_dist_out[id_] = cand_dist[i] + contact_pos_out[id_] = cand_pos[i] + contact_frame_out[id_] = make_frame(cand_nrm[i]) + if geomid >= 0 and geom_type[geomid] == int(GeomType.PLANE): + contact_includemargin_out[id_] = margin - gap + else: + contact_includemargin_out[id_] = margin + contact_friction_out[id_] = friction + contact_solref_out[id_] = solref + contact_solreffriction_out[id_] = solref + contact_solimp_out[id_] = solimp + contact_dim_out[id_] = condim + contact_geom_out[id_] = cand_geom[i] + contact_flex_out[id_] = cand_flex[i] + contact_elem_out[id_] = cand_elem[i] + contact_vert_out[id_] = cand_vert[i] + contact_worldid_out[id_] = cand_worldid[i] + contact_type_out[id_] = cand_type[i] + contact_geomcollisionid_out[id_] = cand_geomcollisionid[i] + + return _write_filtered_contacts + + @wp.kernel -def _filter_flex_candidates( +def _populate_active_sorted( # In: - max_candidates: int, ncand: wp.array[int], - epsilon: float, - cand_dist: wp.array[float], - cand_pos: wp.array[wp.vec3], - cand_geom: wp.array[wp.vec2i], - cand_flex: wp.array[wp.vec2i], - cand_worldid: wp.array[int], + sort_val: wp.array[int], + cand_active: wp.array[int], # Out: - cand_active_out: wp.array[int], + cand_active_sorted_out: wp.array[int], ): - i = wp.tid() - limit = ncand[0] - if i >= limit: - return + si = wp.tid() + ncand_limit = wp.min(ncand[0], sort_val.shape[0]) + if si < ncand_limit: + cand_active_sorted_out[si] = cand_active[sort_val[si]] + else: + if si < cand_active_sorted_out.shape[0]: + cand_active_sorted_out[si] = 0 - geom_i = cand_geom[i][0] - flex_i = cand_flex[i][1] - world_i = cand_worldid[i] - pos_i = cand_pos[i] - dist_i = cand_dist[i] - keep = int(1) - for j in range(max_candidates): - if j >= limit: - break - if j == i: - continue - geom_j = cand_geom[j][0] - if (geom_i >= 0 and geom_j >= 0 and geom_j == geom_i) or (geom_i < 0 and geom_j < 0): - flex_j = cand_flex[j][1] - if flex_j == flex_i: - world_j = cand_worldid[j] - if world_j == world_i: - pos_j = cand_pos[j] - dist_j = cand_dist[j] - - diff = pos_i - pos_j - if wp.dot(diff, diff) < epsilon * epsilon: - if dist_j < dist_i: - keep = 0 - elif dist_j == dist_i and j < i: - keep = 0 +@wp.kernel +def _find_group_starts( + # In: + ncand: wp.array[int], + sort_key: wp.array[wp.int64], + # Out: + flex_group_temp_out: wp.array[int], +): + si = wp.tid() + ncand_limit = wp.min(ncand[0], sort_key.shape[0]) + if si >= ncand_limit: + return - cand_active_out[i] = keep + if si == 0: + flex_group_temp_out[si] = 1 + else: + my_key = sort_key[si] + my_group = my_key >> wp.int64(32) + prev_key = sort_key[si - 1] + prev_group = prev_key >> wp.int64(32) + if my_group != prev_group: + flex_group_temp_out[si] = 1 + else: + flex_group_temp_out[si] = 0 @wp.kernel -def _write_filtered_contacts( - # Model: - opt_warn_overflow: bool, - geom_type: wp.array[int], - geom_condim: wp.array[int], - geom_priority: wp.array[int], - geom_solmix: wp.array2d[float], - geom_solref: wp.array2d[wp.vec2], - geom_solimp: wp.array2d[vec5], - geom_friction: wp.array2d[wp.vec3], - geom_margin: wp.array2d[float], - geom_gap: wp.array2d[float], - flex_condim: wp.array[int], - flex_priority: wp.array[int], - flex_solmix: wp.array[float], - flex_solref: wp.array[wp.vec2], - flex_solimp: wp.array[vec5], - flex_friction: wp.array[wp.vec3], - flex_margin: wp.array[float], - flex_gap: wp.array[float], - flex_dim: wp.array[int], - # Data in: - naconmax_in: int, +def _populate_group_starts( # In: + flex_group_temp_in: wp.array[int], + flex_group_ids_in: wp.array[int], ncand: wp.array[int], - cand_dist: wp.array[float], + # Out: + flex_group_start_indices_out: wp.array[int], + flex_num_groups_out: wp.array[int], +): + si = wp.tid() + ncand_limit = wp.min(ncand[0], flex_group_temp_in.shape[0]) + if si >= ncand_limit: + return + + if flex_group_temp_in[si] == 1: + g = flex_group_ids_in[si] - 1 + if g < flex_group_start_indices_out.shape[0]: + flex_group_start_indices_out[g] = si + else: + wp.printf( + "Error: group index %d out of bounds for group_start_indices (size %d)\n", g, flex_group_start_indices_out.shape[0] + ) + + if si == ncand_limit - 1: + flex_num_groups_out[0] = flex_group_ids_in[si] + + +@wp.func +def _tie_break_fps( + # In: + curr_idx: int, + sel_idx: int, + cand_elem: wp.array[wp.vec2i], +): + if sel_idx < 0: + return True + + elem1_curr = cand_elem[curr_idx][0] + elem2_curr = cand_elem[curr_idx][1] + + elem1_sel = cand_elem[sel_idx][0] + elem2_sel = cand_elem[sel_idx][1] + + if elem1_curr != elem1_sel: + return elem1_curr < elem1_sel + return elem2_curr < elem2_sel + + +@wp.kernel +def _filter_flex_fps( + # In: + flex_group_start_indices_in: wp.array[int], + flex_num_groups_in: wp.array[int], + ncand: wp.array[int], + cand_active_sorted: wp.array[int], + sort_val: wp.array[int], cand_pos: wp.array[wp.vec3], - cand_nrm: wp.array[wp.vec3], - cand_geom: wp.array[wp.vec2i], - cand_flex: wp.array[wp.vec2i], + cand_dist: wp.array[float], cand_elem: wp.array[wp.vec2i], - cand_vert: wp.array[wp.vec2i], - cand_worldid: wp.array[int], - cand_type: wp.array[int], - cand_geomcollisionid: wp.array[int], - cand_active: wp.array[int], - # Data out: - contact_dist_out: wp.array[float], - contact_pos_out: wp.array[wp.vec3], - contact_frame_out: wp.array[wp.mat33], - contact_includemargin_out: wp.array[float], - contact_friction_out: wp.array[vec5], - contact_solref_out: wp.array[wp.vec2], - contact_solreffriction_out: wp.array[wp.vec2], - contact_solimp_out: wp.array[vec5], - contact_dim_out: wp.array[int], - contact_geom_out: wp.array[wp.vec2i], - contact_flex_out: wp.array[wp.vec2i], - contact_elem_out: wp.array[wp.vec2i], - contact_vert_out: wp.array[wp.vec2i], - contact_worldid_out: wp.array[int], - contact_type_out: wp.array[int], - contact_geomcollisionid_out: wp.array[int], - nacon_out: wp.array[int], - # Data out: - overflow_out: wp.array[int], + cand_geom: wp.array[wp.vec2i], + # Out: + fps_min_dist_out: wp.array[float], + cand_active_out: wp.array[int], ): - i = wp.tid() - if i >= ncand[0]: - return + g = wp.tid() - if cand_active[i] == 0: + if g >= flex_num_groups_in[0]: return - geomid = cand_geom[i][0] - worldid = cand_worldid[i] + g_start = flex_group_start_indices_in[g] + first_cand_idx = sort_val[g_start] - condim = int(0) - margin = float(0.0) - gap = float(0.0) - solref = wp.vec2(0.0, 0.0) - solimp = vec5(0.0, 0.0, 0.0, 0.0, 0.0) - friction = vec5(0.0, 0.0, 0.0, 0.0, 0.0) - - if geomid >= 0: - flexid = cand_flex[i][1] - geom_margin_val = geom_margin[worldid % geom_margin.shape[0], geomid] - tri_margin = flex_margin[flexid] - margin = geom_margin_val + tri_margin + # Only perform FPS for flex-flex/self-flex candidate groups (geom_id < 0) + if cand_geom[first_cand_idx][0] >= 0: + return - condim, gap, solref, solimp, friction = _mix_flex_contact_params( - geom_condim[geomid], - geom_priority[geomid], - geom_solmix[worldid % geom_solmix.shape[0], geomid], - geom_solref[worldid % geom_solref.shape[0], geomid], - geom_solimp[worldid % geom_solimp.shape[0], geomid], - geom_friction[worldid % geom_friction.shape[0], geomid], - geom_gap[worldid % geom_gap.shape[0], geomid], - flex_condim[flexid], - flex_priority[flexid], - flex_solmix[flexid], - flex_solref[flexid], - flex_solimp[flexid], - flex_friction[flexid], - flex_gap[flexid], - ) - else: - flex1 = cand_flex[i][0] - flex2 = cand_flex[i][1] - margin = 0.0 - gap = 0.0 - - mixed_condim, _, solref, solimp, friction = _mix_flex_contact_params( - flex_condim[flex1], - flex_priority[flex1], - flex_solmix[flex1], - flex_solref[flex1], - flex_solimp[flex1], - flex_friction[flex1], - 0.0, - flex_condim[flex2], - flex_priority[flex2], - flex_solmix[flex2], - flex_solref[flex2], - flex_solimp[flex2], - flex_friction[flex2], - 0.0, - ) + ncand_limit = wp.min(ncand[0], cand_active_sorted.shape[0]) + g_end = ncand_limit + if g < flex_num_groups_in[0] - 1: + g_end = flex_group_start_indices_in[g + 1] - if cand_vert[i][0] >= 0 and flex_dim[flex1] == 3: - condim = 1 - else: - condim = mixed_condim + # Count active candidates in this group + total_active = int(0) + for si in range(g_start, g_end): + if cand_active_sorted[si] == 1: + total_active += 1 - if cand_dist[i] >= margin: + # If group has <= MJ_MAXCONPAIR candidates, no FPS filtering needed + if total_active <= MJ_MAXCONPAIR: return - id_ = wp.atomic_add(nacon_out, 0, 1) - if id_ >= naconmax_in: - if opt_warn_overflow: - wp.printf( - "flex contact overflow - please increase naconmax to %u\n", - id_ + 1, - ) - wp.atomic_or(overflow_out, worldid, wp.static(OverflowType.NARROWPHASE)) + # Deactivate all candidates in group first + for si in range(g_start, g_end): + if cand_active_sorted[si] == 1: + cand_active_out[sort_val[si]] = 0 + + # 1. Find seed candidate (deepest contact = minimum cand_dist) + min_d = float(1e10) + sel_cand_idx = int(-1) + for si in range(g_start, g_end): + if cand_active_sorted[si] == 1: + c_idx = sort_val[si] + d_val = cand_dist[c_idx] + if d_val < min_d: + min_d = d_val + sel_cand_idx = c_idx + elif d_val == min_d: + if _tie_break_fps(c_idx, sel_cand_idx, cand_elem): + sel_cand_idx = c_idx + + if sel_cand_idx < 0: return - contact_dist_out[id_] = cand_dist[i] - contact_pos_out[id_] = cand_pos[i] - contact_frame_out[id_] = make_frame(cand_nrm[i]) - if geomid >= 0 and geom_type[geomid] == int(GeomType.PLANE): - contact_includemargin_out[id_] = margin - gap - else: - contact_includemargin_out[id_] = margin - contact_friction_out[id_] = friction - contact_solref_out[id_] = solref - contact_solreffriction_out[id_] = solref - contact_solimp_out[id_] = solimp - contact_dim_out[id_] = condim - contact_geom_out[id_] = cand_geom[i] - contact_flex_out[id_] = cand_flex[i] - contact_elem_out[id_] = cand_elem[i] - contact_vert_out[id_] = cand_vert[i] - contact_worldid_out[id_] = cand_worldid[i] - contact_type_out[id_] = cand_type[i] - contact_geomcollisionid_out[id_] = cand_geomcollisionid[i] - - -def flex_broadphase(m: Model, d: Data): + # Mark seed selected + cand_active_out[sel_cand_idx] = 1 + seed_pos = cand_pos[sel_cand_idx] + + # 2. Initialize running minimum distance for all active candidates in group + for si in range(g_start, g_end): + if cand_active_sorted[si] == 1: + c_idx = sort_val[si] + if c_idx == sel_cand_idx: + fps_min_dist_out[c_idx] = -1e10 + else: + fps_min_dist_out[c_idx] = wp.length(cand_pos[c_idx] - seed_pos) + + # 3. Iteratively select remaining MJ_MAXCONPAIR - 1 candidates with max min_dist + for k in range(1, MJ_MAXCONPAIR): + max_min_d = float(-1e10) + sel_cand_idx = int(-1) + for si in range(g_start, g_end): + if cand_active_sorted[si] == 1: + c_idx = sort_val[si] + md = fps_min_dist_out[c_idx] + if md > max_min_d: + max_min_d = md + sel_cand_idx = c_idx + elif md == max_min_d: + if _tie_break_fps(c_idx, sel_cand_idx, cand_elem): + sel_cand_idx = c_idx + + if sel_cand_idx < 0 or max_min_d <= 0.0: + break + + # Mark selected candidate + cand_active_out[sel_cand_idx] = 1 + fps_min_dist_out[sel_cand_idx] = -1e10 + new_pos = cand_pos[sel_cand_idx] + + # Update running min_dist to selected set + for si in range(g_start, g_end): + if cand_active_sorted[si] == 1: + c_idx = sort_val[si] + if fps_min_dist_out[c_idx] > 0.0: + d_new = wp.length(cand_pos[c_idx] - new_pos) + fps_min_dist_out[c_idx] = wp.min(fps_min_dist_out[c_idx], d_new) + + +def flex_broadphase_aabb(m: Model, d: Data): """Precompute dynamic flex object bounding boxes.""" wp.launch( _flex_broadphase_bounds, @@ -3130,7 +3271,6 @@ def flex_collision(m: Model, d: Data, ctx): mesh_nccd = wp.zeros(1, dtype=int) ncollision_dim2 = wp.zeros(1, dtype=int) - ncollision_dim3 = wp.zeros(1, dtype=int) ncollision_plane = wp.zeros(1, dtype=int) if m.has_3d_flex: @@ -3158,7 +3298,6 @@ def flex_collision(m: Model, d: Data, ctx): d.geom_xpos, d.geom_xmat, d.flexvert_xpos, - d.nworld, d.naconmax, ], outputs=[ @@ -3178,24 +3317,18 @@ def flex_collision(m: Model, d: Data, ctx): ) # Update dynamic flex object bounding boxes - flex_broadphase(m, d) + flex_broadphase_aabb(m, d) # 2D Flex Element Collisions if m.flexelem_geom_pair_filtered.shape[0] > 0: wp.launch( - _flex_broadphase_unified, + flex_broadphase_builder(bool(m.opt.warn_overflow)), dim=(d.nworld, m.flexelem_geom_pair_filtered.shape[0]), inputs=[ - m.ngeom, - m.nflex, - m.opt.warn_overflow, m.geom_type, - m.geom_size, m.geom_aabb, - m.geom_rbound, m.geom_margin, m.flex_margin, - m.flex_dim, m.flex_vertadr, m.flex_radius, d.geom_xpos, @@ -3219,33 +3352,16 @@ def flex_collision(m: Model, d: Data, ctx): ) wp.launch( - _flex_narrowphase_unified, + flex_narrowphase_unified_builder(bool(m.opt.warn_overflow)), dim=d.naconmax, inputs=[ - m.ngeom, - m.nflex, m.opt.ccd_tolerance, - m.opt.warn_overflow, + bool(m.opt.warn_overflow), m.geom_type, - m.geom_condim, m.geom_dataid, - m.geom_priority, - m.geom_solmix, - m.geom_solref, - m.geom_solimp, m.geom_size, - m.geom_friction, m.geom_margin, - m.geom_gap, - m.flex_condim, - m.flex_priority, - m.flex_solmix, - m.flex_solref, - m.flex_solimp, - m.flex_friction, m.flex_margin, - m.flex_gap, - m.flex_dim, m.flex_vertadr, m.flex_radius, m.mesh_vertadr, @@ -3263,8 +3379,6 @@ def flex_collision(m: Model, d: Data, ctx): d.geom_xpos, d.geom_xmat, d.flexvert_xpos, - d.nworld, - d.naconmax, d.naccdmax, ncollision_dim2, m.flex_elemadr, @@ -3302,15 +3416,12 @@ def flex_collision(m: Model, d: Data, ctx): # Plane Vertex Collisions if m.flexvert_geom_pair_filtered.shape[0] > 0: wp.launch( - _flex_broadphase_plane, + flex_broadphase_plane_builder(bool(m.opt.warn_overflow)), dim=(d.nworld, m.flexvert_geom_pair_filtered.shape[0]), inputs=[ - m.ngeom, - m.opt.warn_overflow, m.geom_type, m.geom_margin, m.flex_margin, - m.flex_vertadr, m.flex_radius, m.flexvert_geom_pair_filtered, m.flex_vertflexid, @@ -3333,82 +3444,18 @@ def flex_collision(m: Model, d: Data, ctx): _flex_plane_narrowphase, dim=d.naconmax, inputs=[ - m.ngeom, - m.nflexvert, m.geom_type, - m.geom_condim, - m.geom_priority, - m.geom_solmix, - m.geom_solref, - m.geom_solimp, - m.geom_friction, m.geom_margin, - m.geom_gap, - m.flex_condim, - m.flex_priority, - m.flex_solmix, - m.flex_solref, - m.flex_solimp, - m.flex_friction, m.flex_margin, - m.flex_gap, m.flex_vertadr, m.flex_radius, m.flex_vertflexid, d.geom_xpos, d.geom_xmat, d.flexvert_xpos, - d.nworld, - d.naconmax, ncollision_plane, ctx.collision_pair, ctx.collision_worldid, - ], - 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.flex, - d.contact.elem, - d.contact.vert, - d.contact.worldid, - d.contact.type, - d.contact.geomcollisionid, - d.nacon, - ], - ) - - # Geom Vertex Collisions (Candidate-based deduplicated narrowphase) - if m.nflexvert > 0: - wp.launch( - _flex_geom_vertex_narrowphase_detect, - dim=(d.nworld, m.nflexvert), - inputs=[ - m.ngeom, - m.nflexvert, - m.geom_type, - m.geom_contype, - m.geom_conaffinity, - m.geom_size, - m.geom_margin, - m.flex_contype, - m.flex_conaffinity, - m.flex_margin, - m.flex_dim, - m.flex_vertadr, - m.flex_radius, - m.flex_vertflexid, - d.geom_xpos, - d.geom_xmat, - d.flexvert_xpos, - d.nworld, d.naconmax, ], outputs=[ @@ -3427,263 +3474,282 @@ def flex_collision(m: Model, d: Data, ctx): ], ) - if m.nflexevpair > 0: + # Geom Vertex Collisions (Candidate-based deduplicated narrowphase) + wp.launch( + _flex_geom_vertex_narrowphase_detect, + dim=(d.nworld, m.nflexvert), + inputs=[ + m.ngeom, + m.geom_type, + m.geom_contype, + m.geom_conaffinity, + m.geom_size, + m.geom_margin, + m.flex_contype, + m.flex_conaffinity, + m.flex_margin, + m.flex_dim, + m.flex_vertadr, + m.flex_radius, + m.flex_vertflexid, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.naconmax, + ], + outputs=[ + d.overflow, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + wp.launch( + _flex_internal_collisions_detect, + dim=(d.nworld, m.nflexevpair), + inputs=[ + m.flex_margin, + m.flex_internal, + m.flex_dim, + m.flex_vertadr, + m.flex_elemdataadr, + m.flex_elem, + m.flex_evpair, + m.flex_radius, + m.flex_evpairflexid, + d.flexvert_xpos, + d.naconmax, + ], + outputs=[ + d.overflow, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + wp.launch( + _flex_tet_internal_collisions_detect, + dim=(d.nworld, m.nflexelem), + inputs=[ + m.flex_dim, + m.flex_vertadr, + m.flex_elemadr, + m.flex_elemdataadr, + m.flex_elem, + m.flex_radius, + m.flex_elemflexid, + d.flexvert_xpos, + d.naconmax, + ], + outputs=[ + d.overflow, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + + selfcollide_enabled = m.has_flex_selfcollide + run_sap = (selfcollide_enabled and m.nflexelem > NFLEXELEM_SAP) or (m.nflex > 1) + + if run_sap and m.nflexelem > 0: + epa_iterations = m.opt.ccd_iterations + nelem = m.nflexelem + nworldelem = d.nworld * nelem + + # Fixed projection direction (same as sap_broadphase in collision_driver.py) + direction = wp.vec3(0.5935, 0.7790, 0.1235) + direction = wp.normalize(direction) + + # Allocate SAP arrays + sap_lower = wp.empty((d.nworld, nelem, 2), dtype=float) + sap_upper = wp.empty((d.nworld, nelem), dtype=float) + sap_sort_index = wp.empty((d.nworld, nelem, 2), dtype=int) + sap_range_arr = wp.empty((d.nworld, nelem), dtype=int) + sap_cumsum = wp.empty((d.nworld, nelem), dtype=int) + sap_seg_index = wp.empty(d.nworld + 1, dtype=int) + elem_aabb_lower = wp.empty((d.nworld, nelem), dtype=wp.vec3) + elem_aabb_upper = wp.empty((d.nworld, nelem), dtype=wp.vec3) + wp.launch( - _flex_internal_collisions_detect, - dim=(d.nworld, m.nflexevpair), + _flex_sap_project, + dim=(d.nworld, nelem), inputs=[ m.nflex, - m.flex_margin, - m.flex_internal, + m.flex_selfcollide, m.flex_dim, m.flex_vertadr, + m.flex_elemadr, m.flex_elemdataadr, - m.flex_evpairadr, - m.flex_evpairnum, m.flex_elem, - m.flex_evpair, m.flex_radius, - m.flex_evpairflexid, + m.flex_elemflexid, d.flexvert_xpos, - d.naconmax, + d.nworld, + nelem, + direction, ], outputs=[ - d.overflow, - cand_dist, - cand_pos, - cand_nrm, - cand_geom, - cand_flex, - cand_elem, - cand_vert, - cand_worldid, - cand_type, - cand_geomcollisionid, - ncand, + sap_lower.reshape((-1, nelem)), + sap_upper, + sap_sort_index.reshape((-1, nelem)), + elem_aabb_lower, + elem_aabb_upper, + sap_seg_index, ], ) - if m.nflexelem > 0: + # Step 2: Sort + wp.utils.segmented_sort_pairs( + sap_lower.reshape((-1, nelem)), + sap_sort_index.reshape((-1, nelem)), + nworldelem, + sap_seg_index, + ) + + # Step 3: Range wp.launch( - _flex_tet_internal_collisions_detect, - dim=(d.nworld, m.nflexelem), + sap_range, + dim=(d.nworld, nelem), inputs=[ - m.nflex, + nelem, + sap_lower.reshape((-1, nelem)), + sap_upper, + sap_sort_index.reshape((-1, nelem)), + ], + outputs=[ + sap_range_arr, + ], + ) + + # Step 4: Prefix sum for load balancing + wp.utils.array_scan( + sap_range_arr.reshape(-1), + sap_cumsum.reshape(-1), + True, + ) + + # Step 5: SAP sweep - output pairs only (no narrowphase) + nsweep = 5 * nworldelem + + # Step 5a: Generic SAP sweep (shared with geom broadphase) + self_npairs = wp.zeros(1, dtype=int) + self_pair_elem1 = wp.empty(d.naconmax, dtype=int) + self_pair_elem2 = wp.empty(d.naconmax, dtype=int) + self_pair_worldid = wp.empty(d.naconmax, dtype=int) + + ff_npairs = wp.zeros(1, dtype=int) + ff_pair_elem1 = wp.empty(d.naconmax, dtype=int) + ff_pair_elem2 = wp.empty(d.naconmax, dtype=int) + ff_pair_worldid = wp.empty(d.naconmax, dtype=int) + + wp.launch( + flex_sap_sweep_builder(bool(m.opt.warn_overflow)), + dim=nsweep, + inputs=[ + m.flex_contype, + m.flex_conaffinity, + m.flex_selfcollide, m.flex_dim, m.flex_vertadr, m.flex_elemadr, - m.flex_elemnum, m.flex_elemdataadr, + m.flex_vertbodyid, m.flex_elem, - m.flex_radius, m.flex_elemflexid, - d.flexvert_xpos, + nelem, + sap_sort_index.reshape((-1, nelem)), + sap_cumsum.reshape(-1), + nsweep, + elem_aabb_lower, + elem_aabb_upper, + d.naconmax, d.naconmax, ], outputs=[ d.overflow, - cand_dist, - cand_pos, - cand_nrm, - cand_geom, - cand_flex, - cand_elem, - cand_vert, - cand_worldid, - cand_type, - cand_geomcollisionid, - ncand, + self_npairs, + self_pair_elem1, + self_pair_elem2, + self_pair_worldid, + ff_npairs, + ff_pair_elem1, + ff_pair_elem2, + ff_pair_worldid, ], ) - selfcollide_enabled = m.has_flex_selfcollide - - if selfcollide_enabled and m.nflexelem > 0: - epa_iterations = m.opt.ccd_iterations - - # TODO(team): investigate optimal nflexelem threshold for SAP vs brute-force - if m.nflexelem > 32: - # --- SAP broadphase path --- - nelem = m.nflexelem - nworldelem = d.nworld * nelem - - # Fixed projection direction (same as sap_broadphase in collision_driver.py) - # TODO(team): compute optimal SAP direction - direction = wp.vec3(0.5935, 0.7790, 0.1235) - direction = wp.normalize(direction) - - # Allocate SAP arrays - sap_lower = wp.empty((d.nworld, nelem, 2), dtype=float) - sap_upper = wp.empty((d.nworld, nelem), dtype=float) - sap_sort_index = wp.empty((d.nworld, nelem, 2), dtype=int) - sap_range_arr = wp.empty((d.nworld, nelem), dtype=int) - sap_cumsum = wp.empty((d.nworld, nelem), dtype=int) - sap_seg_index = wp.empty(d.nworld + 1, dtype=int) - elem_aabb_lower = wp.empty((d.nworld, nelem), dtype=wp.vec3) - elem_aabb_upper = wp.empty((d.nworld, nelem), dtype=wp.vec3) - - # Step 1: Project element AABBs onto direction - wp.launch( - _flex_sap_project, - dim=(d.nworld, nelem), - inputs=[ - m.nflex, - m.flex_selfcollide, - m.flex_dim, - m.flex_vertadr, - m.flex_elemadr, - m.flex_elemdataadr, - m.flex_elem, - m.flex_radius, - m.flex_elemflexid, - d.flexvert_xpos, - d.nworld, - nelem, - direction, - ], - outputs=[ - sap_lower.reshape((-1, nelem)), - sap_upper, - sap_sort_index.reshape((-1, nelem)), - elem_aabb_lower, - elem_aabb_upper, - sap_seg_index, - ], - ) - - # Step 2: Sort - wp.utils.segmented_sort_pairs( - sap_lower.reshape((-1, nelem)), - sap_sort_index.reshape((-1, nelem)), - nworldelem, - sap_seg_index, - ) - - # Step 3: Range - wp.launch( - sap_range, - dim=(d.nworld, nelem), - inputs=[ - nelem, - sap_lower.reshape((-1, nelem)), - sap_upper, - sap_sort_index.reshape((-1, nelem)), - ], - outputs=[ - sap_range_arr, - ], - ) - - # Step 4: Prefix sum for load balancing - wp.utils.array_scan( - sap_range_arr.reshape(-1), - sap_cumsum.reshape(-1), - True, - ) - - # Step 5: SAP sweep - output pairs only (no narrowphase) - nsweep = 5 * nworldelem - - npairs = wp.zeros(1, dtype=int) - pair_elem1 = wp.empty(d.naconmax, dtype=int) - pair_elem2 = wp.empty(d.naconmax, dtype=int) - pair_worldid = wp.empty(d.naconmax, dtype=int) - - # Step 5a: Generic SAP sweep (shared with geom broadphase) - raw_npairs = wp.zeros(1, dtype=int) - raw_pair_elem1 = wp.empty(d.naconmax, dtype=int) - raw_pair_elem2 = wp.empty(d.naconmax, dtype=int) - raw_pair_worldid = wp.empty(d.naconmax, dtype=int) - - wp.launch( - sap_sweep, - dim=nsweep, - inputs=[ - nelem, - sap_sort_index.reshape((-1, nelem)), - sap_cumsum.reshape(-1), - nsweep, - elem_aabb_lower, - elem_aabb_upper, - d.naconmax, - ], - outputs=[ - raw_npairs, - raw_pair_elem1, - raw_pair_elem2, - raw_pair_worldid, - ], - ) - - # Step 5b: Filter pairs (flex-specific: selfcollide, shared vertices) - wp.launch( - _flex_sap_filter, - dim=d.naconmax, - inputs=[ - m.flex_selfcollide, - m.flex_dim, - m.flex_vertadr, - m.flex_elemadr, - m.flex_elemnum, - m.flex_elemdataadr, - m.flex_vertbodyid, - m.flex_elem, - m.flex_elemflexid, - raw_npairs, - raw_pair_elem1, - raw_pair_elem2, - raw_pair_worldid, - ], - outputs=[ - npairs, - pair_elem1, - pair_elem2, - pair_worldid, - ], - ) - - # Step 6: Narrowphase on actual pairs only - workspace_verts = wp.empty(d.naconmax * 8, dtype=wp.vec3) - - if m.max_flex_dim > 1: - epa_vert = wp.empty(shape=(d.naconmax, 10 + 2 * epa_iterations), dtype=wp.vec3) - epa_vert_index = wp.empty(shape=(d.naconmax, 10 + 2 * epa_iterations), dtype=int) - epa_face = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) - epa_pr = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=wp.vec3) - epa_norm2 = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=float) - epa_horizon = wp.empty(shape=(d.naconmax, MJ_MAX_EPAHORIZON), dtype=int) - else: - epa_vert = wp.empty(shape=(1, 1), dtype=wp.vec3) - epa_vert_index = wp.empty(shape=(1, 1), dtype=int) - epa_face = wp.empty(shape=(1, 1), dtype=int) - epa_pr = wp.empty(shape=(1, 1), dtype=wp.vec3) - epa_norm2 = wp.empty(shape=(1, 1), dtype=float) - epa_horizon = wp.empty(shape=(1, 1), dtype=int) - + # Step 6: Allocate narrowphase workspace, reused by self and flex-flex narrowphases + workspace_verts = wp.empty(d.naconmax * 8, dtype=wp.vec3) + if m.max_flex_dim > 1: + epa_vert = wp.empty(shape=(d.naconmax, 10 + 2 * epa_iterations), dtype=wp.vec3) + epa_vert_index = wp.empty(shape=(d.naconmax, 10 + 2 * epa_iterations), dtype=int) + epa_face = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) + epa_pr = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=wp.vec3) + epa_norm2 = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=float) + epa_horizon = wp.empty(shape=(d.naconmax, MJ_MAX_EPAHORIZON), dtype=int) + else: + epa_vert = wp.empty(shape=(1, 1), dtype=wp.vec3) + epa_vert_index = wp.empty(shape=(1, 1), dtype=int) + epa_face = wp.empty(shape=(1, 1), dtype=int) + epa_pr = wp.empty(shape=(1, 1), dtype=wp.vec3) + epa_norm2 = wp.empty(shape=(1, 1), dtype=float) + epa_horizon = wp.empty(shape=(1, 1), dtype=int) + + # Run self-collision SAP narrowphase if enabled + if selfcollide_enabled and m.nflexelem > NFLEXELEM_SAP: wp.launch( - _flex_selfcollision_narrowphase, + _flex_narrowphase, dim=d.naconmax, inputs=[ - m.nflex, m.opt.ccd_tolerance, - m.opt.warn_overflow, + bool(m.opt.warn_overflow), + m.flex_margin, + m.flex_gap, + m.flex_activelayers, m.flex_dim, m.flex_vertadr, m.flex_elemadr, m.flex_elemdataadr, m.flex_elem, + m.flex_elemlayer, m.flex_radius, m.flex_elemflexid, d.flexvert_xpos, d.naconmax, m.opt.ccd_iterations, epa_iterations, - npairs, - pair_elem1, - pair_elem2, - pair_worldid, + self_npairs, + self_pair_elem1, + self_pair_elem2, + self_pair_worldid, d.naconmax, - m.nflexelem, ], outputs=[ d.overflow, @@ -3708,47 +3774,34 @@ def flex_collision(m: Model, d: Data, ctx): ], ) - else: - # --- Brute-force fallback for small element counts --- - workspace_verts = wp.empty(d.nworld * m.nflexelem * 8, dtype=wp.vec3) - - if m.max_flex_dim > 1: - epa_vert = wp.empty(shape=(d.nworld * m.nflexelem, 10 + 2 * epa_iterations), dtype=wp.vec3) - epa_vert_index = wp.empty(shape=(d.nworld * m.nflexelem, 10 + 2 * epa_iterations), dtype=int) - epa_face = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) - epa_pr = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=wp.vec3) - epa_norm2 = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=float) - epa_horizon = wp.empty(shape=(d.nworld * m.nflexelem, MJ_MAX_EPAHORIZON), dtype=int) - else: - epa_vert = wp.empty(shape=(1, 1), dtype=wp.vec3) - epa_vert_index = wp.empty(shape=(1, 1), dtype=int) - epa_face = wp.empty(shape=(1, 1), dtype=int) - epa_pr = wp.empty(shape=(1, 1), dtype=wp.vec3) - epa_norm2 = wp.empty(shape=(1, 1), dtype=float) - epa_horizon = wp.empty(shape=(1, 1), dtype=int) - + # Run flex-flex SAP narrowphase if there are multiple flexes + if m.nflex > 1: wp.launch( - _flex_active_element_collisions_detect, - dim=(d.nworld, m.nflexelem), + _flex_narrowphase, + dim=d.naconmax, inputs=[ - m.nflex, m.opt.ccd_tolerance, - m.opt.warn_overflow, - m.flex_selfcollide, + bool(m.opt.warn_overflow), + m.flex_margin, + m.flex_gap, + m.flex_activelayers, m.flex_dim, m.flex_vertadr, m.flex_elemadr, - m.flex_elemnum, m.flex_elemdataadr, - m.flex_vertbodyid, m.flex_elem, + m.flex_elemlayer, m.flex_radius, m.flex_elemflexid, d.flexvert_xpos, d.naconmax, m.opt.ccd_iterations, epa_iterations, - m.nflexelem, + ff_npairs, + ff_pair_elem1, + ff_pair_elem2, + ff_pair_worldid, + d.naconmax, ], outputs=[ d.overflow, @@ -3773,10 +3826,77 @@ def flex_collision(m: Model, d: Data, ctx): ], ) + if selfcollide_enabled and m.nflexelem > 0 and m.nflexelem <= NFLEXELEM_SAP: + # --- Brute-force fallback for small element counts --- + epa_iterations = m.opt.ccd_iterations + workspace_verts = wp.empty(d.nworld * m.nflexelem * 8, dtype=wp.vec3) + + if m.max_flex_dim > 1: + epa_vert = wp.empty(shape=(d.nworld * m.nflexelem, 10 + 2 * epa_iterations), dtype=wp.vec3) + epa_vert_index = wp.empty(shape=(d.nworld * m.nflexelem, 10 + 2 * epa_iterations), dtype=int) + epa_face = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) + epa_pr = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=wp.vec3) + epa_norm2 = wp.empty(shape=(d.nworld * m.nflexelem, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=float) + epa_horizon = wp.empty(shape=(d.nworld * m.nflexelem, MJ_MAX_EPAHORIZON), dtype=int) + else: + epa_vert = wp.empty(shape=(1, 1), dtype=wp.vec3) + epa_vert_index = wp.empty(shape=(1, 1), dtype=int) + epa_face = wp.empty(shape=(1, 1), dtype=int) + epa_pr = wp.empty(shape=(1, 1), dtype=wp.vec3) + epa_norm2 = wp.empty(shape=(1, 1), dtype=float) + epa_horizon = wp.empty(shape=(1, 1), dtype=int) + + wp.launch( + _flex_active_element_collisions_detect, + dim=(d.nworld, m.nflexelem), + inputs=[ + m.opt.ccd_tolerance, + bool(m.opt.warn_overflow), + m.flex_selfcollide, + m.flex_activelayers, + m.flex_dim, + m.flex_vertadr, + m.flex_elemadr, + m.flex_elemnum, + m.flex_elemdataadr, + m.flex_vertbodyid, + m.flex_elem, + m.flex_elemlayer, + m.flex_radius, + m.flex_elemflexid, + d.flexvert_xpos, + d.naconmax, + m.opt.ccd_iterations, + epa_iterations, + m.nflexelem, + ], + outputs=[ + d.overflow, + workspace_verts, + epa_vert, + epa_vert_index, + epa_face, + epa_pr, + epa_norm2, + epa_horizon, + cand_dist, + cand_pos, + cand_nrm, + cand_geom, + cand_flex, + cand_elem, + cand_vert, + cand_worldid, + cand_type, + cand_geomcollisionid, + ncand, + ], + ) + # Filter duplicate contacts using sort-based deduplication. # Sort candidates by (worldid, flex_id, geom_id) so duplicates are contiguous, # then compare only within each group. This is O(n log n) vs the naive O(n^2). - filter_key = wp.empty(d.naconmax * 2, dtype=int) + filter_key = wp.empty(d.naconmax * 2, dtype=wp.int64) filter_val = wp.empty(d.naconmax * 2, dtype=int) wp.launch( _compute_filter_key, @@ -3787,6 +3907,7 @@ def flex_collision(m: Model, d: Data, ctx): ncand, cand_geom, cand_flex, + cand_elem, cand_worldid, ], outputs=[ @@ -3807,18 +3928,71 @@ def flex_collision(m: Model, d: Data, ctx): filter_val, cand_dist, cand_pos, + cand_elem, + ], + outputs=[ + cand_active, + ], + ) + + # Farthest point sampling (FPS) filtering to limit flex-flex contacts to MJ_MAXCONPAIR + cand_active_sorted = wp.empty(d.naconmax, dtype=int) + wp.launch( + _populate_active_sorted, + dim=d.naconmax, + inputs=[ncand, filter_val, cand_active], + outputs=[cand_active_sorted], + ) + + flex_group_temp = wp.empty(d.naconmax, dtype=int) + flex_group_ids = wp.empty(d.naconmax, dtype=int) + flex_group_start_indices = wp.empty(d.naconmax, dtype=int) + flex_num_groups = wp.zeros(1, dtype=int) + + wp.launch( + _find_group_starts, + dim=d.naconmax, + inputs=[ncand, filter_key], + outputs=[flex_group_temp], + ) + + wp.utils.array_scan(flex_group_temp, flex_group_ids, True) + + wp.launch( + _populate_group_starts, + dim=d.naconmax, + inputs=[flex_group_temp, flex_group_ids, ncand], + outputs=[flex_group_start_indices, flex_num_groups], + ) + flex_fps_min_dist = wp.empty(d.naconmax, dtype=float) + + world_stride = m.ngeom * m.nflex + m.nflex * m.nflex + nmax_groups = d.nworld * world_stride + wp.launch( + _filter_flex_fps, + dim=nmax_groups, + inputs=[ + flex_group_start_indices, + flex_num_groups, + ncand, + cand_active_sorted, + filter_val, + cand_pos, + cand_dist, + cand_elem, + cand_geom, ], outputs=[ + flex_fps_min_dist, cand_active, ], ) # Copy filtered contacts to the main d.contact array, computing contact parameters on-the-fly wp.launch( - _write_filtered_contacts, + _write_filtered_contacts_builder(bool(m.opt.warn_overflow)), dim=d.naconmax, inputs=[ - m.opt.warn_overflow, m.geom_type, m.geom_condim, m.geom_priority, diff --git a/mujoco_warp/_src/collision_gjk.py b/mujoco_warp/_src/collision_gjk.py index 7feb03790..aeaa68d78 100644 --- a/mujoco_warp/_src/collision_gjk.py +++ b/mujoco_warp/_src/collision_gjk.py @@ -1122,9 +1122,11 @@ def _polytope3( # get normals in both directions n = wp.cross(simplex[1] - simplex[0], simplex[2] - simplex[0]) - if wp.norm_l2(n) < MINVAL: + norm = wp.norm_l2(n) + if norm < MINVAL: pt.status = 2 return pt + n = n / norm pt.vert[0] = simplex1[0] pt.vert[1] = simplex2[0] diff --git a/mujoco_warp/_src/flex_test.py b/mujoco_warp/_src/flex_test.py index 3701e4402..3ef6d4cbc 100644 --- a/mujoco_warp/_src/flex_test.py +++ b/mujoco_warp/_src/flex_test.py @@ -1672,6 +1672,35 @@ def test_plane_cloth_contact_generated(self, nworld): plane_contacts = np.sum(contact_geom[w_indices, 0] == 0) self.assertGreater(plane_contacts, 0, f"Expected at least one contact with the plane in world {w}") + @parameterized.parameters(1, 2) + def test_plane_cloth_no_fps_limit(self, nworld): + """Test that flex-geom contacts are not limited to MJ_MAXCONPAIR (50).""" + xml = """ + + + """ + _, _, m, d = test_data.fixture(xml=xml, nworld=nworld) + + mjw.kinematics(m, d) + mjw.collision(m, d) + + nacon = int(d.nacon.numpy()[0]) + expected_contacts = 100 * nworld + self.assertEqual(nacon, expected_contacts, f"Expected {expected_contacts} contacts, got {nacon}") + @parameterized.parameters(1, 2) def test_mixed_flex_broadphase_and_narrowphase(self, nworld): """Test that broadphase and narrowphase run correctly with mixed 2D and 3D flexes.""" @@ -2594,6 +2623,179 @@ def test_insidesite_flex_body(self, nworld): ) +class FlexFlexCollisionTest(parameterized.TestCase): + """Tests for flex-flex collisions.""" + + def test_flex_flex_collision_1d_1d(self): + """Test collision between two overlapping 1D flexes (capsules).""" + xml = """ + + + + + + + + + + + + """ + _, _, m, d = test_data.fixture(xml=xml) + + self.assertEqual(m.nflex, 2) + + mjw.kinematics(m, d) + mjw.collision(m, d) + + nacon = int(d.nacon.numpy()[0]) + self.assertGreater(nacon, 0, "Expected capsule-capsule flex-flex contacts") + + def test_flex_flex_collision_2d_2d(self): + """Test GJK/EPA collision between two overlapping 2D cloths.""" + xml = """ + + + + + + + + + + + + """ + _, _, m, d = test_data.fixture(xml=xml) + + self.assertEqual(m.nflex, 2) + + mjw.kinematics(m, d) + mjw.collision(m, d) + + nacon = int(d.nacon.numpy()[0]) + self.assertGreater(nacon, 0, "Expected GJK/EPA cloth-cloth contacts") + + def test_flex_flex_collision_1d_2d(self): + """Test GJK/EPA collision between 1D rope and 2D cloth.""" + xml = """ + + + + + + + + + + + + + """ + _, _, m, d = test_data.fixture(xml=xml) + + self.assertEqual(m.nflex, 2) + + mjw.kinematics(m, d) + mjw.collision(m, d) + + nacon = int(d.nacon.numpy()[0]) + self.assertGreater(nacon, 0, "Expected GJK/EPA rope-cloth contacts") + + def test_flex_flex_collision_bitmask_filtering(self): + """Test that flex-flex collisions respect contype/conaffinity bitmasks.""" + xml = """ + + + + + + + + + + + + """ + _, _, m, d = test_data.fixture(xml=xml) + + self.assertEqual(m.nflex, 2) + + mjw.kinematics(m, d) + mjw.collision(m, d) + + nacon = int(d.nacon.numpy()[0]) + self.assertEqual(nacon, 0, "Expected no contacts due to bitmask mismatch") + + def test_flex_flex_collision_shared_body_filtering(self): + """Test that flex-flex collisions exclude elements if vertices share a body.""" + xml = """ + + + + + + + + + + + """ + _, _, m, d = test_data.fixture(xml=xml) + + # First verify we get contacts normally + mjw.kinematics(m, d) + mjw.collision(m, d) + nacon = int(d.nacon.numpy()[0]) + self.assertGreater(nacon, 0, "Expected baseline contacts before weld") + + # Now weld all vertices of cloth1 and cloth2 to the same body ID (e.g. 1) + # So they are treated as sharing bodies. + vertbody = m.flex_vertbodyid.numpy() + vertbody[:] = 1 + m.flex_vertbodyid.assign(vertbody) + + # Run collision again + mjw.collision(m, d) + nacon = int(d.nacon.numpy()[0]) + self.assertEqual(nacon, 0, "Expected 0 contacts due to shared body exclusion") + + def test_flex_flex_collision_3d_3d(self): + """Test GJK/EPA collision between two overlapping 3D blocks.""" + xml = """ + + + + + + + + + + + """ + _, _, m, d = test_data.fixture(xml=xml) + + self.assertEqual(m.nflex, 2) + + mjw.kinematics(m, d) + mjw.flex(m, d) + mjw.collision(m, d) + + nacon = int(d.nacon.numpy()[0]) + self.assertGreater(nacon, 0, "Expected GJK/EPA 3D block-block contacts") + + if __name__ == "__main__": wp.init() absltest.main() diff --git a/mujoco_warp/_src/io.py b/mujoco_warp/_src/io.py index 5a4bb678d..e3e88a207 100644 --- a/mujoco_warp/_src/io.py +++ b/mujoco_warp/_src/io.py @@ -1184,6 +1184,30 @@ def _check_margin(name, t1, t2, margin): m.flexelem_geom_pair_filtered = np.array(flexelem_geom_pairs, dtype=np.int32) m.flexvert_geom_pair_filtered = np.array(flexvert_geom_pairs, dtype=np.int32) + # Precompute flex-flex collision pairs (cross-flex) + # Matching MuJoCo C's canCollide2: check contype/conaffinity bitmask + flexflex_pairs = [] + flexflex_npairs_total = 0 + if mjm.nflex >= 2: + for f1 in range(mjm.nflex): + ct1 = mjm.flex_contype[f1] + ca1 = mjm.flex_conaffinity[f1] + for f2 in range(f1 + 1, mjm.nflex): + ct2 = mjm.flex_contype[f2] + ca2 = mjm.flex_conaffinity[f2] + # Opposite of bitmask filter: can collide if either direction matches + if (ct1 & ca2) or (ct2 & ca1): + pair_count = int(mjm.flex_elemnum[f1]) * int(mjm.flex_elemnum[f2]) + flexflex_pairs.append((f1, f2)) + flexflex_npairs_total += pair_count + + if not flexflex_pairs: + m.flexflex_pair_filtered = np.zeros((0, 2), dtype=np.int32) + else: + m.flexflex_pair_filtered = np.array(flexflex_pairs, dtype=np.int32) + + m.flexflex_npairs_total = flexflex_npairs_total + m.flex_elemflexid = flex_elemflexid m.flex_shellflexid = flex_shellflexid m.flex_evpairflexid = flex_evpairflexid @@ -1524,6 +1548,16 @@ def _eq_bodies(i): _body_pair_nnz(mjm, flex_body_list[idx1], flex_body_list[idx2]), ) + # flex-flex collision + for fj in range(fi + 1, mjm.nflex): + fct_j = mjm.flex_contype[fj] + fca_j = mjm.flex_conaffinity[fj] + if (fct & fca_j) or (fct_j & fca): + dim_i = mjm.flex_dim[fi] + dim_j = mjm.flex_dim[fj] + nnz_flex_flex = (dim_i + dim_j + 2) * 3 + max_contact_nnz = max(max_contact_nnz, nnz_flex_flex) + total_nnz += njmax * max_contact_nnz return int(min(max(total_nnz, 1), njmax * mjm.nv)) diff --git a/mujoco_warp/_src/types.py b/mujoco_warp/_src/types.py index 0e060fd8a..3498ed6a1 100644 --- a/mujoco_warp/_src/types.py +++ b/mujoco_warp/_src/types.py @@ -1164,6 +1164,7 @@ class Model: flex_gap: include in solver if dist