From f703a0e02388d6bc32aac5f704cfd45b567b59b0 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Wed, 7 Jan 2026 11:41:02 -0800 Subject: [PATCH] Import google-deepmind/mujoco_warp from GitHub. PiperOrigin-RevId: 853350801 Change-Id: Idb19c2865b94f4740daf8bc52ed3e7665a109d51 --- .../mjx/third_party/mujoco_warp/__init__.py | 1 + .../mujoco_warp/_src/block_cholesky.py | 4 +- .../mujoco_warp/_src/collision_convex.py | 1177 ++++++++++------- .../mujoco_warp/_src/collision_gjk.py | 78 +- .../mujoco_warp/_src/collision_primitive.py | 18 +- .../mujoco_warp/_src/collision_sdf.py | 7 +- .../mujoco_warp/_src/constraint.py | 2 +- .../mjx/third_party/mujoco_warp/_src/io.py | 57 +- .../mjx/third_party/mujoco_warp/_src/math.py | 32 +- .../third_party/mujoco_warp/_src/passive.py | 6 + .../mjx/third_party/mujoco_warp/_src/ray.py | 276 ++-- .../third_party/mujoco_warp/_src/sensor.py | 22 +- .../third_party/mujoco_warp/_src/solver.py | 150 ++- mjx/mujoco/mjx/warp/forward.py | 2 - mjx/mujoco/mjx/warp/types.py | 1 + 15 files changed, 1173 insertions(+), 660 deletions(-) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py index 744894df..360803e3 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py @@ -46,6 +46,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.io import put_model as put_model from mujoco.mjx.third_party.mujoco_warp._src.io import reset_data as reset_data from mujoco.mjx.third_party.mujoco_warp._src.passive import passive as passive from mujoco.mjx.third_party.mujoco_warp._src.ray import ray as ray +from mujoco.mjx.third_party.mujoco_warp._src.ray import rays as rays from mujoco.mjx.third_party.mujoco_warp._src.sensor import energy_pos as energy_pos from mujoco.mjx.third_party.mujoco_warp._src.sensor import energy_vel as energy_vel from mujoco.mjx.third_party.mujoco_warp._src.sensor import sensor_acc as sensor_acc diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py index 16d67051..36875caf 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py @@ -33,7 +33,7 @@ def create_blocked_cholesky_func(block_size: int): It returns a lower-triangular matrix L such that A = L L^T. """ # TODO(team): remove conditional after mjwarp relies on >= 1.11 - bleeding_edge_warp = wp.static(wp.__version__ >= "1.11") + bleeding_edge_warp = wp.static(tuple(map(int, wp.__version__.split(".")[:2])) >= (1, 11)) # Process the matrix in blocks along its leading dimension. for k in range(0, matrix_size, block_size): @@ -92,7 +92,7 @@ def create_blocked_cholesky_solve_func(block_size: int, matrix_size_static: int) substitution. """ # TODO(team): remove conditional after mjwarp relies on >= 1.11 - bleeding_edge_warp = wp.static(wp.__version__ >= "1.11") + bleeding_edge_warp = wp.static(tuple(map(int, wp.__version__.split(".")[:2])) >= (1, 11)) rhs_tile = wp.tile_load(b, shape=(matrix_size_static, 1), offset=(0, 0), storage="shared", bounds_check=False) # Forward substitution: solve L y = b diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py index de3f485b..e62c6172 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py @@ -191,6 +191,517 @@ def _hfield_filter( return False, xmin, xmax, ymin, ymax, zmin, zmax +@cache_kernel +def ccd_hfield_kernel_builder( + geomtype1: int, + geomtype2: int, + gjk_iterations: int, + epa_iterations: int, +): + """Kernel builder for heightfield CCD collisions (no multiccd args).""" + + # runs convex collision on a set of geom pairs to recover contact info + @nested_kernel(module="unique", enable_backward=False) + def ccd_hfield_kernel( + # Model: + opt_ccd_tolerance: wp.array(dtype=float), + geom_type: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_priority: wp.array(dtype=int), + geom_solmix: wp.array2d(dtype=float), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_solimp: wp.array2d(dtype=vec5), + geom_size: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), + geom_rbound: wp.array2d(dtype=float), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + geom_gap: wp.array2d(dtype=float), + hfield_adr: wp.array(dtype=int), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_data: wp.array(dtype=float), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_graphadr: wp.array(dtype=int), + mesh_graph: wp.array(dtype=int), + mesh_polynum: wp.array(dtype=int), + mesh_polyadr: wp.array(dtype=int), + mesh_polynormal: wp.array(dtype=wp.vec3), + mesh_polyvertadr: wp.array(dtype=int), + mesh_polyvertnum: wp.array(dtype=int), + mesh_polyvert: wp.array(dtype=int), + mesh_polymapadr: wp.array(dtype=int), + mesh_polymapnum: wp.array(dtype=int), + mesh_polymap: wp.array(dtype=int), + pair_dim: wp.array(dtype=int), + pair_solref: wp.array2d(dtype=wp.vec2), + pair_solreffriction: wp.array2d(dtype=wp.vec2), + pair_solimp: wp.array2d(dtype=vec5), + pair_margin: wp.array2d(dtype=float), + pair_gap: wp.array2d(dtype=float), + pair_friction: wp.array2d(dtype=vec5), + # Data in: + naconmax_in: int, + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + collision_pair_in: wp.array(dtype=wp.vec2i), + collision_pairid_in: wp.array(dtype=wp.vec2i), + collision_worldid_in: wp.array(dtype=int), + ncollision_in: wp.array(dtype=int), + # In: + epa_vert_in: wp.array2d(dtype=wp.vec3), + epa_vert1_in: wp.array2d(dtype=wp.vec3), + epa_vert2_in: wp.array2d(dtype=wp.vec3), + epa_vert_index1_in: wp.array2d(dtype=int), + epa_vert_index2_in: wp.array2d(dtype=int), + epa_face_in: wp.array2d(dtype=wp.vec3i), + epa_pr_in: wp.array2d(dtype=wp.vec3), + epa_norm2_in: wp.array2d(dtype=float), + epa_index_in: wp.array2d(dtype=int), + epa_map_in: wp.array2d(dtype=int), + epa_horizon_in: wp.array2d(dtype=int), + # Data out: + nacon_out: wp.array(dtype=int), + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + ): + tid = wp.tid() + if tid >= ncollision_in[0]: + return + + geoms = collision_pair_in[tid] + g1 = geoms[0] + g2 = geoms[1] + + if geom_type[g1] != geomtype1 or geom_type[g2] != geomtype2: + return + + worldid = collision_worldid_in[tid] + + # height field filter + no_hf_collision, xmin, xmax, ymin, ymax, zmin, zmax = _hfield_filter( + geom_dataid, geom_aabb, geom_rbound, geom_margin, hfield_size, geom_xpos_in, geom_xmat_in, worldid, g1, g2 + ) + if no_hf_collision: + return + + _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params( + geom_condim, + geom_priority, + geom_solmix, + geom_solref, + geom_solimp, + geom_friction, + geom_margin, + geom_gap, + pair_dim, + pair_solref, + pair_solreffriction, + pair_solimp, + pair_margin, + pair_gap, + pair_friction, + collision_pair_in, + collision_pairid_in, + tid, + worldid, + ) + + geom1, geom2 = geom_collision_pair( + geom_type, + geom_dataid, + geom_size, + mesh_vertadr, + mesh_vertnum, + mesh_graphadr, + mesh_vert, + mesh_graph, + mesh_polynum, + mesh_polyadr, + mesh_polynormal, + mesh_polyvertadr, + mesh_polyvertnum, + mesh_polyvert, + mesh_polymapadr, + mesh_polymapnum, + mesh_polymap, + geom_xpos_in, + geom_xmat_in, + geoms, + worldid, + ) + + # see MuJoCo mjc_ConvexHField + geom1_dataid = geom_dataid[g1] + + # height field subgrid + nrow = hfield_nrow[geom1_dataid] + ncol = hfield_ncol[geom1_dataid] + size = hfield_size[geom1_dataid] + + # subgrid + x_scale = 0.5 * float(ncol - 1) / size[0] + y_scale = 0.5 * float(nrow - 1) / size[1] + cmin = wp.max(0, int(wp.floor((xmin + size[0]) * x_scale))) + cmax = wp.min(ncol - 1, int(wp.ceil((xmax + size[0]) * x_scale))) + rmin = wp.max(0, int(wp.floor((ymin + size[1]) * y_scale))) + rmax = wp.min(nrow - 1, int(wp.ceil((ymax + size[1]) * y_scale))) + + dx = (2.0 * size[0]) / float(ncol - 1) + dy = (2.0 * size[1]) / float(nrow - 1) + dr = wp.vec2i(1, 0) + + prism = mat63() + + # set zbottom value using base size + prism[0, 2] = -size[3] + prism[1, 2] = -size[3] + prism[2, 2] = -size[3] + + adr = hfield_adr[geom1_dataid] + + hfield_contact_dist = vec_maxconpair() + hfield_contact_pos = mat_maxconpair() + hfield_contact_normal = mat_maxconpair() + min_dist = float(wp.inf) + min_normal = wp.vec3(wp.inf, wp.inf, wp.inf) + min_pos = wp.vec3(wp.inf, wp.inf, wp.inf) + min_id = int(-1) + + # TODO(team): height field margin? + geom1.margin = margin + geom2.margin = margin + + # EPA memory + epa_vert = epa_vert_in[tid] + epa_vert1 = epa_vert1_in[tid] + epa_vert2 = epa_vert2_in[tid] + epa_vert_index1 = epa_vert_index1_in[tid] + epa_vert_index2 = epa_vert_index2_in[tid] + epa_face = epa_face_in[tid] + epa_pr = epa_pr_in[tid] + epa_norm2 = epa_norm2_in[tid] + epa_index = epa_index_in[tid] + epa_map = epa_map_in[tid] + epa_horizon = epa_horizon_in[tid] + + collision_pairid = collision_pairid_in[tid] + + # process all prisms in subgrid + count = int(0) + for r in range(rmin, rmax): + nvert = int(0) + for c in range(cmin, cmax + 1): + # add both triangles from this cell + for i in range(2): + if count >= MJ_MAXCONPAIR: + wp.printf( + "height field collision overflow, number of collisions >= %u - please adjust resolution: \n decrease the number of hfield rows/cols or modify size of colliding geom\n", + MJ_MAXCONPAIR, + ) + continue + + # add vert + x = dx * float(c) - size[0] + y = dy * float(r + dr[i]) - size[1] + z = hfield_data[adr + (r + dr[i]) * ncol + c] * size[2] + margin + + prism[0] = prism[1] + prism[1] = prism[2] + prism[3] = prism[4] + prism[4] = prism[5] + + prism[2, 0] = x + prism[5, 0] = x + prism[2, 1] = y + prism[5, 1] = y + prism[5, 2] = z + + nvert += 1 + + if nvert <= 2: + continue + + # prism height test + if prism[3, 2] < zmin and prism[4, 2] < zmin and prism[5, 2] < zmin: + continue + + geom1.hfprism = prism + + # prism center + x1 = geom1.pos + x1_ = wp.vec3(0.0, 0.0, 0.0) + for i in range(6): + x1_ += prism[i] + x1 += geom1.rot @ (x1_ / 6.0) + + dist, ncontact, w1, w2, idx = ccd( + opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]], + 0.0, + gjk_iterations, + epa_iterations, + geom1, + geom2, + geomtype1, + geomtype2, + x1, + geom2.pos, + epa_vert, + epa_vert1, + epa_vert2, + epa_vert_index1, + epa_vert_index2, + epa_face, + epa_pr, + epa_norm2, + epa_index, + epa_map, + epa_horizon, + ) + + if ncontact == 0: + continue + + # cache contact information + hfield_contact_dist[count] = dist + + pos = 0.5 * (w1 + w2) + hfield_contact_pos[count, 0] = pos[0] + hfield_contact_pos[count, 1] = pos[1] + hfield_contact_pos[count, 2] = pos[2] + + frame = make_frame(w1 - w2) + normal = wp.vec3(frame[0, 0], frame[0, 1], frame[0, 2]) + hfield_contact_normal[count, 0] = normal[0] + hfield_contact_normal[count, 1] = normal[1] + hfield_contact_normal[count, 2] = normal[2] + + # contact with minimum distance + if dist < min_dist: + min_dist = dist + min_normal = normal + min_pos = pos + min_id = count + + count += 1 + + # contact 0: minimum distance + write_contact( + naconmax_in, + 0, + min_dist, + min_pos, + make_frame(min_normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + collision_pairid, + 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_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + + # TODO(team): routine for select subset of contacts + # TODO(team): if use_multiccd? + if wp.static(True): + MIN_DIST_TO_NEXT_CONTACT = 1.0e-3 + + # contact 1: furthest from minimum distance contact + id1 = int(-1) + dist1 = float(-wp.inf) + for i in range(count): + if i == min_id: + continue + + hf_pos = wp.vec3(hfield_contact_pos[i, 0], hfield_contact_pos[i, 1], hfield_contact_pos[i, 2]) + dist = wp.norm_l2(hf_pos - min_pos) + + if dist > dist1: + id1 = i + dist1 = dist + + if id1 == -1 or (0.0 < dist1 and dist1 < MIN_DIST_TO_NEXT_CONTACT): + return + + pos1 = wp.vec3(hfield_contact_pos[id1, 0], hfield_contact_pos[id1, 1], hfield_contact_pos[id1, 2]) + normal1 = wp.vec3(hfield_contact_normal[id1, 0], hfield_contact_normal[id1, 1], hfield_contact_normal[id1, 2]) + + write_contact( + naconmax_in, + 1, + hfield_contact_dist[id1], + pos1, + make_frame(normal1), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + collision_pairid, + 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_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + + # contact 2: point furthest from min_pos - pos1 line + dist_min1 = wp.cross(min_normal, min_pos - pos1) + + id2 = int(-1) + dist_12 = float(-wp.inf) + for i in range(count): + if i == min_id or i == id1: + continue + + hf_pos = wp.vec3(hfield_contact_pos[i, 0], hfield_contact_pos[i, 1], hfield_contact_pos[i, 2]) + dist = wp.abs(wp.dot(hf_pos - min_pos, dist_min1)) + + if dist > dist_12: + id2 = i + dist_12 = dist + + if id2 == -1 or (0.0 < dist_12 and dist_12 < MIN_DIST_TO_NEXT_CONTACT): + return + + pos2 = wp.vec3(hfield_contact_pos[id2, 0], hfield_contact_pos[id2, 1], hfield_contact_pos[id2, 2]) + normal2 = wp.vec3(hfield_contact_normal[id2, 0], hfield_contact_normal[id2, 1], hfield_contact_normal[id2, 2]) + + write_contact( + naconmax_in, + 2, + hfield_contact_dist[id2], + pos2, + make_frame(normal2), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + collision_pairid, + 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_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + + # contact 3: point furthest from other triangle edge + vec_min2 = wp.cross(min_normal, min_pos - pos2) + vec_12 = wp.cross(min_normal, pos1 - pos2) + + id3 = int(-1) + dist3 = float(-wp.inf) + for i in range(count): + if i == min_id or i == id1 or i == id2: + continue + + hf_pos = wp.vec3(hfield_contact_pos[i, 0], hfield_contact_pos[i, 1], hfield_contact_pos[i, 2]) + dist = wp.abs(wp.dot(hf_pos - min_pos, vec_min2)) + wp.abs(wp.dot(pos1 - hf_pos, vec_12)) + + if dist > dist3: + id3 = i + dist3 = dist + + if id3 == -1 or (0.0 < dist3 and dist3 < MIN_DIST_TO_NEXT_CONTACT): + return + + pos3 = wp.vec3(hfield_contact_pos[id3, 0], hfield_contact_pos[id3, 1], hfield_contact_pos[id3, 2]) + normal3 = wp.vec3(hfield_contact_normal[id3, 0], hfield_contact_normal[id3, 1], hfield_contact_normal[id3, 2]) + + write_contact( + naconmax_in, + 3, + hfield_contact_dist[id3], + pos3, + make_frame(normal3), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + collision_pairid, + 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_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + + return ccd_hfield_kernel + + @cache_kernel def ccd_kernel_builder( geomtype1: int, @@ -199,6 +710,8 @@ def ccd_kernel_builder( epa_iterations: int, use_multiccd: bool, ): + """Kernel builder for non-heightfield CCD collisions (no hfield args).""" + @wp.func def eval_ccd_write_contact( # Model: @@ -380,7 +893,7 @@ def ccd_kernel_builder( return ncontact - # runs convex collision on a set of geom pairs to recover contact info + # runs convex collision on a set of geom pairs to recover contact info (non-heightfield) @nested_kernel(module="unique", enable_backward=False) def ccd_kernel( # Model: @@ -393,16 +906,9 @@ def ccd_kernel_builder( geom_solref: wp.array2d(dtype=wp.vec2), geom_solimp: wp.array2d(dtype=vec5), geom_size: wp.array2d(dtype=wp.vec3), - geom_aabb: wp.array3d(dtype=wp.vec3), - geom_rbound: wp.array2d(dtype=float), geom_friction: wp.array2d(dtype=wp.vec3), geom_margin: wp.array2d(dtype=float), geom_gap: wp.array2d(dtype=float), - hfield_adr: wp.array(dtype=int), - hfield_nrow: wp.array(dtype=int), - hfield_ncol: wp.array(dtype=int), - hfield_size: wp.array(dtype=wp.vec4), - hfield_data: wp.array(dtype=float), mesh_vertadr: wp.array(dtype=int), mesh_vertnum: wp.array(dtype=int), mesh_vert: wp.array(dtype=wp.vec3), @@ -484,14 +990,6 @@ def ccd_kernel_builder( worldid = collision_worldid_in[tid] - # height field filter - if wp.static(geomtype1 == GeomType.HFIELD.value): - no_hf_collision, xmin, xmax, ymin, ymax, zmin, zmax = _hfield_filter( - geom_dataid, geom_aabb, geom_rbound, geom_margin, hfield_size, geom_xpos_in, geom_xmat_in, worldid, g1, g2 - ) - if no_hf_collision: - return - _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params( geom_condim, geom_priority, @@ -538,421 +1036,96 @@ def ccd_kernel_builder( worldid, ) - # see MuJoCo mjc_ConvexHField - if wp.static(geomtype1 == GeomType.HFIELD.value): - geom1_dataid = geom_dataid[g1] - - # height field subgrid - nrow = hfield_nrow[geom1_dataid] - ncol = hfield_ncol[geom1_dataid] - size = hfield_size[geom1_dataid] - - # subgrid - x_scale = 0.5 * float(ncol - 1) / size[0] - y_scale = 0.5 * float(nrow - 1) / size[1] - cmin = wp.max(0, int(wp.floor((xmin + size[0]) * x_scale))) - cmax = wp.min(ncol - 1, int(wp.ceil((xmax + size[0]) * x_scale))) - rmin = wp.max(0, int(wp.floor((ymin + size[1]) * y_scale))) - rmax = wp.min(nrow - 1, int(wp.ceil((ymax + size[1]) * y_scale))) - - dx = (2.0 * size[0]) / float(ncol - 1) - dy = (2.0 * size[1]) / float(nrow - 1) - dr = wp.vec2i(1, 0) - - prism = mat63() - - # set zbottom value using base size - prism[0, 2] = -size[3] - prism[1, 2] = -size[3] - prism[2, 2] = -size[3] - - adr = hfield_adr[geom1_dataid] - - hfield_contact_dist = vec_maxconpair() - hfield_contact_pos = mat_maxconpair() - hfield_contact_normal = mat_maxconpair() - min_dist = float(wp.inf) - min_normal = wp.vec3(wp.inf, wp.inf, wp.inf) - min_pos = wp.vec3(wp.inf, wp.inf, wp.inf) - min_id = int(-1) - - # TODO(team): height field margin? - geom1.margin = margin - geom2.margin = margin - - # EPA memory - epa_vert = epa_vert_in[tid] - epa_vert1 = epa_vert1_in[tid] - epa_vert2 = epa_vert2_in[tid] - epa_vert_index1 = epa_vert_index1_in[tid] - epa_vert_index2 = epa_vert_index2_in[tid] - epa_face = epa_face_in[tid] - epa_pr = epa_pr_in[tid] - epa_norm2 = epa_norm2_in[tid] - epa_index = epa_index_in[tid] - epa_map = epa_map_in[tid] - epa_horizon = epa_horizon_in[tid] - - collision_pairid = collision_pairid_in[tid] - - # process all prisms in subgrid - count = int(0) - for r in range(rmin, rmax): - nvert = int(0) - for c in range(cmin, cmax + 1): - # add both triangles from this cell - for i in range(2): - if count >= MJ_MAXCONPAIR: - wp.printf( - "height field collision overflow, number of collisions >= %u - please adjust resolution: \n decrease the number of hfield rows/cols or modify size of colliding geom\n", - MJ_MAXCONPAIR, - ) - continue - - # add vert - x = dx * float(c) - size[0] - y = dy * float(r + dr[i]) - size[1] - z = hfield_data[adr + (r + dr[i]) * ncol + c] * size[2] + margin - - prism[0] = prism[1] - prism[1] = prism[2] - prism[3] = prism[4] - prism[4] = prism[5] - - prism[2, 0] = x - prism[5, 0] = x - prism[2, 1] = y - prism[5, 1] = y - prism[5, 2] = z - - nvert += 1 - - if nvert <= 2: - continue - - # prism height test - if prism[3, 2] < zmin and prism[4, 2] < zmin and prism[5, 2] < zmin: - continue - - geom1.hfprism = prism - - # prism center - x1 = geom1.pos - x1_ = wp.vec3(0.0, 0.0, 0.0) - for i in range(6): - x1_ += prism[i] - x1 += geom1.rot @ (x1_ / 6.0) - - dist, ncontact, w1, w2, idx = ccd( - opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]], - 0.0, - gjk_iterations, - epa_iterations, - geom1, - geom2, - geomtype1, - geomtype2, - x1, - geom2.pos, - epa_vert, - epa_vert1, - epa_vert2, - epa_vert_index1, - epa_vert_index2, - epa_face, - epa_pr, - epa_norm2, - epa_index, - epa_map, - epa_horizon, - ) - - if ncontact == 0: - continue - - # cache contact information - hfield_contact_dist[count] = dist - - pos = 0.5 * (w1 + w2) - hfield_contact_pos[count, 0] = pos[0] - hfield_contact_pos[count, 1] = pos[1] - hfield_contact_pos[count, 2] = pos[2] - - frame = make_frame(w1 - w2) - normal = wp.vec3(frame[0, 0], frame[0, 1], frame[0, 2]) - hfield_contact_normal[count, 0] = normal[0] - hfield_contact_normal[count, 1] = normal[1] - hfield_contact_normal[count, 2] = normal[2] - - # contact with minimum distance - if dist < min_dist: - min_dist = dist - min_normal = normal - min_pos = pos - min_id = count - - count += 1 - - # contact 0: minimum distance - write_contact( - naconmax_in, - 0, - min_dist, - min_pos, - make_frame(min_normal), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - collision_pairid, - 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_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, - ) - - # TODO(team): routine for select subset of contacts - # TODO(team): if use_multiccd? - if wp.static(True): - MIN_DIST_TO_NEXT_CONTACT = 1.0e-3 - - # contact 1: furthest from minimum distance contact - id1 = int(-1) - dist1 = float(-wp.inf) - for i in range(count): - if i == min_id: - continue - - hf_pos = wp.vec3(hfield_contact_pos[i, 0], hfield_contact_pos[i, 1], hfield_contact_pos[i, 2]) - dist = wp.norm_l2(hf_pos - min_pos) - - if dist > dist1: - id1 = i - dist1 = dist - - if id1 == -1 or (0.0 < dist1 and dist1 < MIN_DIST_TO_NEXT_CONTACT): - return - - pos1 = wp.vec3(hfield_contact_pos[id1, 0], hfield_contact_pos[id1, 1], hfield_contact_pos[id1, 2]) - normal1 = wp.vec3(hfield_contact_normal[id1, 0], hfield_contact_normal[id1, 1], hfield_contact_normal[id1, 2]) - - write_contact( - naconmax_in, - 1, - hfield_contact_dist[id1], - pos1, - make_frame(normal1), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - collision_pairid, - 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_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, - ) - - # contact 2: point furthest from min_pos - pos1 line - dist_min1 = wp.cross(min_normal, min_pos - pos1) - - id2 = int(-1) - dist_12 = float(-wp.inf) - for i in range(count): - if i == min_id or i == id1: - continue - - hf_pos = wp.vec3(hfield_contact_pos[i, 0], hfield_contact_pos[i, 1], hfield_contact_pos[i, 2]) - dist = wp.abs(wp.dot(hf_pos - min_pos, dist_min1)) - - if dist > dist_12: - id2 = i - dist_12 = dist - - if id2 == -1 or (0.0 < dist_12 and dist_12 < MIN_DIST_TO_NEXT_CONTACT): - return - - pos2 = wp.vec3(hfield_contact_pos[id2, 0], hfield_contact_pos[id2, 1], hfield_contact_pos[id2, 2]) - normal2 = wp.vec3(hfield_contact_normal[id2, 0], hfield_contact_normal[id2, 1], hfield_contact_normal[id2, 2]) - - write_contact( - naconmax_in, - 2, - hfield_contact_dist[id2], - pos2, - make_frame(normal2), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - collision_pairid, - 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_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, - ) - - # contact 3: point furthest from other triangle edge - vec_min2 = wp.cross(min_normal, min_pos - pos2) - vec_12 = wp.cross(min_normal, pos1 - pos2) - - id3 = int(-1) - dist3 = float(-wp.inf) - for i in range(count): - if i == min_id or i == id1 or i == id2: - continue - - hf_pos = wp.vec3(hfield_contact_pos[i, 0], hfield_contact_pos[i, 1], hfield_contact_pos[i, 2]) - dist = wp.abs(wp.dot(hf_pos - min_pos, vec_min2)) + wp.abs(wp.dot(pos1 - hf_pos, vec_12)) - - if dist > dist3: - id3 = i - dist3 = dist - - if id3 == -1 or (0.0 < dist3 and dist3 < MIN_DIST_TO_NEXT_CONTACT): - return - - pos3 = wp.vec3(hfield_contact_pos[id3, 0], hfield_contact_pos[id3, 1], hfield_contact_pos[id3, 2]) - normal3 = wp.vec3(hfield_contact_normal[id3, 0], hfield_contact_normal[id3, 1], hfield_contact_normal[id3, 2]) - - write_contact( - naconmax_in, - 3, - hfield_contact_dist[id3], - pos3, - make_frame(normal3), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - collision_pairid, - 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_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, - ) - else: - eval_ccd_write_contact( - opt_ccd_tolerance, - geom_type, - naconmax_in, - epa_vert_in, - epa_vert1_in, - epa_vert2_in, - epa_vert_index1_in, - epa_vert_index2_in, - epa_face_in, - epa_pr_in, - epa_norm2_in, - epa_index_in, - epa_map_in, - epa_horizon_in, - multiccd_polygon_in, - multiccd_clipped_in, - multiccd_pnormal_in, - multiccd_pdist_in, - multiccd_idx1_in, - multiccd_idx2_in, - multiccd_n1_in, - multiccd_n2_in, - multiccd_endvert_in, - multiccd_face1_in, - multiccd_face2_in, - geom1, - geom2, - geoms, - worldid, - tid, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geom1.pos, - geom2.pos, - 0, - collision_pairid_in[tid], - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - contact_type_out, - contact_geomcollisionid_out, - nacon_out, - ) + eval_ccd_write_contact( + opt_ccd_tolerance, + geom_type, + naconmax_in, + epa_vert_in, + epa_vert1_in, + epa_vert2_in, + epa_vert_index1_in, + epa_vert_index2_in, + epa_face_in, + epa_pr_in, + epa_norm2_in, + epa_index_in, + epa_map_in, + epa_horizon_in, + multiccd_polygon_in, + multiccd_clipped_in, + multiccd_pnormal_in, + multiccd_pdist_in, + multiccd_idx1_in, + multiccd_idx2_in, + multiccd_n1_in, + multiccd_n2_in, + multiccd_endvert_in, + multiccd_face1_in, + multiccd_face2_in, + geom1, + geom2, + geoms, + worldid, + tid, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geom1.pos, + geom2.pos, + 0, + collision_pairid_in[tid], + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) return ccd_kernel +# Heightfield collision pairs handled by ccd_hfield_kernel_builder +_HFIELD_COLLISION_PAIRS = [ + (GeomType.HFIELD, GeomType.SPHERE), + (GeomType.HFIELD, GeomType.CAPSULE), + (GeomType.HFIELD, GeomType.ELLIPSOID), + (GeomType.HFIELD, GeomType.CYLINDER), + (GeomType.HFIELD, GeomType.BOX), + (GeomType.HFIELD, GeomType.MESH), +] + +# Non-heightfield collision pairs handled by ccd_kernel_builder +_NON_HFIELD_COLLISION_PAIRS = [ + (GeomType.SPHERE, GeomType.ELLIPSOID), + (GeomType.SPHERE, GeomType.MESH), + (GeomType.CAPSULE, GeomType.ELLIPSOID), + (GeomType.CAPSULE, GeomType.CYLINDER), + (GeomType.CAPSULE, GeomType.MESH), + (GeomType.ELLIPSOID, GeomType.ELLIPSOID), + (GeomType.ELLIPSOID, GeomType.CYLINDER), + (GeomType.ELLIPSOID, GeomType.BOX), + (GeomType.ELLIPSOID, GeomType.MESH), + (GeomType.CYLINDER, GeomType.CYLINDER), + (GeomType.CYLINDER, GeomType.BOX), + (GeomType.CYLINDER, GeomType.MESH), + (GeomType.BOX, GeomType.MESH), + (GeomType.MESH, GeomType.MESH), +] + + @event_scope def convex_narrowphase(m: Model, d: Data): """Runs narrowphase collision detection for convex geom pairs. @@ -1006,35 +1179,32 @@ def convex_narrowphase(m: Model, d: Data): epa_map = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) # epa_horizon: index pair (i j) of edges on horizon epa_horizon = wp.empty(shape=(d.naconmax, 2 * MJ_MAX_EPAHORIZON), dtype=int) - # multiccd_polygon: clipped contact surface - multiccd_polygon = wp.empty(shape=(d.naconmax, 2 * nmaxpolygon), dtype=wp.vec3) - # multiccd_clipped: clipped contact surface (intermediate) - multiccd_clipped = wp.empty(shape=(d.naconmax, 2 * nmaxpolygon), dtype=wp.vec3) - # multiccd_pnormal: plane normal of clipping polygon - multiccd_pnormal = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) - # multiccd_pdist: plane distance of clipping polygon - multiccd_pdist = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=float) - # multiccd_idx1: list of normal index candidates for Geom 1 - multiccd_idx1 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=int) - # multiccd_idx2: list of normal index candidates for Geom 2 - multiccd_idx2 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=int) - # multiccd_n1: list of normal candidates for Geom 1 - multiccd_n1 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) - # multiccd_n2: list of normal candidates for Geom 1 - multiccd_n2 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) - # multiccd_endvert: list of edge vertices candidates - multiccd_endvert = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) - # multiccd_face1: contact face - multiccd_face1 = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) - # multiccd_face2: contact face - multiccd_face2 = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) - for geom_pair in _CONVEX_COLLISION_PAIRS: + # Contact outputs + contact_outputs = [ + d.nacon, + 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.worldid, + d.contact.type, + d.contact.geomcollisionid, + ] + + # Launch heightfield collision kernels (no multiccd args, 72 args total) + for geom_pair in _HFIELD_COLLISION_PAIRS: g1 = geom_pair[0].value g2 = geom_pair[1].value if _pair_count(g1, g2): wp.launch( - ccd_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations, use_multiccd), + ccd_hfield_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations), dim=d.naconmax, inputs=[ m.opt.ccd_tolerance, @@ -1095,6 +1265,94 @@ def convex_narrowphase(m: Model, d: Data): epa_index, epa_map, epa_horizon, + ], + outputs=contact_outputs, + ) + + # Allocate multiccd arrays only for non-heightfield collisions + # multiccd_polygon: clipped contact surface + multiccd_polygon = wp.empty(shape=(d.naconmax, 2 * nmaxpolygon), dtype=wp.vec3) + # multiccd_clipped: clipped contact surface (intermediate) + multiccd_clipped = wp.empty(shape=(d.naconmax, 2 * nmaxpolygon), dtype=wp.vec3) + # multiccd_pnormal: plane normal of clipping polygon + multiccd_pnormal = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + # multiccd_pdist: plane distance of clipping polygon + multiccd_pdist = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=float) + # multiccd_idx1: list of normal index candidates for Geom 1 + multiccd_idx1 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=int) + # multiccd_idx2: list of normal index candidates for Geom 2 + multiccd_idx2 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=int) + # multiccd_n1: list of normal candidates for Geom 1 + multiccd_n1 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + # multiccd_n2: list of normal candidates for Geom 1 + multiccd_n2 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + # multiccd_endvert: list of edge vertices candidates + multiccd_endvert = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + # multiccd_face1: contact face + multiccd_face1 = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + # multiccd_face2: contact face + multiccd_face2 = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + + # Launch non-heightfield collision kernels (no hfield args, 78 args total) + for geom_pair in _NON_HFIELD_COLLISION_PAIRS: + g1 = geom_pair[0].value + g2 = geom_pair[1].value + if _pair_count(g1, g2): + wp.launch( + ccd_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations, use_multiccd), + dim=d.naconmax, + inputs=[ + m.opt.ccd_tolerance, + 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.mesh_vertadr, + m.mesh_vertnum, + m.mesh_vert, + m.mesh_graphadr, + m.mesh_graph, + m.mesh_polynum, + m.mesh_polyadr, + m.mesh_polynormal, + m.mesh_polyvertadr, + m.mesh_polyvertnum, + m.mesh_polyvert, + m.mesh_polymapadr, + m.mesh_polymapnum, + m.mesh_polymap, + m.pair_dim, + m.pair_solref, + m.pair_solreffriction, + m.pair_solimp, + m.pair_margin, + m.pair_gap, + m.pair_friction, + d.naconmax, + d.geom_xpos, + d.geom_xmat, + d.collision_pair, + d.collision_pairid, + d.collision_worldid, + d.ncollision, + epa_vert, + epa_vert1, + epa_vert2, + epa_vert_index1, + epa_vert_index2, + epa_face, + epa_pr, + epa_norm2, + epa_index, + epa_map, + epa_horizon, multiccd_polygon, multiccd_clipped, multiccd_pnormal, @@ -1107,20 +1365,5 @@ def convex_narrowphase(m: Model, d: Data): multiccd_face1, multiccd_face2, ], - outputs=[ - d.nacon, - 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.worldid, - d.contact.type, - d.contact.geomcollisionid, - ], + outputs=contact_outputs, ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py index bc2e2513..e8ca7e6f 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py @@ -618,8 +618,7 @@ def gjk( return result elif cutoff < FLOAT_MAX: vs = wp.dot(x_k, simplex[n]) - vv = wp.dot(x_k, x_k) - if wp.dot(x_k, simplex[n]) > 0 and (vs * vs / vv) >= cutoff2: + if wp.dot(x_k, simplex[n]) > 0 and (vs * vs / xnorm) >= cutoff2: result = GJKResult() result.dim = 0 result.dist = FLOAT_MAX @@ -906,22 +905,34 @@ def _epa_witness( if geomtype1 == GeomType.HFIELD and (i1 != i2 or i1 != i3): # TODO(kbayes): Fix case where geom2 is near bottom of height field or "extreme" prism heights n = geom1.rot[:, 2] - a = geom1.hfprism[3] - b = geom1.hfprism[4] - c = geom1.hfprism[5] - x2 = wp.normalize(x2) + + # height field prism vertices in global frame + a = geom1.pos + geom1.rot @ geom1.hfprism[3] + b = geom1.pos + geom1.rot @ geom1.hfprism[4] + c = geom1.pos + geom1.rot @ geom1.hfprism[5] # TODO(kbayes): Support cases where geom2 is larger than the height field - sp = _support(geom2, geomtype2, x2) - x2 = sp.point + if geomtype2 == GeomType.CAPSULE or geomtype2 == GeomType.SPHERE: + radius = geom2.size[0] + margin = geom2.margin + geom2.margin = 0.0 + geom2.size = wp.vec3(0.0, geom2.size[1], geom2.size[2]) + sp = _support(geom2, geomtype2, x2) + x2 = sp.point - (0.5 * margin + radius) * n + geom2.size[0] = radius + geom2.margin = margin + else: + x2 = wp.normalize(x2) + sp = _support(geom2, geomtype2, x2) + x2 = sp.point coordinates2 = _tri_affine_coord(a, b, c, x2) if coordinates2[0] > 0 and coordinates2[1] > 0 and coordinates2[2] > 0: - x1 = coordinates[0] * a + coordinates[1] * b + coordinates[2] * c + x1 = coordinates2[0] * a + coordinates2[1] * b + coordinates2[2] * c else: p = c - p = wp.where(coordinates[1] > 0, b, p) - p = wp.where(coordinates[0] > 0, a, p) + p = wp.where(coordinates2[1] > 0, b, p) + p = wp.where(coordinates2[0] > 0, a, p) x1 = x2 - wp.dot(x2 - p, n) * n return x1, x2, -wp.norm_l2(x1 - x2) @@ -1256,13 +1267,14 @@ def _epa( # compute support point w from the closest face's normal lower = wp.sqrt(lower2) wi = pt.nvert - i1, i2 = _epa_support(pt, wi, geom1, geom2, geomtype1, geomtype2, pt.face_pr[idx] / lower) + face_pr_normalized = pt.face_pr[idx] / lower + i1, i2 = _epa_support(pt, wi, geom1, geom2, geomtype1, geomtype2, face_pr_normalized) geom1.index = i1 geom2.index = i2 pt.nvert += 1 # upper bound for kth iteration - upper_k = wp.dot(pt.face_pr[idx], pt.vert[wi]) / lower + upper_k = wp.dot(face_pr_normalized, pt.vert[wi]) if upper_k < upper: upper = upper_k upper2 = upper * upper @@ -2162,7 +2174,40 @@ def multicontact( @wp.func -def _inflate(dist: float, x1: wp.vec3, x2: wp.vec3, margin1: float, margin2: float) -> Tuple[float, wp.vec3, wp.vec3]: +def _inflate( + result: GJKResult, geom1: Geom, geom2: Geom, geomtype1: int, geomtype2: int, margin1: float, margin2: float +) -> Tuple[float, wp.vec3, wp.vec3]: + dist = result.dist + x1 = result.x1 + x2 = result.x2 + + if geomtype1 == GeomType.HFIELD: + v = result.simplex_index1[0] + is_side = bool(False) + for i in range(result.dim): + if result.simplex_index1[i] != v: + is_side = True + break + + if is_side: + n = geom1.rot[:, 2] + sp = _support(geom2, geomtype2, x2) + x2 = sp.point - margin2 * n + + a = geom1.hfprism[3] + b = geom1.hfprism[4] + c = geom1.hfprism[5] + coordinates = _tri_affine_coord(a, b, c, x2) + if coordinates[0] > 0 and coordinates[1] > 0 and coordinates[2] > 0: + x1 = coordinates[0] * a + coordinates[1] * b + coordinates[2] * c + else: + p = c + p = wp.where(coordinates[1] > 0, b, p) + p = wp.where(coordinates[0] > 0, a, p) + x1 = x2 - wp.dot(x2 - p, n) * n + dist = -wp.norm_l2(x1 - x2) + return dist, x1, x2 + n = wp.normalize(x2 - x1) if margin1 > 0.0: x1 += margin1 * n @@ -2213,8 +2258,7 @@ def ccd( geom1.margin = 0.0 geom1.size = wp.vec3(0.0, geom1.size[1], geom1.size[2]) - # TODO(kbayes): support gjk margin trick with height fields - if geomtype1 != GeomType.HFIELD and (geomtype2 == GeomType.SPHERE or geomtype2 == GeomType.CAPSULE): + if geomtype2 == GeomType.SPHERE or geomtype2 == GeomType.CAPSULE: size2 = geom2.size[0] full_margin2 = size2 + 0.5 * geom2.margin geom2.margin = 0.0 @@ -2229,7 +2273,7 @@ def ccd( if result.dist > tolerance: if result.dist == FLOAT_MAX: return result.dist, 1, result.x1, result.x2, -1 - dist, x1, x2 = _inflate(result.dist, result.x1, result.x2, full_margin1, full_margin2) + dist, x1, x2 = _inflate(result, geom1, geom2, geomtype1, geomtype2, full_margin1, full_margin2) return dist, 1, x1, x2, -1 # deep penetration, reset initial conditions and rerun GJK + EPA diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py index 287becad..5cf03c7d 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py @@ -535,11 +535,11 @@ def contact_params( max_geom_friction = wp.max(geom_friction[friction_id, g1], geom_friction[friction_id, g2]) friction = vec5( - wp.max(MJ_MINMU, max_geom_friction[0]), - wp.max(MJ_MINMU, max_geom_friction[0]), - wp.max(MJ_MINMU, max_geom_friction[1]), - wp.max(MJ_MINMU, max_geom_friction[2]), - wp.max(MJ_MINMU, max_geom_friction[2]), + max_geom_friction[0], + max_geom_friction[0], + max_geom_friction[1], + max_geom_friction[2], + max_geom_friction[2], ) if geom_solref[solref_id, g1][0] > 0.0 and geom_solref[solref_id, g2][0] > 0.0: @@ -553,6 +553,14 @@ def contact_params( margin = wp.max(geom_margin[margin_id, g1], geom_margin[margin_id, g2]) gap = wp.max(geom_gap[gap_id, g1], geom_gap[gap_id, g2]) + friction = vec5( + wp.max(MJ_MINMU, friction[0]), + wp.max(MJ_MINMU, friction[1]), + wp.max(MJ_MINMU, friction[2]), + wp.max(MJ_MINMU, friction[3]), + wp.max(MJ_MINMU, friction[4]), + ) + return geoms, margin, gap, condim, friction, solref, solreffriction, solimp diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py index 4e855036..f9392cad 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py @@ -367,7 +367,7 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: Volume elif type == GeomType.MESH and mesh_data.valid: mesh_data.pnt = p mesh_data.vec = -wp.normalize(p) - dist = ray_mesh( + dist, normal = ray_mesh( mesh_data.nmeshface, mesh_data.mesh_vertadr, mesh_data.mesh_faceadr, @@ -380,7 +380,7 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: Volume mesh_data.vec, ) if dist > wp.norm_l2(p): - return -ray_mesh( + dist, normal = ray_mesh( mesh_data.nmeshface, mesh_data.mesh_vertadr, mesh_data.mesh_faceadr, @@ -392,6 +392,7 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: Volume mesh_data.pnt, -mesh_data.vec, ) + return -dist return dist elif type == GeomType.SDF: if sdf_type == -1: @@ -416,7 +417,7 @@ def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: V elif type == GeomType.MESH and mesh_data.valid: mesh_data.pnt = p mesh_data.vec = -wp.normalize(p) - dist = ray_mesh( + dist, normal = ray_mesh( mesh_data.nmeshface, mesh_data.mesh_vertadr, mesh_data.mesh_faceadr, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py index 210c7541..13468c8c 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py @@ -333,7 +333,7 @@ def _efc_equality_joint( # Two joint constraint qposadr2 = jnt_qposadr[jntid_2] dofadr2 = jnt_dofadr[jntid_2] - dif = qpos_in[worldid, qposadr2] - qpos0[worldid, qposadr2] + dif = qpos_in[worldid, qposadr2] - qpos0[qpos0_id, qposadr2] # Horner's method for polynomials rhs = data[0] + dif * (data[1] + dif * (data[2] + dif * (data[3] + dif * data[4]))) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py index faedac78..a902247e 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -560,6 +560,32 @@ def _get_padded_sizes(nv: int, njmax: int, is_sparse: bool, tile_size: int): return njmax_padded, nv_padded +def _default_nconmax(mjm: mujoco.MjModel, mjd: Optional[mujoco.MjData] = None) -> int: + """Returns a default guess for an ideal nconmax given a Model and optional Data. + + This guess is based off a very simple heuristic, and may need to be manually raised if MJWarp + reports ncon overflow, or lowered in order to get the very best performance. + """ + valid_sizes = (2 + (np.arange(19) % 2)) * (2 ** (np.arange(19) // 2 + 3)) # 16, 24, 32, 48, ... 8192 + has_sdf = (mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF).any() + has_flex = mjm.nflex > 0 + nconmax = max(mjm.nv * 0.35 * (mjm.nhfield > 0) * 10 + 45, 256 * has_flex, 64 * has_sdf, mjd.ncon if mjd else 0) + return int(valid_sizes[np.searchsorted(valid_sizes, nconmax)]) + + +def _default_njmax(mjm: mujoco.MjModel, mjd: Optional[mujoco.MjData] = None) -> int: + """Returns a default guess for an ideal njmax given a Model and optional Data. + + This guess is based off a very simple heuristic, and may need to be manually raised if MJWarp + reports ncon overflow, or lowered in order to get the very best performance. + """ + valid_sizes = (2 + (np.arange(19) % 2)) * (2 ** (np.arange(19) // 2 + 3)) # 16, 24, 32, 48, ... 8192 + has_sdf = (mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF).any() + has_flex = mjm.nflex > 0 + njmax = max(mjm.nv * 2.26 * (mjm.nhfield > 0) * 18 + 53, 512 * has_flex, 256 * has_sdf, mjd.nefc if mjd else 0) + return int(valid_sizes[np.searchsorted(valid_sizes, njmax)]) + + def make_data( mjm: mujoco.MjModel, nworld: int = 1, @@ -582,9 +608,11 @@ def make_data( The data object containing the current state and output arrays (device). """ # TODO(team): move nconmax, njmax to Model? - # TODO(team): improve heuristic for nconmax and njmax - nconmax = nconmax or 20 - njmax = njmax or nconmax * 6 + if nconmax is None: + nconmax = _default_nconmax(mjm) + + if njmax is None: + njmax = _default_njmax(mjm) if nworld < 1: raise ValueError(f"nworld must be >= 1") @@ -592,7 +620,7 @@ def make_data( if naconmax is None: if nconmax < 0: raise ValueError("nconmax must be >= 0") - naconmax = max(512, nworld * nconmax) + naconmax = nworld * nconmax elif naconmax < 0: raise ValueError("naconmax must be >= 0") @@ -689,9 +717,11 @@ def put_data( # TODO(team): decide what to do about uninitialized warp-only fields created by put_data # we need to ensure these are only workspace fields and don't carry state - # TODO(team): better heuristic for nconmax and njmax - nconmax = nconmax or max(5, 4 * mjd.ncon) - njmax = njmax or max(5, 4 * mjd.nefc) + if nconmax is None: + nconmax = _default_nconmax(mjm, mjd) + + if njmax is None: + njmax = _default_njmax(mjm, mjd) if nworld < 1: raise ValueError(f"nworld must be >= 1") @@ -699,11 +729,9 @@ def put_data( if naconmax is None: if nconmax < 0: raise ValueError("nconmax must be >= 0") - if mjd.ncon > nconmax: raise ValueError(f"nconmax overflow (nconmax must be >= {mjd.ncon})") - - naconmax = max(512, nworld * nconmax) + naconmax = nworld * nconmax elif naconmax < mjd.ncon * nworld: raise ValueError(f"naconmax overflow (naconmax must be >= {mjd.ncon * nworld})") @@ -845,8 +873,9 @@ def put_data( d.actuator_moment = wp.array(np.full((nworld, mjm.nu, mjm.nv), actuator_moment), dtype=float) d.nacon = wp.array([mjd.ncon * nworld], dtype=int) - d.ne_connect = wp.full(nworld, 3 * np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_CONNECT) & mjd.eq_active), dtype=int) - d.ne_weld = wp.full(nworld, 6 * np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_WELD) & mjd.eq_active), dtype=int) + + d.ne_connect = wp.full(nworld, 3 * int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_CONNECT) & mjd.eq_active)), dtype=int) + d.ne_weld = wp.full(nworld, 6 * int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_WELD) & mjd.eq_active)), dtype=int) d.ne_jnt = wp.full(nworld, np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_JOINT) & mjd.eq_active), dtype=int) d.ne_ten = wp.full(nworld, np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_TENDON) & mjd.eq_active), dtype=int) d.ne_flex = wp.full(nworld, np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_FLEX) & mjd.eq_active), dtype=int) @@ -1324,6 +1353,7 @@ def override_model(model: Union[types.Model, mujoco.MjModel], overrides: Union[d } mjw_only_fields = {"opt.broadphase", "opt.broadphase_filter", "opt.ls_parallel", "opt.graph_conditional"} mj_only_fields = {"opt.jacobian"} + readonly_fields = {"opt.is_sparse"} if not isinstance(overrides, dict): overrides_dict = {} @@ -1341,6 +1371,9 @@ def override_model(model: Union[types.Model, mujoco.MjModel], overrides: Union[d if key in mj_only_fields and isinstance(model, types.Model): continue + if key in readonly_fields and isinstance(model, types.Model): + raise ValueError(f"Cannot override {key} on mjw.Model: field affects model initialization and has side effects") + obj, attrs = model, key.split(".") for i, attr in enumerate(attrs): if not hasattr(obj, attr): diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py index d3b7d063..ec49041e 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py @@ -58,20 +58,28 @@ def axis_angle_to_quat(axis: wp.vec3, angle: float) -> wp.quat: @wp.func def quat_to_mat(quat: wp.quat) -> wp.mat33: - """Converts a quaternion into a 9-dimensional rotation matrix.""" - vec = wp.vec4(quat[0], quat[1], quat[2], quat[3]) - q = wp.outer(vec, vec) + """Converts a quaternion into 3x3 rotation matrix.""" + q00 = quat[0] * quat[0] + q01 = quat[0] * quat[1] + q02 = quat[0] * quat[2] + q03 = quat[0] * quat[3] + q11 = quat[1] * quat[1] + q12 = quat[1] * quat[2] + q13 = quat[1] * quat[3] + q22 = quat[2] * quat[2] + q23 = quat[2] * quat[3] + q33 = quat[3] * quat[3] return wp.mat33( - q[0, 0] + q[1, 1] - q[2, 2] - q[3, 3], - 2.0 * (q[1, 2] - q[0, 3]), - 2.0 * (q[1, 3] + q[0, 2]), - 2.0 * (q[1, 2] + q[0, 3]), - q[0, 0] - q[1, 1] + q[2, 2] - q[3, 3], - 2.0 * (q[2, 3] - q[0, 1]), - 2.0 * (q[1, 3] - q[0, 2]), - 2.0 * (q[2, 3] + q[0, 1]), - q[0, 0] - q[1, 1] - q[2, 2] + q[3, 3], + q00 + q11 - q22 - q33, + 2.0 * (q12 - q03), + 2.0 * (q13 + q02), + 2.0 * (q12 + q03), + q00 - q11 + q22 - q33, + 2.0 * (q23 - q01), + 2.0 * (q13 - q02), + 2.0 * (q23 + q01), + q00 - q11 - q22 + q33, ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py index 61127a8b..3ee4ffb3 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -296,6 +296,12 @@ def _fluid_force( fluid_applied_out[worldid, bodyid] = zero_force return + # skip bodies with negligible mass + mass = body_mass[worldid % body_mass.shape[0], bodyid] + if mass < MJ_MINVAL: + fluid_applied_out[worldid, bodyid] = zero_force + return + wind = opt_wind[worldid % opt_wind.shape[0]] density = opt_density[worldid % opt_density.shape[0]] viscosity = opt_viscosity[worldid % opt_viscosity.shape[0]] diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py index 15959db7..41267c73 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py @@ -105,7 +105,7 @@ def _ray_quad(a: float, b: float, c: float) -> Tuple[float, wp.vec2]: """Compute solutions from quadratic: a*x^2 + 2*b*x + c = 0.""" det = b * b - a * c if det < MJ_MINVAL: - return wp.inf, wp.vec2(wp.inf, wp.inf) + return -1.0, wp.vec2(-1.0, -1.0) det = wp.sqrt(det) # compute the two solutions @@ -120,12 +120,14 @@ def _ray_quad(a: float, b: float, c: float) -> Tuple[float, wp.vec2]: elif x1 >= 0.0: return x1, x else: - return wp.inf, x + return -1.0, x @wp.func -def _ray_triangle(v0: wp.vec3, v1: wp.vec3, v2: wp.vec3, pnt: wp.vec3, vec: wp.vec3, b0: wp.vec3, b1: wp.vec3) -> float: - """Returns the distance at which a ray intersects with a triangle.""" +def _ray_triangle( + v0: wp.vec3, v1: wp.vec3, v2: wp.vec3, pnt: wp.vec3, vec: wp.vec3, b0: wp.vec3, b1: wp.vec3 +) -> Tuple[float, wp.vec3]: + """Returns the distance and normal at which a ray intersects with a triangle.""" dif0 = v0 - pnt dif1 = v1 - pnt dif2 = v2 - pnt @@ -145,7 +147,7 @@ def _ray_triangle(v0: wp.vec3, v1: wp.vec3, v2: wp.vec3, pnt: wp.vec3, vec: wp.v or (planar_01 > 0.0 and planar_11 > 0.0 and planar_21 > 0.0) or (planar_01 < 0.0 and planar_11 < 0.0 and planar_21 < 0.0) ): - return float(wp.inf) + return -1.0, wp.vec3() # determine if origin is inside planar projection of triangle # A = (p0-p2, p1-p2), b = -p2, solve A*t = b @@ -158,14 +160,14 @@ def _ray_triangle(v0: wp.vec3, v1: wp.vec3, v2: wp.vec3, pnt: wp.vec3, vec: wp.v det = A00 * A11 - A10 * A01 if wp.abs(det) < MJ_MINVAL: - return float(wp.inf) + return -1.0, wp.vec3() t0 = (A11 * b[0] - A10 * b[1]) / det t1 = (-A01 * b[0] + A00 * b[1]) / det # check if outside if t0 < 0.0 or t1 < 0.0 or t0 + t1 > 1.0: - return float(wp.inf) + return -1.0, wp.vec3() # intersect ray with plane of triangle dif0 = v0 - v2 @@ -174,39 +176,39 @@ def _ray_triangle(v0: wp.vec3, v1: wp.vec3, v2: wp.vec3, pnt: wp.vec3, vec: wp.v nrm = wp.cross(dif0, dif1) # normal to triangle plane denom = wp.dot(vec, nrm) if wp.abs(denom) < MJ_MINVAL: - return float(wp.inf) + return -1.0, wp.vec3() dist = -wp.dot(dif2, nrm) / denom - return wp.where(dist >= 0.0, dist, float(wp.inf)) + return wp.where(dist >= 0.0, dist, -1.0), wp.normalize(nrm) @wp.func -def _ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: - """Returns the distance at which a ray intersects with a plane.""" +def _ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: + """Returns the distance and normal at which a ray intersects with a plane.""" # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) # z-vec not pointing towards front face: reject if lvec[2] > -MJ_MINVAL: - return wp.inf + return -1.0, wp.vec3() # intersection with plane x = -lpnt[2] / lvec[2] if x < 0.0: - return wp.inf + return -1.0, wp.vec3() p = wp.vec2(lpnt[0] + x * lvec[0], lpnt[1] + x * lvec[1]) # accept only within rendered rectangle if (size[0] <= 0.0 or wp.abs(p[0]) <= size[0]) and (size[1] <= 0.0 or wp.abs(p[1]) <= size[1]): - return x + return x, wp.vec3(mat[0, 2], mat[1, 2], mat[2, 2]) else: - return wp.inf + return -1.0, wp.vec3() @wp.func -def _ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> float: - """Returns the distance at which a ray intersects with a sphere.""" +def _ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: + """Returns the distance and normal at which a ray intersects with a sphere.""" dif = pnt - pos a = wp.dot(vec, vec) @@ -214,16 +216,21 @@ def _ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> fl c = wp.dot(dif, dif) - dist_sqr sol, _ = _ray_quad(a, b, c) - return sol + normal = wp.vec3() + if sol >= 0: + s = pnt + vec * sol + normal = wp.normalize(s - pos) + return sol, normal @wp.func -def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: - """Returns the distance at which a ray intersects with a capsule.""" +def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: + """Returns the distance and normal at which a ray intersects with a capsule.""" # bounding sphere test ssz = size[0] + size[1] - if _ray_sphere(pos, ssz * ssz, pnt, vec) < 0.0: - return wp.inf + dist_sphere, normal_sphere = _ray_sphere(pos, ssz * ssz, pnt, vec) + if dist_sphere < 0: + return -1.0, wp.vec3() # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) @@ -239,6 +246,7 @@ def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: # solve a * x^2 + 2 * b * x + c = 0 sol, xx = _ray_quad(a, b, c) + part = 0 # -1: bottom, 0: cylinder, 1: top # make sure round solution is between flat sides if sol >= 0.0 and wp.abs(lpnt[2] + sol * vec[2]) <= size[1]: @@ -257,6 +265,7 @@ def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: if xx[i] >= 0.0 and lpnt[2] + xx[i] * lvec[2] >= size[1]: if x < 0.0 or xx[i] < x: x = xx[i] + part = 1 # bottom cap ldif = wp.vec3(ldif[0], ldif[1], lpnt[2] + size[1]) @@ -269,13 +278,27 @@ def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: if xx[i] >= 0.0 and lpnt[2] + xx[i] * lvec[2] <= -size[1]: if x < 0.0 or xx[i] < x: x = xx[i] + part = -1 - return x + normal = wp.vec3() + if x >= 0: + normal[0] = lpnt[0] + lvec[0] * x + normal[1] = lpnt[1] + lvec[1] * x + if part == 0: + normal[2] = 0.0 + else: + normal[2] = lpnt[2] + lvec[2] * x - size[1] * float(part) + + # normalize, rotate into global frame + normal = wp.normalize(normal) + normal = mat @ normal + + return x, normal @wp.func -def _ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: - """Returns the distance at which a ray intersects with an ellipsoid.""" +def _ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: + """Returns the distance and normal at which a ray intersects with an ellipsoid.""" # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) @@ -290,22 +313,35 @@ def _ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec # solve a * x^2 + 2 * b * x + c = 0 sol, _ = _ray_quad(a, b, c) - return sol + + normal = wp.vec3() + if sol >= 0: + # surface intersection (local frame) + l = lpnt + lvec * sol + + # gradient of ellipsoid function + normal = wp.cw_mul(s, l) + normal = wp.normalize(normal) + normal = mat @ normal + + return sol, normal @wp.func -def _ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: - """Returns the distance at which a ray intersects with a cylinder.""" +def _ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: + """Returns the distance and normal at which a ray intersects with a cylinder.""" # bounding sphere test ssz = size[0] * size[0] + size[1] * size[1] - if _ray_sphere(pos, ssz, pnt, vec) < 0.0: - return wp.inf + dist_sphere, normal_sphere = _ray_sphere(pos, ssz, pnt, vec) + if dist_sphere < 0: + return -1.0, wp.vec3() # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) # init solution - x = wp.inf + x = -1.0 + part = 0 # -1: bottom, 0: cylinder, 1: top # flat sides if wp.abs(lvec[2]) > MJ_MINVAL: @@ -322,6 +358,7 @@ def _ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: if wp.dot(p, p) <= size[0] * size[0]: if x < 0.0 or sol < x: x = sol + part = side # (x * lvec + lpnt)' * (x * lvec + lpnt) = size[0] * size[0] a = lvec[0] * lvec[0] + lvec[1] * lvec[1] @@ -335,28 +372,43 @@ def _ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: if sol >= 0.0 and wp.abs(lpnt[2] + sol * lvec[2]) <= size[1]: if x < 0.0 or sol < x: x = sol + part = 0 - return x + normal = wp.vec3() + if x >= 0: + if part == 0: + normal = lpnt + lvec * x + normal[2] = 0.0 + normal = wp.normalize(normal) + else: + normal = wp.vec3(0.0, 0.0, float(part)) + + normal = mat @ normal + + return x, normal _IFACE = wp.types.matrix((3, 2), dtype=int)(1, 2, 0, 2, 0, 1) @wp.func -def _ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, vec6]: - """Returns the distance at which a ray intersects with a box.""" +def _ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, vec6, wp.vec3]: + """Returns distance, per side information, and normal at which a ray intersects with a box.""" all = vec6(-1.0, -1.0, -1.0, -1.0, -1.0, -1.0) # bounding sphere test ssz = wp.dot(size, size) - if _ray_sphere(pos, ssz, pnt, vec) < 0.0: - return wp.inf, all + dist_sphere, _ = _ray_sphere(pos, ssz, pnt, vec) + if dist_sphere < 0: + return -1.0, all, wp.vec3() # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) # init solution - x = wp.inf + x = float(-1.0) + face_side = -1 + face_axis = -1 # loop over axes with non-zero vec for i in range(3): @@ -377,13 +429,20 @@ def _ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.v # accept within rectangle if (wp.abs(p0) <= size[id0]) and (wp.abs(p1) <= size[id1]): # update - if (x < 0.0) or (sol < x): + if x < 0.0 or sol < x: x = sol + face_axis = i + face_side = side # save in all all[2 * i + (side + 1) // 2] = sol - return x, all + normal = wp.vec3() + if x >= 0: + normal[face_axis] = float(face_side) + normal = mat @ normal + + return x, all, normal @wp.func @@ -402,10 +461,10 @@ def _ray_hfield( pnt: wp.vec3, vec: wp.vec3, id: int, -): +) -> Tuple[float, wp.vec3]: # check geom type if geom_type[id] != GeomType.HFIELD: - return wp.inf + return -1.0, wp.vec3() # hfield id and dimensions hid = geom_dataid[id] @@ -420,7 +479,7 @@ def _ray_hfield( # compute size and pos of base box base_scale = size[3] * 0.5 base_size = wp.vec3(size[0], size[1], base_scale) - base_pos = pos + mat_col * base_scale + base_pos = pos - mat_col * base_scale # compute size and pos of top box top_scale = size[2] * 0.5 @@ -428,13 +487,13 @@ def _ray_hfield( top_pos = pos + mat_col * top_scale # init: intersection with base box - x, _ = _ray_box(base_pos, mat, base_size, pnt, vec) + x, _, normal_base = _ray_box(base_pos, mat, base_size, pnt, vec) # check top box: done if no intersection - top_intersect, all = _ray_box(top_pos, mat, top_size, pnt, vec) + top_intersect, all, normal_top = _ray_box(top_pos, mat, top_size, pnt, vec) if top_intersect < 0.0: - return x + return x, normal_base # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) @@ -451,8 +510,8 @@ def _ray_hfield( b1 = b0 + lvec * -safe_div(wp.dot(lvec, b0), wp.dot(lvec, lvec)) b1 = wp.normalize(b1) - b2 = wp.cross(b1, lvec) - b2 = wp.normalize(b2) + b0 = wp.cross(b1, lvec) + b0 = wp.normalize(b0) # find ray segment intersecting top box seg = wp.vec2(0.0, top_intersect) @@ -464,8 +523,14 @@ def _ray_hfield( # project segment endpoints in horizontal plane, discretize dx = safe_div(2.0 * size[0], float(ncol - 1)) dy = safe_div(2.0 * size[1], float(nrow - 1)) - SX = wp.vec2(safe_div(lpnt[0] * seg[0] * lvec[0] + size[0], dx), safe_div(lpnt[0] * seg[1] * lvec[0] + size[0], dx)) - SY = wp.vec2(safe_div(lpnt[1] + seg[0] * lvec[1] + size[1], dy), safe_div(lpnt[1] + seg[1] * lvec[1] + size[1], dy)) + SX = wp.vec2( + safe_div(lpnt[0] + seg[0] * lvec[0] + size[0], dx), + safe_div(lpnt[0] + seg[1] * lvec[0] + size[0], dx), + ) + SY = wp.vec2( + safe_div(lpnt[1] + seg[0] * lvec[1] + size[1], dy), + safe_div(lpnt[1] + seg[1] * lvec[1] + size[1], dy), + ) # compute ranges, with +1 padding cmin = wp.max(0, int(wp.floor(wp.min(SX[0], SX[1])) - 1.0)) @@ -473,18 +538,33 @@ def _ray_hfield( rmin = wp.max(0, int(wp.floor(wp.min(SY[0], SY[1])) - 1.0)) rmax = wp.min(nrow - 1, int(wp.ceil(wp.max(SY[0], SY[1])) + 1.0)) + normal_local = wp.vec3() + if x >= 0: + normal_local = wp.transpose(mat) @ normal_base + # check triangles within bounds for r in range(rmin, rmax): for c in range(cmin, cmax): # first triangle - v0 = wp.vec3(dx * float(c) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + c] * size[2]) - v1 = wp.vec3( - dx * float(c + 1) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + (c + 1)] * size[2] + v0 = wp.vec3( + dx * float(c) - size[0], + dy * float(r) - size[1], + hfield_data[adr + r * ncol + c] * size[2], ) - v2 = wp.vec3(dx * float(c + 1) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + (c + 1)] * size[2]) - sol = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1) + v1 = wp.vec3( + dx * float(c + 1) - size[0], + dy * float(r) - size[1], + hfield_data[adr + r * ncol + (c + 1)] * size[2], + ) + v2 = wp.vec3( + dx * float(c + 1) - size[0], + dy * float(r + 1) - size[1], + hfield_data[adr + (r + 1) * ncol + (c + 1)] * size[2], + ) + sol, normal_tri = _ray_triangle(v0, v1, v2, lpnt, lvec, b0, b1) if sol >= 0.0 and (x < 0.0 or sol < x): x = sol + normal_local = normal_tri # second triangle v0 = wp.vec3(dx * float(c) - size[0], dy * float(r) - size[1], hfield_data[adr + r * ncol + c] * size[2]) @@ -492,9 +572,10 @@ def _ray_hfield( dx * float(c + 1) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + (c + 1)] * size[2] ) v2 = wp.vec3(dx * float(c) - size[0], dy * float(r + 1) - size[1], hfield_data[adr + (r + 1) * ncol + c] * size[2]) - sol = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1) + sol, normal_tri = _ray_triangle(v0, v1, v2, lpnt, lvec, b0, b1) if sol >= 0.0 and (x < 0.0 or sol < x): x = sol + normal_local = normal_tri # check viable sides of top box for i in range(4): @@ -527,8 +608,11 @@ def _ray_hfield( # check if point is below line segments if z < z0 * (y0 + 1.0 - y) + z1 * (y - y0): x = all[i] + normal_local = wp.vec3(float(i == 1) - float(i == 0), float(i == 3) - float(i == 2), 0.0) - return x + if x >= 0: + normal_local = mat @ normal_local + return x, normal_local @wp.func @@ -545,8 +629,8 @@ def ray_mesh( mat: wp.mat33, pnt: wp.vec3, vec: wp.vec3, -) -> float: - """Returns the distance and geomid for ray mesh intersections.""" +) -> Tuple[float, wp.vec3]: + """Returns the distance and normal for ray mesh intersections.""" pnt, vec = _ray_map(pos, mat, pnt, vec) # compute orthogonal basis vectors @@ -568,7 +652,8 @@ def ray_mesh( b1 = wp.cross(vec, b0) b1 = wp.normalize(b1) - min_dist = float(wp.inf) + x = float(-1.0) + normal = wp.vec3() # get mesh vertex data range vert_start = mesh_vertadr[data_id] @@ -592,16 +677,22 @@ def ray_mesh( v2 = mesh_vert[vert_start + v_idx.z] # calculate intersection - dist = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1) - if dist < min_dist: - min_dist = dist + dist, normal_tri = _ray_triangle(v0, v1, v2, pnt, vec, b0, b1) + if dist >= 0 and (x < 0 or dist < x): + x = dist + normal = normal_tri - return min_dist + normal = mat @ normal + + return x, normal @wp.func -def ray_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3, geomtype: int) -> float: - """Returns distance along ray to intersection with geom, or infinity if none.""" +def ray_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3, geomtype: int) -> Tuple[float, wp.vec3]: + """Returns distance along ray to intersection with geom and normal at intersection point. + + If no intersection is found, returns -1 and zero vector. + """ # TODO(team): static loop unrolling to remove unnecessary branching if geomtype == GeomType.PLANE: return _ray_plane(pos, mat, size, pnt, vec) @@ -614,10 +705,10 @@ def ray_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.v elif geomtype == GeomType.CYLINDER: return _ray_cylinder(pos, mat, size, pnt, vec) elif geomtype == GeomType.BOX: - dist, _ = _ray_box(pos, mat, size, pnt, vec) - return dist + dist, _, normal = _ray_box(pos, mat, size, pnt, vec) + return dist, normal else: - return wp.inf + return -1.0, wp.vec3() @wp.func @@ -653,7 +744,7 @@ def _ray_geom_mesh( flg_static: bool, bodyexclude: int, geomid: int, -) -> float: +) -> Tuple[float, wp.vec3]: if not _ray_eliminate( body_weldid, geom_bodyid, @@ -701,7 +792,7 @@ def _ray_geom_mesh( else: return ray_geom(pos, mat, geom_size[worldid % geom_size.shape[0], geomid], pnt, vec, type) else: - return wp.inf + return -1.0, wp.vec3() @wp.kernel @@ -737,8 +828,9 @@ def _ray( flg_static: bool, bodyexclude: wp.array(dtype=int), # Out: - dist_out: wp.array(dtype=float, ndim=2), - geomid_out: wp.array(dtype=int, ndim=2), + dist_out: wp.array2d(dtype=float), + geomid_out: wp.array2d(dtype=int), + normal_out: wp.array2d(dtype=wp.vec3), ): worldid, rayid, tid = wp.tid() @@ -746,11 +838,12 @@ def _ray( min_dist = float(wp.inf) min_geomid = int(-1) + min_normal = wp.vec3() upper = ((ngeom + num_threads - 1) // num_threads) * num_threads for geomid in range(tid, upper, num_threads): if geomid < ngeom: - dist = _ray_geom_mesh( + dist, normal = _ray_geom_mesh( nmeshface, body_weldid, geom_type, @@ -780,24 +873,30 @@ def _ray( bodyexclude[rayid], geomid, ) + if dist < 0: + dist = wp.inf else: dist = wp.inf + normal = wp.vec3() tile_dist = wp.tile(dist) local_min_geomid = wp.tile_argmin(tile_dist) local_min_dist = tile_dist[local_min_geomid[0]] tile_geomid = wp.tile(geomid) + tile_normal = wp.tile(normal, preserve_type=True) if local_min_dist < min_dist: min_dist = local_min_dist min_geomid = tile_geomid[local_min_geomid[0]] + min_normal = tile_normal[local_min_geomid[0]] if wp.isinf(min_dist): dist_out[worldid, rayid] = -1.0 else: dist_out[worldid, rayid] = min_dist geomid_out[worldid, rayid] = min_geomid + normal_out[worldid, rayid] = min_normal def ray( @@ -808,7 +907,7 @@ def ray( geomgroup: Optional[vec6] = None, flg_static: bool = True, bodyexclude: int = -1, -) -> Tuple[wp.array, wp.array]: +) -> Tuple[wp.array, wp.array, wp.array]: """Returns the distance at which rays intersect with primitive geoms. Args: @@ -821,7 +920,8 @@ def ray( bodyexclude: Ignore geoms on specified body id (-1 to disable). Returns: - Distances from ray origins to geom surfaces and IDs of intersected geoms (-1 if none). + Distances from ray origins to geom surfaces, IDs of intersected geoms (-1 if none), + and normals at intersection points. """ assert pnt.shape[0] == 1 assert pnt.shape[0] == vec.shape[0] @@ -833,10 +933,11 @@ def ray( ray_bodyexclude.fill_(bodyexclude) ray_dist = wp.empty((d.nworld, 1), dtype=float) ray_geomid = wp.empty((d.nworld, 1), dtype=int) + ray_normal = wp.empty((d.nworld, 1), dtype=wp.vec3) - rays(m, d, pnt, vec, geomgroup, flg_static, ray_bodyexclude, ray_dist, ray_geomid) + rays(m, d, pnt, vec, geomgroup, flg_static, ray_bodyexclude, ray_dist, ray_geomid, ray_normal) - return ray_dist, ray_geomid + return ray_dist, ray_geomid, ray_normal def rays( @@ -847,9 +948,27 @@ def rays( geomgroup: vec6, flg_static: bool, bodyexclude: wp.array(dtype=int), - dist: wp.array2d(dtype=wp.vec3), + dist: wp.array2d(dtype=float), geomid: wp.array2d(dtype=int), + normal: wp.array2d(dtype=wp.vec3), ): + """Ray intersection for multiple worlds and multiple rays. + + Args: + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + pnt: Ray origin points, shape (nworld, nray). + vec: Ray directions, shape (nworld, nray). + geomgroup: Group inclusion/exclusion mask. Set all elements to -1 to ignore. + flg_static: If True, allows rays to intersect with static geoms. + bodyexclude: Per-ray body exclusion array of shape (nray,). Geoms on the + specified body ids are ignored (-1 to disable for that ray). + dist: Output array for distances from ray origins to geom surfaces, shape + (nworld, nray). -1 indicates no intersection. + geomid: Output array for IDs of intersected geoms, shape (nworld, nray). -1 + indicates no intersection. + normal: Output array for normals at intersection points, shape (nworld, nray). + """ wp.launch_tiled( _ray, dim=(d.nworld, pnt.shape[1]), @@ -883,6 +1002,7 @@ def rays( bodyexclude, dist, geomid, + normal, ], block_dim=m.block_dim.ray, ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py index 36c50c6e..41ff3008 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -766,6 +766,7 @@ def sensor_pos(m: Model, d: Data): rangefinder_pnt = wp.empty((d.nworld, m.nrangefinder), dtype=wp.vec3) rangefinder_vec = wp.empty((d.nworld, m.nrangefinder), dtype=wp.vec3) rangefinder_geomid = wp.empty((d.nworld, m.nrangefinder), dtype=int) + rangefinder_normal = wp.empty((d.nworld, m.nrangefinder), dtype=wp.vec3) # get position and direction wp.launch( @@ -786,6 +787,7 @@ def sensor_pos(m: Model, d: Data): m.sensor_rangefinder_bodyid, rangefinder_dist, rangefinder_geomid, + rangefinder_normal, ) if m.sensor_e_potential: @@ -2065,17 +2067,15 @@ def _sensor_touch( conray = -conray # add if ray-zone intersection (always true when contact.pos inside zone) - if ( - ray.ray_geom( - site_xpos_in[worldid, objid], - site_xmat_in[worldid, objid], - site_size[objid], - contact_pos_in[conid], - conray, - site_type[objid], - ) - >= 0.0 - ): + dist, normal = ray.ray_geom( + site_xpos_in[worldid, objid], + site_xmat_in[worldid, objid], + site_size[objid], + contact_pos_in[conid], + conray, + site_type[objid], + ) + if dist >= 0.0: adr = sensor_adr[sensorid] wp.atomic_add(sensordata_out[worldid], adr, normalforce) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py index d8853115..1398976a 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py @@ -738,33 +738,52 @@ def linesearch_jv_fused(nv: int, dofs_per_thread: int): return kernel -@wp.kernel -def linesearch_prepare_gauss( - # Model: - nv: int, - # Data in: - qfrc_smooth_in: wp.array2d(dtype=float), - efc_Ma_in: wp.array2d(dtype=float), - efc_search_in: wp.array2d(dtype=float), - efc_gauss_in: wp.array(dtype=float), - efc_mv_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_quad_gauss_out: wp.array(dtype=wp.vec3), -): - worldid = wp.tid() - if efc_done_in[worldid]: - return +@cache_kernel +def linesearch_prepare_gauss(nv: int, dofs_per_thread: int): + @nested_kernel(module="unique", enable_backward=False) + def kernel( + # Data in: + qfrc_smooth_in: wp.array2d(dtype=float), + efc_Ma_in: wp.array2d(dtype=float), + efc_search_in: wp.array2d(dtype=float), + efc_gauss_in: wp.array(dtype=float), + efc_mv_in: wp.array2d(dtype=float), + efc_done_in: wp.array(dtype=bool), + # Data out: + efc_quad_gauss_out: wp.array(dtype=wp.vec3), + ): + worldid, dofstart = wp.tid() - quad_gauss_0 = efc_gauss_in[worldid] - quad_gauss_1 = float(0.0) - quad_gauss_2 = float(0.0) - for i in range(nv): - search = efc_search_in[worldid, i] - quad_gauss_1 += search * (efc_Ma_in[worldid, i] - qfrc_smooth_in[worldid, i]) - quad_gauss_2 += 0.5 * search * efc_mv_in[worldid, i] + if efc_done_in[worldid]: + return - efc_quad_gauss_out[worldid] = wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2) + quad_gauss_1 = float(0.0) + quad_gauss_2 = float(0.0) + + if wp.static(dofs_per_thread >= nv): + for i in range(wp.static(nv)): + search = efc_search_in[worldid, i] + quad_gauss_1 += search * (efc_Ma_in[worldid, i] - qfrc_smooth_in[worldid, i]) + quad_gauss_2 += 0.5 * search * efc_mv_in[worldid, i] + + quad_gauss_0 = efc_gauss_in[worldid] + efc_quad_gauss_out[worldid] = wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2) + + else: + for i in range(wp.static(dofs_per_thread)): + ii = dofstart * wp.static(dofs_per_thread) + i + if ii < nv: + search = efc_search_in[worldid, ii] + quad_gauss_1 += search * (efc_Ma_in[worldid, ii] - qfrc_smooth_in[worldid, ii]) + quad_gauss_2 += 0.5 * search * efc_mv_in[worldid, ii] + + if dofstart == 0: + quad_gauss_0 = efc_gauss_in[worldid] + wp.atomic_add(efc_quad_gauss_out, worldid, wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2)) + else: + wp.atomic_add(efc_quad_gauss_out, worldid, wp.vec3(0.0, quad_gauss_1, quad_gauss_2)) + + return kernel @wp.kernel @@ -939,10 +958,13 @@ def _linesearch(m: types.Model, d: types.Data, cost: wp.array2d(dtype=float)): # prepare quadratics # quad_gauss = [gauss, search.T @ Ma - search.T @ qfrc_smooth, 0.5 * search.T @ mv] + if threads_per_efc > 1: + d.efc.quad_gauss.zero_() + wp.launch( - linesearch_prepare_gauss, - dim=(d.nworld), - inputs=[m.nv, d.qfrc_smooth, d.efc.Ma, d.efc.search, d.efc.gauss, d.efc.mv, d.efc.done], + linesearch_prepare_gauss(m.nv, dofs_per_thread), + dim=(d.nworld, threads_per_efc), + inputs=[d.qfrc_smooth, d.efc.Ma, d.efc.search, d.efc.gauss, d.efc.mv, d.efc.done], outputs=[d.efc.quad_gauss], ) @@ -1002,28 +1024,42 @@ def solve_init_efc( efc_search_dot_out[worldid] = 0.0 -@wp.kernel -def solve_init_jaref( - # Model: - nv: int, - # Data in: - nefc_in: wp.array(dtype=int), - qacc_in: wp.array2d(dtype=float), - efc_J_in: wp.array3d(dtype=float), - efc_aref_in: wp.array2d(dtype=float), - # Data out: - efc_Jaref_out: wp.array2d(dtype=float), -): - worldid, efcid = wp.tid() +@cache_kernel +def solve_init_jaref(nv: int, dofs_per_thread: int): + @nested_kernel(module="unique", enable_backward=False) + def kernel( + # Data in: + nefc_in: wp.array(dtype=int), + qacc_in: wp.array2d(dtype=float), + efc_J_in: wp.array3d(dtype=float), + efc_aref_in: wp.array2d(dtype=float), + # Data out: + efc_Jaref_out: wp.array2d(dtype=float), + ): + worldid, efcid, dofstart = wp.tid() - if efcid >= nefc_in[worldid]: - return + if efcid >= nefc_in[worldid]: + return - jaref = float(0.0) - for i in range(nv): - jaref += efc_J_in[worldid, efcid, i] * qacc_in[worldid, i] + jaref = float(0.0) - efc_Jaref_out[worldid, efcid] = jaref - efc_aref_in[worldid, efcid] + if wp.static(dofs_per_thread >= nv): + for i in range(wp.static(min(dofs_per_thread, nv))): + jaref += efc_J_in[worldid, efcid, i] * qacc_in[worldid, i] + efc_Jaref_out[worldid, efcid] = jaref - efc_aref_in[worldid, efcid] + + else: + for i in range(wp.static(dofs_per_thread)): + ii = dofstart * wp.static(dofs_per_thread) + i + if ii < nv: + jaref += efc_J_in[worldid, efcid, ii] * qacc_in[worldid, ii] + + if dofstart == 0: + wp.atomic_add(efc_Jaref_out, worldid, efcid, jaref - efc_aref_in[worldid, efcid]) + else: + wp.atomic_add(efc_Jaref_out, worldid, efcid, jaref) + + return kernel @wp.kernel @@ -2043,10 +2079,24 @@ def create_context( ) # jaref = d.efc_J @ d.qacc - d.efc_aref + + # if we are only using 1 thread, it makes sense to do more dofs as we can also skip the + # init kernel. For more than 1 thread, dofs_per_thread is lower for better load balancing. + + if m.nv > 50: + dofs_per_thread = 20 + else: + dofs_per_thread = 50 + + threads_per_efc = ceil(m.nv / dofs_per_thread) + # we need to clear the jaref array if we're doing atomic adds. + if threads_per_efc > 1: + d.efc.Jaref.zero_() + wp.launch( - solve_init_jaref, - dim=(d.nworld, d.njmax), - inputs=[m.nv, d.nefc, d.qacc, d.efc.J, d.efc.aref], + solve_init_jaref(m.nv, dofs_per_thread), + dim=(d.nworld, d.njmax, threads_per_efc), + inputs=[d.nefc, d.qacc, d.efc.J, d.efc.aref], outputs=[d.efc.Jaref], ) diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index e1702192..f3ab0ac6 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -42,7 +42,6 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) - @ffi.format_args_for_warp def _forward_shim( # Model @@ -1837,7 +1836,6 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) - @ffi.format_args_for_warp def _step_shim( # Model diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index df02a52b..89467e4e 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -35,6 +35,7 @@ class TileSet: adr: address of each tile in the set size: size of all the tiles in this set """ + adr: np.ndarray size: int