From e9de329e4ebc1692f1052505b8b3f9be3df0f4b4 Mon Sep 17 00:00:00 2001 From: Taylor Howell Date: Wed, 1 Apr 2026 09:32:47 -0700 Subject: [PATCH] Import google-deepmind/mujoco_warp from GitHub. PiperOrigin-RevId: 892971818 Change-Id: Ibce0c41294a4fcf2825a246a9533e8927fdc7e82 --- .../mjx/third_party/mujoco_warp/__init__.py | 2 + .../mjx/third_party/mujoco_warp/_src/bvh.py | 602 +-- .../mujoco_warp/_src/collision_convex.py | 27 +- .../mujoco_warp/_src/collision_core.py | 213 +- .../mujoco_warp/_src/collision_driver.py | 31 +- .../mujoco_warp/_src/collision_flex.py | 834 +++++ .../mujoco_warp/_src/collision_gjk.py | 31 +- .../mujoco_warp/_src/collision_primitive.py | 38 +- .../_src/collision_primitive_core.py | 504 +++ .../mujoco_warp/_src/collision_sdf.py | 441 +-- .../mujoco_warp/_src/constraint.py | 3258 +++++++++-------- .../mujoco_warp/_src/derivative.py | 195 +- .../third_party/mujoco_warp/_src/forward.py | 421 +-- .../mjx/third_party/mujoco_warp/_src/io.py | 874 +++-- .../third_party/mujoco_warp/_src/island.py | 33 +- .../mjx/third_party/mujoco_warp/_src/math.py | 29 + .../third_party/mujoco_warp/_src/passive.py | 141 +- .../mjx/third_party/mujoco_warp/_src/ray.py | 22 +- .../third_party/mujoco_warp/_src/render.py | 384 +- .../mujoco_warp/_src/render_util.py | 38 + .../third_party/mujoco_warp/_src/sensor.py | 272 +- .../third_party/mujoco_warp/_src/smooth.py | 1352 ++++--- .../third_party/mujoco_warp/_src/solver.py | 1180 +++--- .../third_party/mujoco_warp/_src/support.py | 34 + .../mjx/third_party/mujoco_warp/_src/types.py | 281 +- .../third_party/mujoco_warp/_src/util_pkg.py | 20 +- .../third_party/mujoco_warp/_src/warp_util.py | 4 +- .../third_party/mujoco_warp/pyproject.toml | 4 +- .../mjx/third_party/mujoco_warp/viewer.py | 5 +- mjx/mujoco/mjx/warp/bvh.py | 31 +- mjx/mujoco/mjx/warp/collision_driver.py | 113 +- mjx/mujoco/mjx/warp/forward.py | 290 +- mjx/mujoco/mjx/warp/forward_test.py | 11 +- mjx/mujoco/mjx/warp/render.py | 16 +- mjx/mujoco/mjx/warp/smooth.py | 23 +- mjx/mujoco/mjx/warp/types.py | 82 +- 36 files changed, 7231 insertions(+), 4605 deletions(-) create mode 100644 mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py index 1ff05ff6..40653b62 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py @@ -64,6 +64,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.ray import rays as rays from mujoco.mjx.third_party.mujoco_warp._src.render import render as render from mujoco.mjx.third_party.mujoco_warp._src.render_util import get_depth as get_depth from mujoco.mjx.third_party.mujoco_warp._src.render_util import get_rgb as get_rgb +from mujoco.mjx.third_party.mujoco_warp._src.render_util import get_segmentation as get_segmentation 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 @@ -92,6 +93,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.support import xfrc_accumulate as x from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType as BiasType from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseFilter as BroadphaseFilter from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType as BroadphaseType +from mujoco.mjx.third_party.mujoco_warp._src.types import Callback as Callback from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType as ConeType from mujoco.mjx.third_party.mujoco_warp._src.types import Constraint as Constraint from mujoco.mjx.third_party.mujoco_warp._src.types import Contact as Contact diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py index c40fcfc9..58f7b221 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py @@ -189,12 +189,12 @@ def _compute_bvh_bounds( upper_out: wp.array(dtype=wp.vec3), group_out: wp.array(dtype=int), ): - world_id, geom_local_id = wp.tid() + worldid, geom_local_id = wp.tid() geom_id = enabled_geom_ids[geom_local_id] - pos = geom_xpos_in[world_id, geom_id] - rot = geom_xmat_in[world_id, geom_id] - size = geom_size[world_id % geom_size.shape[0], geom_id] + pos = geom_xpos_in[worldid, geom_id] + rot = geom_xmat_in[worldid, geom_id] + size = geom_size[worldid % geom_size.shape[0], geom_id] type = geom_type[geom_id] # TODO: Investigate branch elimination with static loop unrolling @@ -218,9 +218,9 @@ def _compute_bvh_bounds( hfield_center = pos + rot[:, 2] * size[2] lower_bound, upper_bound = _compute_box_bounds(hfield_center, rot, size) - lower_out[world_id * bvh_ngeom + geom_local_id] = lower_bound - upper_out[world_id * bvh_ngeom + geom_local_id] = upper_bound - group_out[world_id * bvh_ngeom + geom_local_id] = world_id + lower_out[worldid * bvh_ngeom + geom_local_id] = lower_bound + upper_out[worldid * bvh_ngeom + geom_local_id] = upper_bound + group_out[worldid * bvh_ngeom + geom_local_id] = worldid @wp.kernel @@ -235,14 +235,70 @@ def compute_bvh_group_roots( group_root_out[tid] = root +@wp.kernel +def _compute_flex_bvh_bounds( + # Model: + flex_vertadr: wp.array(dtype=int), + flex_vertnum: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flex_radius: wp.array(dtype=float), + # Data in: + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + # In: + flex_geom_flexid: wp.array(dtype=int), + flex_geom_edgeid: wp.array(dtype=int), + bvh_ngeom: int, + total_bvh_size: int, + # Out: + lower_out: wp.array(dtype=wp.vec3), + upper_out: wp.array(dtype=wp.vec3), + group_out: wp.array(dtype=int), +): + worldid, flexlocalid = wp.tid() + + flex_id = flex_geom_flexid[flexlocalid] + edge_id = flex_geom_edgeid[flexlocalid] + out_idx = worldid * total_bvh_size + bvh_ngeom + flexlocalid + radius = flex_radius[flex_id] + inflate = wp.vec3(radius, radius, radius) + + if edge_id >= 0: # capsule (1D edge) + edge = flex_edge[edge_id] + vert_adr = flex_vertadr[flex_id] + v0 = flexvert_xpos_in[worldid, vert_adr + edge[0]] + v1 = flexvert_xpos_in[worldid, vert_adr + edge[1]] + lower_out[out_idx] = wp.min(v0, v1) - inflate + upper_out[out_idx] = wp.max(v0, v1) + inflate + else: # mesh (2D/3D) + vert_adr = flex_vertadr[flex_id] + nvert = flex_vertnum[flex_id] + min_bound = wp.vec3(MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL) + max_bound = wp.vec3(-MJ_MAXVAL, -MJ_MAXVAL, -MJ_MAXVAL) + for i in range(nvert): + v = flexvert_xpos_in[worldid, vert_adr + i] + min_bound = wp.min(min_bound, v) + max_bound = wp.max(max_bound, v) + lower_out[out_idx] = min_bound - inflate + upper_out[out_idx] = max_bound + inflate + + group_out[out_idx] = worldid + + def build_scene_bvh(mjm: mujoco.MjModel, mjd: mujoco.MjData, rc: RenderContext, nworld: int): """Build a global BVH for all geometries in all worlds.""" + total_bvh_size = rc.bvh_ngeom + rc.bvh_nflexgeom + geom_type = wp.array(mjm.geom_type, dtype=int) geom_dataid = wp.array(mjm.geom_dataid, dtype=int) geom_size = wp.array(np.tile(mjm.geom_size[np.newaxis, :, :], (nworld, 1, 1)), dtype=wp.vec3) geom_xpos = wp.array(np.tile(mjd.geom_xpos[np.newaxis, :, :], (nworld, 1, 1)), dtype=wp.vec3) geom_xmat = wp.array(np.tile(mjd.geom_xmat.reshape(mjm.ngeom, 3, 3)[np.newaxis, :, :, :], (nworld, 1, 1, 1)), dtype=wp.mat33) + flex_vertadr = wp.array(mjm.flex_vertadr, dtype=int) + flex_vertnum = wp.array(mjm.flex_vertnum, dtype=int) + flex_edge = wp.array(mjm.flex_edge, dtype=wp.vec2i) + flex_radius = wp.array(mjm.flex_radius, dtype=float) + wp.launch( kernel=_compute_bvh_bounds, dim=(nworld, rc.bvh_ngeom), @@ -252,7 +308,7 @@ def build_scene_bvh(mjm: mujoco.MjModel, mjd: mujoco.MjData, rc: RenderContext, geom_size, geom_xpos, geom_xmat, - rc.bvh_ngeom, + total_bvh_size, rc.enabled_geom_ids, rc.mesh_bounds_size, rc.hfield_bounds_size, @@ -262,6 +318,26 @@ def build_scene_bvh(mjm: mujoco.MjModel, mjd: mujoco.MjData, rc: RenderContext, ], ) + flexvert_xpos = wp.array(np.tile(mjd.flexvert_xpos[np.newaxis, :, :], (nworld, 1, 1)), dtype=wp.vec3) + wp.launch( + kernel=_compute_flex_bvh_bounds, + dim=(nworld, rc.bvh_nflexgeom), + inputs=[ + flex_vertadr, + flex_vertnum, + flex_edge, + flex_radius, + flexvert_xpos, + rc.flex_geom_flexid, + rc.flex_geom_edgeid, + rc.bvh_ngeom, + total_bvh_size, + rc.lower, + rc.upper, + rc.group, + ], + ) + bvh = wp.Bvh(rc.lower, rc.upper, groups=rc.group, constructor="sah") # BVH handle must be stored to avoid garbage collection @@ -277,6 +353,8 @@ def build_scene_bvh(mjm: mujoco.MjModel, mjd: mujoco.MjData, rc: RenderContext, def refit_scene_bvh(m: Model, d: Data, rc: RenderContext): + total_bvh_size = rc.bvh_ngeom + rc.bvh_nflexgeom + wp.launch( kernel=_compute_bvh_bounds, dim=(d.nworld, rc.bvh_ngeom), @@ -286,7 +364,7 @@ def refit_scene_bvh(m: Model, d: Data, rc: RenderContext): m.geom_size, d.geom_xpos, d.geom_xmat, - rc.bvh_ngeom, + total_bvh_size, rc.enabled_geom_ids, rc.mesh_bounds_size, rc.hfield_bounds_size, @@ -296,6 +374,26 @@ def refit_scene_bvh(m: Model, d: Data, rc: RenderContext): ], ) + if rc.bvh_nflexgeom > 0: + wp.launch( + kernel=_compute_flex_bvh_bounds, + dim=(d.nworld, rc.bvh_nflexgeom), + inputs=[ + m.flex_vertadr, + m.flex_vertnum, + m.flex_edge, + m.flex_radius, + d.flexvert_xpos, + rc.flex_geom_flexid, + rc.flex_geom_edgeid, + rc.bvh_ngeom, + total_bvh_size, + rc.lower, + rc.upper, + rc.group, + ], + ) + rc.bvh.refit() @@ -500,6 +598,12 @@ def build_hfield_bvh( @wp.kernel def accumulate_flex_vertex_normals( # Model: + nflex: int, + flex_dim: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_elemadr: wp.array(dtype=int), + flex_elemnum: wp.array(dtype=int), + flex_elemdataadr: wp.array(dtype=int), flex_elem: wp.array(dtype=int), # Data in: flexvert_xpos_in: wp.array2d(dtype=wp.vec3), @@ -509,10 +613,22 @@ def accumulate_flex_vertex_normals( """Accumulate per-vertex normals by summing adjacent face normals.""" worldid, elemid = wp.tid() - elem_base = elemid * 3 - i0 = flex_elem[elem_base + 0] - i1 = flex_elem[elem_base + 1] - i2 = flex_elem[elem_base + 2] + for i in range(nflex): + locid = elemid - flex_elemadr[i] + if locid >= 0 and locid < flex_elemnum[i]: + f = i + break + + if flex_dim[f] == 1 or flex_dim[f] == 3: + return + + local_elemid = elemid - flex_elemadr[f] + elem_adr = flex_elemdataadr[f] + vert_adr = flex_vertadr[f] + elem_base = elem_adr + local_elemid * 3 + i0 = vert_adr + flex_elem[elem_base + 0] + i1 = vert_adr + flex_elem[elem_base + 1] + i2 = vert_adr + flex_elem[elem_base + 2] v0 = flexvert_xpos_in[worldid, i0] v1 = flexvert_xpos_in[worldid, i1] @@ -611,11 +727,12 @@ def _build_flex_2d_elements( @wp.kernel def _build_flex_2d_sides( + # Model: + flex_shell: wp.array(dtype=int), # Data in: flexvert_xpos_in: wp.array2d(dtype=wp.vec3), # In: flexvert_norm_in: wp.array2d(dtype=wp.vec3), - flex_shell_in: wp.array(dtype=int), shell_adr: int, vert_adr: int, face_offset: int, @@ -635,8 +752,8 @@ def _build_flex_2d_sides( worldid, shellid = wp.tid() base = shell_adr + 2 * shellid - i0 = vert_adr + flex_shell_in[base + 0] - i1 = vert_adr + flex_shell_in[base + 1] + i0 = vert_adr + flex_shell[base + 0] + i1 = vert_adr + flex_shell[base + 1] v0 = flexvert_xpos_in[worldid, i0] v1 = flexvert_xpos_in[worldid, i1] @@ -672,10 +789,11 @@ def _build_flex_2d_sides( @wp.kernel def _build_flex_3d_shells( + # Model: + flex_shell: wp.array(dtype=int), # Data in: flexvert_xpos_in: wp.array2d(dtype=wp.vec3), # In: - flex_shell_in: wp.array(dtype=int), shell_adr: int, vert_adr: int, face_offset: int, @@ -693,9 +811,9 @@ def _build_flex_3d_shells( worldid, shellid = wp.tid() base = shell_adr + shellid * 3 - i0 = vert_adr + flex_shell_in[base + 0] - i1 = vert_adr + flex_shell_in[base + 1] - i2 = vert_adr + flex_shell_in[base + 2] + i0 = vert_adr + flex_shell[base + 0] + i1 = vert_adr + flex_shell[base + 1] + i2 = vert_adr + flex_shell[base + 2] face_id = worldid * nface + face_offset + shellid base = face_id * 3 @@ -716,163 +834,163 @@ def _build_flex_3d_shells( @wp.kernel -def _update_flex_face_points( +def _update_flex_2d_face_points( # Model: - nflex: int, - flex_dim: wp.array(dtype=int), flex_vertadr: wp.array(dtype=int), flex_elemnum: wp.array(dtype=int), + flex_elemdataadr: wp.array(dtype=int), + flex_shelldataadr: wp.array(dtype=int), flex_elem: wp.array(dtype=int), + flex_shell: wp.array(dtype=int), + flex_radius: wp.array(dtype=float), # Data in: flexvert_xpos_in: wp.array2d(dtype=wp.vec3), # In: - flex_shell_in: wp.array(dtype=int), flexvert_norm_in: wp.array2d(dtype=wp.vec3), - flex_elemdataadr: wp.array(dtype=int), - flex_shelldataadr: wp.array(dtype=int), - flex_faceadr: wp.array(dtype=int), - flex_radius: wp.array(dtype=float), - flex_workadr: wp.array(dtype=int), - flex_worknum: wp.array(dtype=int), - nfaces: int, + flex_id: int, + nface: int, smooth: bool, # Out: face_point_out: wp.array(dtype=wp.vec3), ): worldid, workid = wp.tid() - # identify which flex this work item belongs to - f = int(0) - locid = int(0) - for i in range(nflex): - locid = workid - flex_workadr[i] - if locid >= 0 and locid < flex_worknum[i]: - f = i - break + elem_adr = flex_elemdataadr[flex_id] + vert_adr = flex_vertadr[flex_id] + radius = flex_radius[flex_id] + nelem = flex_elemnum[flex_id] + world_face_offset = worldid * nface - dim = flex_dim[f] - face_offset = flex_faceadr[f] - world_face_offset = worldid * nfaces - vert_adr = flex_vertadr[f] - - if dim == 2: - radius = flex_radius[f] - elem_count = flex_elemnum[f] - - if locid < elem_count: - # 2D element faces - elemid = locid - elem_adr = flex_elemdataadr[f] - ebase = elem_adr + elemid * 3 - i0 = vert_adr + flex_elem[ebase + 0] - i1 = vert_adr + flex_elem[ebase + 1] - i2 = vert_adr + flex_elem[ebase + 2] - - v0 = flexvert_xpos_in[worldid, i0] - v1 = flexvert_xpos_in[worldid, i1] - v2 = flexvert_xpos_in[worldid, i2] - - # TODO: Use static conditional - if smooth: - n0 = flexvert_norm_in[worldid, i0] - n1 = flexvert_norm_in[worldid, i1] - n2 = flexvert_norm_in[worldid, i2] - else: - face_nrm = wp.cross(v1 - v0, v2 - v0) - face_nrm = wp.normalize(face_nrm) - n0 = face_nrm - n1 = face_nrm - n2 = face_nrm - - p0_pos = v0 + radius * n0 - p1_pos = v1 + radius * n1 - p2_pos = v2 + radius * n2 - - p0_neg = v0 - radius * n0 - p1_neg = v1 - radius * n1 - p2_neg = v2 - radius * n2 - - face_id0 = world_face_offset + face_offset + (2 * elemid) - base0 = face_id0 * 3 - face_point_out[base0 + 0] = p0_pos - face_point_out[base0 + 1] = p1_pos - face_point_out[base0 + 2] = p2_pos - - face_id1 = world_face_offset + face_offset + (2 * elemid + 1) - base1 = face_id1 * 3 - face_point_out[base1 + 0] = p0_neg - face_point_out[base1 + 1] = p1_neg - face_point_out[base1 + 2] = p2_neg - else: - # 2D shell faces - shellid = locid - elem_count - shell_adr = flex_shelldataadr[f] - sbase = shell_adr + 2 * shellid - i0 = vert_adr + flex_shell_in[sbase + 0] - i1 = vert_adr + flex_shell_in[sbase + 1] - - v0 = flexvert_xpos_in[worldid, i0] - v1 = flexvert_xpos_in[worldid, i1] - - n0 = flexvert_norm_in[worldid, i0] - n1 = flexvert_norm_in[worldid, i1] - - shell_face_offset = face_offset + (2 * elem_count) - face_id0 = world_face_offset + shell_face_offset + (2 * shellid) - base0 = face_id0 * 3 - face_point_out[base0 + 0] = v0 + radius * n0 - face_point_out[base0 + 1] = v1 - radius * n1 - face_point_out[base0 + 2] = v1 + radius * n1 - - face_id1 = world_face_offset + shell_face_offset + (2 * shellid + 1) - base1 = face_id1 * 3 - face_point_out[base1 + 0] = v1 - radius * n1 - face_point_out[base1 + 1] = v0 + radius * n0 - face_point_out[base1 + 2] = v0 - radius * n0 - else: - # 3D shell faces - shellid = locid - shell_adr = flex_shelldataadr[f] - sbase = shell_adr + shellid * 3 - i0 = vert_adr + flex_shell_in[sbase + 0] - i1 = vert_adr + flex_shell_in[sbase + 1] - i2 = vert_adr + flex_shell_in[sbase + 2] + if workid < nelem: + # 2D element faces + elemid = workid + ebase = elem_adr + elemid * 3 + i0 = vert_adr + flex_elem[ebase + 0] + i1 = vert_adr + flex_elem[ebase + 1] + i2 = vert_adr + flex_elem[ebase + 2] v0 = flexvert_xpos_in[worldid, i0] v1 = flexvert_xpos_in[worldid, i1] v2 = flexvert_xpos_in[worldid, i2] - face_id = world_face_offset + face_offset + shellid - fbase = face_id * 3 + # TODO: Use static conditional + if smooth: + n0 = flexvert_norm_in[worldid, i0] + n1 = flexvert_norm_in[worldid, i1] + n2 = flexvert_norm_in[worldid, i2] + else: + face_nrm = wp.cross(v1 - v0, v2 - v0) + face_nrm = wp.normalize(face_nrm) + n0 = face_nrm + n1 = face_nrm + n2 = face_nrm - face_point_out[fbase + 0] = v0 - face_point_out[fbase + 1] = v1 - face_point_out[fbase + 2] = v2 + p0_pos = v0 + radius * n0 + p1_pos = v1 + radius * n1 + p2_pos = v2 + radius * n2 + + p0_neg = v0 - radius * n0 + p1_neg = v1 - radius * n1 + p2_neg = v2 - radius * n2 + + face_id0 = world_face_offset + (2 * elemid) + base0 = face_id0 * 3 + face_point_out[base0 + 0] = p0_pos + face_point_out[base0 + 1] = p1_pos + face_point_out[base0 + 2] = p2_pos + + face_id1 = world_face_offset + (2 * elemid + 1) + base1 = face_id1 * 3 + face_point_out[base1 + 0] = p0_neg + face_point_out[base1 + 1] = p1_neg + face_point_out[base1 + 2] = p2_neg + else: + # 2D shell faces + shell_adr = flex_shelldataadr[flex_id] + shellid = workid - nelem + sbase = shell_adr + 2 * shellid + i0 = vert_adr + flex_shell[sbase + 0] + i1 = vert_adr + flex_shell[sbase + 1] + + v0 = flexvert_xpos_in[worldid, i0] + v1 = flexvert_xpos_in[worldid, i1] + + n0 = flexvert_norm_in[worldid, i0] + n1 = flexvert_norm_in[worldid, i1] + + shell_face_offset = 2 * nelem + face_id0 = world_face_offset + shell_face_offset + (2 * shellid) + base0 = face_id0 * 3 + face_point_out[base0 + 0] = v0 + radius * n0 + face_point_out[base0 + 1] = v1 - radius * n1 + face_point_out[base0 + 2] = v1 + radius * n1 + + face_id1 = world_face_offset + shell_face_offset + (2 * shellid + 1) + base1 = face_id1 * 3 + face_point_out[base1 + 0] = v1 - radius * n1 + face_point_out[base1 + 1] = v0 + radius * n0 + face_point_out[base1 + 2] = v0 - radius * n0 + + +@wp.kernel +def _update_flex_3d_face_points( + # Model: + flex_vertadr: wp.array(dtype=int), + flex_shelldataadr: wp.array(dtype=int), + flex_shell: wp.array(dtype=int), + # Data in: + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + # In: + flex_id: int, + nface: int, + # Out: + face_point_out: wp.array(dtype=wp.vec3), +): + worldid, shellid = wp.tid() + + shell_adr = flex_shelldataadr[flex_id] + vert_adr = flex_vertadr[flex_id] + + face_id = worldid * nface + shellid + fbase = face_id * 3 + + sbase = shell_adr + shellid * 3 + i0 = vert_adr + flex_shell[sbase + 0] + i1 = vert_adr + flex_shell[sbase + 1] + i2 = vert_adr + flex_shell[sbase + 2] + + face_point_out[fbase + 0] = flexvert_xpos_in[worldid, i0] + face_point_out[fbase + 1] = flexvert_xpos_in[worldid, i1] + face_point_out[fbase + 2] = flexvert_xpos_in[worldid, i2] def build_flex_bvh( - mjm: mujoco.MjModel, mjd: mujoco.MjData, nworld: int, constructor: str = "sah", leaf_size: int = 2 -) -> tuple[wp.Mesh, wp.array, wp.array, wp.array, wp.array, wp.array, int]: - """Create a Warp mesh BVH from flex data.""" - if (mjm.flex_dim == 1).any(): - raise ValueError("1D Flex objects are not currently supported.") - - nflex = mjm.nflex + mjm: mujoco.MjModel, + mjd: mujoco.MjData, + nworld: int, + flex_id: int, + constructor: str = "sah", + leaf_size: int = 2, +) -> tuple[wp.Mesh, wp.array, wp.array, wp.array, int]: + """Create a Warp mesh BVH for a single 2D or 3D flex.""" nflexvert = mjm.nflexvert - nflexelemdata = len(mjm.flex_elem) + flex_dim = wp.array(mjm.flex_dim, dtype=int) + flex_elemadr = wp.array(mjm.flex_elemadr, dtype=int) + flex_elemnum = wp.array(mjm.flex_elemnum, dtype=int) flex_elem = wp.array(mjm.flex_elem, dtype=int) + flex_elemdataadr = wp.array(mjm.flex_elemdataadr, dtype=int) + flex_vertadr = wp.array(mjm.flex_vertadr, dtype=int) flexvert_xpos = wp.array(np.tile(mjd.flexvert_xpos[np.newaxis, :, :], (nworld, 1, 1)), dtype=wp.vec3) - flex_faceadr = [0] - for f in range(nflex): - if mjm.flex_dim[f] == 2: - flex_faceadr.append(flex_faceadr[-1] + 2 * mjm.flex_elemnum[f] + 2 * mjm.flex_shellnum[f]) - elif mjm.flex_dim[f] == 3: - flex_faceadr.append(flex_faceadr[-1] + mjm.flex_shellnum[f]) + dim = int(mjm.flex_dim[flex_id]) + nelem = int(mjm.flex_elemnum[flex_id]) + nshell = int(mjm.flex_shellnum[flex_id]) - nface = int(flex_faceadr[-1]) - flex_faceadr = flex_faceadr[:-1] + if dim == 2: + nface = 2 * nelem + 2 * nshell + else: + nface = nshell face_point = wp.empty(nface * 3 * nworld, dtype=wp.vec3) face_index = wp.empty(nface * 3 * nworld, dtype=wp.int32) @@ -883,8 +1001,8 @@ def build_flex_bvh( wp.launch( kernel=accumulate_flex_vertex_normals, - dim=(nworld, nflexelemdata // 3), - inputs=[flex_elem, flexvert_xpos], + dim=(nworld, mjm.nflexelem), + inputs=[mjm.nflex, flex_dim, flex_vertadr, flex_elemadr, flex_elemnum, flex_elemdataadr, flex_elem, flexvert_xpos], outputs=[flexvert_norm], ) @@ -894,60 +1012,56 @@ def build_flex_bvh( inputs=[flexvert_norm], ) - for f in range(nflex): - dim = mjm.flex_dim[f] - elem_adr = mjm.flex_elemdataadr[f] - nelem = mjm.flex_elemnum[f] - shell_adr = mjm.flex_shelldataadr[f] - nshell = mjm.flex_shellnum[f] - vert_adr = mjm.flex_vertadr[f] + elem_adr = mjm.flex_elemdataadr[flex_id] + shell_adr = mjm.flex_shelldataadr[flex_id] + vert_adr = mjm.flex_vertadr[flex_id] - if dim == 2: - wp.launch( - kernel=_build_flex_2d_elements, - dim=(nworld, nelem), - inputs=[ - flex_elem, - flexvert_xpos, - flexvert_norm, - elem_adr, - vert_adr, - flex_faceadr[f], - mjm.flex_radius[f], - nface, - ], - outputs=[face_point, face_index, group], - ) + if dim == 2: + wp.launch( + kernel=_build_flex_2d_elements, + dim=(nworld, nelem), + inputs=[ + flex_elem, + flexvert_xpos, + flexvert_norm, + elem_adr, + vert_adr, + 0, # face_offset + mjm.flex_radius[flex_id], + nface, + ], + outputs=[face_point, face_index, group], + ) - wp.launch( - kernel=_build_flex_2d_sides, - dim=(nworld, nshell), - inputs=[ - flexvert_xpos, - flexvert_norm, - flex_shell, - shell_adr, - vert_adr, - flex_faceadr[f] + 2 * nelem, - mjm.flex_radius[f], - nface, - ], - outputs=[face_point, face_index, group], - ) - elif dim == 3: - wp.launch( - kernel=_build_flex_3d_shells, - dim=(nworld, nshell), - inputs=[ - flexvert_xpos, - flex_shell, - shell_adr, - vert_adr, - flex_faceadr[f], - nface, - ], - outputs=[face_point, face_index, group], - ) + wp.launch( + kernel=_build_flex_2d_sides, + dim=(nworld, nshell), + inputs=[ + flex_shell, + flexvert_xpos, + flexvert_norm, + shell_adr, + vert_adr, + 2 * nelem, # face_offset + mjm.flex_radius[flex_id], + nface, + ], + outputs=[face_point, face_index, group], + ) + elif dim == 3: + wp.launch( + kernel=_build_flex_3d_shells, + dim=(nworld, nshell), + inputs=[ + flex_shell, + flexvert_xpos, + shell_adr, + vert_adr, + 0, # face_offset + nface, + ], + outputs=[face_point, face_index, group], + ) flex_mesh = wp.Mesh( points=face_point, @@ -965,24 +1079,23 @@ def build_flex_bvh( outputs=[group_root], ) - return ( - flex_mesh, - face_point, - group_root, - flex_shell, - flex_faceadr, - nface, - ) + return flex_mesh, group_root def refit_flex_bvh(m: Model, d: Data, rc: RenderContext): - """Refit the flex BVH.""" + """Refit per-flex BVHs.""" flexvert_norm = wp.zeros(d.flexvert_xpos.shape, dtype=wp.vec3) wp.launch( kernel=accumulate_flex_vertex_normals, - dim=(d.nworld, m.nflexelemdata // 3), + dim=(d.nworld, m.nflexelem), inputs=[ + m.nflex, + m.flex_dim, + m.flex_vertadr, + m.flex_elemadr, + m.flex_elemnum, + m.flex_elemdataadr, m.flex_elem, d.flexvert_xpos, ], @@ -991,32 +1104,49 @@ def refit_flex_bvh(m: Model, d: Data, rc: RenderContext): wp.launch( kernel=normalize_vertex_normals, - dim=(d.nworld, m.nflexvert), + dim=(d.nworld, d.flexvert_xpos.shape[1]), inputs=[flexvert_norm], ) - wp.launch( - kernel=_update_flex_face_points, - dim=(d.nworld, rc.flex_nwork), - inputs=[ - m.nflex, - m.flex_dim, - m.flex_vertadr, - m.flex_elemnum, - m.flex_elem, - d.flexvert_xpos, - rc.flex_shell, - flexvert_norm, - rc.flex_elemdataadr, - rc.flex_shelldataadr, - rc.flex_faceadr, - rc.flex_radius, - rc.flex_workadr, - rc.flex_worknum, - rc.flex_nface, - rc.flex_render_smooth, - ], - outputs=[rc.flex_face_point], - ) + for i in range(m.nflex): + if rc.flex_dim_np[i] == 1: + continue + mesh = rc.flex_mesh_registry[i] + nface = mesh.points.shape[0] // (3 * d.nworld) - rc.flex_mesh.refit() + if rc.flex_dim_np[i] == 2: + wp.launch( + kernel=_update_flex_2d_face_points, + dim=(d.nworld, nface // 2), + inputs=[ + m.flex_vertadr, + m.flex_elemnum, + m.flex_elemdataadr, + m.flex_shelldataadr, + m.flex_elem, + m.flex_shell, + m.flex_radius, + d.flexvert_xpos, + flexvert_norm, + i, + nface, + rc.flex_render_smooth, + ], + outputs=[mesh.points], + ) + else: + wp.launch( + kernel=_update_flex_3d_face_points, + dim=(d.nworld, nface), + inputs=[ + m.flex_vertadr, + m.flex_shelldataadr, + m.flex_shell, + d.flexvert_xpos, + i, + nface, + ], + outputs=[mesh.points], + ) + + mesh.refit() 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 69ea9cbe..12d9824a 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 @@ -15,34 +15,35 @@ from typing import Tuple +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src.collision_core import CollisionContext -from mujoco.mjx.third_party.mujoco_warp._src.collision_core import contact_params from mujoco.mjx.third_party.mujoco_warp._src.collision_core import Geom +from mujoco.mjx.third_party.mujoco_warp._src.collision_core import contact_params from mujoco.mjx.third_party.mujoco_warp._src.collision_core import geom_collision_pair from mujoco.mjx.third_party.mujoco_warp._src.collision_core import write_contact from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import multicontact from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import support -from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom +from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom_collision_pair from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index -from mujoco.mjx.third_party.mujoco_warp._src.types import Data -from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit -from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType -from mujoco.mjx.third_party.mujoco_warp._src.types import mat43 -from mujoco.mjx.third_party.mujoco_warp._src.types import mat63 from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAX_EPAFACES from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAX_EPAHORIZON from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import mat43 +from mujoco.mjx.third_party.mujoco_warp._src.types import mat63 from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp # TODO(team): improve compile time to enable backward pass wp.set_module_options({"enable_backward": False}) @@ -233,6 +234,7 @@ def ccd_hfield_kernel_builder( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -516,6 +518,7 @@ def ccd_hfield_kernel_builder( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -572,6 +575,7 @@ def ccd_hfield_kernel_builder( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -626,6 +630,7 @@ def ccd_hfield_kernel_builder( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -681,6 +686,7 @@ def ccd_hfield_kernel_builder( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -751,6 +757,7 @@ def ccd_kernel_builder( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -871,6 +878,7 @@ def ccd_kernel_builder( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -956,6 +964,7 @@ def ccd_kernel_builder( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -1070,6 +1079,7 @@ def ccd_kernel_builder( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -1156,6 +1166,7 @@ def convex_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table d.contact.solimp, d.contact.dim, d.contact.geom, + d.contact.efc_address, d.contact.worldid, d.contact.type, d.contact.geomcollisionid, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_core.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_core.py index f8294f65..b7affa7d 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_core.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_core.py @@ -18,14 +18,15 @@ import dataclasses from typing import Tuple +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src.math import safe_div +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import mat63 -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -63,30 +64,30 @@ class Geom: @wp.func def geom_collision_pair( - # Model: - geom_type: wp.array(dtype=int), - geom_dataid: wp.array(dtype=int), - geom_size: wp.array2d(dtype=wp.vec3), - mesh_vertadr: wp.array(dtype=int), - mesh_vertnum: wp.array(dtype=int), - mesh_graphadr: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), - 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), - # Data in: - geom_xpos_in: wp.array2d(dtype=wp.vec3), - geom_xmat_in: wp.array2d(dtype=wp.mat33), - # In: - geoms: wp.vec2i, - worldid: int, + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_graphadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + 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), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + geoms: wp.vec2i, + worldid: int, ) -> Tuple[Geom, Geom]: geom1 = Geom() geom2 = Geom() @@ -155,38 +156,39 @@ def geom_collision_pair( @wp.func def write_contact( - # Data in: - naconmax_in: int, - # In: - id_: int, - dist_in: float, - pos_in: wp.vec3, - frame_in: wp.mat33, - margin_in: float, - gap_in: float, - condim_in: int, - friction_in: vec5, - solref_in: wp.vec2, - solreffriction_in: wp.vec2, - solimp_in: vec5, - geoms_in: wp.vec2i, - pairid_in: wp.vec2i, - worldid_in: int, - # Data out: - contact_dist_out: wp.array(dtype=float), - contact_pos_out: wp.array(dtype=wp.vec3), - contact_frame_out: wp.array(dtype=wp.mat33), - contact_includemargin_out: wp.array(dtype=float), - contact_friction_out: wp.array(dtype=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), - nacon_out: wp.array(dtype=int), + # Data in: + naconmax_in: int, + # In: + id_: int, + dist_in: float, + pos_in: wp.vec3, + frame_in: wp.mat33, + margin_in: float, + gap_in: float, + condim_in: int, + friction_in: vec5, + solref_in: wp.vec2, + solreffriction_in: wp.vec2, + solimp_in: vec5, + geoms_in: wp.vec2i, + pairid_in: wp.vec2i, + worldid_in: int, + # Data out: + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=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_efc_address_out: wp.array2d(dtype=int), + contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ) -> int: """Atomically write a detected contact into the contact output arrays. @@ -222,33 +224,35 @@ def write_contact( contact_solimp_out[cid] = solimp_in contact_type_out[cid] = contact_type contact_geomcollisionid_out[cid] = id_ + for i in range(contact_efc_address_out.shape[1]): + contact_efc_address_out[cid, i] = -1 return int(active) return 0 @wp.func def contact_params( - # Model: - geom_condim: 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_friction: wp.array2d(dtype=wp.vec3), - geom_margin: wp.array2d(dtype=float), - geom_gap: wp.array2d(dtype=float), - 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), - # In: - collision_pair_in: wp.array(dtype=wp.vec2i), - collision_pairid_in: wp.array(dtype=wp.vec2i), - cid: int, - worldid: int, + # Model: + geom_condim: 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_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + geom_gap: wp.array2d(dtype=float), + 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), + # In: + collision_pair_in: wp.array(dtype=wp.vec2i), + collision_pairid_in: wp.array(dtype=wp.vec2i), + cid: int, + worldid: int, ): """Resolve contact parameters for a collision pair. @@ -267,9 +271,7 @@ def contact_params( condim = pair_dim[pairid] friction = pair_friction[worldid % pair_friction.shape[0], pairid] solref = pair_solref[worldid % pair_solref.shape[0], pairid] - solreffriction = pair_solreffriction[ - worldid % pair_solreffriction.shape[0], pairid - ] + solreffriction = pair_solreffriction[worldid % pair_solreffriction.shape[0], pairid] solimp = pair_solimp[worldid % pair_solimp.shape[0], pairid] else: g1 = geoms[0] @@ -305,44 +307,33 @@ def contact_params( mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 >= MJ_MINVAL), 0.0, mix) mix = wp.where((solmix1 >= MJ_MINVAL) and (solmix2 < MJ_MINVAL), 1.0, mix) condim = wp.max(condim1, condim2) - max_geom_friction = wp.max( - geom_friction[friction_id, g1], geom_friction[friction_id, g2] - ) + max_geom_friction = wp.max(geom_friction[friction_id, g1], geom_friction[friction_id, g2]) friction = vec5( - max_geom_friction[0], - max_geom_friction[0], - max_geom_friction[1], - max_geom_friction[2], - 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 - ): - solref = ( - mix * geom_solref[solref_id, g1] - + (1.0 - mix) * geom_solref[solref_id, g2] - ) + if geom_solref[solref_id, g1][0] > 0.0 and geom_solref[solref_id, g2][0] > 0.0: + solref = mix * geom_solref[solref_id, g1] + (1.0 - mix) * geom_solref[solref_id, g2] else: solref = wp.min(geom_solref[solref_id, g1], geom_solref[solref_id, g2]) solreffriction = wp.vec2(0.0, 0.0) - solimp = ( - mix * geom_solimp[solimp_id, g1] - + (1.0 - mix) * geom_solimp[solimp_id, g2] - ) + solimp = mix * geom_solimp[solimp_id, g1] + (1.0 - mix) * geom_solimp[solimp_id, g2] # geom priority is ignored margin = geom_margin[margin_id, g1] + geom_margin[margin_id, g2] gap = 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]), + 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 @@ -366,7 +357,7 @@ class CollisionContext: def create_collision_context(naconmax: int) -> CollisionContext: """Create a CollisionContext with allocated arrays.""" return CollisionContext( - collision_pair=wp.empty(naconmax, dtype=wp.vec2i), - collision_pairid=wp.empty(naconmax, dtype=wp.vec2i), - collision_worldid=wp.empty(naconmax, dtype=int), + collision_pair=wp.empty(naconmax, dtype=wp.vec2i), + collision_pairid=wp.empty(naconmax, dtype=wp.vec2i), + collision_worldid=wp.empty(naconmax, dtype=int), ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py index 36cdaea1..786ee9b8 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_driver.py @@ -15,25 +15,27 @@ from typing import Any +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src.collision_convex import convex_narrowphase from mujoco.mjx.third_party.mujoco_warp._src.collision_core import CollisionContext from mujoco.mjx.third_party.mujoco_warp._src.collision_core import create_collision_context +from mujoco.mjx.third_party.mujoco_warp._src.collision_flex import flex_narrowphase from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import primitive_narrowphase from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf_narrowphase from mujoco.mjx.third_party.mujoco_warp._src.math import upper_tri_index +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseFilter from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseType from mujoco.mjx.third_party.mujoco_warp._src.types import CollisionType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model from mujoco.mjx.third_party.mujoco_warp._src.types import mat23 from mujoco.mjx.third_party.mujoco_warp._src.types import mat63 -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL -from mujoco.mjx.third_party.mujoco_warp._src.types import Model from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -290,25 +292,13 @@ def _broadphase_filter(opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound # 8: obb aabb_id = worldid % ngeom_aabb if wp.static(ngeom_aabb > 1) else 0 - center1, center2 = ( - geom_aabb[aabb_id, geom1, 0], - geom_aabb[aabb_id, geom2, 0], - ) # kernel_analyzer: ignore - size1, size2 = ( - geom_aabb[aabb_id, geom1, 1], - geom_aabb[aabb_id, geom2, 1], - ) # kernel_analyzer: ignore + center1, center2 = geom_aabb[aabb_id, geom1, 0], geom_aabb[aabb_id, geom2, 0] # kernel_analyzer: ignore + size1, size2 = geom_aabb[aabb_id, geom1, 1], geom_aabb[aabb_id, geom2, 1] # kernel_analyzer: ignore rbound_id = worldid % ngeom_rbound if wp.static(ngeom_rbound > 1) else 0 - rbound1, rbound2 = ( - geom_rbound[rbound_id, geom1], - geom_rbound[rbound_id, geom2], - ) # kernel_analyzer: ignore + rbound1, rbound2 = geom_rbound[rbound_id, geom1], geom_rbound[rbound_id, geom2] # kernel_analyzer: ignore margin_id = worldid % ngeom_margin if wp.static(ngeom_margin > 1) else 0 - margin1, margin2 = ( - geom_margin[margin_id, geom1], - geom_margin[margin_id, geom2], - ) # kernel_analyzer: ignore + margin1, margin2 = geom_margin[margin_id, geom1], geom_margin[margin_id, geom2] # kernel_analyzer: ignore xpos1, xpos2 = geom_xpos_in[worldid, geom1], geom_xpos_in[worldid, geom2] xmat1, xmat2 = geom_xmat_in[worldid, geom1], geom_xmat_in[worldid, geom2] @@ -757,6 +747,9 @@ def _narrowphase(m: Model, d: Data, ctx: CollisionContext): if m.has_sdf_geom: sdf_narrowphase(m, d, ctx) + if m.nflex > 0: + flex_narrowphase(m, d) + @event_scope def collision(m: Model, d: Data): diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py new file mode 100644 index 00000000..215423e9 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_flex.py @@ -0,0 +1,834 @@ +# Copyright 2026 The Newton Developers +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Flex collision detection (geom vs flex triangles).""" + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import collision_primitive_core +from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope + +wp.set_module_options({"enable_backward": False}) + + +@wp.func +def _write_flex_contact( + # Data in: + naconmax_in: int, + # In: + dist: float, + pos: wp.vec3, + frame: wp.mat33, + margin: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solimp: vec5, + geom: int, + flexid: int, + vertid: int, + worldid: int, + # Data out: + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=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_flex_out: wp.array(dtype=wp.vec2i), + contact_vert_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), + nacon_out: wp.array(dtype=int), +): + if dist >= margin or dist >= MJ_MAXVAL: + return + + id_ = wp.atomic_add(nacon_out, 0, 1) + if id_ >= naconmax_in: + return + + contact_dist_out[id_] = dist + contact_pos_out[id_] = pos + contact_frame_out[id_] = frame + contact_includemargin_out[id_] = margin + contact_friction_out[id_] = friction + contact_solref_out[id_] = solref + contact_solreffriction_out[id_] = wp.vec2(0.0, 0.0) + contact_solimp_out[id_] = solimp + contact_dim_out[id_] = condim + contact_geom_out[id_] = wp.vec2i(geom, -1) + contact_flex_out[id_] = wp.vec2i(-1, flexid) + contact_vert_out[id_] = wp.vec2i(-1, vertid) + contact_worldid_out[id_] = worldid + contact_type_out[id_] = 1 + contact_geomcollisionid_out[id_] = 0 + + +@wp.func +def _collide_geom_triangle( + # Data in: + naconmax_in: int, + # In: + gtype: int, + pos: wp.vec3, + rot: wp.mat33, + size_val: wp.vec3, + t1: wp.vec3, + t2: wp.vec3, + t3: wp.vec3, + tri_radius: float, + margin: float, + condim: int, + friction: vec5, + solref: wp.vec2, + solimp: vec5, + geomid: int, + flexid: int, + vertex_id: int, + worldid: int, + # Data out: + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=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_flex_out: wp.array(dtype=wp.vec2i), + contact_vert_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), + nacon_out: wp.array(dtype=int), +): + if gtype == int(GeomType.SPHERE): + 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_flex_contact( + naconmax_in, + dist, + contact_pos, + make_frame(nrm), + margin, + condim, + friction, + solref, + solimp, + geomid, + flexid, + vertex_id, + 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_vert_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + return + + # Capsule, box, cylinder all return up to 2 contacts - compute then share writing code + dists = wp.vec2(collision_primitive_core.MJ_MAXVAL, collision_primitive_core.MJ_MAXVAL) + poss = collision_primitive_core.mat23f(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + nrms = collision_primitive_core.mat23f(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + + if gtype == int(GeomType.CAPSULE): + cap_radius = size_val[0] + cap_half_len = size_val[1] + cap_axis = wp.vec3(rot[0, 2], rot[1, 2], rot[2, 2]) + dists, poss, nrms = collision_primitive_core.capsule_triangle( + pos, cap_axis, cap_radius, cap_half_len, t1, t2, t3, tri_radius + ) + elif gtype == int(GeomType.BOX): + dists, poss, nrms = collision_primitive_core.box_triangle(pos, rot, size_val, t1, t2, t3, tri_radius) + elif gtype == int(GeomType.CYLINDER): + cyl_radius = size_val[0] + cyl_half_height = size_val[1] + cyl_axis = wp.vec3(rot[0, 2], rot[1, 2], rot[2, 2]) + dists, poss, nrms = collision_primitive_core.cylinder_triangle( + pos, cyl_axis, cyl_radius, cyl_half_height, t1, t2, t3, tri_radius + ) + + # Write up to 2 contacts (shared code for capsule/box/cylinder) + 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_flex_contact( + naconmax_in, + dists[0], + p1, + make_frame(n1), + margin, + condim, + friction, + solref, + solimp, + geomid, + flexid, + vertex_id, + 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_vert_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + 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_flex_contact( + naconmax_in, + dists[1], + p2, + make_frame(n2), + margin, + condim, + friction, + solref, + solimp, + geomid, + flexid, + vertex_id, + 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_vert_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + + +@wp.kernel +def _flex_plane_narrowphase( + # Model: + ngeom: int, + nflexvert: int, + geom_type: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_solimp: wp.array2d(dtype=vec5), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + flex_condim: wp.array(dtype=int), + flex_friction: wp.array(dtype=wp.vec3), + flex_margin: wp.array(dtype=float), + flex_vertadr: wp.array(dtype=int), + flex_radius: wp.array(dtype=float), + flex_vertflexid: wp.array(dtype=int), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + nworld_in: int, + naconmax_in: int, + # Data out: + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=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_flex_out: wp.array(dtype=wp.vec2i), + contact_vert_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), + nacon_out: wp.array(dtype=int), +): + worldid, vertid = wp.tid() + + flexid = flex_vertflexid[vertid] + radius = flex_radius[flexid] + flex_margin_val = flex_margin[flexid] + flex_condim_val = flex_condim[flexid] + flex_fric = flex_friction[flexid] + # Convert global vertid to local vertex index within this flex + local_vertid = vertid - flex_vertadr[flexid] + + vert = flexvert_xpos_in[worldid, vertid] + + # TODO: Add a broadphase + for geomid in range(ngeom): + gtype = geom_type[geomid] + if gtype != int(GeomType.PLANE): + continue + + plane_pos = geom_xpos_in[worldid, geomid] + plane_rot = geom_xmat_in[worldid, geomid] + plane_normal = wp.vec3(plane_rot[0, 2], plane_rot[1, 2], plane_rot[2, 2]) + + margin = geom_margin[worldid % geom_margin.shape[0], geomid] + flex_margin_val + + diff = vert - plane_pos + signed_dist = wp.dot(diff, plane_normal) + dist = signed_dist - radius + + if dist < margin: + geom_condim_val = geom_condim[geomid] + condim = wp.max(geom_condim_val, flex_condim_val) + solref = geom_solref[worldid % geom_solref.shape[0], geomid] + solimp = geom_solimp[worldid % geom_solimp.shape[0], geomid] + geom_fric = geom_friction[worldid % geom_friction.shape[0], geomid] + fric0 = wp.max(geom_fric[0], flex_fric[0]) + fric1 = wp.max(geom_fric[1], flex_fric[1]) + fric2 = wp.max(geom_fric[2], flex_fric[2]) + friction = vec5( + wp.max(MJ_MINMU, fric0), + wp.max(MJ_MINMU, fric0), + wp.max(MJ_MINMU, fric1), + wp.max(MJ_MINMU, fric2), + wp.max(MJ_MINMU, fric2), + ) + + contact_pos = vert - plane_normal * (dist * 0.5 + radius) + _write_flex_contact( + naconmax_in, + dist, + contact_pos, + make_frame(plane_normal), + margin, + condim, + friction, + solref, + solimp, + geomid, + flexid, + 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_vert_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + + +@wp.kernel +def _flex_narrowphase_dim2( + # Model: + ngeom: int, + nflex: int, + geom_type: wp.array(dtype=int), + geom_contype: wp.array(dtype=int), + geom_conaffinity: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_solimp: wp.array2d(dtype=vec5), + geom_size: wp.array2d(dtype=wp.vec3), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + flex_contype: wp.array(dtype=int), + flex_conaffinity: wp.array(dtype=int), + flex_margin: wp.array(dtype=float), + flex_dim: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_elemadr: wp.array(dtype=int), + flex_elemnum: wp.array(dtype=int), + flex_elemdataadr: wp.array(dtype=int), + flex_elem: wp.array(dtype=int), + flex_radius: wp.array(dtype=float), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + nworld_in: int, + naconmax_in: int, + # Data out: + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=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_flex_out: wp.array(dtype=wp.vec2i), + contact_vert_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), + nacon_out: wp.array(dtype=int), +): + worldid, elemid = wp.tid() + + flexid = int(-1) + for i in range(nflex): + if flex_dim[i] != 2: + continue + elem_adr = flex_elemadr[i] + elem_num = flex_elemnum[i] + if elemid >= elem_adr and elemid < elem_adr + elem_num: + flexid = i + break + + if flexid < 0: + return + + vert_adr = flex_vertadr[flexid] + tri_radius = flex_radius[flexid] + tri_margin = flex_margin[flexid] + + elem_data_idx = flex_elemdataadr[flexid] + (elemid - flex_elemadr[flexid]) * 3 + v0_local = flex_elem[elem_data_idx] + v1_local = flex_elem[elem_data_idx + 1] + v2_local = flex_elem[elem_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] + + # TODO: Add a broadphase + for geomid in range(ngeom): + gtype = geom_type[geomid] + if ( + gtype != int(GeomType.SPHERE) + and gtype != int(GeomType.CAPSULE) + and gtype != int(GeomType.BOX) + and gtype != int(GeomType.CYLINDER) + ): + continue + + g_contype = geom_contype[geomid] + g_conaffinity = geom_conaffinity[geomid] + f_contype = flex_contype[flexid] + f_conaffinity = flex_conaffinity[flexid] + if not ((g_contype & f_conaffinity) or (f_contype & g_conaffinity)): + continue + + 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] + + condim = geom_condim[geomid] + gf = geom_friction[worldid % geom_friction.shape[0], geomid] + friction = vec5( + wp.max(MJ_MINMU, gf[0]), + wp.max(MJ_MINMU, gf[0]), + wp.max(MJ_MINMU, gf[1]), + wp.max(MJ_MINMU, gf[2]), + wp.max(MJ_MINMU, gf[2]), + ) + solref = geom_solref[worldid % geom_solref.shape[0], geomid] + solimp = geom_solimp[worldid % geom_solimp.shape[0], geomid] + + _collide_geom_triangle( + naconmax_in, + gtype, + geom_pos, + geom_rot, + geom_size_val, + t1, + t2, + t3, + tri_radius, + margin, + condim, + friction, + solref, + solimp, + geomid, + flexid, + v0_local, + 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_vert_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + + +@wp.kernel +def _flex_narrowphase_dim3( + # Model: + ngeom: int, + nflex: int, + geom_type: wp.array(dtype=int), + geom_contype: wp.array(dtype=int), + geom_conaffinity: wp.array(dtype=int), + geom_condim: wp.array(dtype=int), + geom_solref: wp.array2d(dtype=wp.vec2), + geom_solimp: wp.array2d(dtype=vec5), + geom_size: wp.array2d(dtype=wp.vec3), + geom_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + flex_contype: wp.array(dtype=int), + flex_conaffinity: wp.array(dtype=int), + flex_margin: wp.array(dtype=float), + flex_dim: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_shellnum: wp.array(dtype=int), + flex_shelldataadr: wp.array(dtype=int), + flex_shell: wp.array(dtype=int), + flex_radius: wp.array(dtype=float), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + nworld_in: int, + naconmax_in: int, + # Data out: + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=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_flex_out: wp.array(dtype=wp.vec2i), + contact_vert_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), + nacon_out: wp.array(dtype=int), +): + worldid, shellid = wp.tid() + + flexid = int(-1) + shell_offset = int(0) + for i in range(nflex): + if flex_dim[i] != 3: + continue + shell_num = flex_shellnum[i] + if shellid >= shell_offset and shellid < shell_offset + shell_num: + flexid = i + break + shell_offset += shell_num + + if flexid < 0: + return + + vert_adr = flex_vertadr[flexid] + tri_radius = flex_radius[flexid] + tri_margin = flex_margin[flexid] + + shell_adr = flex_shelldataadr[flexid] + local_shellid = shellid - shell_offset + shell_data_idx = shell_adr + local_shellid * 3 + + v0_local = flex_shell[shell_data_idx] + v1_local = flex_shell[shell_data_idx + 1] + v2_local = flex_shell[shell_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] + + # TODO: Add a broadphase + for geomid in range(ngeom): + gtype = geom_type[geomid] + if ( + gtype != int(GeomType.SPHERE) + and gtype != int(GeomType.CAPSULE) + and gtype != int(GeomType.BOX) + and gtype != int(GeomType.CYLINDER) + ): + continue + + g_contype = geom_contype[geomid] + g_conaffinity = geom_conaffinity[geomid] + f_contype = flex_contype[flexid] + f_conaffinity = flex_conaffinity[flexid] + if not ((g_contype & f_conaffinity) or (f_contype & g_conaffinity)): + continue + + 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] + + condim = geom_condim[geomid] + gf = geom_friction[worldid % geom_friction.shape[0], geomid] + friction = vec5( + wp.max(MJ_MINMU, gf[0]), + wp.max(MJ_MINMU, gf[0]), + wp.max(MJ_MINMU, gf[1]), + wp.max(MJ_MINMU, gf[2]), + wp.max(MJ_MINMU, gf[2]), + ) + solref = geom_solref[worldid % geom_solref.shape[0], geomid] + solimp = geom_solimp[worldid % geom_solimp.shape[0], geomid] + + _collide_geom_triangle( + naconmax_in, + gtype, + geom_pos, + geom_rot, + geom_size_val, + t1, + t2, + t3, + tri_radius, + margin, + condim, + friction, + solref, + solimp, + geomid, + flexid, + v0_local, + 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_vert_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) + + +@event_scope +def flex_narrowphase(m: Model, d: Data): + """Runs collision detection between geoms and flex elements.""" + if m.nflex == 0: + return + + wp.launch( + _flex_narrowphase_dim2, + dim=(d.nworld, m.nflexelem), + inputs=[ + m.ngeom, + m.nflex, + m.geom_type, + m.geom_contype, + m.geom_conaffinity, + m.geom_condim, + m.geom_solref, + m.geom_solimp, + m.geom_size, + m.geom_friction, + m.geom_margin, + m.flex_contype, + m.flex_conaffinity, + m.flex_margin, + m.flex_dim, + m.flex_vertadr, + m.flex_elemadr, + m.flex_elemnum, + m.flex_elemdataadr, + m.flex_elem, + m.flex_radius, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.nworld, + d.naconmax, + ], + 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.vert, + d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + d.nacon, + ], + ) + + wp.launch( + _flex_narrowphase_dim3, + dim=(d.nworld, m.nflexshelldata // 3), + inputs=[ + m.ngeom, + m.nflex, + m.geom_type, + m.geom_contype, + m.geom_conaffinity, + m.geom_condim, + m.geom_solref, + m.geom_solimp, + m.geom_size, + m.geom_friction, + m.geom_margin, + m.flex_contype, + m.flex_conaffinity, + m.flex_margin, + m.flex_dim, + m.flex_vertadr, + m.flex_shellnum, + m.flex_shelldataadr, + m.flex_shell, + m.flex_radius, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.nworld, + d.naconmax, + ], + 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.vert, + d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + d.nacon, + ], + ) + + wp.launch( + _flex_plane_narrowphase, + dim=(d.nworld, m.nflexvert), + inputs=[ + m.ngeom, + m.nflexvert, + m.geom_type, + m.geom_condim, + m.geom_solref, + m.geom_solimp, + m.geom_friction, + m.geom_margin, + m.flex_condim, + m.flex_friction, + m.flex_margin, + m.flex_vertadr, + m.flex_radius, + m.flex_vertflexid, + d.geom_xpos, + d.geom_xmat, + d.flexvert_xpos, + d.nworld, + d.naconmax, + ], + 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.vert, + d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + d.nacon, + ], + ) 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 7e4948ed..fe1c4445 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 @@ -16,11 +16,12 @@ import math from typing import Tuple +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src.collision_core import Geom from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import mat43 from mujoco.mjx.third_party.mujoco_warp._src.types import mat63 -import warp as wp # TODO(team): improve compile time to enable backward pass wp.set_module_options({"enable_backward": False}) @@ -581,16 +582,19 @@ def gjk( simplex_index2 = wp.vec4i() n = int(0) coordinates = wp.vec4() # barycentric coordinates - epsilon = wp.where(is_discrete, 0.0, 0.5 * tolerance * tolerance) + tol2 = tolerance * tolerance + epsilon = wp.where(is_discrete, 0.0, 0.5 * tol2) # set initial guess x_k = x1_0 - x2_0 + xnorm_old = FLOAT_MAX - for k in range(gjk_iterations): + for _ in range(gjk_iterations): xnorm = wp.dot(x_k, x_k) # TODO(kbayes): determine new constant here - if xnorm < 1e-12: + if xnorm < tol2 or wp.abs(xnorm_old - xnorm) < tol2: break + xnorm_old = xnorm dir_neg = x_k / wp.sqrt(xnorm) # compute kth support point in geom1 @@ -663,13 +667,6 @@ def gjk( if n == 4: break - if k == gjk_iterations - 1: - wp.printf( - "Warning: opt.ccd_iterations, currently set to %d, needs to be" - " increased.\n", - gjk_iterations, - ) - result = GJKResult() # compute the approximate witness points @@ -1205,7 +1202,6 @@ def _is_invalid_face(face: int) -> bool: def _epa( # In: tolerance: float, - gjk_iterations: int, epa_iterations: int, pt: Polytope, geom1: Geom, @@ -1226,7 +1222,7 @@ def _epa( # so iterations must be cap to limit the number of generated vertices # (one new vertex per iteration) epa_iterations = wp.min(epa_iterations, 1000) - for k in range(epa_iterations): + for _ in range(epa_iterations): pidx = idx idx = int(-1) lower2 = float(FLOAT_MAX) @@ -1325,13 +1321,6 @@ def _epa( # clear horizon pt.nhorizon = 0 - if k == epa_iterations - 1: - wp.printf( - "Warning: opt.ccd_iterations, currently set to %d, needs to be" - " increased.\n", - gjk_iterations, - ) - # return from valid face if idx > -1: x1, x2, dist = _epa_witness(pt, geom1, geom2, geomtype1, geomtype2, idx) @@ -2347,7 +2336,7 @@ def ccd( if pt.status: return result.dist, 1, result.x1, result.x2, -1 - dist, x1, x2, idx = _epa(tolerance, gjk_iterations, epa_iterations, pt, geom1, geom2, geomtype1, geomtype2, is_discrete) + dist, x1, x2, idx = _epa(tolerance, epa_iterations, pt, geom1, geom2, geomtype1, geomtype2, is_discrete) if idx == -1: return FLOAT_MAX, 0, wp.vec3(), wp.vec3(), -1 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 c7a40514..f1829de4 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 @@ -15,9 +15,11 @@ from typing import Tuple +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src.collision_core import CollisionContext -from mujoco.mjx.third_party.mujoco_warp._src.collision_core import contact_params from mujoco.mjx.third_party.mujoco_warp._src.collision_core import Geom +from mujoco.mjx.third_party.mujoco_warp._src.collision_core import contact_params from mujoco.mjx.third_party.mujoco_warp._src.collision_core import geom_collision_pair from mujoco.mjx.third_party.mujoco_warp._src.collision_core import write_contact from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import box_box @@ -34,15 +36,14 @@ from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import sph from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import sphere_sphere from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType -from mujoco.mjx.third_party.mujoco_warp._src.types import mat43 -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import mat43 from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -304,6 +305,7 @@ def plane_sphere_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -339,6 +341,7 @@ def plane_sphere_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -374,6 +377,7 @@ def sphere_sphere_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -408,6 +412,7 @@ def sphere_sphere_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -443,6 +448,7 @@ def sphere_capsule_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -480,6 +486,7 @@ def sphere_capsule_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -515,6 +522,7 @@ def capsule_capsule_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -564,6 +572,7 @@ def capsule_capsule_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -599,6 +608,7 @@ def plane_capsule_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -644,6 +654,7 @@ def plane_capsule_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -679,6 +690,7 @@ def plane_ellipsoid_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -713,6 +725,7 @@ def plane_ellipsoid_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -748,6 +761,7 @@ def plane_box_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -784,6 +798,7 @@ def plane_box_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -819,6 +834,7 @@ def plane_convex_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -855,6 +871,7 @@ def plane_convex_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -890,6 +907,7 @@ def sphere_cylinder_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -934,6 +952,7 @@ def sphere_cylinder_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -969,6 +988,7 @@ def plane_cylinder_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -1015,6 +1035,7 @@ def plane_cylinder_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -1050,6 +1071,7 @@ def sphere_box_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -1083,6 +1105,7 @@ def sphere_box_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -1118,6 +1141,7 @@ def capsule_box_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -1166,6 +1190,7 @@ def capsule_box_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -1201,6 +1226,7 @@ def box_box_wrapper( contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -1245,6 +1271,7 @@ def box_box_wrapper( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -1327,6 +1354,7 @@ def _primitive_narrowphase(primitive_collisions_types, primitive_collisions_func contact_solimp_out: wp.array(dtype=vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), @@ -1416,6 +1444,7 @@ def _primitive_narrowphase(primitive_collisions_types, primitive_collisions_func contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -1511,6 +1540,7 @@ def primitive_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_ta d.contact.solimp, d.contact.dim, d.contact.geom, + d.contact.efc_address, d.contact.worldid, d.contact.type, d.contact.geomcollisionid, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py index 306a301e..8a85d241 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive_core.py @@ -1489,3 +1489,507 @@ def capsule_box( mat23f(pos1[0], pos1[1], pos1[2], pos2[0], pos2[1], pos2[2]), mat23f(normal1[0], normal1[1], normal1[2], normal2[0], normal2[1], normal2[2]), ) + + +@wp.func +def _tri_area_sign(p1: wp.vec2, p2: wp.vec2, p3: wp.vec2) -> float: + """Sign of (signed) area of planar triangle.""" + return wp.sign((p1[0] - p3[0]) * (p2[1] - p3[1]) - (p2[0] - p3[0]) * (p1[1] - p3[1])) + + +@wp.func +def _tri_point_segment(p: wp.vec2, u: wp.vec2, v: wp.vec2) -> wp.vec2: + """Find nearest point to p within line segment (u, v).""" + uv = v - u + up = p - u + + denom = wp.max(MJ_MINVAL, wp.dot(uv, uv)) + a = wp.dot(uv, up) / denom + + if a <= 0.0: + return u + elif a >= 1.0: + return v + else: + return u + a * uv + + +@wp.func +def sphere_triangle( + sphere_pos: wp.vec3, + sphere_radius: float, + t1: wp.vec3, + t2: wp.vec3, + t3: wp.vec3, + tri_radius: float, +) -> Tuple[float, wp.vec3, wp.vec3]: + """Core contact geometry calculation for sphere-triangle collision. + + Port of mjraw_SphereTriangle from engine_collision_primitive.c + + Args: + sphere_pos: Center position of the sphere. + sphere_radius: Radius of the sphere. + t1: Triangle vertex positions. + t2: Triangle vertex positions. + t3: Triangle vertex positions. + tri_radius: Triangle (flex element) radius. + + Returns: + - Contact distance (MJ_MAXVAL if no collision). + - Contact position. + - Contact normal vector. + """ + S = sphere_pos - t1 + A = t2 - t1 + B = t3 - t1 + + N = wp.normalize(wp.cross(A, B)) + + dstS = wp.dot(N, S) + + P = S - dstS * N + + V1 = wp.normalize(A) + lenA = wp.length(A) + V2 = wp.normalize(wp.cross(N, A)) + + o = wp.vec2(0.0, 0.0) + a = wp.vec2(lenA, 0.0) + b = wp.vec2(wp.dot(V1, B), wp.dot(V2, B)) + p = wp.vec2(wp.dot(V1, P), wp.dot(V2, P)) + + sign1 = _tri_area_sign(p, o, a) + sign2 = _tri_area_sign(p, a, b) + sign3 = _tri_area_sign(p, b, o) + + X = wp.vec3(0.0) + if sign1 == sign2 and sign2 == sign3: + X = P + else: + x0 = _tri_point_segment(p, o, a) + x1 = _tri_point_segment(p, a, b) + x2 = _tri_point_segment(p, b, o) + + d0 = wp.length(p - x0) + d1 = wp.length(p - x1) + d2 = wp.length(p - x2) + + if d0 < d1 and d0 < d2: + X = x0[0] * V1 + x0[1] * V2 + elif d1 < d2: + X = x1[0] * V1 + x1[1] * V2 + else: + X = x2[0] * V1 + x2[1] * V2 + + nrm = X - S + dst = wp.length(nrm) + + if dst > MJ_MINVAL: + nrm = nrm / dst + else: + nrm = N + + dist = dst - sphere_radius - tri_radius + pos = sphere_pos + nrm * (sphere_radius + 0.5 * dist) + + return dist, pos, nrm + + +@wp.func +def box_triangle( + box_pos: wp.vec3, + box_rot: wp.mat33, + box_size: wp.vec3, + t1: wp.vec3, + t2: wp.vec3, + t3: wp.vec3, + tri_radius: float, +) -> Tuple[wp.vec2, mat23f, mat23f]: + """Core contact geometry calculation for box-triangle collision. + + Port of mjraw_BoxTriangle from engine_collision_primitive.c + + Args: + box_pos: Center position of the box. + box_rot: Orientation matrix of the box. + box_size: Half-sizes of the box. + t1: Triangle vertex positions. + t2: Triangle vertex positions. + t3: Triangle vertex positions. + tri_radius: Triangle (flex element) radius. + + Returns: + - wp.vec2 of distances for up to 2 contacts (MJ_MAXVAL if no collision). + - mat23f of contact positions (2 x vec3). + - mat23f of contact normals (2 x vec3). + """ + dist1 = MJ_MAXVAL + dist2 = MJ_MAXVAL + pos1 = wp.vec3(0.0) + pos2 = wp.vec3(0.0) + nrm1 = wp.vec3(0.0) + nrm2 = wp.vec3(0.0) + cnt = 0 + + box_rotT = wp.transpose(box_rot) + + for vi in range(3): + vert = wp.vec3(0.0) + if vi == 0: + vert = t1 + elif vi == 1: + vert = t2 + else: + vert = t3 + + diff = vert - box_pos + local = box_rotT @ diff + + maxaxis = 0 + maxval = wp.abs(local[0]) - box_size[0] + for j in range(1, 3): + val = wp.abs(local[j]) - box_size[j] + if val > maxval: + maxval = val + maxaxis = j + + inside = True + for j in range(3): + if wp.abs(local[j]) > box_size[j] + tri_radius: + inside = False + + if inside and cnt < 2: + nrm_local = wp.vec3(0.0) + if maxaxis == 0: + nrm_local = wp.vec3(wp.sign(local[0]), 0.0, 0.0) + elif maxaxis == 1: + nrm_local = wp.vec3(0.0, wp.sign(local[1]), 0.0) + else: + nrm_local = wp.vec3(0.0, 0.0, wp.sign(local[2])) + + nrm_global = box_rot @ nrm_local + d = maxval - tri_radius + offset = tri_radius + d * 0.5 + p = vert - nrm_global * offset + + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = nrm_global + else: + dist2 = d + pos2 = p + nrm2 = nrm_global + cnt += 1 + + for i in range(8): + if cnt >= 2: + break + + vec = wp.vec3( + wp.where(i & 1, box_size[0], -box_size[0]), + wp.where(i & 2, box_size[1], -box_size[1]), + wp.where(i & 4, box_size[2], -box_size[2]), + ) + corner = box_rot @ vec + box_pos + + d, p, n = sphere_triangle(corner, 0.0, t1, t2, t3, tri_radius) + if d < MJ_MAXVAL: + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = n + elif cnt == 1: + dist2 = d + pos2 = p + nrm2 = n + cnt += 1 + + return ( + wp.vec2(dist1, dist2), + mat23f(pos1[0], pos1[1], pos1[2], pos2[0], pos2[1], pos2[2]), + mat23f(nrm1[0], nrm1[1], nrm1[2], nrm2[0], nrm2[1], nrm2[2]), + ) + + +@wp.func +def capsule_triangle( + capsule_pos: wp.vec3, + capsule_axis: wp.vec3, + capsule_radius: float, + capsule_half_length: float, + t1: wp.vec3, + t2: wp.vec3, + t3: wp.vec3, + tri_radius: float, +) -> Tuple[wp.vec2, mat23f, mat23f]: + """Core contact geometry calculation for capsule-triangle collision. + + Port of mjraw_CapsuleTriangle from engine_collision_primitive.c + + Args: + capsule_pos: Center position of the capsule. + capsule_axis: Unit axis direction of the capsule. + capsule_radius: Radius of the capsule. + capsule_half_length: Half-length of the capsule cylinder. + t1: Triangle vertex positions. + t2: Triangle vertex positions. + t3: Triangle vertex positions. + tri_radius: Triangle (flex element) radius. + + Returns: + - wp.vec2 of distances for up to 2 contacts (MJ_MAXVAL if no collision). + - mat23f of contact positions (2 x vec3). + - mat23f of contact normals (2 x vec3). + """ + dist1 = MJ_MAXVAL + dist2 = MJ_MAXVAL + pos1 = wp.vec3(0.0) + pos2 = wp.vec3(0.0) + nrm1 = wp.vec3(0.0) + nrm2 = wp.vec3(0.0) + cnt = 0 + + p1 = capsule_pos - capsule_axis * capsule_half_length + p2 = capsule_pos + capsule_axis * capsule_half_length + + d, p, n = sphere_triangle(p1, capsule_radius, t1, t2, t3, tri_radius) + if d < MJ_MAXVAL: + dist1 = d + pos1 = p + nrm1 = n + cnt = 1 + + d, p, n = sphere_triangle(p2, capsule_radius, t1, t2, t3, tri_radius) + if d < MJ_MAXVAL and cnt < 2: + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = n + else: + dist2 = d + pos2 = p + nrm2 = n + cnt += 1 + + ab = p2 - p1 + ab_len_sq = 4.0 * capsule_half_length * capsule_half_length + + for vi in range(3): + if cnt >= 2: + break + + vert = wp.vec3(0.0) + if vi == 0: + vert = t1 + elif vi == 1: + vert = t2 + else: + vert = t3 + + vec = vert - p1 + t_param = wp.dot(vec, ab) / wp.max(MJ_MINVAL, ab_len_sq) + + if t_param > MJ_MINVAL and t_param < 1.0 - MJ_MINVAL: + closest = p1 + ab * t_param + diff = vert - closest + dist_raw = wp.length(diff) + + if dist_raw > MJ_MINVAL: + nrm = diff / dist_raw + d = dist_raw - capsule_radius - tri_radius + p = (closest + vert + nrm * (capsule_radius - tri_radius)) * 0.5 + + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = nrm + else: + dist2 = d + pos2 = p + nrm2 = nrm + cnt += 1 + + return ( + wp.vec2(dist1, dist2), + mat23f(pos1[0], pos1[1], pos1[2], pos2[0], pos2[1], pos2[2]), + mat23f(nrm1[0], nrm1[1], nrm1[2], nrm2[0], nrm2[1], nrm2[2]), + ) + + +@wp.func +def cylinder_triangle( + cylinder_pos: wp.vec3, + cylinder_axis: wp.vec3, + cylinder_radius: float, + cylinder_half_height: float, + t1: wp.vec3, + t2: wp.vec3, + t3: wp.vec3, + tri_radius: float, +) -> Tuple[wp.vec2, mat23f, mat23f]: + """Core contact geometry calculation for cylinder-triangle collision. + + Args: + cylinder_pos: Center position of the cylinder. + cylinder_axis: Unit axis direction of the cylinder. + cylinder_radius: Radius of the cylinder. + cylinder_half_height: Half-height of the cylinder. + t1: Triangle vertex positions. + t2: Triangle vertex positions. + t3: Triangle vertex positions. + tri_radius: Triangle (flex element) radius. + + Returns: + - wp.vec2 of distances for up to 2 contacts (MJ_MAXVAL if no collision). + - mat23f of contact positions (2 x vec3). + - mat23f of contact normals (2 x vec3). + """ + dist1 = MJ_MAXVAL + dist2 = MJ_MAXVAL + pos1 = wp.vec3(0.0) + pos2 = wp.vec3(0.0) + nrm1 = wp.vec3(0.0) + nrm2 = wp.vec3(0.0) + cnt = int(0) + + p1 = cylinder_pos - cylinder_axis * cylinder_half_height + p2 = cylinder_pos + cylinder_axis * cylinder_half_height + + ab = p2 - p1 + ab_len_sq = 4.0 * cylinder_half_height * cylinder_half_height + + for vi in range(3): + if cnt >= 2: + break + + vert = wp.vec3(0.0) + if vi == 0: + vert = t1 + elif vi == 1: + vert = t2 + else: + vert = t3 + + vec = vert - p1 + t_param = wp.dot(vec, ab) / wp.max(MJ_MINVAL, ab_len_sq) + + if t_param > MJ_MINVAL and t_param < 1.0 - MJ_MINVAL: + closest = p1 + ab * t_param + diff = vert - closest + dist_raw = wp.length(diff) + + if dist_raw < cylinder_radius + tri_radius: + if dist_raw > MJ_MINVAL: + nrm = diff / dist_raw + d = dist_raw - cylinder_radius - tri_radius + p = (closest + vert + nrm * (cylinder_radius - tri_radius)) * 0.5 + else: + dist_to_side = cylinder_radius + dist_to_p2 = (1.0 - t_param) * wp.sqrt(ab_len_sq) + dist_to_p1 = t_param * wp.sqrt(ab_len_sq) + + if dist_to_p2 < dist_to_side and dist_to_p2 < dist_to_p1: + nrm = cylinder_axis + d = -dist_to_p2 - tri_radius + p = vert + elif dist_to_p1 < dist_to_side: + nrm = -cylinder_axis + d = -dist_to_p1 - tri_radius + p = vert + else: + tri_normal = wp.normalize(wp.cross(t2 - t1, t3 - t1)) + nrm = tri_normal + d = -cylinder_radius - tri_radius + p = closest + + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = nrm + else: + dist2 = d + pos2 = p + nrm2 = nrm + cnt += 1 + elif t_param <= MJ_MINVAL: + diff = vert - p1 + signed_dist = wp.dot(diff, cylinder_axis) + perp = diff - cylinder_axis * signed_dist + perp_len = wp.length(perp) + + if perp_len < cylinder_radius: + d = -signed_dist - tri_radius + nrm = -cylinder_axis + p = vert - nrm * (tri_radius + d * 0.5) + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = nrm + else: + dist2 = d + pos2 = p + nrm2 = nrm + cnt += 1 + elif perp_len < cylinder_radius + tri_radius: + edge_dir = perp / perp_len + edge_point = p1 + edge_dir * cylinder_radius + diff_to_edge = vert - edge_point + dist_raw = wp.length(diff_to_edge) + if dist_raw > MJ_MINVAL: + nrm = diff_to_edge / dist_raw + d = dist_raw - tri_radius + p = vert - nrm * (tri_radius + d * 0.5) + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = nrm + else: + dist2 = d + pos2 = p + nrm2 = nrm + cnt += 1 + else: + diff = vert - p2 + signed_dist = wp.dot(diff, cylinder_axis) + perp = diff - cylinder_axis * signed_dist + perp_len = wp.length(perp) + + if perp_len < cylinder_radius: + d = signed_dist - tri_radius + nrm = cylinder_axis + p = vert - nrm * (tri_radius + d * 0.5) + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = nrm + else: + dist2 = d + pos2 = p + nrm2 = nrm + cnt += 1 + elif perp_len < cylinder_radius + tri_radius: + edge_dir = perp / perp_len + edge_point = p2 + edge_dir * cylinder_radius + diff_to_edge = vert - edge_point + dist_raw = wp.length(diff_to_edge) + if dist_raw > MJ_MINVAL: + nrm = diff_to_edge / dist_raw + d = dist_raw - tri_radius + p = vert - nrm * (tri_radius + d * 0.5) + if cnt == 0: + dist1 = d + pos1 = p + nrm1 = nrm + else: + dist2 = d + pos2 = p + nrm2 = nrm + cnt += 1 + + return ( + wp.vec2(dist1, dist2), + mat23f(pos1[0], pos1[1], pos1[2], pos2[0], pos2[1], pos2[2]), + mat23f(nrm1[0], nrm1[1], nrm1[2], nrm2[0], nrm2[1], nrm2[2]), + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_sdf.py index c84b67c8..4decf33c 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 @@ -15,6 +15,8 @@ from typing import Tuple +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src.collision_core import CollisionContext from mujoco.mjx.third_party.mujoco_warp._src.collision_core import contact_params from mujoco.mjx.third_party.mujoco_warp._src.collision_core import geom_collision_pair @@ -27,9 +29,9 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import Model from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 from mujoco.mjx.third_party.mujoco_warp._src.types import vec8 from mujoco.mjx.third_party.mujoco_warp._src.types import vec8i +from mujoco.mjx.third_party.mujoco_warp._src.types import vec_pluginattr from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -38,8 +40,8 @@ wp.set_module_options({"enable_backward": False}) class OptimizationParams: rel_mat: wp.mat33 rel_pos: wp.vec3 - attr1: wp.vec3 - attr2: wp.vec3 + attr1: vec_pluginattr + attr2: vec_pluginattr @wp.struct @@ -77,20 +79,24 @@ class MeshData: @wp.func def get_sdf_params( - # Model: - oct_child: wp.array(dtype=vec8i), - oct_aabb: wp.array2d(dtype=wp.vec3), - oct_coeff: wp.array(dtype=vec8), - mesh_octadr: wp.array(dtype=int), - plugin: wp.array(dtype=int), - plugin_attr: wp.array(dtype=wp.vec3f), - # In: - g_type: int, - g_size: wp.vec3, - plugin_id: int, - mesh_id: int, -) -> Tuple[wp.vec3, int, VolumeData, MeshData]: - attributes = g_size + # Model: + oct_child: wp.array(dtype=vec8i), + oct_aabb: wp.array2d(dtype=wp.vec3), + oct_coeff: wp.array(dtype=vec8), + mesh_octadr: wp.array(dtype=int), + plugin: wp.array(dtype=int), + plugin_attr: wp.array(dtype=vec_pluginattr), + # In: + g_type: int, + g_size: wp.vec3, + plugin_id: int, + mesh_id: int, +) -> Tuple[vec_pluginattr, int, VolumeData, MeshData]: + # default attributes from geom size, first 3 values copied + attributes = vec_pluginattr() + attributes[0] = g_size[0] + attributes[1] = g_size[1] + attributes[2] = g_size[2] plugin_index = -1 volume_data = VolumeData() @@ -108,6 +114,16 @@ def get_sdf_params( volume_data.oct_coeff = oct_coeff volume_data.valid = True + elif g_type == GeomType.MESH and mesh_id != -1 and mesh_octadr[mesh_id] != -1: + octadr = mesh_octadr[mesh_id] + volume_data.center = oct_aabb[octadr, 0] + volume_data.half_size = oct_aabb[octadr, 1] + volume_data.root = octadr + volume_data.oct_aabb = oct_aabb + volume_data.oct_child = oct_child + volume_data.oct_coeff = oct_coeff + volume_data.valid = True + return attributes, plugin_index, volume_data, MeshData() @@ -215,24 +231,28 @@ def grad_ellipsoid(p: wp.vec3, size: wp.vec3) -> wp.vec3: @wp.func -def user_sdf(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> float: +def user_sdf(p: wp.vec3, attr: vec_pluginattr, sdf_type: int) -> float: + """User-defined SDF function. + + Access attributes via attr[i] where i is the attribute index (0 to _NPLUGINATTR-1). + """ wp.printf("ERROR: user_sdf function must be implemented by user code\n") return 0.0 @wp.func -def user_sdf_grad(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3: +def user_sdf_grad(p: wp.vec3, attr: vec_pluginattr, sdf_type: int) -> wp.vec3: + """User-defined SDF gradient function. + + Access attributes via attr[i] where i is the attribute index (0 to _NPLUGINATTR-1). + """ wp.printf("ERROR: user_sdf_grad function must be implemented by user code\n") return wp.vec3(0.0) @wp.func def find_oct( - oct_child: wp.array(dtype=vec8i), - oct_aabb: wp.array2d(dtype=wp.vec3), - p: wp.vec3, - grad: bool, - root: int, + oct_child: wp.array(dtype=vec8i), oct_aabb: wp.array2d(dtype=wp.vec3), p: wp.vec3, grad: bool, root: int ) -> Tuple[int, Tuple[vec8, vec8, vec8]]: stack = root niter = int(100) @@ -268,14 +288,14 @@ def find_oct( # child indices are relative to root (mesh_octadr offset) child0 = oct_child[node][0] if ( - child0 == -1 - and oct_child[node][1] == -1 - and oct_child[node][2] == -1 - and oct_child[node][3] == -1 - and oct_child[node][4] == -1 - and oct_child[node][5] == -1 - and oct_child[node][6] == -1 - and oct_child[node][7] == -1 + child0 == -1 + and oct_child[node][1] == -1 + and oct_child[node][2] == -1 + and oct_child[node][3] == -1 + and oct_child[node][4] == -1 + and oct_child[node][5] == -1 + and oct_child[node][6] == -1 + and oct_child[node][7] == -1 ): for j in range(8): if not grad: @@ -342,13 +362,7 @@ def box_project(center: wp.vec3, half_size: wp.vec3, xyz: wp.vec3) -> Tuple[floa @wp.func def sample_volume_sdf(xyz: wp.vec3, volume_data: VolumeData) -> float: dist0, point = box_project(volume_data.center, volume_data.half_size, xyz) - node, weights = find_oct( - volume_data.oct_child, - volume_data.oct_aabb, - point, - grad=False, - root=volume_data.root, - ) + node, weights = find_oct(volume_data.oct_child, volume_data.oct_aabb, point, grad=False, root=volume_data.root) return dist0 + wp.dot(weights[0], volume_data.oct_coeff[node]) @@ -365,13 +379,7 @@ def sample_volume_grad(xyz: wp.vec3, volume_data: VolumeData) -> wp.vec3: grad_y = (sample_volume_sdf(xyz + dy, volume_data) - f) / h grad_z = (sample_volume_sdf(xyz + dz, volume_data) - f) / h return wp.vec3(grad_x, grad_y, grad_z) - node, weights = find_oct( - volume_data.oct_child, - volume_data.oct_aabb, - point, - grad=True, - root=volume_data.root, - ) + node, weights = find_oct(volume_data.oct_child, volume_data.oct_aabb, point, grad=True, root=volume_data.root) grad_x = wp.dot(weights[0], volume_data.oct_coeff[node]) grad_y = wp.dot(weights[1], volume_data.oct_coeff[node]) grad_z = wp.dot(weights[2], volume_data.oct_coeff[node]) @@ -379,15 +387,17 @@ def sample_volume_grad(xyz: wp.vec3, volume_data: VolumeData) -> wp.vec3: @wp.func -def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: VolumeData, mesh_data: MeshData) -> float: +def sdf(type: int, p: wp.vec3, attr: vec_pluginattr, sdf_type: int, volume_data: VolumeData, mesh_data: MeshData) -> float: + # extract first 3 elements as vec3 for primitive sdf functions + attr_vec3 = wp.vec3(attr[0], attr[1], attr[2]) if type == GeomType.PLANE: return p[2] elif type == GeomType.SPHERE: - return sphere(p, attr) + return sphere(p, attr_vec3) elif type == GeomType.BOX: - return box(p, attr) + return box(p, attr_vec3) elif type == GeomType.ELLIPSOID: - return ellipsoid(p, attr) + return ellipsoid(p, attr_vec3) elif type == GeomType.MESH and mesh_data.valid: mesh_data.pnt = p mesh_data.vec = -wp.normalize(p) @@ -425,21 +435,27 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: Volume return sample_volume_sdf(p, volume_data) else: return user_sdf(p, attr, sdf_type) + elif type == GeomType.MESH and volume_data.valid: + return sample_volume_sdf(p, volume_data) wp.printf("ERROR: SDF type not implemented\n") return 0.0 @wp.func -def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: VolumeData, mesh_data: MeshData) -> wp.vec3: +def sdf_grad( + type: int, p: wp.vec3, attr: vec_pluginattr, sdf_type: int, volume_data: VolumeData, mesh_data: MeshData +) -> wp.vec3: + # extract first 3 elements as vec3 for primitive sdf functions + attr_vec3 = wp.vec3(attr[0], attr[1], attr[2]) if type == GeomType.PLANE: grad = wp.vec3(0.0, 0.0, 1.0) return grad elif type == GeomType.SPHERE: return grad_sphere(p) elif type == GeomType.BOX: - return grad_box(p, attr) + return grad_box(p, attr_vec3) elif type == GeomType.ELLIPSOID: - return grad_ellipsoid(p, attr) + return grad_ellipsoid(p, attr_vec3) elif type == GeomType.MESH and mesh_data.valid: mesh_data.pnt = p mesh_data.vec = -wp.normalize(p) @@ -466,6 +482,8 @@ def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: V return sample_volume_grad(p, volume_data) else: return user_sdf_grad(p, attr, sdf_type) + elif type == GeomType.MESH and volume_data.valid: + return sample_volume_grad(p, volume_data) wp.printf("ERROR: SDF grad type not implemented\n") return wp.vec3(0.0) @@ -476,8 +494,8 @@ def clearance( type1: int, p1: wp.vec3, p2: wp.vec3, - s1: wp.vec3, - s2: wp.vec3, + s1: vec_pluginattr, + s2: vec_pluginattr, sdf_type1: int, sdf_type2: int, sfd_intersection: bool, @@ -606,8 +624,8 @@ def gradient_descent( # In: type1: int, x0_initial: wp.vec3, - attr1: wp.vec3, - attr2: wp.vec3, + attr1: vec_pluginattr, + attr2: vec_pluginattr, pos1: wp.vec3, rot1: wp.mat33, pos2: wp.vec3, @@ -645,76 +663,77 @@ def gradient_descent( @wp.kernel def _sdf_narrowphase( - # Model: - nmeshface: int, - oct_child: wp.array(dtype=vec8i), - oct_aabb: wp.array2d(dtype=wp.vec3), - oct_coeff: wp.array(dtype=vec8), - 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_friction: wp.array2d(dtype=wp.vec3), - geom_margin: wp.array2d(dtype=float), - geom_gap: wp.array2d(dtype=float), - mesh_vertadr: wp.array(dtype=int), - mesh_vertnum: wp.array(dtype=int), - mesh_faceadr: wp.array(dtype=int), - mesh_octadr: wp.array(dtype=int), - mesh_graphadr: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), - mesh_face: wp.array(dtype=wp.vec3i), - 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), - plugin: wp.array(dtype=int), - plugin_attr: wp.array(dtype=wp.vec3f), - geom_plugin_index: wp.array(dtype=int), - # Data in: - geom_xpos_in: wp.array2d(dtype=wp.vec3), - geom_xmat_in: wp.array2d(dtype=wp.mat33), - naconmax_in: int, - ncollision_in: wp.array(dtype=int), - # In: - collision_pair_in: wp.array(dtype=wp.vec2i), - collision_pairid_in: wp.array(dtype=wp.vec2i), - collision_worldid_in: wp.array(dtype=int), - sdf_initpoints: int, - sdf_iterations: int, - # Data out: - contact_dist_out: wp.array(dtype=float), - contact_pos_out: wp.array(dtype=wp.vec3), - contact_frame_out: wp.array(dtype=wp.mat33), - contact_includemargin_out: wp.array(dtype=float), - contact_friction_out: wp.array(dtype=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), - nacon_out: wp.array(dtype=int), + # Model: + nmeshface: int, + oct_child: wp.array(dtype=vec8i), + oct_aabb: wp.array2d(dtype=wp.vec3), + oct_coeff: wp.array(dtype=vec8), + 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_friction: wp.array2d(dtype=wp.vec3), + geom_margin: wp.array2d(dtype=float), + geom_gap: wp.array2d(dtype=float), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_faceadr: wp.array(dtype=int), + mesh_octadr: wp.array(dtype=int), + mesh_graphadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_face: wp.array(dtype=wp.vec3i), + 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), + plugin: wp.array(dtype=int), + plugin_attr: wp.array(dtype=vec_pluginattr), + geom_plugin_index: wp.array(dtype=int), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + naconmax_in: int, + ncollision_in: wp.array(dtype=int), + # In: + collision_pair_in: wp.array(dtype=wp.vec2i), + collision_pairid_in: wp.array(dtype=wp.vec2i), + collision_worldid_in: wp.array(dtype=int), + sdf_initpoints: int, + sdf_iterations: int, + # Data out: + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=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_efc_address_out: wp.array2d(dtype=int), + contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): i, contact_tid = wp.tid() if i >= sdf_initpoints: @@ -799,29 +818,11 @@ def _sdf_narrowphase( rot1 = geom1.rot attr1, g1_plugin_id, volume_data1, mesh_data1 = get_sdf_params( - oct_child, - oct_aabb, - oct_coeff, - mesh_octadr, - plugin, - plugin_attr, - type1, - geom1.size, - g1_plugin, - geom_dataid[g1], + oct_child, oct_aabb, oct_coeff, mesh_octadr, plugin, plugin_attr, type1, geom1.size, g1_plugin, geom_dataid[g1] ) attr2, g2_plugin_id, volume_data2, mesh_data2 = get_sdf_params( - oct_child, - oct_aabb, - oct_coeff, - mesh_octadr, - plugin, - plugin_attr, - type2, - geom2.size, - g2_plugin, - geom_dataid[g2], + oct_child, oct_aabb, oct_coeff, mesh_octadr, plugin, plugin_attr, type2, geom2.size, g2_plugin, geom_dataid[g2] ) mesh_data1.nmeshface = nmeshface @@ -900,6 +901,7 @@ def _sdf_narrowphase( contact_solimp_out, contact_dim_out, contact_geom_out, + contact_efc_address_out, contact_worldid_out, contact_type_out, contact_geomcollisionid_out, @@ -910,76 +912,77 @@ def _sdf_narrowphase( @event_scope def sdf_narrowphase(m: Model, d: Data, ctx: CollisionContext): wp.launch( - _sdf_narrowphase, - dim=(m.opt.sdf_initpoints, d.naconmax), - inputs=[ - m.nmeshface, - m.oct_child, - m.oct_aabb, - m.oct_coeff, - 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_aabb, - m.geom_friction, - m.geom_margin, - m.geom_gap, - m.mesh_vertadr, - m.mesh_vertnum, - m.mesh_faceadr, - m.mesh_octadr, - m.mesh_graphadr, - m.mesh_vert, - m.mesh_face, - 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, - m.plugin, - m.plugin_attr, - m.geom_plugin_index, - d.geom_xpos, - d.geom_xmat, - d.naconmax, - d.ncollision, - ctx.collision_pair, - ctx.collision_pairid, - ctx.collision_worldid, - m.opt.sdf_initpoints, - m.opt.sdf_iterations, - ], - 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.worldid, - d.contact.type, - d.contact.geomcollisionid, - d.nacon, - ], + _sdf_narrowphase, + dim=(m.opt.sdf_initpoints, d.naconmax), + inputs=[ + m.nmeshface, + m.oct_child, + m.oct_aabb, + m.oct_coeff, + 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_aabb, + m.geom_friction, + m.geom_margin, + m.geom_gap, + m.mesh_vertadr, + m.mesh_vertnum, + m.mesh_faceadr, + m.mesh_octadr, + m.mesh_graphadr, + m.mesh_vert, + m.mesh_face, + 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, + m.plugin, + m.plugin_attr, + m.geom_plugin_index, + d.geom_xpos, + d.geom_xmat, + d.naconmax, + d.ncollision, + ctx.collision_pair, + ctx.collision_pairid, + ctx.collision_worldid, + m.opt.sdf_initpoints, + m.opt.sdf_iterations, + ], + outputs=[ + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.includemargin, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.dim, + d.contact.geom, + d.contact.efc_address, + d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + d.nacon, + ], ) 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 228006e8..eec47583 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py @@ -13,18 +13,18 @@ # limitations under the License. # ============================================================================== +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src import math from mujoco.mjx.third_party.mujoco_warp._src import support from mujoco.mjx.third_party.mujoco_warp._src import types from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit -from mujoco.mjx.third_party.mujoco_warp._src.types import SPARSE_CONSTRAINT_JACOBIAN -from mujoco.mjx.third_party.mujoco_warp._src.types import vec11 from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.types import vec11 from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -36,6 +36,8 @@ def _zero_constraint_counts( nf_out: wp.array(dtype=int), nl_out: wp.array(dtype=int), nefc_out: wp.array(dtype=int), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid = wp.tid() @@ -44,35 +46,36 @@ def _zero_constraint_counts( nf_out[worldid] = 0 nl_out[worldid] = 0 nefc_out[worldid] = 0 + efc_nnz_out[worldid] = 0 @wp.func def _efc_row( - # Model: - opt_disableflags: int, - # In: - worldid: int, - timestep: float, - efcid: int, - pos_aref: float, - pos_imp: float, - invweight: float, - solref: wp.vec2, - solimp: vec5, - margin: float, - vel: float, - frictionloss: float, - type: int, - id: int, - # Out: - type_out: wp.array2d(dtype=int), - id_out: wp.array2d(dtype=int), - pos_out: wp.array2d(dtype=float), - margin_out: wp.array2d(dtype=float), - D_out: wp.array2d(dtype=float), - vel_out: wp.array2d(dtype=float), - aref_out: wp.array2d(dtype=float), - frictionloss_out: wp.array2d(dtype=float), + # Model: + opt_disableflags: int, + # In: + worldid: int, + timestep: float, + efcid: int, + pos_aref: float, + pos_imp: float, + invweight: float, + solref: wp.vec2, + solimp: vec5, + margin: float, + vel: float, + frictionloss: float, + type: int, + id: int, + # Out: + type_out: wp.array2d(dtype=int), + id_out: wp.array2d(dtype=int), + pos_out: wp.array2d(dtype=float), + margin_out: wp.array2d(dtype=float), + D_out: wp.array2d(dtype=float), + vel_out: wp.array2d(dtype=float), + aref_out: wp.array2d(dtype=float), + frictionloss_out: wp.array2d(dtype=float), ): # calculate kbi timeconst = solref[0] @@ -108,9 +111,7 @@ def _efc_row( imp = wp.where(imp_x > 1.0, dmax, imp) # set outputs - D_out[worldid, efcid] = 1.0 / wp.max( - invweight * (1.0 - imp) / imp, types.MJ_MINVAL - ) + D_out[worldid, efcid] = 1.0 / wp.max(invweight * (1.0 - imp) / imp, types.MJ_MINVAL) vel_out[worldid, efcid] = vel aref_out[worldid, efcid] = -k * imp * pos_aref - b * vel pos_out[worldid, efcid] = pos_aref + margin @@ -122,52 +123,55 @@ def _efc_row( @wp.kernel def _equality_connect( - # Model: - nv: int, - nsite: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - body_parentid: wp.array(dtype=int), - body_rootid: wp.array(dtype=int), - body_weldid: wp.array(dtype=int), - body_dofnum: wp.array(dtype=int), - body_dofadr: wp.array(dtype=int), - body_invweight0: wp.array2d(dtype=wp.vec2), - dof_bodyid: wp.array(dtype=int), - dof_parentid: wp.array(dtype=int), - site_bodyid: wp.array(dtype=int), - eq_obj1id: wp.array(dtype=int), - eq_obj2id: wp.array(dtype=int), - eq_objtype: wp.array(dtype=int), - eq_solref: wp.array2d(dtype=wp.vec2), - eq_solimp: wp.array2d(dtype=vec5), - eq_data: wp.array2d(dtype=vec11), - is_sparse: bool, - eq_connect_adr: wp.array(dtype=int), - # Data in: - qvel_in: wp.array2d(dtype=float), - eq_active_in: wp.array2d(dtype=bool), - xpos_in: wp.array2d(dtype=wp.vec3), - xmat_in: wp.array2d(dtype=wp.mat33), - site_xpos_in: wp.array2d(dtype=wp.vec3), - subtree_com_in: wp.array2d(dtype=wp.vec3), - cdof_in: wp.array2d(dtype=wp.spatial_vector), - njmax_in: int, - # Data out: - ne_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + nsite: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_weldid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + body_invweight0: wp.array2d(dtype=wp.vec2), + dof_bodyid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_objtype: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_data: wp.array2d(dtype=vec11), + is_sparse: bool, + eq_connect_adr: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + eq_active_in: wp.array2d(dtype=bool), + xpos_in: wp.array2d(dtype=wp.vec3), + xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): """Calculates constraint rows for connect equality constraints.""" worldid, eqconnectid = wp.tid() @@ -182,6 +186,10 @@ def _equality_connect( if efcid >= njmax_in - 3: return + efcid0 = efcid + 0 + efcid1 = efcid + 1 + efcid2 = efcid + 2 + data = eq_data[worldid % eq_data.shape[0], eqid] anchor1 = wp.vec3f(data[0], data[1], data[2]) anchor2 = wp.vec3f(data[3], data[4], data[5]) @@ -207,26 +215,39 @@ def _equality_connect( Jqvel = wp.vec3f(0.0, 0.0, 0.0) if is_sparse: + # TODO(team): pre-compute number of non-zeros body1 = body_weldid[body1] body2 = body_weldid[body2] da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) - efcid0 = efcid + 0 - efcid1 = efcid + 1 - efcid2 = efcid + 2 - - rowadr0 = efcid0 * nv - rowadr1 = efcid1 * nv - rowadr2 = efcid2 * nv - - efc_J_rowadr_out[worldid, efcid0] = rowadr0 - efc_J_rowadr_out[worldid, efcid1] = rowadr1 - efc_J_rowadr_out[worldid, efcid2] = rowadr2 - + # count non-zeros + pda1 = da1 + pda2 = da2 rownnz = int(0) + while pda1 >= 0 or pda2 >= 0: + da = wp.max(pda1, pda2) + if pda1 == da: + pda1 = dof_parentid[pda1] + if pda2 == da: + pda2 = dof_parentid[pda2] + rownnz += 1 + # get rowadr + rowadr = wp.atomic_add(efc_nnz_out, worldid, 3 * rownnz) + if rowadr + 3 * rownnz > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid0] = rowadr + efc_J_rowadr_out[worldid, efcid1] = rowadr + rownnz + efc_J_rowadr_out[worldid, efcid2] = rowadr + 2 * rownnz + + efc_J_rownnz_out[worldid, efcid0] = rownnz + efc_J_rownnz_out[worldid, efcid1] = rownnz + efc_J_rownnz_out[worldid, efcid2] = rownnz + + # compute J and colind + nnz = int(0) while da1 >= 0 or da2 >= 0: da = wp.max(da1, da2) if da1 == da: @@ -235,32 +256,32 @@ def _equality_connect( da2 = dof_parentid[da2] jacp1, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - pos1, - body1, - da, - worldid, + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos1, + body1, + da, + worldid, ) jacp2, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - pos2, - body2, - da, - worldid, + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos2, + body2, + da, + worldid, ) j1mj2 = jacp1 - jacp2 - sparseid0 = rowadr0 + rownnz - sparseid1 = rowadr1 + rownnz - sparseid2 = rowadr2 + rownnz + sparseid0 = rowadr + nnz + sparseid1 = rowadr + rownnz + nnz + sparseid2 = rowadr + 2 * rownnz + nnz efc_J_colind_out[worldid, 0, sparseid0] = da efc_J_colind_out[worldid, 0, sparseid1] = da @@ -272,49 +293,42 @@ def _equality_connect( Jqvel += j1mj2 * qvel_in[worldid, da] - rownnz += 1 - - efc_J_rownnz_out[worldid, efcid0] = rownnz - efc_J_rownnz_out[worldid, efcid1] = rownnz - efc_J_rownnz_out[worldid, efcid2] = rownnz + nnz += 1 else: # TODO(team): dof tree traversal for dofid in range(nv): jacp1, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - pos1, - body1, - dofid, - worldid, + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos1, + body1, + dofid, + worldid, ) jacp2, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - pos2, - body2, - dofid, - worldid, + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos2, + body2, + dofid, + worldid, ) j1mj2 = jacp1 - jacp2 - efc_J_out[worldid, efcid + 0, dofid] = j1mj2[0] - efc_J_out[worldid, efcid + 1, dofid] = j1mj2[1] - efc_J_out[worldid, efcid + 2, dofid] = j1mj2[2] + efc_J_out[worldid, efcid0, dofid] = j1mj2[0] + efc_J_out[worldid, efcid1, dofid] = j1mj2[1] + efc_J_out[worldid, efcid2, dofid] = j1mj2[2] Jqvel += j1mj2 * qvel_in[worldid, dofid] body_invweight0_id = worldid % body_invweight0.shape[0] - invweight = ( - body_invweight0[body_invweight0_id, body1][0] - + body_invweight0[body_invweight0_id, body2][0] - ) + invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] pos_imp = wp.length(pos) solref = eq_solref[worldid % eq_solref.shape[0], eqid] @@ -325,68 +339,71 @@ def _equality_connect( efcidi = efcid + i _efc_row( - opt_disableflags, - worldid, - timestep, - efcidi, - pos[i], - pos_imp, - invweight, - solref, - solimp, - 0.0, - Jqvel[i], - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + timestep, + efcidi, + pos[i], + pos_imp, + invweight, + solref, + solimp, + 0.0, + Jqvel[i], + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _equality_joint( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - qpos0: wp.array2d(dtype=float), - jnt_qposadr: wp.array(dtype=int), - jnt_dofadr: wp.array(dtype=int), - dof_invweight0: wp.array2d(dtype=float), - eq_obj1id: wp.array(dtype=int), - eq_obj2id: wp.array(dtype=int), - eq_solref: wp.array2d(dtype=wp.vec2), - eq_solimp: wp.array2d(dtype=vec5), - eq_data: wp.array2d(dtype=vec11), - is_sparse: bool, - eq_jnt_adr: wp.array(dtype=int), - # Data in: - qpos_in: wp.array2d(dtype=float), - qvel_in: wp.array2d(dtype=float), - eq_active_in: wp.array2d(dtype=bool), - njmax_in: int, - # Data out: - ne_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + qpos0: wp.array2d(dtype=float), + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_invweight0: wp.array2d(dtype=float), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_data: wp.array2d(dtype=vec11), + is_sparse: bool, + eq_jnt_adr: wp.array(dtype=int), + # Data in: + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + eq_active_in: wp.array2d(dtype=bool), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, eqjntid = wp.tid() eqid = eq_jnt_adr[eqjntid] @@ -414,7 +431,9 @@ def _equality_joint( else: rownnz = 1 efc_J_rownnz_out[worldid, efcid] = rownnz - rowadr = efcid * nv + rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) + if rowadr + rownnz > njmax_nnz_in: + return efc_J_rowadr_out[worldid, efcid] = rowadr efc_J_colind_out[worldid, 0, rowadr] = dofadr1 efc_J_out[worldid, 0, rowadr] = 1.0 @@ -430,19 +449,12 @@ def _equality_joint( 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])) - ) - deriv_2 = data[1] + dif * ( - 2.0 * data[2] + dif * (3.0 * data[3] + dif * 4.0 * data[4]) - ) + rhs = data[0] + dif * (data[1] + dif * (data[2] + dif * (data[3] + dif * data[4]))) + deriv_2 = data[1] + dif * (2.0 * data[2] + dif * (3.0 * data[3] + dif * 4.0 * data[4])) pos = qpos_in[worldid, qposadr1] - qpos0[qpos0_id, qposadr1] - rhs Jqvel = qvel_in[worldid, dofadr1] - qvel_in[worldid, dofadr2] * deriv_2 - invweight = ( - dof_invweight0[dof_invweight0_id, dofadr1] - + dof_invweight0[dof_invweight0_id, dofadr2] - ) + invweight = dof_invweight0[dof_invweight0_id, dofadr1] + dof_invweight0[dof_invweight0_id, dofadr2] if is_sparse: sparseid = rowadr + 1 @@ -458,67 +470,73 @@ def _equality_joint( # Update constraint parameters _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - invweight, - eq_solref[worldid % eq_solref.shape[0], eqid], - eq_solimp[worldid % eq_solimp.shape[0], eqid], - 0.0, - Jqvel, - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + invweight, + eq_solref[worldid % eq_solref.shape[0], eqid], + eq_solimp[worldid % eq_solimp.shape[0], eqid], + 0.0, + Jqvel, + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _equality_tendon( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - eq_obj1id: wp.array(dtype=int), - eq_obj2id: wp.array(dtype=int), - eq_solref: wp.array2d(dtype=wp.vec2), - eq_solimp: wp.array2d(dtype=vec5), - eq_data: wp.array2d(dtype=vec11), - tendon_length0: wp.array2d(dtype=float), - tendon_invweight0: wp.array2d(dtype=float), - is_sparse: bool, - eq_ten_adr: wp.array(dtype=int), - # Data in: - qvel_in: wp.array2d(dtype=float), - eq_active_in: wp.array2d(dtype=bool), - ten_J_in: wp.array3d(dtype=float), - ten_length_in: wp.array2d(dtype=float), - njmax_in: int, - # Data out: - ne_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_data: wp.array2d(dtype=vec11), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + tendon_length0: wp.array2d(dtype=float), + tendon_invweight0: wp.array2d(dtype=float), + is_sparse: bool, + eq_ten_adr: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + eq_active_in: wp.array2d(dtype=bool), + ten_J_in: wp.array2d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, eqtenid = wp.tid() eqid = eq_ten_adr[eqtenid] @@ -540,89 +558,118 @@ def _equality_tendon( solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] tendon_length0_id = worldid % tendon_length0.shape[0] tendon_invweight0_id = worldid % tendon_invweight0.shape[0] - pos1 = ( - ten_length_in[worldid, obj1id] - tendon_length0[tendon_length0_id, obj1id] - ) - jac1 = ten_J_in[worldid, obj1id] + pos1 = ten_length_in[worldid, obj1id] - tendon_length0[tendon_length0_id, obj1id] if obj2id > -1: - invweight = ( - tendon_invweight0[tendon_invweight0_id, obj1id] - + tendon_invweight0[tendon_invweight0_id, obj2id] - ) + invweight = tendon_invweight0[tendon_invweight0_id, obj1id] + tendon_invweight0[tendon_invweight0_id, obj2id] - pos2 = ( - ten_length_in[worldid, obj2id] - - tendon_length0[tendon_length0_id, obj2id] - ) - jac2 = ten_J_in[worldid, obj2id] + pos2 = ten_length_in[worldid, obj2id] - tendon_length0[tendon_length0_id, obj2id] dif = pos2 dif2 = dif * dif dif3 = dif2 * dif dif4 = dif3 * dif - pos = pos1 - ( - data[0] - + data[1] * dif - + data[2] * dif2 - + data[3] * dif3 - + data[4] * dif4 - ) - deriv = ( - data[1] - + 2.0 * data[2] * dif - + 3.0 * data[3] * dif2 - + 4.0 * data[4] * dif3 - ) + pos = pos1 - (data[0] + data[1] * dif + data[2] * dif2 + data[3] * dif3 + data[4] * dif4) + deriv = data[1] + 2.0 * data[2] * dif + 3.0 * data[3] * dif2 + 4.0 * data[4] * dif3 else: invweight = tendon_invweight0[tendon_invweight0_id, obj1id] pos = pos1 - data[0] deriv = 0.0 - Jqvel = float(0.0) + rownnz1 = ten_J_rownnz[obj1id] + rowadr1 = ten_J_rowadr[obj1id] + rownnz2 = 0 + rowadr2 = 0 + + if deriv != 0.0: + rownnz2 = ten_J_rownnz[obj2id] + rowadr2 = ten_J_rowadr[obj2id] - # TODO(team): sparse tendon jacobian if is_sparse: - rowadr = efcid * nv - efc_J_rownnz_out[worldid, efcid] = nv + # TODO(team): pre-compute rownnz + # count unique dofs + p1, p2 = int(0), int(0) + rownnz = int(0) + while p1 < rownnz1 or p2 < rownnz2: + col1 = nv + col2 = nv + if p1 < rownnz1: + col1 = ten_J_colind[rowadr1 + p1] + if p2 < rownnz2: + col2 = ten_J_colind[rowadr2 + p2] + if col1 <= col2: + p1 += 1 + if col2 <= col1: + p2 += 1 + rownnz += 1 + + rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) + if rowadr + rownnz > njmax_nnz_in: + return efc_J_rowadr_out[worldid, efcid] = rowadr + ptr1 = int(0) + ptr2 = int(0) + + Jqvel = float(0.0) + + nnz = int(0) for i in range(nv): + J1 = float(0.0) + if ptr1 < rownnz1: + sparseid1 = rowadr1 + ptr1 + if ten_J_colind[sparseid1] == i: + J1 = ten_J_in[worldid, sparseid1] + ptr1 += 1 + + J = J1 if deriv != 0.0: - J = jac1[i] + jac2[i] * -deriv - else: - J = jac1[i] + J2 = float(0.0) + if ptr2 < rownnz2: + sparseid2 = rowadr2 + ptr2 + if ten_J_colind[sparseid2] == i: + J2 = ten_J_in[worldid, sparseid2] + ptr2 += 1 + J += J2 * -deriv + if is_sparse: - efc_J_colind_out[worldid, 0, rowadr + i] = i - efc_J_out[worldid, 0, rowadr + i] = J + if J != 0.0: + sparseid = rowadr + nnz + efc_J_colind_out[worldid, 0, sparseid] = i + efc_J_out[worldid, 0, sparseid] = J + nnz += 1 else: efc_J_out[worldid, efcid, i] = J + Jqvel += J * qvel_in[worldid, i] + if is_sparse: + efc_J_rownnz_out[worldid, efcid] = nnz + _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - invweight, - solref, - solimp, - 0.0, - Jqvel, - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + invweight, + solref, + solimp, + 0.0, + Jqvel, + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @@ -630,41 +677,50 @@ def _equality_tendon( def _equality_flex(is_sparse: bool): @wp.kernel(module="unique", enable_backward=False) def kernel( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - flexedge_length0: wp.array(dtype=float), - flexedge_invweight0: wp.array(dtype=float), - flexedge_J_rownnz: wp.array(dtype=int), - flexedge_J_rowadr: wp.array(dtype=int), - flexedge_J_colind: wp.array(dtype=int), - eq_solref: wp.array2d(dtype=wp.vec2), - eq_solimp: wp.array2d(dtype=vec5), - eq_flex_adr: wp.array(dtype=int), - # Data in: - qvel_in: wp.array2d(dtype=float), - flexedge_J_in: wp.array2d(dtype=float), - flexedge_length_in: wp.array2d(dtype=float), - njmax_in: int, - # Data out: - ne_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + flex_edgeadr: wp.array(dtype=int), + flex_edgenum: wp.array(dtype=int), + flexedge_length0: wp.array(dtype=float), + flexedge_invweight0: wp.array(dtype=float), + flexedge_J_rownnz: wp.array(dtype=int), + flexedge_J_rowadr: wp.array(dtype=int), + flexedge_J_colind: wp.array(dtype=int), + eq_obj1id: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_flex_adr: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + flexedge_J_in: wp.array2d(dtype=float), + flexedge_length_in: wp.array2d(dtype=float), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, eqflexid, edgeid = wp.tid() eqid = eq_flex_adr[eqflexid] + flexid = eq_obj1id[eqid] + if edgeid < flex_edgeadr[flexid] or edgeid >= flex_edgeadr[flexid] + flex_edgenum[flexid]: + return wp.atomic_add(ne_out, worldid, 1) efcid = wp.atomic_add(nefc_out, worldid, 1) @@ -683,7 +739,9 @@ def _equality_flex(is_sparse: bool): if wp.static(is_sparse): efc_J_rownnz_out[worldid, efcid] = rownnz - efc_rowadr = efcid * nv + efc_rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) + if efc_rowadr + rownnz > njmax_nnz_in: + return efc_J_rowadr_out[worldid, efcid] = efc_rowadr for i in range(rownnz): flex_sparseid = flex_rowadr + i @@ -704,28 +762,28 @@ def _equality_flex(is_sparse: bool): Jqvel += J * qvel_in[worldid, colind] _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - flexedge_invweight0[edgeid], - solref, - solimp, - 0.0, - Jqvel, - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + flexedge_invweight0[edgeid], + solref, + solimp, + 0.0, + Jqvel, + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) return kernel @@ -733,54 +791,57 @@ def _equality_flex(is_sparse: bool): @wp.kernel def _equality_weld( - # Model: - nv: int, - nsite: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - body_parentid: wp.array(dtype=int), - body_rootid: wp.array(dtype=int), - body_weldid: wp.array(dtype=int), - body_dofnum: wp.array(dtype=int), - body_dofadr: wp.array(dtype=int), - body_invweight0: wp.array2d(dtype=wp.vec2), - dof_bodyid: wp.array(dtype=int), - dof_parentid: wp.array(dtype=int), - site_bodyid: wp.array(dtype=int), - site_quat: wp.array2d(dtype=wp.quat), - eq_obj1id: wp.array(dtype=int), - eq_obj2id: wp.array(dtype=int), - eq_objtype: wp.array(dtype=int), - eq_solref: wp.array2d(dtype=wp.vec2), - eq_solimp: wp.array2d(dtype=vec5), - eq_data: wp.array2d(dtype=vec11), - is_sparse: bool, - eq_wld_adr: wp.array(dtype=int), - # Data in: - qvel_in: wp.array2d(dtype=float), - eq_active_in: wp.array2d(dtype=bool), - xpos_in: wp.array2d(dtype=wp.vec3), - xquat_in: wp.array2d(dtype=wp.quat), - xmat_in: wp.array2d(dtype=wp.mat33), - site_xpos_in: wp.array2d(dtype=wp.vec3), - subtree_com_in: wp.array2d(dtype=wp.vec3), - cdof_in: wp.array2d(dtype=wp.spatial_vector), - njmax_in: int, - # Data out: - ne_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + nsite: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_weldid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + body_invweight0: wp.array2d(dtype=wp.vec2), + dof_bodyid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + site_quat: wp.array2d(dtype=wp.quat), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_objtype: wp.array(dtype=int), + eq_solref: wp.array2d(dtype=wp.vec2), + eq_solimp: wp.array2d(dtype=vec5), + eq_data: wp.array2d(dtype=vec11), + is_sparse: bool, + eq_wld_adr: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + eq_active_in: wp.array2d(dtype=bool), + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + xmat_in: wp.array2d(dtype=wp.mat33), + site_xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + ne_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, eqweldid = wp.tid() eqid = eq_wld_adr[eqweldid] @@ -794,6 +855,13 @@ def _equality_weld( if efcid >= njmax_in - 6: return + efcid0 = efcid + 0 + efcid1 = efcid + 1 + efcid2 = efcid + 2 + efcid3 = efcid + 3 + efcid4 = efcid + 4 + efcid5 = efcid + 5 + is_site = eq_objtype[eqid] == types.ObjType.SITE and nsite > 0 obj1id = eq_obj1id[eqid] @@ -812,12 +880,8 @@ def _equality_weld( pos2 = site_xpos_in[worldid, obj2id] site_quat_id = worldid % site_quat.shape[0] - quat = math.mul_quat( - xquat_in[worldid, body1], site_quat[site_quat_id, obj1id] - ) - quat1 = math.quat_inv( - math.mul_quat(xquat_in[worldid, body2], site_quat[site_quat_id, obj2id]) - ) + quat = math.mul_quat(xquat_in[worldid, body1], site_quat[site_quat_id, obj1id]) + quat1 = math.quat_inv(math.mul_quat(xquat_in[worldid, body2], site_quat[site_quat_id, obj2id])) else: body1 = obj1id @@ -833,35 +897,45 @@ def _equality_weld( Jqvelr = wp.vec3f(0.0, 0.0, 0.0) if is_sparse: + # TODO(team): pre-compute number of non-zeros body1 = body_weldid[body1] body2 = body_weldid[body2] da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) - efcid0 = efcid + 0 - efcid1 = efcid + 1 - efcid2 = efcid + 2 - efcid3 = efcid + 3 - efcid4 = efcid + 4 - efcid5 = efcid + 5 - - rowadr0 = efcid0 * nv - rowadr1 = efcid1 * nv - rowadr2 = efcid2 * nv - rowadr3 = efcid3 * nv - rowadr4 = efcid4 * nv - rowadr5 = efcid5 * nv - - efc_J_rowadr_out[worldid, efcid0] = rowadr0 - efc_J_rowadr_out[worldid, efcid1] = rowadr1 - efc_J_rowadr_out[worldid, efcid2] = rowadr2 - efc_J_rowadr_out[worldid, efcid3] = rowadr3 - efc_J_rowadr_out[worldid, efcid4] = rowadr4 - efc_J_rowadr_out[worldid, efcid5] = rowadr5 - + # count non-zeros + pda1 = da1 + pda2 = da2 rownnz = int(0) + while pda1 >= 0 or pda2 >= 0: + da = wp.max(pda1, pda2) + if pda1 == da: + pda1 = dof_parentid[da] + if pda2 == da: + pda2 = dof_parentid[da] + rownnz += 1 + # get rowadr + rowadr = wp.atomic_add(efc_nnz_out, worldid, 6 * rownnz) + if rowadr + 6 * rownnz > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid0] = rowadr + efc_J_rowadr_out[worldid, efcid1] = rowadr + rownnz + efc_J_rowadr_out[worldid, efcid2] = rowadr + 2 * rownnz + efc_J_rowadr_out[worldid, efcid3] = rowadr + 3 * rownnz + efc_J_rowadr_out[worldid, efcid4] = rowadr + 4 * rownnz + efc_J_rowadr_out[worldid, efcid5] = rowadr + 5 * rownnz + + efc_J_rownnz_out[worldid, efcid0] = rownnz + efc_J_rownnz_out[worldid, efcid1] = rownnz + efc_J_rownnz_out[worldid, efcid2] = rownnz + efc_J_rownnz_out[worldid, efcid3] = rownnz + efc_J_rownnz_out[worldid, efcid4] = rownnz + efc_J_rownnz_out[worldid, efcid5] = rownnz + + # compute J and colind + nnz = int(0) while da1 >= 0 or da2 >= 0: da = wp.max(da1, da2) if da1 == da: @@ -870,26 +944,26 @@ def _equality_weld( da2 = dof_parentid[da] jacp1, jacr1 = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - pos1, - body1, - da, - worldid, + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos1, + body1, + da, + worldid, ) jacp2, jacr2 = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - pos2, - body2, - da, - worldid, + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos2, + body2, + da, + worldid, ) jacdifp = jacp1 - jacp2 @@ -898,12 +972,12 @@ def _equality_weld( jacdifrq = math.mul_quat(math.quat_mul_axis(quat1, jacdifr), quat) jacdifr = 0.5 * wp.vec3(jacdifrq[1], jacdifrq[2], jacdifrq[3]) - sparseid0 = rowadr0 + rownnz - sparseid1 = rowadr1 + rownnz - sparseid2 = rowadr2 + rownnz - sparseid3 = rowadr3 + rownnz - sparseid4 = rowadr4 + rownnz - sparseid5 = rowadr5 + rownnz + sparseid0 = rowadr + nnz + sparseid1 = rowadr + rownnz + nnz + sparseid2 = rowadr + 2 * rownnz + nnz + sparseid3 = rowadr + 3 * rownnz + nnz + sparseid4 = rowadr + 4 * rownnz + nnz + sparseid5 = rowadr + 5 * rownnz + nnz efc_J_colind_out[worldid, 0, sparseid0] = da efc_J_colind_out[worldid, 0, sparseid1] = da @@ -922,50 +996,45 @@ def _equality_weld( Jqvelp += jacdifp * qvel_in[worldid, da] Jqvelr += jacdifr * qvel_in[worldid, da] - rownnz += 1 - - efc_J_rownnz_out[worldid, efcid0] = rownnz - efc_J_rownnz_out[worldid, efcid1] = rownnz - efc_J_rownnz_out[worldid, efcid2] = rownnz - efc_J_rownnz_out[worldid, efcid3] = rownnz - efc_J_rownnz_out[worldid, efcid4] = rownnz - efc_J_rownnz_out[worldid, efcid5] = rownnz + nnz += 1 else: for dofid in range(nv): jacp1, jacr1 = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - pos1, - body1, - dofid, - worldid, + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos1, + body1, + dofid, + worldid, ) jacp2, jacr2 = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - pos2, - body2, - dofid, - worldid, + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + pos2, + body2, + dofid, + worldid, ) jacdifp = jacp1 - jacp2 - for i in range(3): - efc_J_out[worldid, efcid + i, dofid] = jacdifp[i] + efc_J_out[worldid, efcid0, dofid] = jacdifp[0] + efc_J_out[worldid, efcid1, dofid] = jacdifp[1] + efc_J_out[worldid, efcid2, dofid] = jacdifp[2] jacdifr = (jacr1 - jacr2) * torquescale jacdifrq = math.mul_quat(math.quat_mul_axis(quat1, jacdifr), quat) jacdifr = 0.5 * wp.vec3(jacdifrq[1], jacdifrq[2], jacdifrq[3]) - for i in range(3): - efc_J_out[worldid, efcid + 3 + i, dofid] = jacdifr[i] + efc_J_out[worldid, efcid3, dofid] = jacdifr[0] + efc_J_out[worldid, efcid4, dofid] = jacdifr[1] + efc_J_out[worldid, efcid5, dofid] = jacdifr[2] Jqvelp += jacdifp * qvel_in[worldid, dofid] Jqvelr += jacdifr * qvel_in[worldid, dofid] @@ -977,10 +1046,7 @@ def _equality_weld( crot = wp.vec3(crotq[1], crotq[2], crotq[3]) * torquescale body_invweight0_id = worldid % body_invweight0.shape[0] - invweight_t = ( - body_invweight0[body_invweight0_id, body1][0] - + body_invweight0[body_invweight0_id, body2][0] - ) + invweight_t = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] pos_imp = wp.sqrt(wp.length_sq(cpos) + wp.length_sq(crot)) @@ -991,91 +1057,91 @@ def _equality_weld( for i in range(3): _efc_row( - opt_disableflags, - worldid, - timestep, - efcid + i, - cpos[i], - pos_imp, - invweight_t, - solref, - solimp, - 0.0, - Jqvelp[i], - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + timestep, + efcid + i, + cpos[i], + pos_imp, + invweight_t, + solref, + solimp, + 0.0, + Jqvelp[i], + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) - invweight_r = ( - body_invweight0[body_invweight0_id, body1][1] - + body_invweight0[body_invweight0_id, body2][1] - ) + invweight_r = body_invweight0[body_invweight0_id, body1][1] + body_invweight0[body_invweight0_id, body2][1] for i in range(3): _efc_row( - opt_disableflags, - worldid, - timestep, - efcid + 3 + i, - crot[i], - pos_imp, - invweight_r, - solref, - solimp, - 0.0, - Jqvelr[i], - 0.0, - ConstraintType.EQUALITY, - eqid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + timestep, + efcid + 3 + i, + crot[i], + pos_imp, + invweight_r, + solref, + solimp, + 0.0, + Jqvelr[i], + 0.0, + ConstraintType.EQUALITY, + eqid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _friction_dof( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - dof_solref: wp.array2d(dtype=wp.vec2), - dof_solimp: wp.array2d(dtype=vec5), - dof_frictionloss: wp.array2d(dtype=float), - dof_invweight0: wp.array2d(dtype=float), - is_sparse: bool, - # Data in: - qvel_in: wp.array2d(dtype=float), - njmax_in: int, - # Data out: - nf_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + dof_solref: wp.array2d(dtype=wp.vec2), + dof_solimp: wp.array2d(dtype=vec5), + dof_frictionloss: wp.array2d(dtype=float), + dof_invweight0: wp.array2d(dtype=float), + is_sparse: bool, + # Data in: + qvel_in: wp.array2d(dtype=float), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nf_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, dofid = wp.tid() @@ -1092,7 +1158,9 @@ def _friction_dof( if is_sparse: efc_J_rownnz_out[worldid, efcid] = 1 - rowadr = efcid * nv + rowadr = wp.atomic_add(efc_nnz_out, worldid, 1) + if rowadr + 1 > njmax_nnz_in: + return efc_J_rowadr_out[worldid, efcid] = rowadr efc_J_colind_out[worldid, 0, rowadr] = dofid efc_J_out[worldid, 0, rowadr] = 1.0 @@ -1107,61 +1175,67 @@ def _friction_dof( dof_solref_id = worldid % dof_solref.shape[0] dof_solimp_id = worldid % dof_solimp.shape[0] _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - 0.0, - 0.0, - dof_invweight0[dof_invweight0_id, dofid], - dof_solref[dof_solref_id, dofid], - dof_solimp[dof_solimp_id, dofid], - 0.0, - Jqvel, - dof_frictionloss[dof_frictionloss_id, dofid], - ConstraintType.FRICTION_DOF, - dofid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + 0.0, + 0.0, + dof_invweight0[dof_invweight0_id, dofid], + dof_solref[dof_solref_id, dofid], + dof_solimp[dof_solimp_id, dofid], + 0.0, + Jqvel, + dof_frictionloss[dof_frictionloss_id, dofid], + ConstraintType.FRICTION_DOF, + dofid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _friction_tendon( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - tendon_solref_fri: wp.array2d(dtype=wp.vec2), - tendon_solimp_fri: wp.array2d(dtype=vec5), - tendon_frictionloss: wp.array2d(dtype=float), - tendon_invweight0: wp.array2d(dtype=float), - is_sparse: bool, - # Data in: - qvel_in: wp.array2d(dtype=float), - ten_J_in: wp.array3d(dtype=float), - njmax_in: int, - # Data out: - nf_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + tendon_solref_fri: wp.array2d(dtype=wp.vec2), + tendon_solimp_fri: wp.array2d(dtype=vec5), + tendon_frictionloss: wp.array2d(dtype=float), + tendon_invweight0: wp.array2d(dtype=float), + is_sparse: bool, + # Data in: + qvel_in: wp.array2d(dtype=float), + ten_J_in: wp.array2d(dtype=float), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nf_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, tenid = wp.tid() @@ -1179,86 +1253,103 @@ def _friction_tendon( Jqvel = float(0.0) - # TODO(team): sparse tendon jacobian + rownnz_tenJ = ten_J_rownnz[tenid] + rowadr_tenJ = ten_J_rowadr[tenid] if is_sparse: - rowadr = efcid * nv - efc_J_rownnz_out[worldid, efcid] = nv - efc_J_rowadr_out[worldid, efcid] = rowadr + efc_J_rownnz_out[worldid, efcid] = rownnz_tenJ + rowadr_efc = wp.atomic_add(efc_nnz_out, worldid, rownnz_tenJ) + if rowadr_efc + rownnz_tenJ > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr_efc - for i in range(nv): - # TODO(team): sparse ten_J - J = ten_J_in[worldid, tenid, i] - if is_sparse: - efc_J_colind_out[worldid, 0, rowadr + i] = i - efc_J_out[worldid, 0, rowadr + i] = J - else: - efc_J_out[worldid, efcid, i] = J - - Jqvel += J * qvel_in[worldid, i] + for i in range(rownnz_tenJ): + sparseid_ten = rowadr_tenJ + i + sparseid_efc = rowadr_efc + i + colind = ten_J_colind[sparseid_ten] + J = ten_J_in[worldid, sparseid_ten] + efc_J_colind_out[worldid, 0, sparseid_efc] = colind + efc_J_out[worldid, 0, sparseid_efc] = J + Jqvel += J * qvel_in[worldid, colind] + else: + nnz = int(0) + colind = ten_J_colind[rowadr_tenJ] + for i in range(nv): + if nnz < rownnz_tenJ and i == colind: + J = ten_J_in[worldid, rowadr_tenJ + nnz] + efc_J_out[worldid, efcid, i] = J + Jqvel += J * qvel_in[worldid, i] + nnz += 1 + if nnz < rownnz_tenJ: + colind = ten_J_colind[rowadr_tenJ + nnz] + else: + efc_J_out[worldid, efcid, i] = 0.0 tendon_invweight0_id = worldid % tendon_invweight0.shape[0] tendon_solref_fri_id = worldid % tendon_solref_fri.shape[0] tendon_solimp_fri_id = worldid % tendon_solimp_fri.shape[0] _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - 0.0, - 0.0, - tendon_invweight0[tendon_invweight0_id, tenid], - tendon_solref_fri[tendon_solref_fri_id, tenid], - tendon_solimp_fri[tendon_solimp_fri_id, tenid], - 0.0, - Jqvel, - frictionloss, - ConstraintType.FRICTION_TENDON, - tenid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + 0.0, + 0.0, + tendon_invweight0[tendon_invweight0_id, tenid], + tendon_solref_fri[tendon_solref_fri_id, tenid], + tendon_solimp_fri[tendon_solimp_fri_id, tenid], + 0.0, + Jqvel, + frictionloss, + ConstraintType.FRICTION_TENDON, + tenid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _limit_slide_hinge( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - jnt_qposadr: wp.array(dtype=int), - jnt_dofadr: wp.array(dtype=int), - jnt_solref: wp.array2d(dtype=wp.vec2), - jnt_solimp: wp.array2d(dtype=vec5), - jnt_range: wp.array2d(dtype=wp.vec2), - jnt_margin: wp.array2d(dtype=float), - dof_invweight0: wp.array2d(dtype=float), - is_sparse: bool, - jnt_limited_slide_hinge_adr: wp.array(dtype=int), - # Data in: - qpos_in: wp.array2d(dtype=float), - qvel_in: wp.array2d(dtype=float), - njmax_in: int, - # Data out: - nl_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + jnt_solref: wp.array2d(dtype=wp.vec2), + jnt_solimp: wp.array2d(dtype=vec5), + jnt_range: wp.array2d(dtype=wp.vec2), + jnt_margin: wp.array2d(dtype=float), + dof_invweight0: wp.array2d(dtype=float), + is_sparse: bool, + jnt_limited_slide_hinge_adr: wp.array(dtype=int), + # Data in: + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nl_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, jntlimitedid = wp.tid() jntid = jnt_limited_slide_hinge_adr[jntlimitedid] @@ -1285,7 +1376,9 @@ def _limit_slide_hinge( if is_sparse: efc_J_rownnz_out[worldid, efcid] = 1 - rowadr = efcid * nv + rowadr = wp.atomic_add(efc_nnz_out, worldid, 1) + if rowadr + 1 > njmax_nnz_in: + return efc_J_rowadr_out[worldid, efcid] = rowadr efc_J_colind_out[worldid, 0, rowadr] = dofadr efc_J_out[worldid, 0, rowadr] = J @@ -1300,65 +1393,68 @@ def _limit_slide_hinge( jnt_solref_id = worldid % jnt_solref.shape[0] jnt_solimp_id = worldid % jnt_solimp.shape[0] _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - dof_invweight0[dof_invweight0_id, dofadr], - jnt_solref[jnt_solref_id, jntid], - jnt_solimp[jnt_solimp_id, jntid], - jntmargin, - Jqvel, - 0.0, - ConstraintType.LIMIT_JOINT, - jntid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + dof_invweight0[dof_invweight0_id, dofadr], + jnt_solref[jnt_solref_id, jntid], + jnt_solimp[jnt_solimp_id, jntid], + jntmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_JOINT, + jntid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _limit_ball( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - jnt_qposadr: wp.array(dtype=int), - jnt_dofadr: wp.array(dtype=int), - jnt_solref: wp.array2d(dtype=wp.vec2), - jnt_solimp: wp.array2d(dtype=vec5), - jnt_range: wp.array2d(dtype=wp.vec2), - jnt_margin: wp.array2d(dtype=float), - dof_invweight0: wp.array2d(dtype=float), - is_sparse: bool, - jnt_limited_ball_adr: wp.array(dtype=int), - # Data in: - qpos_in: wp.array2d(dtype=float), - qvel_in: wp.array2d(dtype=float), - njmax_in: int, - # Data out: - nl_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + jnt_solref: wp.array2d(dtype=wp.vec2), + jnt_solimp: wp.array2d(dtype=vec5), + jnt_range: wp.array2d(dtype=wp.vec2), + jnt_margin: wp.array2d(dtype=float), + dof_invweight0: wp.array2d(dtype=float), + is_sparse: bool, + jnt_limited_ball_adr: wp.array(dtype=int), + # Data in: + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nl_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, jntlimitedid = wp.tid() jntid = jnt_limited_ball_adr[jntlimitedid] @@ -1391,7 +1487,9 @@ def _limit_ball( if is_sparse: efc_J_rownnz_out[worldid, efcid] = 3 - rowadr = efcid * nv + rowadr = wp.atomic_add(efc_nnz_out, worldid, 3) + if rowadr + 3 > njmax_nnz_in: + return efc_J_rowadr_out[worldid, efcid] = rowadr sparseid0 = rowadr + 0 @@ -1420,69 +1518,70 @@ def _limit_ball( jnt_solref_id = worldid % jnt_solref.shape[0] jnt_solimp_id = worldid % jnt_solimp.shape[0] _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - dof_invweight0[dof_invweight0_id, dofadr], - jnt_solref[jnt_solref_id, jntid], - jnt_solimp[jnt_solimp_id, jntid], - jntmargin, - Jqvel, - 0.0, - ConstraintType.LIMIT_JOINT, - jntid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + dof_invweight0[dof_invweight0_id, dofadr], + jnt_solref[jnt_solref_id, jntid], + jnt_solimp[jnt_solimp_id, jntid], + jntmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_JOINT, + jntid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _limit_tendon( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - jnt_dofadr: wp.array(dtype=int), - tendon_adr: wp.array(dtype=int), - tendon_num: wp.array(dtype=int), - tendon_solref_lim: wp.array2d(dtype=wp.vec2), - tendon_solimp_lim: wp.array2d(dtype=vec5), - tendon_range: wp.array2d(dtype=wp.vec2), - tendon_margin: wp.array2d(dtype=float), - tendon_invweight0: wp.array2d(dtype=float), - wrap_type: wp.array(dtype=int), - wrap_objid: wp.array(dtype=int), - is_sparse: bool, - tendon_limited_adr: wp.array(dtype=int), - # Data in: - qvel_in: wp.array2d(dtype=float), - ten_J_in: wp.array3d(dtype=float), - ten_length_in: wp.array2d(dtype=float), - njmax_in: int, - # Data out: - nl_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + tendon_solref_lim: wp.array2d(dtype=wp.vec2), + tendon_solimp_lim: wp.array2d(dtype=vec5), + tendon_range: wp.array2d(dtype=wp.vec2), + tendon_margin: wp.array2d(dtype=float), + tendon_invweight0: wp.array2d(dtype=float), + is_sparse: bool, + tendon_limited_adr: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + ten_J_in: wp.array2d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + njmax_in: int, + njmax_nnz_in: int, + # Data out: + nl_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): worldid, tenlimitedid = wp.tid() tenid = tendon_limited_adr[tenlimitedid] @@ -1506,122 +1605,123 @@ def _limit_tendon( Jqvel = float(0.0) scl = float(dist_min < dist_max) * 2.0 - 1.0 - # TODO(team): sparse tendon jacobian + rownnz_tenJ = ten_J_rownnz[tenid] + rowadr_tenJ = ten_J_rowadr[tenid] if is_sparse: - rowadr = efcid * nv - efc_J_rownnz_out[worldid, efcid] = nv - efc_J_rowadr_out[worldid, efcid] = rowadr - for i in range(nv): - efc_J_colind_out[worldid, 0, rowadr + i] = i - efc_J_out[worldid, 0, rowadr + i] = 0.0 + efc_J_rownnz_out[worldid, efcid] = rownnz_tenJ + rowadr_efc = wp.atomic_add(efc_nnz_out, worldid, rownnz_tenJ) + if rowadr_efc + rownnz_tenJ > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr_efc - adr = tendon_adr[tenid] - if wrap_type[adr] == types.WrapType.JOINT: - if not is_sparse: - for i in range(nv): - efc_J_out[worldid, efcid, i] = 0.0 - - ten_num = tendon_num[tenid] - for i in range(ten_num): - dofadr = jnt_dofadr[wrap_objid[adr + i]] - J = scl * ten_J_in[worldid, tenid, dofadr] - - if is_sparse: - efc_J_out[worldid, 0, rowadr + dofadr] = J - else: - efc_J_out[worldid, efcid, dofadr] = J - - Jqvel += J * qvel_in[worldid, dofadr] + for i in range(rownnz_tenJ): + sparseid_ten = rowadr_tenJ + i + sparseid_efc = rowadr_efc + i + colind = ten_J_colind[sparseid_ten] + J = scl * ten_J_in[worldid, sparseid_ten] + efc_J_colind_out[worldid, 0, sparseid_efc] = colind + efc_J_out[worldid, 0, sparseid_efc] = J + Jqvel += J * qvel_in[worldid, colind] else: + nnz = int(0) + colind = ten_J_colind[rowadr_tenJ] for i in range(nv): - J = scl * ten_J_in[worldid, tenid, i] - - if is_sparse: - efc_J_out[worldid, 0, rowadr + i] = J - else: + if nnz < rownnz_tenJ and i == colind: + J = scl * ten_J_in[worldid, rowadr_tenJ + nnz] efc_J_out[worldid, efcid, i] = J - - Jqvel += J * qvel_in[worldid, i] + Jqvel += J * qvel_in[worldid, i] + nnz += 1 + if nnz < rownnz_tenJ: + colind = ten_J_colind[rowadr_tenJ + nnz] + else: + efc_J_out[worldid, efcid, i] = 0.0 tendon_invweight0_id = worldid % tendon_invweight0.shape[0] tendon_solref_lim_id = worldid % tendon_solref_lim.shape[0] tendon_solimp_lim_id = worldid % tendon_solimp_lim.shape[0] _efc_row( - opt_disableflags, - worldid, - opt_timestep[worldid % opt_timestep.shape[0]], - efcid, - pos, - pos, - tendon_invweight0[tendon_invweight0_id, tenid], - tendon_solref_lim[tendon_solref_lim_id, tenid], - tendon_solimp_lim[tendon_solimp_lim_id, tenid], - tenmargin, - Jqvel, - 0.0, - ConstraintType.LIMIT_TENDON, - tenid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + opt_timestep[worldid % opt_timestep.shape[0]], + efcid, + pos, + pos, + tendon_invweight0[tendon_invweight0_id, tenid], + tendon_solref_lim[tendon_solref_lim_id, tenid], + tendon_solimp_lim[tendon_solimp_lim_id, tenid], + tenmargin, + Jqvel, + 0.0, + ConstraintType.LIMIT_TENDON, + tenid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _contact_pyramidal( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - opt_impratio_invsqrt: wp.array(dtype=float), - body_parentid: wp.array(dtype=int), - body_rootid: wp.array(dtype=int), - body_weldid: wp.array(dtype=int), - body_dofnum: wp.array(dtype=int), - body_dofadr: wp.array(dtype=int), - body_invweight0: wp.array2d(dtype=wp.vec2), - dof_bodyid: wp.array(dtype=int), - dof_parentid: wp.array(dtype=int), - geom_bodyid: wp.array(dtype=int), - is_sparse: bool, - # Data in: - qvel_in: wp.array2d(dtype=float), - subtree_com_in: wp.array2d(dtype=wp.vec3), - cdof_in: wp.array2d(dtype=wp.spatial_vector), - njmax_in: int, - nacon_in: wp.array(dtype=int), - # In: - dist_in: wp.array(dtype=float), - condim_in: wp.array(dtype=int), - includemargin_in: wp.array(dtype=float), - worldid_in: wp.array(dtype=int), - geom_in: wp.array(dtype=wp.vec2i), - pos_in: wp.array(dtype=wp.vec3), - frame_in: wp.array(dtype=wp.mat33), - friction_in: wp.array(dtype=vec5), - solref_in: wp.array(dtype=wp.vec2), - solimp_in: wp.array(dtype=vec5), - type_in: wp.array(dtype=int), - # Data out: - nefc_out: wp.array(dtype=int), - contact_efc_address_out: wp.array2d(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + opt_impratio_invsqrt: wp.array(dtype=float), + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_weldid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + body_invweight0: wp.array2d(dtype=wp.vec2), + dof_bodyid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + is_sparse: bool, + # Data in: + qvel_in: wp.array2d(dtype=float), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, + njmax_nnz_in: int, + nacon_in: wp.array(dtype=int), + # In: + dist_in: wp.array(dtype=float), + condim_in: wp.array(dtype=int), + includemargin_in: wp.array(dtype=float), + worldid_in: wp.array(dtype=int), + geom_in: wp.array(dtype=wp.vec2i), + flex_in: wp.array(dtype=wp.vec2i), + vert_in: wp.array(dtype=wp.vec2i), + pos_in: wp.array(dtype=wp.vec3), + frame_in: wp.array(dtype=wp.mat33), + friction_in: wp.array(dtype=vec5), + solref_in: wp.array(dtype=wp.vec2), + solimp_in: wp.array(dtype=vec5), + type_in: wp.array(dtype=int), + # Data out: + nefc_out: wp.array(dtype=int), + contact_efc_address_out: wp.array2d(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): conid, dimid = wp.tid() @@ -1655,8 +1755,20 @@ def _contact_pyramidal( contact_efc_address_out[conid, dimid] = efcid geom = geom_in[conid] - body1 = geom_bodyid[geom[0]] - body2 = geom_bodyid[geom[1]] + + if geom[0] >= 0: + body1 = geom_bodyid[geom[0]] + else: + flex = flex_in[conid] + vert = vert_in[conid] + body1 = flex_vertbodyid[flex_vertadr[flex[0]] + vert[0]] + + if geom[1] >= 0: + body2 = geom_bodyid[geom[1]] + else: + flex = flex_in[conid] + vert = vert_in[conid] + body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] con_pos = pos_in[conid] frame = frame_in[conid] @@ -1674,29 +1786,48 @@ def _contact_pyramidal( invweight = invweight + fri0 * fri0 * invweight invweight = invweight * 2.0 * fri0 * fri0 * impratio_invsqrt * impratio_invsqrt - if is_sparse: - rowadr = efcid * nv - efc_J_rowadr_out[worldid, efcid] = rowadr - Jqvel = float(0.0) # skip fixed bodies body1 = body_weldid[body1] body2 = body_weldid[body2] - da1 = body_dofadr[body1] + body_dofnum[body1] - 1 - da2 = body_dofadr[body2] + body_dofnum[body2] - 1 + da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) + da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) + + if is_sparse: + pda1 = da1 + pda2 = da2 + rownnz = int(0) + while pda1 >= 0 or pda2 >= 0: + da = wp.max(pda1, pda2) + # skip common dofs + if pda1 == da and pda2 == da: + break + if pda1 == da: + pda1 = dof_parentid[pda1] + if pda2 == da: + pda2 = dof_parentid[pda2] + rownnz += 1 + + # get rowadr + rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) + if rowadr + rownnz > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr + efc_J_rownnz_out[worldid, efcid] = rownnz + da = wp.max(da1, da2) if is_sparse: - rownnz = int(0) + nnz = int(0) dofid = int(da) else: dofid = int(nv - 1) while True: if is_sparse: - if da1 < 0 and da2 < 0: + if nnz >= rownnz: break else: if dofid < 0: @@ -1749,13 +1880,15 @@ def _contact_pyramidal( J -= Ji * frii if is_sparse: - sparseid = rowadr + rownnz + sparseid = rowadr + nnz efc_J_colind_out[worldid, 0, sparseid] = dofid efc_J_out[worldid, 0, sparseid] = J - rownnz += 1 + nnz += 1 else: efc_J_out[worldid, efcid, dofid] = J Jqvel += J * qvel_in[worldid, dofid] + if is_sparse and nnz >= rownnz: + break # Advance tree pointers and recompute da for next iteration if da1 == da: @@ -1772,91 +1905,95 @@ def _contact_pyramidal( efc_J_out[worldid, efcid, dofid] = 0.0 dofid -= 1 - if is_sparse: - efc_J_rownnz_out[worldid, efcid] = rownnz - if condim == 1: efc_type = ConstraintType.CONTACT_FRICTIONLESS else: efc_type = ConstraintType.CONTACT_PYRAMIDAL _efc_row( - opt_disableflags, - worldid, - timestep, - efcid, - pos, - pos, - invweight, - solref_in[conid], - solimp_in[conid], - includemargin, - Jqvel, - 0.0, - efc_type, - conid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + timestep, + efcid, + pos, + pos, + invweight, + solref_in[conid], + solimp_in[conid], + includemargin, + Jqvel, + 0.0, + efc_type, + conid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @wp.kernel def _contact_elliptic( - # Model: - nv: int, - opt_timestep: wp.array(dtype=float), - opt_disableflags: int, - opt_impratio_invsqrt: wp.array(dtype=float), - body_parentid: wp.array(dtype=int), - body_rootid: wp.array(dtype=int), - body_weldid: wp.array(dtype=int), - body_dofnum: wp.array(dtype=int), - body_dofadr: wp.array(dtype=int), - body_invweight0: wp.array2d(dtype=wp.vec2), - dof_bodyid: wp.array(dtype=int), - dof_parentid: wp.array(dtype=int), - geom_bodyid: wp.array(dtype=int), - is_sparse: bool, - # Data in: - qvel_in: wp.array2d(dtype=float), - subtree_com_in: wp.array2d(dtype=wp.vec3), - cdof_in: wp.array2d(dtype=wp.spatial_vector), - njmax_in: int, - nacon_in: wp.array(dtype=int), - # In: - dist_in: wp.array(dtype=float), - condim_in: wp.array(dtype=int), - includemargin_in: wp.array(dtype=float), - worldid_in: wp.array(dtype=int), - geom_in: wp.array(dtype=wp.vec2i), - pos_in: wp.array(dtype=wp.vec3), - frame_in: wp.array(dtype=wp.mat33), - friction_in: wp.array(dtype=vec5), - solref_in: wp.array(dtype=wp.vec2), - solreffriction_in: wp.array(dtype=wp.vec2), - solimp_in: wp.array(dtype=vec5), - type_in: wp.array(dtype=int), - # Data out: - nefc_out: wp.array(dtype=int), - contact_efc_address_out: wp.array2d(dtype=int), - efc_type_out: wp.array2d(dtype=int), - efc_id_out: wp.array2d(dtype=int), - efc_J_rownnz_out: wp.array2d(dtype=int), - efc_J_rowadr_out: wp.array2d(dtype=int), - efc_J_colind_out: wp.array3d(dtype=int), - efc_J_out: wp.array3d(dtype=float), - efc_pos_out: wp.array2d(dtype=float), - efc_margin_out: wp.array2d(dtype=float), - efc_D_out: wp.array2d(dtype=float), - efc_vel_out: wp.array2d(dtype=float), - efc_aref_out: wp.array2d(dtype=float), - efc_frictionloss_out: wp.array2d(dtype=float), + # Model: + nv: int, + opt_timestep: wp.array(dtype=float), + opt_disableflags: int, + opt_impratio_invsqrt: wp.array(dtype=float), + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_weldid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + body_invweight0: wp.array2d(dtype=wp.vec2), + dof_bodyid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + is_sparse: bool, + # Data in: + qvel_in: wp.array2d(dtype=float), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, + njmax_nnz_in: int, + nacon_in: wp.array(dtype=int), + # In: + dist_in: wp.array(dtype=float), + condim_in: wp.array(dtype=int), + includemargin_in: wp.array(dtype=float), + worldid_in: wp.array(dtype=int), + geom_in: wp.array(dtype=wp.vec2i), + flex_in: wp.array(dtype=wp.vec2i), + vert_in: wp.array(dtype=wp.vec2i), + pos_in: wp.array(dtype=wp.vec3), + frame_in: wp.array(dtype=wp.mat33), + friction_in: wp.array(dtype=vec5), + solref_in: wp.array(dtype=wp.vec2), + solreffriction_in: wp.array(dtype=wp.vec2), + solimp_in: wp.array(dtype=vec5), + type_in: wp.array(dtype=int), + # Data out: + nefc_out: wp.array(dtype=int), + contact_efc_address_out: wp.array2d(dtype=int), + efc_type_out: wp.array2d(dtype=int), + efc_id_out: wp.array2d(dtype=int), + efc_J_rownnz_out: wp.array2d(dtype=int), + efc_J_rowadr_out: wp.array2d(dtype=int), + efc_J_colind_out: wp.array3d(dtype=int), + efc_J_out: wp.array3d(dtype=float), + efc_pos_out: wp.array2d(dtype=float), + efc_margin_out: wp.array2d(dtype=float), + efc_D_out: wp.array2d(dtype=float), + efc_vel_out: wp.array2d(dtype=float), + efc_aref_out: wp.array2d(dtype=float), + efc_frictionloss_out: wp.array2d(dtype=float), + # Out: + efc_nnz_out: wp.array(dtype=int), ): conid, dimid = wp.tid() @@ -1888,35 +2025,67 @@ def _contact_elliptic( contact_efc_address_out[conid, dimid] = efcid geom = geom_in[conid] - body1 = geom_bodyid[geom[0]] - body2 = geom_bodyid[geom[1]] + + if geom[0] >= 0: + body1 = geom_bodyid[geom[0]] + else: + flex = flex_in[conid] + vert = vert_in[conid] + body1 = flex_vertbodyid[flex_vertadr[flex[0]] + vert[0]] + + if geom[1] >= 0: + body2 = geom_bodyid[geom[1]] + else: + flex = flex_in[conid] + vert = vert_in[conid] + body2 = flex_vertbodyid[flex_vertadr[flex[1]] + vert[1]] con_pos = pos_in[conid] frame = frame_in[conid] - if is_sparse: - rowadr = efcid * nv - efc_J_rowadr_out[worldid, efcid] = rowadr - Jqvel = float(0.0) # skip fixed bodies body1 = body_weldid[body1] body2 = body_weldid[body2] - da1 = body_dofadr[body1] + body_dofnum[body1] - 1 - da2 = body_dofadr[body2] + body_dofnum[body2] - 1 + da1 = int(body_dofadr[body1] + body_dofnum[body1] - 1) + da2 = int(body_dofadr[body2] + body_dofnum[body2] - 1) + + if is_sparse: + # count non-zeros + pda1 = da1 + pda2 = da2 + rownnz = int(0) + while pda1 >= 0 or pda2 >= 0: + da = wp.max(pda1, pda2) + # skip common dofs + if pda1 == da and pda2 == da: + break + if pda1 == da: + pda1 = dof_parentid[pda1] + if pda2 == da: + pda2 = dof_parentid[pda2] + rownnz += 1 + + # get rowadr + rowadr = wp.atomic_add(efc_nnz_out, worldid, rownnz) + if rowadr + rownnz > njmax_nnz_in: + return + efc_J_rowadr_out[worldid, efcid] = rowadr + efc_J_rownnz_out[worldid, efcid] = rownnz + da = wp.max(da1, da2) if is_sparse: - rownnz = int(0) + nnz = int(0) dofid = int(da) else: dofid = int(nv - 1) while True: if is_sparse: - if da1 < 0 and da2 < 0: + if nnz >= rownnz: break else: if dofid < 0: @@ -1957,13 +2126,15 @@ def _contact_elliptic( J += frame[dimid - 3, xyz] * jac_dif if is_sparse: - sparseid = rowadr + rownnz + sparseid = rowadr + nnz efc_J_colind_out[worldid, 0, sparseid] = dofid efc_J_out[worldid, 0, sparseid] = J - rownnz += 1 + nnz += 1 else: efc_J_out[worldid, efcid, dofid] = J Jqvel += J * qvel_in[worldid, dofid] + if is_sparse and nnz >= rownnz: + break # Advance tree pointers and recompute da for next iteration if da1 == da: @@ -1980,9 +2151,6 @@ def _contact_elliptic( efc_J_out[worldid, efcid, dofid] = 0.0 dofid -= 1 - if is_sparse: - efc_J_rownnz_out[worldid, efcid] = rownnz - body_invweight0_id = worldid % body_invweight0.shape[0] invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] @@ -2013,561 +2181,599 @@ def _contact_elliptic( efc_type = ConstraintType.CONTACT_ELLIPTIC _efc_row( - opt_disableflags, - worldid, - timestep, - efcid, - pos_aref, - pos, - invweight, - ref, - solimp_in[conid], - includemargin, - Jqvel, - 0.0, - efc_type, - conid, - efc_type_out, - efc_id_out, - efc_pos_out, - efc_margin_out, - efc_D_out, - efc_vel_out, - efc_aref_out, - efc_frictionloss_out, + opt_disableflags, + worldid, + timestep, + efcid, + pos_aref, + pos, + invweight, + ref, + solimp_in[conid], + includemargin, + Jqvel, + 0.0, + efc_type, + conid, + efc_type_out, + efc_id_out, + efc_pos_out, + efc_margin_out, + efc_D_out, + efc_vel_out, + efc_aref_out, + efc_frictionloss_out, ) @event_scope def make_constraint(m: types.Model, d: types.Data): """Creates constraint jacobians and other supporting data.""" + efc_nnz = wp.empty((d.nworld,), dtype=int) + wp.launch( _zero_constraint_counts, dim=d.nworld, - inputs=[d.ne, d.nf, d.nl, d.nefc], + inputs=[d.ne, d.nf, d.nl, d.nefc, efc_nnz], ) - if types.SPARSE_CONSTRAINT_JACOBIAN: - d.contact.efc_address.fill_(-1) - if not (m.opt.disableflags & types.DisableBit.CONSTRAINT): if not (m.opt.disableflags & types.DisableBit.EQUALITY): wp.launch( - _equality_connect, - dim=(d.nworld, m.eq_connect_adr.size), - inputs=[ - m.nv, - m.nsite, - m.opt.timestep, - m.opt.disableflags, - m.body_parentid, - m.body_rootid, - m.body_weldid, - m.body_dofnum, - m.body_dofadr, - m.body_invweight0, - m.dof_bodyid, - m.dof_parentid, - m.site_bodyid, - m.eq_obj1id, - m.eq_obj2id, - m.eq_objtype, - m.eq_solref, - m.eq_solimp, - m.eq_data, - SPARSE_CONSTRAINT_JACOBIAN, - m.eq_connect_adr, - d.qvel, - d.eq_active, - d.xpos, - d.xmat, - d.site_xpos, - d.subtree_com, - d.cdof, - d.njmax, - ], - outputs=[ - d.ne, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _equality_connect, + dim=(d.nworld, m.eq_connect_adr.size), + inputs=[ + m.nv, + m.nsite, + m.opt.timestep, + m.opt.disableflags, + m.body_parentid, + m.body_rootid, + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.body_invweight0, + m.dof_bodyid, + m.dof_parentid, + m.site_bodyid, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + m.eq_solref, + m.eq_solimp, + m.eq_data, + m.is_sparse, + m.eq_connect_adr, + d.qvel, + d.eq_active, + d.xpos, + d.xmat, + d.site_xpos, + d.subtree_com, + d.cdof, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.ne, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) wp.launch( - _equality_weld, - dim=(d.nworld, m.eq_wld_adr.size), - inputs=[ - m.nv, - m.nsite, - m.opt.timestep, - m.opt.disableflags, - m.body_parentid, - m.body_rootid, - m.body_weldid, - m.body_dofnum, - m.body_dofadr, - m.body_invweight0, - m.dof_bodyid, - m.dof_parentid, - m.site_bodyid, - m.site_quat, - m.eq_obj1id, - m.eq_obj2id, - m.eq_objtype, - m.eq_solref, - m.eq_solimp, - m.eq_data, - SPARSE_CONSTRAINT_JACOBIAN, - m.eq_wld_adr, - d.qvel, - d.eq_active, - d.xpos, - d.xquat, - d.xmat, - d.site_xpos, - d.subtree_com, - d.cdof, - d.njmax, - ], - outputs=[ - d.ne, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _equality_weld, + dim=(d.nworld, m.eq_wld_adr.size), + inputs=[ + m.nv, + m.nsite, + m.opt.timestep, + m.opt.disableflags, + m.body_parentid, + m.body_rootid, + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.body_invweight0, + m.dof_bodyid, + m.dof_parentid, + m.site_bodyid, + m.site_quat, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + m.eq_solref, + m.eq_solimp, + m.eq_data, + m.is_sparse, + m.eq_wld_adr, + d.qvel, + d.eq_active, + d.xpos, + d.xquat, + d.xmat, + d.site_xpos, + d.subtree_com, + d.cdof, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.ne, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) wp.launch( - _equality_joint, - dim=(d.nworld, m.eq_jnt_adr.size), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.qpos0, - m.jnt_qposadr, - m.jnt_dofadr, - m.dof_invweight0, - m.eq_obj1id, - m.eq_obj2id, - m.eq_solref, - m.eq_solimp, - m.eq_data, - SPARSE_CONSTRAINT_JACOBIAN, - m.eq_jnt_adr, - d.qpos, - d.qvel, - d.eq_active, - d.njmax, - ], - outputs=[ - d.ne, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _equality_joint, + dim=(d.nworld, m.eq_jnt_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.qpos0, + m.jnt_qposadr, + m.jnt_dofadr, + m.dof_invweight0, + m.eq_obj1id, + m.eq_obj2id, + m.eq_solref, + m.eq_solimp, + m.eq_data, + m.is_sparse, + m.eq_jnt_adr, + d.qpos, + d.qvel, + d.eq_active, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.ne, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) wp.launch( - _equality_tendon, - dim=(d.nworld, m.eq_ten_adr.size), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.eq_obj1id, - m.eq_obj2id, - m.eq_solref, - m.eq_solimp, - m.eq_data, - m.tendon_length0, - m.tendon_invweight0, - SPARSE_CONSTRAINT_JACOBIAN, - m.eq_ten_adr, - d.qvel, - d.eq_active, - d.ten_J, - d.ten_length, - d.njmax, - ], - outputs=[ - d.ne, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _equality_tendon, + dim=(d.nworld, m.eq_ten_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.eq_obj1id, + m.eq_obj2id, + m.eq_solref, + m.eq_solimp, + m.eq_data, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, + m.tendon_length0, + m.tendon_invweight0, + m.is_sparse, + m.eq_ten_adr, + d.qvel, + d.eq_active, + d.ten_J, + d.ten_length, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.ne, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) wp.launch( - _equality_flex(SPARSE_CONSTRAINT_JACOBIAN), - dim=(d.nworld, m.eq_flex_adr.size, m.nflexedge), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.flexedge_length0, - m.flexedge_invweight0, - m.flexedge_J_rownnz, - m.flexedge_J_rowadr, - m.flexedge_J_colind, - m.eq_solref, - m.eq_solimp, - m.eq_flex_adr, - d.qvel, - d.flexedge_J, - d.flexedge_length, - d.njmax, - ], - outputs=[ - d.ne, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _equality_flex(m.is_sparse), + dim=(d.nworld, m.eq_flex_adr.size, m.nflexedge), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.flex_edgeadr, + m.flex_edgenum, + m.flexedge_length0, + m.flexedge_invweight0, + m.flexedge_J_rownnz, + m.flexedge_J_rowadr, + m.flexedge_J_colind, + m.eq_obj1id, + m.eq_solref, + m.eq_solimp, + m.eq_flex_adr, + d.qvel, + d.flexedge_J, + d.flexedge_length, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.ne, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) if not (m.opt.disableflags & types.DisableBit.FRICTIONLOSS): wp.launch( - _friction_dof, - dim=(d.nworld, m.nv), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.dof_solref, - m.dof_solimp, - m.dof_frictionloss, - m.dof_invweight0, - SPARSE_CONSTRAINT_JACOBIAN, - d.qvel, - d.njmax, - ], - outputs=[ - d.nf, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _friction_dof, + dim=(d.nworld, m.nv), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.dof_solref, + m.dof_solimp, + m.dof_frictionloss, + m.dof_invweight0, + m.is_sparse, + d.qvel, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.nf, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) wp.launch( - _friction_tendon, - dim=(d.nworld, m.ntendon), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.tendon_solref_fri, - m.tendon_solimp_fri, - m.tendon_frictionloss, - m.tendon_invweight0, - SPARSE_CONSTRAINT_JACOBIAN, - d.qvel, - d.ten_J, - d.njmax, - ], - outputs=[ - d.nf, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _friction_tendon, + dim=(d.nworld, m.ntendon), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, + m.tendon_solref_fri, + m.tendon_solimp_fri, + m.tendon_frictionloss, + m.tendon_invweight0, + m.is_sparse, + d.qvel, + d.ten_J, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.nf, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) # limit if not (m.opt.disableflags & types.DisableBit.LIMIT): wp.launch( - _limit_ball, - dim=(d.nworld, m.jnt_limited_ball_adr.size), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.jnt_qposadr, - m.jnt_dofadr, - m.jnt_solref, - m.jnt_solimp, - m.jnt_range, - m.jnt_margin, - m.dof_invweight0, - SPARSE_CONSTRAINT_JACOBIAN, - m.jnt_limited_ball_adr, - d.qpos, - d.qvel, - d.njmax, - ], - outputs=[ - d.nl, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _limit_ball, + dim=(d.nworld, m.jnt_limited_ball_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.jnt_qposadr, + m.jnt_dofadr, + m.jnt_solref, + m.jnt_solimp, + m.jnt_range, + m.jnt_margin, + m.dof_invweight0, + m.is_sparse, + m.jnt_limited_ball_adr, + d.qpos, + d.qvel, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.nl, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) wp.launch( - _limit_slide_hinge, - dim=(d.nworld, m.jnt_limited_slide_hinge_adr.size), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.jnt_qposadr, - m.jnt_dofadr, - m.jnt_solref, - m.jnt_solimp, - m.jnt_range, - m.jnt_margin, - m.dof_invweight0, - SPARSE_CONSTRAINT_JACOBIAN, - m.jnt_limited_slide_hinge_adr, - d.qpos, - d.qvel, - d.njmax, - ], - outputs=[ - d.nl, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _limit_slide_hinge, + dim=(d.nworld, m.jnt_limited_slide_hinge_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.jnt_qposadr, + m.jnt_dofadr, + m.jnt_solref, + m.jnt_solimp, + m.jnt_range, + m.jnt_margin, + m.dof_invweight0, + m.is_sparse, + m.jnt_limited_slide_hinge_adr, + d.qpos, + d.qvel, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.nl, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) wp.launch( - _limit_tendon, - dim=(d.nworld, m.tendon_limited_adr.size), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.jnt_dofadr, - m.tendon_adr, - m.tendon_num, - m.tendon_solref_lim, - m.tendon_solimp_lim, - m.tendon_range, - m.tendon_margin, - m.tendon_invweight0, - m.wrap_type, - m.wrap_objid, - SPARSE_CONSTRAINT_JACOBIAN, - m.tendon_limited_adr, - d.qvel, - d.ten_J, - d.ten_length, - d.njmax, - ], - outputs=[ - d.nl, - d.nefc, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _limit_tendon, + dim=(d.nworld, m.tendon_limited_adr.size), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, + m.tendon_solref_lim, + m.tendon_solimp_lim, + m.tendon_range, + m.tendon_margin, + m.tendon_invweight0, + m.is_sparse, + m.tendon_limited_adr, + d.qvel, + d.ten_J, + d.ten_length, + d.njmax, + d.njmax_nnz, + ], + outputs=[ + d.nl, + d.nefc, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) # contact if not (m.opt.disableflags & types.DisableBit.CONTACT): if m.opt.cone == types.ConeType.PYRAMIDAL: wp.launch( - _contact_pyramidal, - dim=(d.naconmax, m.nmaxpyramid), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.opt.impratio_invsqrt, - m.body_parentid, - m.body_rootid, - m.body_weldid, - m.body_dofnum, - m.body_dofadr, - m.body_invweight0, - m.dof_bodyid, - m.dof_parentid, - m.geom_bodyid, - SPARSE_CONSTRAINT_JACOBIAN, - d.qvel, - d.subtree_com, - d.cdof, - d.njmax, - d.nacon, - d.contact.dist, - d.contact.dim, - d.contact.includemargin, - d.contact.worldid, - d.contact.geom, - d.contact.pos, - d.contact.frame, - d.contact.friction, - d.contact.solref, - d.contact.solimp, - d.contact.type, - ], - outputs=[ - d.nefc, - d.contact.efc_address, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _contact_pyramidal, + dim=(d.naconmax, m.nmaxpyramid), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.opt.impratio_invsqrt, + m.body_parentid, + m.body_rootid, + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.body_invweight0, + m.dof_bodyid, + m.dof_parentid, + m.geom_bodyid, + m.flex_vertadr, + m.flex_vertbodyid, + m.is_sparse, + d.qvel, + d.subtree_com, + d.cdof, + d.njmax, + d.njmax_nnz, + d.nacon, + d.contact.dist, + d.contact.dim, + d.contact.includemargin, + d.contact.worldid, + d.contact.geom, + d.contact.flex, + d.contact.vert, + d.contact.pos, + d.contact.frame, + d.contact.friction, + d.contact.solref, + d.contact.solimp, + d.contact.type, + ], + outputs=[ + d.nefc, + d.contact.efc_address, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) elif m.opt.cone == types.ConeType.ELLIPTIC: wp.launch( - _contact_elliptic, - dim=(d.naconmax, m.nmaxcondim), - inputs=[ - m.nv, - m.opt.timestep, - m.opt.disableflags, - m.opt.impratio_invsqrt, - m.body_parentid, - m.body_rootid, - m.body_weldid, - m.body_dofnum, - m.body_dofadr, - m.body_invweight0, - m.dof_bodyid, - m.dof_parentid, - m.geom_bodyid, - SPARSE_CONSTRAINT_JACOBIAN, - d.qvel, - d.subtree_com, - d.cdof, - d.njmax, - d.nacon, - d.contact.dist, - d.contact.dim, - d.contact.includemargin, - d.contact.worldid, - d.contact.geom, - d.contact.pos, - d.contact.frame, - d.contact.friction, - d.contact.solref, - d.contact.solreffriction, - d.contact.solimp, - d.contact.type, - ], - outputs=[ - d.nefc, - d.contact.efc_address, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.pos, - d.efc.margin, - d.efc.D, - d.efc.vel, - d.efc.aref, - d.efc.frictionloss, - ], + _contact_elliptic, + dim=(d.naconmax, m.nmaxcondim), + inputs=[ + m.nv, + m.opt.timestep, + m.opt.disableflags, + m.opt.impratio_invsqrt, + m.body_parentid, + m.body_rootid, + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.body_invweight0, + m.dof_bodyid, + m.dof_parentid, + m.geom_bodyid, + m.flex_vertadr, + m.flex_vertbodyid, + m.is_sparse, + d.qvel, + d.subtree_com, + d.cdof, + d.njmax, + d.njmax_nnz, + d.nacon, + d.contact.dist, + d.contact.dim, + d.contact.includemargin, + d.contact.worldid, + d.contact.geom, + d.contact.flex, + d.contact.vert, + d.contact.pos, + d.contact.frame, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.type, + ], + outputs=[ + d.nefc, + d.contact.efc_address, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.pos, + d.efc.margin, + d.efc.D, + d.efc.vel, + d.efc.aref, + d.efc.frictionloss, + efc_nnz, + ], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py index 406c36c4..20da751a 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py @@ -15,6 +15,7 @@ import warp as wp +from mujoco.mjx.third_party.mujoco_warp._src.support import next_act from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit @@ -30,18 +31,24 @@ wp.set_module_options({"enable_backward": False}) @wp.kernel def _qderiv_actuator_passive_vel( # Model: + opt_timestep: wp.array(dtype=float), actuator_dyntype: wp.array(dtype=int), actuator_gaintype: wp.array(dtype=int), actuator_biastype: wp.array(dtype=int), actuator_actadr: wp.array(dtype=int), actuator_actnum: wp.array(dtype=int), actuator_forcelimited: wp.array(dtype=bool), + actuator_actlimited: wp.array(dtype=bool), + actuator_dynprm: wp.array2d(dtype=vec10f), actuator_gainprm: wp.array2d(dtype=vec10f), actuator_biasprm: wp.array2d(dtype=vec10f), + actuator_actearly: wp.array(dtype=bool), actuator_forcerange: wp.array2d(dtype=wp.vec2), + actuator_actrange: wp.array2d(dtype=wp.vec2), # Data in: act_in: wp.array2d(dtype=float), ctrl_in: wp.array2d(dtype=float), + act_dot_in: wp.array2d(dtype=float), actuator_force_in: wp.array2d(dtype=float), # Out: vel_out: wp.array2d(dtype=float), @@ -76,9 +83,24 @@ def _qderiv_actuator_passive_vel( vel = float(bias) if actuator_dyntype[actid] != DynType.NONE: if gain != 0.0: - act_first = actuator_actadr[actid] - act_last = act_first + actuator_actnum[actid] - 1 - vel += gain * act_in[worldid, act_last] + act_adr = actuator_actadr[actid] + actuator_actnum[actid] - 1 + + # use next activation if actearly is set (matching forward pass) + if actuator_actearly[actid]: + act = next_act( + opt_timestep[worldid % opt_timestep.shape[0]], + actuator_dyntype[actid], + actuator_dynprm[worldid % actuator_dynprm.shape[0], actid], + actuator_actrange[worldid % actuator_actrange.shape[0], actid], + act_in[worldid, act_adr], + act_dot_in[worldid, act_adr], + 1.0, + actuator_actlimited[actid], + ) + else: + act = act_in[worldid, act_adr] + + vel += gain * act else: if gain != 0.0: vel += gain * ctrl_in[worldid, actid] @@ -95,21 +117,20 @@ def _nonzero_mask(x: float) -> float: @wp.kernel -def _qderiv_actuator_passive_actuation_sparse( - # Model: - nu: int, - is_sparse: bool, - # Data in: - moment_rownnz_in: wp.array2d(dtype=int), - moment_rowadr_in: wp.array2d(dtype=int), - moment_colind_in: wp.array2d(dtype=int), - actuator_moment_in: wp.array2d(dtype=float), - # In: - vel_in: wp.array2d(dtype=float), - qMi: wp.array(dtype=int), - qMj: wp.array(dtype=int), - # Out: - qDeriv_out: wp.array3d(dtype=float), +def _qderiv_actuator_passive_actuation_dense( + # Model: + nu: int, + # Data in: + moment_rownnz_in: wp.array2d(dtype=int), + moment_rowadr_in: wp.array2d(dtype=int), + moment_colind_in: wp.array2d(dtype=int), + actuator_moment_in: wp.array2d(dtype=float), + # In: + vel_in: wp.array2d(dtype=float), + qMi: wp.array(dtype=int), + qMj: wp.array(dtype=int), + # Out: + qDeriv_out: wp.array3d(dtype=float), ): worldid, elemid = wp.tid() @@ -142,12 +163,63 @@ def _qderiv_actuator_passive_actuation_sparse( qderiv_contrib += moment_i * moment_j * vel - if is_sparse: - qDeriv_out[worldid, 0, elemid] = qderiv_contrib - else: - qDeriv_out[worldid, dofiid, dofjid] = qderiv_contrib - if dofiid != dofjid: - qDeriv_out[worldid, dofjid, dofiid] = qderiv_contrib + qDeriv_out[worldid, dofiid, dofjid] = qderiv_contrib + if dofiid != dofjid: + qDeriv_out[worldid, dofjid, dofiid] = qderiv_contrib + + +@wp.kernel +def _qderiv_actuator_passive_actuation_sparse( + # Model: + M_rownnz: wp.array(dtype=int), + M_rowadr: wp.array(dtype=int), + # Data in: + moment_rownnz_in: wp.array2d(dtype=int), + moment_rowadr_in: wp.array2d(dtype=int), + moment_colind_in: wp.array2d(dtype=int), + actuator_moment_in: wp.array2d(dtype=float), + # In: + vel_in: wp.array2d(dtype=float), + qMj: wp.array(dtype=int), + # Out: + qDeriv_out: wp.array3d(dtype=float), +): + worldid, actid = wp.tid() + + vel = vel_in[worldid, actid] + if vel == 0.0: + return + + rownnz = moment_rownnz_in[worldid, actid] + rowadr = moment_rowadr_in[worldid, actid] + + for i in range(rownnz): + rowadri = rowadr + i + moment_i = actuator_moment_in[worldid, rowadri] + if moment_i == 0.0: + continue + dofi = moment_colind_in[worldid, rowadri] + + for j in range(i + 1): + rowadrj = rowadr + j + moment_j = actuator_moment_in[worldid, rowadrj] + if moment_j == 0.0: + continue + dofj = moment_colind_in[worldid, rowadrj] + + contrib = moment_i * moment_j * vel + + # Search the corresponding elemid + # TODO: This could be precalculated for improved performance + row = dofi + col = dofj + row_startk = M_rowadr[row] - 1 + row_nnz = M_rownnz[row] + for k in range(row_nnz): + row_startk += 1 + if qMj[row_startk] == col: + wp.atomic_add(qDeriv_out[worldid, 0], row_startk, contrib) + break @wp.kernel @@ -176,7 +248,7 @@ def _qderiv_actuator_passive( else: qderiv = qDeriv_in[worldid, dofiid, dofjid] - if not opt_disableflags & DisableBit.DAMPER and dofiid == dofjid: + if not (opt_disableflags & DisableBit.DAMPER) and dofiid == dofjid: qderiv -= dof_damping[worldid % dof_damping.shape[0], dofiid] qderiv *= opt_timestep[worldid % opt_timestep.shape[0]] @@ -196,10 +268,13 @@ def _qderiv_tendon_damping( # Model: ntendon: int, opt_timestep: wp.array(dtype=float), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), tendon_damping: wp.array2d(dtype=float), is_sparse: bool, # Data in: - ten_J_in: wp.array3d(dtype=float), + ten_J_in: wp.array2d(dtype=float), # In: qMi: wp.array(dtype=int), qMj: wp.array(dtype=int), @@ -213,7 +288,24 @@ def _qderiv_tendon_damping( qderiv = float(0.0) tendon_damping_id = worldid % tendon_damping.shape[0] for tenid in range(ntendon): - qderiv -= ten_J_in[worldid, tenid, dofiid] * ten_J_in[worldid, tenid, dofjid] * tendon_damping[tendon_damping_id, tenid] + damping = tendon_damping[tendon_damping_id, tenid] + if damping == 0.0: + continue + + rownnz = ten_J_rownnz[tenid] + rowadr = ten_J_rowadr[tenid] + Ji = float(0.0) + Jj = float(0.0) + for k in range(rownnz): + if Ji != 0.0 and Jj != 0.0: + break + sparseid = rowadr + k + colind = ten_J_colind[sparseid] + if colind == dofiid: + Ji = ten_J_in[worldid, sparseid] + if colind == dofjid: + Jj = ten_J_in[worldid, sparseid] + qderiv -= Ji * Jj * damping qderiv *= opt_timestep[worldid % opt_timestep.shape[0]] @@ -242,43 +334,47 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d(dtype=float)): if ~(m.opt.disableflags & (DisableBit.ACTUATION | DisableBit.DAMPER)): # TODO(team): only clear elements not set by _qderiv_actuator_passive out.zero_() - if m.nu > 0 and not m.opt.disableflags & DisableBit.ACTUATION: + if m.nu > 0 and not (m.opt.disableflags & DisableBit.ACTUATION): vel = wp.empty((d.nworld, m.nu), dtype=float) wp.launch( _qderiv_actuator_passive_vel, dim=(d.nworld, m.nu), inputs=[ + m.opt.timestep, m.actuator_dyntype, m.actuator_gaintype, m.actuator_biastype, m.actuator_actadr, m.actuator_actnum, m.actuator_forcelimited, + m.actuator_actlimited, + m.actuator_dynprm, m.actuator_gainprm, m.actuator_biasprm, + m.actuator_actearly, m.actuator_forcerange, + m.actuator_actrange, d.act, d.ctrl, + d.act_dot, d.actuator_force, ], outputs=[vel], ) - wp.launch( + if m.is_sparse: + wp.launch( _qderiv_actuator_passive_actuation_sparse, - dim=(d.nworld, qMi.size), - inputs=[ - m.nu, - m.is_sparse, - d.moment_rownnz, - d.moment_rowadr, - d.moment_colind, - d.actuator_moment, - vel, - qMi, - qMj, - ], + dim=(d.nworld, m.nu), + inputs=[m.M_rownnz, m.M_rowadr, d.moment_rownnz, d.moment_rowadr, d.moment_colind, d.actuator_moment, vel, qMj], outputs=[out], - ) + ) + else: + wp.launch( + _qderiv_actuator_passive_actuation_dense, + dim=(d.nworld, qMi.size), + inputs=[m.nu, d.moment_rownnz, d.moment_rowadr, d.moment_colind, d.actuator_moment, vel, qMi, qMj], + outputs=[out], + ) wp.launch( _qderiv_actuator_passive, dim=(d.nworld, qMi.size), @@ -298,11 +394,22 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d(dtype=float)): # TODO(team): directly utilize qM for these settings wp.copy(out, d.qM) - if not m.opt.disableflags & DisableBit.DAMPER: + if not (m.opt.disableflags & DisableBit.DAMPER): wp.launch( _qderiv_tendon_damping, dim=(d.nworld, qMi.size), - inputs=[m.ntendon, m.opt.timestep, m.tendon_damping, m.is_sparse, d.ten_J, qMi, qMj], + inputs=[ + m.ntendon, + m.opt.timestep, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, + m.tendon_damping, + m.is_sparse, + d.ten_J, + qMi, + qMj, + ], outputs=[out], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py index 0dc3de14..64bdd91f 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py @@ -15,6 +15,8 @@ from typing import Optional +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src import collision_driver from mujoco.mjx.third_party.mujoco_warp._src import constraint from mujoco.mjx.third_party.mujoco_warp._src import derivative @@ -25,7 +27,9 @@ from mujoco.mjx.third_party.mujoco_warp._src import sensor from mujoco.mjx.third_party.mujoco_warp._src import smooth from mujoco.mjx.third_party.mujoco_warp._src import solver from mujoco.mjx.third_party.mujoco_warp._src import util_misc +from mujoco.mjx.third_party.mujoco_warp._src.support import next_act from mujoco.mjx.third_party.mujoco_warp._src.support import xfrc_accumulate +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit @@ -34,14 +38,12 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import EnableBit from mujoco.mjx.third_party.mujoco_warp._src.types import GainType from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType from mujoco.mjx.third_party.mujoco_warp._src.types import JointType -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import Model from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -127,55 +129,24 @@ def _next_velocity( qvel_out[worldid, dofid] = qvel_in[worldid, dofid] + qacc_scale_in * qacc_in[worldid, dofid] * timestep -# TODO(team): kernel analyzer array slice? -@wp.func -def _next_act( - # Model: - opt_timestep: float, # kernel_analyzer: ignore - actuator_dyntype: int, # kernel_analyzer: ignore - actuator_dynprm: vec10f, # kernel_analyzer: ignore - actuator_actrange: wp.vec2, # kernel_analyzer: ignore - # Data In: - act_in: float, # kernel_analyzer: ignore - act_dot_in: float, # kernel_analyzer: ignore - # In: - act_dot_scale: float, - clamp: bool, -) -> float: - # advance actuation - if actuator_dyntype == DynType.FILTEREXACT: - tau = wp.max(MJ_MINVAL, actuator_dynprm[0]) - act = act_in + act_dot_scale * act_dot_in * tau * (1.0 - wp.exp(-opt_timestep / tau)) - elif actuator_dyntype == DynType.USER: - return act_in - else: - act = act_in + act_dot_scale * act_dot_in * opt_timestep - - # clamp to actrange - if clamp: - act = wp.clamp(act, actuator_actrange[0], actuator_actrange[1]) - - return act - - @wp.kernel def _next_activation( - # Model: - opt_timestep: wp.array(dtype=float), - actuator_dyntype: wp.array(dtype=int), - actuator_actadr: wp.array(dtype=int), - actuator_actnum: wp.array(dtype=int), - actuator_actlimited: wp.array(dtype=bool), - actuator_dynprm: wp.array2d(dtype=vec10f), - actuator_actrange: wp.array2d(dtype=wp.vec2), - # Data in: - act_in: wp.array2d(dtype=float), - act_dot_in: wp.array2d(dtype=float), - # In: - act_dot_scale: float, - limit: bool, - # Data out: - act_out: wp.array2d(dtype=float), + # Model: + opt_timestep: wp.array(dtype=float), + actuator_dyntype: wp.array(dtype=int), + actuator_actadr: wp.array(dtype=int), + actuator_actnum: wp.array(dtype=int), + actuator_actlimited: wp.array(dtype=bool), + actuator_dynprm: wp.array2d(dtype=vec10f), + actuator_actrange: wp.array2d(dtype=wp.vec2), + # Data in: + act_in: wp.array2d(dtype=float), + act_dot_in: wp.array2d(dtype=float), + # In: + act_dot_scale: float, + limit: bool, + # Data out: + act_out: wp.array2d(dtype=float), ): worldid, uid = wp.tid() opt_timestep_id = worldid % opt_timestep.shape[0] @@ -184,15 +155,15 @@ def _next_activation( actadr = actuator_actadr[uid] actnum = actuator_actnum[uid] for j in range(actadr, actadr + actnum): - act = _next_act( - opt_timestep[opt_timestep_id], - actuator_dyntype[uid], - actuator_dynprm[actuator_dynprm_id, uid], - actuator_actrange[actuator_actrange_id, uid], - act_in[worldid, j], - act_dot_in[worldid, j], - act_dot_scale, - limit and actuator_actlimited[uid], + act = next_act( + opt_timestep[opt_timestep_id], + actuator_dyntype[uid], + actuator_dynprm[actuator_dynprm_id, uid], + actuator_actrange[actuator_actrange_id, uid], + act_in[worldid, j], + act_dot_in[worldid, j], + act_dot_scale, + limit and actuator_actlimited[uid], ) act_out[worldid, j] = act @@ -201,12 +172,16 @@ def _next_activation( def _next_time( # Model: opt_timestep: wp.array(dtype=float), + is_sparse: bool, # Data in: nefc_in: wp.array(dtype=int), time_in: wp.array(dtype=float), + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), nworld_in: int, naconmax_in: int, njmax_in: int, + njmax_nnz_in: int, nacon_in: wp.array(dtype=int), ncollision_in: wp.array(dtype=int), # Data out: @@ -218,6 +193,11 @@ def _next_time( if nefc > njmax_in: wp.printf("nefc overflow - please increase njmax to %u\n", nefc) + elif nefc > 0 and is_sparse: + efcid = wp.min(nefc, njmax_in) - 1 + efc_nnz = efc_J_rowadr_in[worldid, efcid] + efc_J_rownnz_in[worldid, efcid] + if efc_nnz > njmax_nnz_in: + wp.printf("njmax_nnz overflow - please increase njmax_nnz to %u\n", efc_nnz) if worldid == 0: ncollision = ncollision_in[0] @@ -236,22 +216,22 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None) # advance activations wp.launch( - _next_activation, - dim=(d.nworld, m.nu), - inputs=[ - m.opt.timestep, - m.actuator_dyntype, - m.actuator_actadr, - m.actuator_actnum, - m.actuator_actlimited, - m.actuator_dynprm, - m.actuator_actrange, - d.act, - d.act_dot, - 1.0, - True, - ], - outputs=[d.act], + _next_activation, + dim=(d.nworld, m.nu), + inputs=[ + m.opt.timestep, + m.actuator_dyntype, + m.actuator_actadr, + m.actuator_actnum, + m.actuator_actlimited, + m.actuator_dynprm, + m.actuator_actrange, + d.act, + d.act_dot, + 1.0, + True, + ], + outputs=[d.act], ) wp.launch( @@ -274,7 +254,20 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None) wp.launch( _next_time, dim=d.nworld, - inputs=[m.opt.timestep, d.nefc, d.time, d.nworld, d.naconmax, d.njmax, d.nacon, d.ncollision], + inputs=[ + m.opt.timestep, + m.is_sparse, + d.nefc, + d.time, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.nworld, + d.naconmax, + d.njmax, + d.njmax_nnz, + d.nacon, + d.ncollision, + ], outputs=[d.time], ) @@ -294,9 +287,7 @@ def _euler_damp_qfrc_sparse( timestep = opt_timestep[worldid % opt_timestep.shape[0]] adr = dof_Madr[tid] - qM_integration_out[worldid, 0, adr] += ( - timestep * dof_damping[worldid % dof_damping.shape[0], tid] - ) + qM_integration_out[worldid, 0, adr] += timestep * dof_damping[worldid % dof_damping.shape[0], tid] @cache_kernel @@ -336,7 +327,7 @@ def _tile_euler_dense(tile: TileSet): def euler(m: Model, d: Data): """Euler integrator, semi-implicit in velocity.""" # integrate damping implicitly - if not m.opt.disableflags & (DisableBit.EULERDAMP | DisableBit.DAMPER): + if not (m.opt.disableflags & (DisableBit.EULERDAMP | DisableBit.DAMPER)): qacc = wp.empty((d.nworld, m.nv), dtype=float) if m.is_sparse: qM = wp.clone(d.qM) @@ -390,22 +381,22 @@ def _rk_perturb_state( # activation if m.na and act_t0 is not None: wp.launch( - _next_activation, - dim=(d.nworld, m.nu), - inputs=[ - m.opt.timestep, - m.actuator_dyntype, - m.actuator_actadr, - m.actuator_actnum, - m.actuator_actlimited, - m.actuator_dynprm, - m.actuator_actrange, - act_t0, - d.act_dot, - scale, - False, - ], - outputs=[d.act], + _next_activation, + dim=(d.nworld, m.nu), + inputs=[ + m.opt.timestep, + m.actuator_dyntype, + m.actuator_actadr, + m.actuator_actnum, + m.actuator_actlimited, + m.actuator_dynprm, + m.actuator_actrange, + act_t0, + d.act_dot, + scale, + False, + ], + outputs=[d.act], ) @@ -548,14 +539,14 @@ def fwd_position(m: Model, d: Data, factorize: bool = True): @wp.kernel def _actuator_velocity( - # Data in: - qvel_in: wp.array2d(dtype=float), - moment_rownnz_in: wp.array2d(dtype=int), - moment_rowadr_in: wp.array2d(dtype=int), - moment_colind_in: wp.array2d(dtype=int), - actuator_moment_in: wp.array2d(dtype=float), - # Data out: - actuator_velocity_out: wp.array2d(dtype=float), + # Data in: + qvel_in: wp.array2d(dtype=float), + moment_rownnz_in: wp.array2d(dtype=int), + moment_rowadr_in: wp.array2d(dtype=int), + moment_colind_in: wp.array2d(dtype=int), + actuator_moment_in: wp.array2d(dtype=float), + # Data out: + actuator_velocity_out: wp.array2d(dtype=float), ): worldid, actid = wp.tid() @@ -571,50 +562,49 @@ def _actuator_velocity( actuator_velocity_out[worldid, actid] = vel -@cache_kernel -def _tendon_velocity(nv: int): - @wp.kernel(module="unique", enable_backward=False) - def tendon_velocity( - # Data in: - qvel_in: wp.array2d(dtype=float), - ten_J_in: wp.array3d(dtype=float), - # Data out: - ten_velocity_out: wp.array2d(dtype=float), - ): - worldid, tenid = wp.tid() - ten_J_tile = wp.tile_load(ten_J_in[worldid, tenid], shape=wp.static(nv)) - qvel_tile = wp.tile_load(qvel_in[worldid], shape=wp.static(nv)) - ten_J_qvel_tile = wp.tile_map(wp.mul, ten_J_tile, qvel_tile) - ten_velocity_tile = wp.tile_reduce(wp.add, ten_J_qvel_tile) - ten_velocity_out[worldid, tenid] = ten_velocity_tile[0] +@wp.kernel +def _tendon_velocity( + # Model: + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + ten_J_in: wp.array2d(dtype=float), + # Data out: + ten_velocity_out: wp.array2d(dtype=float), +): + worldid, tenid = wp.tid() - return tendon_velocity + velocity = float(0.0) + rownnz = ten_J_rownnz[tenid] + rowadr = ten_J_rowadr[tenid] + for i in range(rownnz): + sparseid = rowadr + i + J = ten_J_in[worldid, sparseid] + if J != 0.0: + colind = ten_J_colind[sparseid] + velocity += J * qvel_in[worldid, colind] + + ten_velocity_out[worldid, tenid] = velocity @event_scope def fwd_velocity(m: Model, d: Data): """Velocity-dependent computations.""" - wp.launch_tiled( - _actuator_velocity, - dim=(d.nworld, m.nu), - inputs=[ - d.qvel, - d.moment_rownnz, - d.moment_rowadr, - d.moment_colind, - d.actuator_moment, - ], - outputs=[d.actuator_velocity], - block_dim=m.block_dim.actuator_velocity, + wp.launch( + _actuator_velocity, + dim=(d.nworld, m.nu), + inputs=[d.qvel, d.moment_rownnz, d.moment_rowadr, d.moment_colind, d.actuator_moment], + outputs=[d.actuator_velocity], + block_dim=m.block_dim.actuator_velocity, ) - # TODO(team): sparse version - wp.launch_tiled( - _tendon_velocity(m.nv), + wp.launch( + _tendon_velocity, dim=(d.nworld, m.ntendon), - inputs=[d.qvel, d.ten_J], + inputs=[m.ten_J_rownnz, m.ten_J_rowadr, m.ten_J_colind, d.qvel, d.ten_J], outputs=[d.ten_velocity], - block_dim=m.block_dim.tendon_velocity, ) smooth.com_vel(m, d) @@ -625,36 +615,36 @@ def fwd_velocity(m: Model, d: Data): @wp.kernel def _actuator_force( - # Model: - na: int, - opt_timestep: wp.array(dtype=float), - actuator_dyntype: wp.array(dtype=int), - actuator_gaintype: wp.array(dtype=int), - actuator_biastype: wp.array(dtype=int), - actuator_actadr: wp.array(dtype=int), - actuator_actnum: wp.array(dtype=int), - actuator_ctrllimited: wp.array(dtype=bool), - actuator_forcelimited: wp.array(dtype=bool), - actuator_actlimited: wp.array(dtype=bool), - actuator_dynprm: wp.array2d(dtype=vec10f), - actuator_gainprm: wp.array2d(dtype=vec10f), - actuator_biasprm: wp.array2d(dtype=vec10f), - actuator_actearly: wp.array(dtype=bool), - actuator_ctrlrange: wp.array2d(dtype=wp.vec2), - actuator_forcerange: wp.array2d(dtype=wp.vec2), - actuator_actrange: wp.array2d(dtype=wp.vec2), - actuator_acc0: wp.array2d(dtype=float), - actuator_lengthrange: wp.array2d(dtype=wp.vec2), - # Data in: - act_in: wp.array2d(dtype=float), - ctrl_in: wp.array2d(dtype=float), - actuator_length_in: wp.array2d(dtype=float), - actuator_velocity_in: wp.array2d(dtype=float), - # In: - dsbl_clampctrl: int, - # Data out: - act_dot_out: wp.array2d(dtype=float), - actuator_force_out: wp.array2d(dtype=float), + # Model: + na: int, + opt_timestep: wp.array(dtype=float), + actuator_dyntype: wp.array(dtype=int), + actuator_gaintype: wp.array(dtype=int), + actuator_biastype: wp.array(dtype=int), + actuator_actadr: wp.array(dtype=int), + actuator_actnum: wp.array(dtype=int), + actuator_ctrllimited: wp.array(dtype=bool), + actuator_forcelimited: wp.array(dtype=bool), + actuator_actlimited: wp.array(dtype=bool), + actuator_dynprm: wp.array2d(dtype=vec10f), + actuator_gainprm: wp.array2d(dtype=vec10f), + actuator_biasprm: wp.array2d(dtype=vec10f), + actuator_actearly: wp.array(dtype=bool), + actuator_ctrlrange: wp.array2d(dtype=wp.vec2), + actuator_forcerange: wp.array2d(dtype=wp.vec2), + actuator_actrange: wp.array2d(dtype=wp.vec2), + actuator_acc0: wp.array2d(dtype=float), + actuator_lengthrange: wp.array2d(dtype=wp.vec2), + # Data in: + act_in: wp.array2d(dtype=float), + ctrl_in: wp.array2d(dtype=float), + actuator_length_in: wp.array2d(dtype=float), + actuator_velocity_in: wp.array2d(dtype=float), + # In: + dsbl_clampctrl: int, + # Data out: + act_dot_out: wp.array2d(dtype=float), + actuator_force_out: wp.array2d(dtype=float), ): worldid, uid = wp.tid() @@ -693,7 +683,7 @@ def _actuator_force( if dyntype == DynType.INTEGRATOR or dyntype == DynType.NONE: act = act_in[worldid, act_last] - ctrl_act = _next_act( + ctrl_act = next_act( opt_timestep[worldid % opt_timestep.shape[0]], dyntype, dynprm, @@ -720,9 +710,7 @@ def _actuator_force( gain = gainprm[0] + gainprm[1] * length + gainprm[2] * velocity elif gaintype == GainType.MUSCLE: acc0 = actuator_acc0[worldid % actuator_acc0.shape[0], uid] - lengthrange = actuator_lengthrange[ - worldid % actuator_lengthrange.shape[0], uid - ] + lengthrange = actuator_lengthrange[worldid % actuator_lengthrange.shape[0], uid] gain = util_misc.muscle_gain(length, velocity, lengthrange, acc0, gainprm) # GainType.USER: gain stays 0, modified by act_gain_callback @@ -735,9 +723,7 @@ def _actuator_force( bias = biasprm[0] + biasprm[1] * length + biasprm[2] * velocity elif biastype == BiasType.MUSCLE: acc0 = actuator_acc0[worldid % actuator_acc0.shape[0], uid] - lengthrange = actuator_lengthrange[ - worldid % actuator_lengthrange.shape[0], uid - ] + lengthrange = actuator_lengthrange[worldid % actuator_lengthrange.shape[0], uid] bias = util_misc.muscle_bias(length, lengthrange, acc0, biasprm) force = gain * ctrl_act + bias @@ -795,14 +781,14 @@ def _tendon_actuator_force_clamp( @wp.kernel def _qfrc_actuator( - # Data in: - moment_rownnz_in: wp.array2d(dtype=int), - moment_rowadr_in: wp.array2d(dtype=int), - moment_colind_in: wp.array2d(dtype=int), - actuator_moment_in: wp.array2d(dtype=float), - actuator_force_in: wp.array2d(dtype=float), - # Data out: - qfrc_actuator_out: wp.array2d(dtype=float), + # Data in: + moment_rownnz_in: wp.array2d(dtype=int), + moment_rowadr_in: wp.array2d(dtype=int), + moment_colind_in: wp.array2d(dtype=int), + actuator_moment_in: wp.array2d(dtype=float), + actuator_force_in: wp.array2d(dtype=float), + # Data out: + qfrc_actuator_out: wp.array2d(dtype=float), ): worldid, actid = wp.tid() @@ -812,26 +798,23 @@ def _qfrc_actuator( for i in range(rownnz): sparseid = rowadr + i colind = moment_colind_in[worldid, sparseid] - qfrc = ( - actuator_moment_in[worldid, sparseid] - * actuator_force_in[worldid, actid] - ) + qfrc = actuator_moment_in[worldid, sparseid] * actuator_force_in[worldid, actid] wp.atomic_add(qfrc_actuator_out[worldid], colind, qfrc) @wp.kernel def _qfrc_actuator_gravcomp_limits( - # Model: - ngravcomp: int, - jnt_actfrclimited: wp.array(dtype=bool), - jnt_actgravcomp: wp.array(dtype=int), - jnt_actfrcrange: wp.array2d(dtype=wp.vec2), - dof_jntid: wp.array(dtype=int), - # Data in: - qfrc_gravcomp_in: wp.array2d(dtype=float), - qfrc_actuator_in: wp.array2d(dtype=float), - # Data out: - qfrc_actuator_out: wp.array2d(dtype=float), + # Model: + ngravcomp: int, + jnt_actfrclimited: wp.array(dtype=bool), + jnt_actgravcomp: wp.array(dtype=int), + jnt_actfrcrange: wp.array2d(dtype=wp.vec2), + dof_jntid: wp.array(dtype=int), + # Data in: + qfrc_gravcomp_in: wp.array2d(dtype=float), + qfrc_actuator_in: wp.array2d(dtype=float), + # Data out: + qfrc_actuator_out: wp.array2d(dtype=float), ): worldid, dofid = wp.tid() jntid = dof_jntid[dofid] @@ -917,30 +900,30 @@ def fwd_actuation(m: Model, d: Data): # TODO(team): optimize performance d.qfrc_actuator.zero_() wp.launch( - _qfrc_actuator, - dim=(d.nworld, m.nu), - inputs=[ - d.moment_rownnz, - d.moment_rowadr, - d.moment_colind, - d.actuator_moment, - d.actuator_force, - ], - outputs=[d.qfrc_actuator], + _qfrc_actuator, + dim=(d.nworld, m.nu), + inputs=[ + d.moment_rownnz, + d.moment_rowadr, + d.moment_colind, + d.actuator_moment, + d.actuator_force, + ], + outputs=[d.qfrc_actuator], ) wp.launch( - _qfrc_actuator_gravcomp_limits, - dim=(d.nworld, m.nv), - inputs=[ - m.ngravcomp, - m.jnt_actfrclimited, - m.jnt_actgravcomp, - m.jnt_actfrcrange, - m.dof_jntid, - d.qfrc_gravcomp, - d.qfrc_actuator, - ], - outputs=[d.qfrc_actuator], + _qfrc_actuator_gravcomp_limits, + dim=(d.nworld, m.nv), + inputs=[ + m.ngravcomp, + m.jnt_actfrclimited, + m.jnt_actgravcomp, + m.jnt_actfrcrange, + m.dof_jntid, + d.qfrc_gravcomp, + d.qfrc_actuator, + ], + outputs=[d.qfrc_actuator], ) 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 e516739c..0b53094b 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -14,24 +14,24 @@ # ============================================================================== import dataclasses -from typing import Any, Optional, Sequence import warnings +from typing import Any, Optional, Sequence import mujoco +import numpy as np +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src import bvh from mujoco.mjx.third_party.mujoco_warp._src import math as mjmath from mujoco.mjx.third_party.mujoco_warp._src import render_util from mujoco.mjx.third_party.mujoco_warp._src import smooth from mujoco.mjx.third_party.mujoco_warp._src import types from mujoco.mjx.third_party.mujoco_warp._src import warp_util -from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL -from mujoco.mjx.third_party.mujoco_warp._src.types import SPARSE_CONSTRAINT_JACOBIAN +from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType from mujoco.mjx.third_party.mujoco_warp._src.types import vec10 from mujoco.mjx.third_party.mujoco_warp._src.util_pkg import check_version -import numpy as np -import warp as wp def _create_array(data: Any, spec: wp.array, sizes: dict[str, int]) -> wp.array | None: @@ -114,9 +114,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: if unsupported: raise NotImplementedError(f"{mj_type(unsupported).name} is unsupported.") - if ((mjm.flex_contype != 0) | (mjm.flex_conaffinity != 0)).any(): - raise NotImplementedError("Flex collisions are not implemented.") - if mjm.opt.noslip_iterations > 0: raise NotImplementedError(f"noslip solver not implemented.") @@ -226,6 +223,8 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.is_sparse = is_sparse(mjm) m.has_fluid = mjm.opt.wind.any() or mjm.opt.density > 0 or mjm.opt.viscosity > 0 + m.max_ten_J_rownnz = int(mjm.ten_J_rownnz.max()) if mjm.ntendon else 0 + # body ids grouped by tree level (depth-based traversal) bodies, body_depth = {}, np.zeros(mjm.nbody, dtype=int) - 1 for i in range(mjm.nbody): @@ -364,6 +363,46 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: ) ) + # check for unsupported margin + multicontact / box-box CCD combinations + use_multiccd = mjm.opt.enableflags & types.EnableBit.MULTICCD + nativeccd_disabled = mjm.opt.disableflags & types.DisableBit.NATIVECCD + BOX = int(mujoco.mjtGeom.mjGEOM_BOX) + MESH = int(mujoco.mjtGeom.mjGEOM_MESH) + + has_boxbox = m.geom_pair_type_count[geom_trid_index(BOX, BOX)] > 0 + has_multiccd_pairs = has_boxbox or ( + use_multiccd + and (m.geom_pair_type_count[geom_trid_index(BOX, MESH)] > 0 or m.geom_pair_type_count[geom_trid_index(MESH, MESH)] > 0) + ) + + if has_multiccd_pairs: + + def _check_margin(name, t1, t2, margin): + if use_multiccd: + raise NotImplementedError( + f"{name} has non-zero margin ({margin}) with MULTICCD enabled. Set margin to 0 or disable MULTICCD." + ) + if t1 == BOX and t2 == BOX and not nativeccd_disabled: + raise NotImplementedError( + f"{name} has non-zero margin ({margin}) with NATIVECCD enabled. Set margin to 0 or disable NATIVECCD." + ) + + geom_name = lambda g: mujoco.mj_id2name(mjm, mujoco.mjtObj.mjOBJ_GEOM, g) or str(g) + + for idx in np.nonzero(nxn_include & (nxn_pairid_contact == -1))[0]: + g1, g2 = int(geom1[idx]), int(geom2[idx]) + t1, t2 = int(mjm.geom_type[g1]), int(mjm.geom_type[g2]) + m1, m2 = float(mjm.geom_margin[g1]), float(mjm.geom_margin[g2]) + if (m1 or m2) and t1 in (BOX, MESH) and t2 in (BOX, MESH): + _check_margin(f"geom pair ({geom_name(g1)}, {geom_name(g2)})", t1, t2, (m1, m2)) + + for pid in range(mjm.npair): + g1, g2 = int(mjm.pair_geom1[pid]), int(mjm.pair_geom2[pid]) + t1, t2 = int(mjm.geom_type[g1]), int(mjm.geom_type[g2]) + pm = float(mjm.pair_margin[pid]) + if pm and t1 in (BOX, MESH) and t2 in (BOX, MESH): + _check_margin(f"pair {pid} ({geom_name(g1)}, {geom_name(g2)})", t1, t2, pm) + m.nmaxpolygon = np.append(mjm.mesh_polyvertnum, 0).max() m.nmaxmeshdeg = np.append(mjm.mesh_polymapnum, 0).max() @@ -390,9 +429,11 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: current = [] else: current.append(v) - # Pad with zeros if less than 3 - attr_values += [0.0] * (3 - len(attr_values)) - m.plugin_attr.append(attr_values[:3]) + if len(attr_values) > types._NPLUGINATTR: + raise ValueError(f"Plugin has {len(attr_values)} attributes, which exceeds the maximum of {types._NPLUGINATTR}. ") + # pad with zeros to _NPLUGINATTR + attr_values += [0.0] * (types._NPLUGINATTR - len(attr_values)) + m.plugin_attr.append(attr_values[: types._NPLUGINATTR]) # equality constraint addresses m.eq_connect_adr = np.nonzero(mjm.eq_type == types.EqType.CONNECT)[0] @@ -542,6 +583,15 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: Madr_ki -= 1 m.qLD_updates = tuple(wp.array(qLD_updates[i], dtype=wp.vec3i) for i in sorted(qLD_updates)) + # Build concatenated updates for fused kernel + all_updates_flat = [] + level_offsets = [0] + for level in sorted(qLD_updates): + all_updates_flat.extend(qLD_updates[level]) + level_offsets.append(len(all_updates_flat)) + m.qLD_all_updates = all_updates_flat if all_updates_flat else [(0, 0, 0)] + m.qLD_level_offsets = level_offsets + # indices for sparse qM_fullm (used in solver) m.qM_fullm_i, m.qM_fullm_j = [], [] for i in range(mjm.nv): @@ -631,9 +681,168 @@ def _default_njmax(mjm: mujoco.MjModel, mjd: Optional[mujoco.MjData] = None) -> return int(valid_sizes[np.searchsorted(valid_sizes, njmax)]) -def _resolve_batch_size( - na: int | None, n: int | None, nworld: int, default: int -) -> int: +def _body_pair_nnz(mjm: mujoco.MjModel, body1: int, body2: int) -> int: + """Returns the number of unique DOFs in the kinematic tree union of two bodies.""" + body1 = mjm.body_weldid[body1] + body2 = mjm.body_weldid[body2] + da1 = mjm.body_dofadr[body1] + mjm.body_dofnum[body1] - 1 + da2 = mjm.body_dofadr[body2] + mjm.body_dofnum[body2] - 1 + nnz = 0 + while da1 >= 0 or da2 >= 0: + da = max(da1, da2) + if da1 == da: + da1 = mjm.dof_parentid[da1] + if da2 == da: + da2 = mjm.dof_parentid[da2] + nnz += 1 + return nnz + + +def _default_njmax_nnz(mjm: mujoco.MjModel, nconmax: int, njmax: int) -> int: + """Returns a heuristic estimate for the number of non-zeros in the sparse constraint Jacobian. + + Assumes all equality, friction, and limit constraints are active and computes + their non-zeros. For contacts, assumes njmax contact rows at the maximum + body-pair non-zeros from all enabled collision pairs. + + Args: + mjm: The model containing kinematic and dynamic information (host). + nconmax: Maximum number of contacts per world. + njmax: Maximum number of constraint rows per world. + + Returns: + Estimated number of non-zeros in the constraint Jacobian. + """ + total_nnz = 0 + + def _eq_bodies(i): + """Returns body pair for equality constraint i.""" + obj1id, obj2id = mjm.eq_obj1id[i], mjm.eq_obj2id[i] + if mjm.eq_objtype[i] == mujoco.mjtObj.mjOBJ_SITE: + return mjm.site_bodyid[obj1id], mjm.site_bodyid[obj2id] + return obj1id, obj2id + + # equality constraints (assume all active) + for i in range(mjm.neq): + eq_type = mjm.eq_type[i] + + if eq_type == mujoco.mjtEq.mjEQ_CONNECT: + total_nnz += 3 * _body_pair_nnz(mjm, *_eq_bodies(i)) + + elif eq_type == mujoco.mjtEq.mjEQ_WELD: + total_nnz += 6 * _body_pair_nnz(mjm, *_eq_bodies(i)) + + elif eq_type == mujoco.mjtEq.mjEQ_JOINT: + total_nnz += 2 if mjm.eq_obj2id[i] >= 0 else 1 + + elif eq_type == mujoco.mjtEq.mjEQ_TENDON: + obj1id = mjm.eq_obj1id[i] + obj2id = mjm.eq_obj2id[i] + rownnz1 = mjm.ten_J_rownnz[obj1id] if obj1id < mjm.ntendon else 0 + if obj2id >= 0 and obj2id < mjm.ntendon: + rowadr1 = mjm.ten_J_rowadr[obj1id] + rowadr2 = mjm.ten_J_rowadr[obj2id] + rownnz2 = mjm.ten_J_rownnz[obj2id] + cols = set() + for j in range(rownnz1): + cols.add(mjm.ten_J_colind[rowadr1 + j]) + for j in range(rownnz2): + cols.add(mjm.ten_J_colind[rowadr2 + j]) + total_nnz += len(cols) + else: + total_nnz += rownnz1 + + elif eq_type == mujoco.mjtEq.mjEQ_FLEX: + obj1id = mjm.eq_obj1id[i] + if obj1id < mjm.nflex: + edge_start = mjm.flex_edgeadr[obj1id] + edge_count = mjm.flex_edgenum[obj1id] + for e in range(edge_count): + total_nnz += mjm.flexedge_J_rownnz[edge_start + e] + + # friction constraints + total_nnz += (mjm.dof_frictionloss > 0).sum() + for i in range(mjm.ntendon): + if mjm.tendon_frictionloss[i] > 0: + total_nnz += mjm.ten_J_rownnz[i] + + # limit constraints (assume all active) + for i in range(mjm.njnt): + if mjm.jnt_limited[i]: + jnt_type = mjm.jnt_type[i] + if jnt_type == mujoco.mjtJoint.mjJNT_BALL: + total_nnz += 3 + elif jnt_type in (mujoco.mjtJoint.mjJNT_SLIDE, mujoco.mjtJoint.mjJNT_HINGE): + total_nnz += 1 + for i in range(mjm.ntendon): + if mjm.tendon_limited[i]: + total_nnz += mjm.ten_J_rownnz[i] + + # contact constraints: njmax rows at max body-pair non-zeros + max_contact_nnz = 0 + + # contact pairs + for i in range(mjm.npair): + g1, g2 = mjm.pair_geom1[i], mjm.pair_geom2[i] + b1, b2 = mjm.geom_bodyid[g1], mjm.geom_bodyid[g2] + max_contact_nnz = max(max_contact_nnz, _body_pair_nnz(mjm, b1, b2)) + + # filter geom-geom pairs (unique body pairs, filtered) + body_pair_seen = set() + for i in range(mjm.ngeom): + bi = mjm.geom_bodyid[i] + cti, cai = mjm.geom_contype[i], mjm.geom_conaffinity[i] + for j in range(i + 1, mjm.ngeom): + bj = mjm.geom_bodyid[j] + if bi == bj: + continue + if mjm.body_weldid[bi] == 0 and mjm.body_weldid[bj] == 0: + continue + bp = (min(bi, bj), max(bi, bj)) + if bp in body_pair_seen: + continue + ctj, caj = mjm.geom_contype[j], mjm.geom_conaffinity[j] + if not ((cti & caj) or (ctj & cai)): + continue + body_pair_seen.add(bp) + max_contact_nnz = max(max_contact_nnz, _body_pair_nnz(mjm, bi, bj)) + + # flex vertex contacts + for fi in range(mjm.nflex): + fct = mjm.flex_contype[fi] + fca = mjm.flex_conaffinity[fi] + + vert_start = mjm.flex_vertadr[fi] + vert_count = mjm.flex_vertnum[fi] + flex_bodies = {mjm.flex_vertbodyid[vert_start + v] for v in range(vert_count)} + + geom_bodies = set() + for g in range(mjm.ngeom): + ct, ca = mjm.geom_contype[g], mjm.geom_conaffinity[g] + if (fct & ca) or (ct & fca): + geom_bodies.add(mjm.geom_bodyid[g]) + + for fb in flex_bodies: + for gb in geom_bodies: + if fb != gb: + max_contact_nnz = max(max_contact_nnz, _body_pair_nnz(mjm, fb, gb)) + + # flex self-collision + if mjm.flex_selfcollide[fi]: + flex_body_list = sorted(flex_bodies) + for idx1 in range(len(flex_body_list)): + for idx2 in range(idx1 + 1, len(flex_body_list)): + max_contact_nnz = max( + max_contact_nnz, + _body_pair_nnz(mjm, flex_body_list[idx1], flex_body_list[idx2]), + ) + + total_nnz += njmax * max_contact_nnz + + return int(min(max(total_nnz, 1), njmax * mjm.nv)) + + +def _resolve_batch_size(na: int | None, n: int | None, nworld: int, default: int) -> int: if na is not None: return na if n is not None: @@ -647,6 +856,7 @@ def make_data( nconmax: Optional[int] = None, nccdmax: Optional[int] = None, njmax: Optional[int] = None, + njmax_nnz: Optional[int] = None, naconmax: Optional[int] = None, naccdmax: Optional[int] = None, ) -> types.Data: @@ -660,6 +870,7 @@ def make_data( nccdmax: Number of CCD contacts to allocate per world. Same semantics as nconmax. njmax: Number of constraints to allocate per world. Constraint arrays are batched by world: no world may have more than njmax constraints. + njmax_nnz: Number of non-zeros in constraint Jacobian (sparse). Defaults to njmax * nv. naconmax: Number of contacts to allocate for all worlds. Overrides nconmax. naccdmax: Maximum number of CCD contacts. Defaults to naconmax. @@ -709,20 +920,32 @@ def make_data( sizes["naconmax"] = naconmax sizes["njmax"] = njmax + if njmax_nnz is None: + if is_sparse(mjm): + njmax_nnz = _default_njmax_nnz(mjm, nconmax, njmax) + else: + njmax_nnz = njmax * mjm.nv + contact = types.Contact(**{f.name: _create_array(None, f.type, sizes) for f in dataclasses.fields(types.Contact)}) + contact.efc_address = wp.array(np.full((naconmax, sizes["nmaxpyramid"]), -1, dtype=int), dtype=int) efc = types.Constraint(**{f.name: _create_array(None, f.type, sizes) for f in dataclasses.fields(types.Constraint)}) - if SPARSE_CONSTRAINT_JACOBIAN: + if is_sparse(mjm): efc.J_rownnz = wp.zeros((nworld, njmax), dtype=int) efc.J_rowadr = wp.zeros((nworld, njmax), dtype=int) - efc.J_colind = wp.zeros((nworld, 1, njmax * mjm.nv), dtype=int) - efc.J = wp.zeros((nworld, 1, njmax * mjm.nv), dtype=float) + efc.J_colind = wp.zeros((nworld, 1, njmax_nnz), dtype=int) + efc.J = wp.zeros((nworld, 1, njmax_nnz), dtype=float) else: efc.J_rownnz = wp.zeros((nworld, 0), dtype=int) efc.J_rowadr = wp.zeros((nworld, 0), dtype=int) efc.J_colind = wp.zeros((nworld, 0, 0), dtype=int) efc.J = wp.zeros((nworld, sizes["njmax_pad"], sizes["nv_pad"]), dtype=float) + contact_kwargs = {} + for f in dataclasses.fields(types.Contact): + contact_kwargs[f.name] = _create_array(None, f.type, sizes) + contact = types.Contact(**contact_kwargs) + # world body and static geom (attached to the world) poses are precomputed # this speeds up scenes with many static geoms (e.g. terrains) # TODO(team): remove this when we introduce dof islands + sleeping @@ -734,65 +957,34 @@ def make_data( mocap_id = mjm.body_mocapid[mocap_body] d_kwargs = { - "qpos": wp.array( - np.tile(mjm.qpos0, nworld), shape=(nworld, mjm.nq), dtype=float - ), - "contact": contact, - "efc": efc, - "nworld": nworld, - "naconmax": naconmax, - "naccdmax": naccdmax, - "njmax": njmax, - "njmax_pad": sizes["njmax_pad"], - "qM": None, - "qLD": None, - # world body - "xquat": wp.array( - np.tile(mjd.xquat, (nworld, 1)), - shape=(nworld, mjm.nbody), - dtype=wp.quat, - ), - "xmat": wp.array( - np.tile(mjd.xmat, (nworld, 1)), - shape=(nworld, mjm.nbody), - dtype=wp.mat33, - ), - "ximat": wp.array( - np.tile(mjd.ximat, (nworld, 1)), - shape=(nworld, mjm.nbody), - dtype=wp.mat33, - ), - # static geoms - "geom_xpos": wp.array( - np.tile(mjd.geom_xpos, (nworld, 1)), - shape=(nworld, mjm.ngeom), - dtype=wp.vec3, - ), - "geom_xmat": wp.array( - np.tile(mjd.geom_xmat, (nworld, 1)), - shape=(nworld, mjm.ngeom), - dtype=wp.mat33, - ), - # mocap - "mocap_pos": wp.array( - np.tile(mjm.body_pos[mocap_body[mocap_id]], (nworld, 1)), - shape=(nworld, mjm.nmocap), - dtype=wp.vec3, - ), - "mocap_quat": wp.array( - np.tile(mjm.body_quat[mocap_body[mocap_id]], (nworld, 1)), - shape=(nworld, mjm.nmocap), - dtype=wp.quat, - ), - # equality constraints - "eq_active": wp.array( - np.tile(mjm.eq_active0.astype(bool), (nworld, 1)), - shape=(nworld, mjm.neq), - dtype=bool, - ), - # island arrays - "nisland": None, - "tree_island": None, + "qpos": wp.array(np.tile(mjm.qpos0, nworld), shape=(nworld, mjm.nq), dtype=float), + "contact": contact, + "efc": efc, + "nworld": nworld, + "naconmax": naconmax, + "naccdmax": naccdmax, + "njmax": njmax, + "njmax_pad": sizes["njmax_pad"], + "njmax_nnz": njmax_nnz, + "qM": None, + "qLD": None, + # world body + "xquat": wp.array(np.tile(mjd.xquat, (nworld, 1)), shape=(nworld, mjm.nbody), dtype=wp.quat), + "xmat": wp.array(np.tile(mjd.xmat, (nworld, 1)), shape=(nworld, mjm.nbody), dtype=wp.mat33), + "ximat": wp.array(np.tile(mjd.ximat, (nworld, 1)), shape=(nworld, mjm.nbody), dtype=wp.mat33), + # static geoms + "geom_xpos": wp.array(np.tile(mjd.geom_xpos, (nworld, 1)), shape=(nworld, mjm.ngeom), dtype=wp.vec3), + "geom_xmat": wp.array(np.tile(mjd.geom_xmat, (nworld, 1)), shape=(nworld, mjm.ngeom), dtype=wp.mat33), + # mocap + "mocap_pos": wp.array(np.tile(mjm.body_pos[mocap_body[mocap_id]], (nworld, 1)), shape=(nworld, mjm.nmocap), dtype=wp.vec3), + "mocap_quat": wp.array( + np.tile(mjm.body_quat[mocap_body[mocap_id]], (nworld, 1)), shape=(nworld, mjm.nmocap), dtype=wp.quat + ), + # equality constraints + "eq_active": wp.array(np.tile(mjm.eq_active0.astype(bool), (nworld, 1)), shape=(nworld, mjm.neq), dtype=bool), + # island arrays + "nisland": None, + "tree_island": None, } for f in dataclasses.fields(types.Data): if f.name in d_kwargs: @@ -822,6 +1014,7 @@ def put_data( nconmax: Optional[int] = None, nccdmax: Optional[int] = None, njmax: Optional[int] = None, + njmax_nnz: Optional[int] = None, naconmax: Optional[int] = None, naccdmax: Optional[int] = None, ) -> types.Data: @@ -836,6 +1029,7 @@ def put_data( nccdmax: Number of CCD contacts to allocate per world. Same semantics as nconmax. njmax: Number of constraints to allocate per world. Constraint arrays are batched by world: no world may have more than njmax constraints. + njmax_nnz: Number of non-zeros in constraint Jacobian (sparse). Defaults to njmax * nv. naconmax: Number of contacts to allocate for all worlds. Overrides nconmax. naccdmax: Maximum number of CCD contacts. Defaults to naconmax. @@ -898,6 +1092,12 @@ def put_data( sizes["naconmax"] = naconmax sizes["njmax"] = njmax + if njmax_nnz is None: + if is_sparse(mjm): + njmax_nnz = _default_njmax_nnz(mjm, nconmax, njmax) + else: + njmax_nnz = njmax * mjm.nv + # ensure static geom positions are computed # TODO: remove once MjData creation semantics are fixed mujoco.mj_kinematics(mjm, mjd) @@ -915,7 +1115,7 @@ def put_data( contact = types.Contact(**contact_kwargs) - contact.efc_address = np.zeros((naconmax, sizes["nmaxpyramid"]), dtype=int) + contact.efc_address = np.full((naconmax, sizes["nmaxpyramid"]), -1, dtype=int) for i in range(mjd.ncon): efc_address = mjd.contact.efc_address[i] if efc_address == -1: @@ -945,43 +1145,28 @@ def put_data( efc = types.Constraint(**efc_kwargs) - if SPARSE_CONSTRAINT_JACOBIAN: - # TODO(team): process efc_J sparsity structure for nv row shift - efc.J_rownnz = wp.array( - np.full((nworld, njmax), mjm.nv, dtype=int), dtype=int - ) - efc.J_rowadr = wp.array( - np.tile( - np.arange(0, njmax * mjm.nv, mjm.nv) - if mjm.nv - else np.zeros(njmax, dtype=int), - (nworld, 1), - ), - dtype=int, - ) - efc.J_colind = wp.array( - np.tile(np.arange(mjm.nv), (nworld, njmax)).reshape((nworld, 1, -1)), - dtype=int, - ) - - mj_efc_J = np.zeros((mjd.nefc, mjm.nv)) + if is_sparse(mjm): + J_rownnz = np.zeros(njmax, dtype=np.int32) + J_rowadr = np.zeros(njmax, dtype=np.int32) + J_colind = np.zeros(njmax_nnz, dtype=np.int32) + J = np.zeros(njmax_nnz, dtype=np.float64) if mjd.nefc: if mujoco.mj_isSparse(mjm): - mujoco.mju_sparse2dense( - mj_efc_J, - mjd.efc_J, - mjd.efc_J_rownnz, - mjd.efc_J_rowadr, - mjd.efc_J_colind, - ) + J_rownnz[: mjd.nefc] = mjd.efc_J_rownnz[: mjd.nefc] + J_rowadr[: mjd.nefc] = mjd.efc_J_rowadr[: mjd.nefc] + nnz = int(mjd.efc_J_rownnz[: mjd.nefc].sum()) + J_colind[:nnz] = mjd.efc_J_colind[:nnz] + J[:nnz] = mjd.efc_J[:nnz] else: - mj_efc_J = mjd.efc_J.reshape((mjd.nefc, mjm.nv)) - efc_J = np.zeros((njmax, mjm.nv), dtype=float) - efc_J[: mjd.nefc, : mjm.nv] = mj_efc_J - efc.J = wp.array( - np.tile(efc_J.reshape(-1), (nworld, 1, 1)).reshape((nworld, 1, -1)), - dtype=float, - ) + dense_J = mjd.efc_J.reshape((-1, mjm.nv))[: mjd.nefc] + mujoco.mju_dense2sparse( + J[: mjd.nefc * mjm.nv], dense_J, J_rownnz[: mjd.nefc], J_rowadr[: mjd.nefc], J_colind[: mjd.nefc * mjm.nv] + ) + + efc.J_rownnz = wp.array(np.tile(J_rownnz, (nworld, 1)), dtype=int) + efc.J_rowadr = wp.array(np.tile(J_rowadr, (nworld, 1)), dtype=int) + efc.J_colind = wp.array(np.tile(J_colind, (nworld, 1)).reshape((nworld, 1, -1)), dtype=int) + efc.J = wp.array(np.tile(J, (nworld, 1)).reshape((nworld, 1, -1)), dtype=float) else: efc.J_rownnz = wp.zeros((nworld, 0), dtype=int) efc.J_rowadr = wp.zeros((nworld, 0), dtype=int) @@ -990,13 +1175,7 @@ def put_data( mj_efc_J = np.zeros((mjd.nefc, mjm.nv)) if mjd.nefc: if mujoco.mj_isSparse(mjm): - mujoco.mju_sparse2dense( - mj_efc_J, - mjd.efc_J, - mjd.efc_J_rownnz, - mjd.efc_J_rowadr, - mjd.efc_J_colind, - ) + mujoco.mju_sparse2dense(mj_efc_J, mjd.efc_J, mjd.efc_J_rownnz, mjd.efc_J_rowadr, mjd.efc_J_colind) else: mj_efc_J = mjd.efc_J.reshape((mjd.nefc, mjm.nv)) efc_J = np.zeros((nworld, sizes["njmax_pad"], sizes["nv_pad"]), dtype=float) @@ -1005,22 +1184,22 @@ def put_data( # create data d_kwargs = { - "contact": contact, - "efc": efc, - "nworld": nworld, - "naconmax": naconmax, - "naccdmax": naccdmax, - "njmax": njmax, - "njmax_pad": sizes["njmax_pad"], - # fields set after initialization: - "solver_niter": None, - "qM": None, - "qLD": None, - "ten_J": None, - "nacon": None, - # island arrays - "nisland": None, - "tree_island": None, + "contact": contact, + "efc": efc, + "nworld": nworld, + "naconmax": naconmax, + "naccdmax": naccdmax, + "njmax": njmax, + "njmax_pad": sizes["njmax_pad"], + "njmax_nnz": njmax_nnz, + # fields set after initialization: + "solver_niter": None, + "qM": None, + "qLD": None, + "nacon": None, + # island arrays + "nisland": None, + "tree_island": None, } for f in dataclasses.fields(types.Data): if f.name in d_kwargs: @@ -1050,29 +1229,6 @@ def put_data( d.nisland = wp.array(np.full(nworld, mjd.nisland), dtype=int) d.tree_island = wp.array(np.tile(mjd.tree_island, (nworld, 1)), dtype=int) - ten_J = np.zeros((mjm.ntendon, mjm.nv)) - if mujoco.mj_isSparse(mjm) or check_version("mujoco>=3.5.1.dev872479828"): - if mjm.ntendon: - if check_version("mujoco>=3.5.1.dev875093374"): - mujoco.mju_sparse2dense( - ten_J, - mjd.ten_J.reshape(-1), - mjm.ten_J_rownnz, - mjm.ten_J_rowadr, - mjm.ten_J_colind.reshape(-1), - ) - else: - mujoco.mju_sparse2dense( - ten_J, - mjd.ten_J.reshape(-1), - mjd.ten_J_rownnz, - mjd.ten_J_rowadr, - mjd.ten_J_colind.reshape(-1), - ) - else: - ten_J = mjd.ten_J.reshape((mjm.ntendon, mjm.nv)) - d.ten_J = wp.array(np.full((nworld, mjm.ntendon, mjm.nv), ten_J), dtype=float) - d.nacon = wp.array([mjd.ncon * nworld], dtype=int) return d @@ -1233,14 +1389,14 @@ def get_data_into( mujoco.mj_factorM(mjm, result) if nefc > 0: - if SPARSE_CONSTRAINT_JACOBIAN: + if is_sparse(mjm): efc_J = np.zeros((nefc, mjm.nv)) mujoco.mju_sparse2dense( - efc_J, - d.efc.J.numpy()[world_id, 0], - d.efc.J_rownnz.numpy()[world_id, :nefc], - d.efc.J_rowadr.numpy()[world_id, :nefc], - d.efc.J_colind.numpy()[world_id, 0], + efc_J, + d.efc.J.numpy()[world_id, 0], + d.efc.J_rownnz.numpy()[world_id, :nefc], + d.efc.J_rowadr.numpy()[world_id, :nefc], + d.efc.J_colind.numpy()[world_id, 0], ) else: efc_J = d.efc.J.numpy()[world_id, :nefc, : mjm.nv] @@ -1248,11 +1404,11 @@ def get_data_into( # write to mujoco result (format depends on mj_isSparse) if mujoco.mj_isSparse(mjm): mujoco.mju_dense2sparse( - result.efc_J, - efc_J[efc_idx], - result.efc_J_rownnz, - result.efc_J_rowadr, - result.efc_J_colind, + result.efc_J, + efc_J[efc_idx], + result.efc_J_rownnz, + result.efc_J_rowadr, + result.efc_J_colind, ) else: result.efc_J[: nefc * mjm.nv] = efc_J[efc_idx].flatten() @@ -1276,24 +1432,7 @@ def get_data_into( # tendon result.ten_length[:] = d.ten_length.numpy()[world_id] - if check_version("mujoco>=3.5.1.dev869712136"): - ten_J = d.ten_J.numpy()[world_id] - if check_version("mujoco>=3.5.1.dev875093374"): - ten_J_rownnz = mjm.ten_J_rownnz - ten_J_rowadr = mjm.ten_J_rowadr - ten_J_colind = mjm.ten_J_colind.reshape(-1) - else: - ten_J_rownnz = result.ten_J_rownnz - ten_J_rowadr = result.ten_J_rowadr - ten_J_colind = result.ten_J_colind.reshape(-1) - mujoco.mju_dense2sparse( - result.ten_J, - ten_J, - ten_J_rownnz, - ten_J_rowadr, - ten_J_colind, - ) - else: + if mjm.ntendon > 0: result.ten_J[:] = d.ten_J.numpy()[world_id] result.ten_wrapadr[:] = d.ten_wrapadr.numpy()[world_id] result.ten_wrapnum[:] = d.ten_wrapnum.numpy()[world_id] @@ -1428,12 +1567,8 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): mocapid = body_mocapid[bodyid] if mocapid >= 0: - mocap_pos_out[worldid, mocapid] = body_pos[ - worldid % body_pos.shape[0], bodyid - ] - mocap_quat_out[worldid, mocapid] = body_quat[ - worldid % body_quat.shape[0], bodyid - ] + mocap_pos_out[worldid, mocapid] = body_pos[worldid % body_pos.shape[0], bodyid] + mocap_quat_out[worldid, mocapid] = body_quat[worldid % body_quat.shape[0], bodyid] @wp.kernel(module="unique", enable_backward=False) def reset_contact( @@ -1453,6 +1588,8 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): contact_solimp_out: wp.array(dtype=types.vec5), contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), + contact_flex_out: wp.array(dtype=wp.vec2i), + contact_vert_out: wp.array(dtype=wp.vec2i), contact_efc_address_out: wp.array2d(dtype=int), contact_worldid_out: wp.array(dtype=int), contact_type_out: wp.array(dtype=int), @@ -1479,8 +1616,10 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): contact_solimp_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0) contact_dim_out[conid] = 0 contact_geom_out[conid] = wp.vec2i(0, 0) + contact_flex_out[conid] = wp.vec2i(0, 0) + contact_vert_out[conid] = wp.vec2i(0, 0) for i in range(nefcaddress): - contact_efc_address_out[conid, i] = 0 + contact_efc_address_out[conid, i] = -1 contact_worldid_out[conid] = 0 contact_type_out[conid] = 0 contact_geomcollisionid_out[conid] = 0 @@ -1519,6 +1658,8 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): d.contact.solimp, d.contact.dim, d.contact.geom, + d.contact.flex, + d.contact.vert, d.contact.efc_address, d.contact.worldid, d.contact.type, @@ -1812,28 +1953,45 @@ def _finalize_body_invweight0( @wp.kernel def _copy_tendon_jacobian( tenid_target: int, - ten_J_in: wp.array3d(dtype=float), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + ten_J_in: wp.array2d(dtype=float), ten_J_vec_out: wp.array2d(dtype=float), ): worldid = wp.tid() nv = ten_J_in.shape[2] - for i in range(nv): - ten_J_vec_out[worldid, i] = ten_J_in[worldid, tenid_target, i] + rownnz = ten_J_rownnz[tenid_target] + rowadr = ten_J_rowadr[tenid_target] + for i in range(rownnz): + colind = ten_J_colind[rowadr + i] + ten_J_vec_out[worldid, colind] = ten_J_in[worldid, rowadr + i] @wp.kernel def _compute_tendon_dot_product( + # Model: + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + # In: tenid_target: int, - nv: int, - ten_J_in: wp.array3d(dtype=float), + ten_J_in: wp.array2d(dtype=float), result_vec_in: wp.array2d(dtype=float), + # Out: tendon_invweight0_out: wp.array2d(dtype=float), ): worldid = wp.tid() tendon_invweight0_id = worldid % tendon_invweight0_out.shape[0] dot_prod = float(0.0) - for i in range(nv): - dot_prod += ten_J_in[worldid, tenid_target, i] * result_vec_in[worldid, i] + + rownnz = ten_J_rownnz[tenid_target] + rowadr = ten_J_rowadr[tenid_target] + for i in range(rownnz): + sparseid = rowadr + i + colind = ten_J_colind[sparseid] + dot_prod += ten_J_in[worldid, sparseid] * result_vec_in[worldid, colind] + tendon_invweight0_out[tendon_invweight0_id, tenid_target] = dot_prod @@ -1891,12 +2049,12 @@ def _compute_light_pos0( @wp.kernel def _copy_actuator_moment( - actid_target: int, - moment_rownnz_in: wp.array2d(dtype=int), - moment_rowadr_in: wp.array2d(dtype=int), - moment_colind_in: wp.array2d(dtype=int), - actuator_moment_in: wp.array2d(dtype=float), - act_moment_vec_out: wp.array2d(dtype=float), + actid_target: int, + moment_rownnz_in: wp.array2d(dtype=int), + moment_rowadr_in: wp.array2d(dtype=int), + moment_colind_in: wp.array2d(dtype=int), + actuator_moment_in: wp.array2d(dtype=float), + act_moment_vec_out: wp.array2d(dtype=float), ): worldid = wp.tid() nv = act_moment_vec_out.shape[1] @@ -1912,10 +2070,10 @@ def _copy_actuator_moment( @wp.kernel def _compute_actuator_acc0( - actid_target: int, - nv: int, - result_vec_in: wp.array2d(dtype=float), - actuator_acc0_out: wp.array2d(dtype=float), + actid_target: int, + nv: int, + result_vec_in: wp.array2d(dtype=float), + actuator_acc0_out: wp.array2d(dtype=float), ): worldid = wp.tid() norm_sq = float(0.0) @@ -1926,11 +2084,11 @@ def _compute_actuator_acc0( @wp.kernel def _compute_dof_M0( - dof_bodyid: wp.array(dtype=int), - dof_armature: wp.array2d(dtype=float), - cdof_in: wp.array2d(dtype=wp.spatial_vector), - crb_in: wp.array2d(dtype=vec10), - dof_M0_out: wp.array2d(dtype=float), + dof_bodyid: wp.array(dtype=int), + dof_armature: wp.array2d(dtype=float), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + crb_in: wp.array2d(dtype=vec10), + dof_M0_out: wp.array2d(dtype=float), ): worldid, dofid = wp.tid() bodyid = dof_bodyid[dofid] @@ -1941,15 +2099,15 @@ def _compute_dof_M0( @wp.kernel def _resolve_dampratio( - actuator_biastype: wp.array(dtype=int), - actuator_gainprm: wp.array2d(dtype=types.vec10f), - moment_rownnz_in: wp.array2d(dtype=int), - moment_rowadr_in: wp.array2d(dtype=int), - moment_colind_in: wp.array2d(dtype=int), - actuator_moment_in: wp.array2d(dtype=float), - dof_M0_in: wp.array2d(dtype=float), - nv: int, - actuator_biasprm: wp.array2d(dtype=types.vec10f), + actuator_biastype: wp.array(dtype=int), + actuator_gainprm: wp.array2d(dtype=types.vec10f), + moment_rownnz_in: wp.array2d(dtype=int), + moment_rowadr_in: wp.array2d(dtype=int), + moment_colind_in: wp.array2d(dtype=int), + actuator_moment_in: wp.array2d(dtype=float), + dof_M0_in: wp.array2d(dtype=float), + nv: int, + actuator_biasprm: wp.array2d(dtype=types.vec10f), ): worldid, actid = wp.tid() biastype = actuator_biastype[actid] @@ -1992,15 +2150,15 @@ def _resolve_dampratio( @wp.kernel def _set_length_range( - actuator_trntype: wp.array(dtype=int), - actuator_trnid: wp.array(dtype=wp.vec2i), - actuator_gear: wp.array2d(dtype=wp.spatial_vector), - jnt_limited: wp.array(dtype=int), - jnt_range: wp.array2d(dtype=wp.vec2), - tendon_limited: wp.array(dtype=int), - tendon_range: wp.array2d(dtype=wp.vec2), - ntendon: int, - actuator_lengthrange_out: wp.array2d(dtype=wp.vec2), + actuator_trntype: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + actuator_gear: wp.array2d(dtype=wp.spatial_vector), + jnt_limited: wp.array(dtype=int), + jnt_range: wp.array2d(dtype=wp.vec2), + tendon_limited: wp.array(dtype=int), + tendon_range: wp.array2d(dtype=wp.vec2), + ntendon: int, + actuator_lengthrange_out: wp.array2d(dtype=wp.vec2), ): worldid, actid = wp.tid() trntype = actuator_trntype[actid] @@ -2167,16 +2325,22 @@ def set_const_0(m: types.Model, d: types.Data): # tendon_invweight0[t] = J_t * inv(M) * J_t' if m.ntendon > 0: - ten_J_vec = wp.zeros((d.nworld, m.nv), dtype=float) - ten_result_vec = wp.zeros((d.nworld, m.nv), dtype=float) + ten_J_vec = wp.empty((d.nworld, m.nv), dtype=float) + ten_result_vec = wp.empty((d.nworld, m.nv), dtype=float) for tenid in range(m.ntendon): - wp.launch(_copy_tendon_jacobian, dim=d.nworld, inputs=[tenid, d.ten_J], outputs=[ten_J_vec]) + ten_J_vec.zero_() + wp.launch( + _copy_tendon_jacobian, + dim=d.nworld, + inputs=[tenid, m.ten_J_rownnz, m.ten_J_rowadr, m.ten_J_colind, d.ten_J], + outputs=[ten_J_vec], + ) smooth.solve_m(m, d, ten_result_vec, ten_J_vec) wp.launch( _compute_tendon_dot_product, dim=d.nworld, - inputs=[tenid, m.nv, d.ten_J, ten_result_vec], + inputs=[m.ten_J_rownnz, m.ten_J_rowadr, m.ten_J_colind, tenid, d.ten_J, ten_result_vec], outputs=[m.tendon_invweight0], ) @@ -2201,16 +2365,10 @@ def set_const_0(m: types.Model, d: types.Data): for actid in range(m.nu): wp.launch( - _copy_actuator_moment, - dim=d.nworld, - inputs=[ - actid, - d.moment_rownnz, - d.moment_rowadr, - d.moment_colind, - d.actuator_moment, - ], - outputs=[act_moment_vec], + _copy_actuator_moment, + dim=d.nworld, + inputs=[actid, d.moment_rownnz, d.moment_rowadr, d.moment_colind, d.actuator_moment], + outputs=[act_moment_vec], ) smooth.solve_m(m, d, act_result_vec, act_moment_vec) wp.launch(_compute_actuator_acc0, dim=d.nworld, inputs=[actid, m.nv, act_result_vec], outputs=[m.actuator_acc0]) @@ -2219,25 +2377,25 @@ def set_const_0(m: types.Model, d: types.Data): if m.nu > 0 and m.nv > 0: dof_M0 = wp.zeros((d.nworld, m.nv), dtype=float) wp.launch( - _compute_dof_M0, - dim=(d.nworld, m.nv), - inputs=[m.dof_bodyid, m.dof_armature, d.cdof, d.crb], - outputs=[dof_M0], + _compute_dof_M0, + dim=(d.nworld, m.nv), + inputs=[m.dof_bodyid, m.dof_armature, d.cdof, d.crb], + outputs=[dof_M0], ) wp.launch( - _resolve_dampratio, - dim=(d.nworld, m.nu), - inputs=[ - m.actuator_biastype, - m.actuator_gainprm, - d.moment_rownnz, - d.moment_rowadr, - d.moment_colind, - d.actuator_moment, - dof_M0, - m.nv, - ], - outputs=[m.actuator_biasprm], + _resolve_dampratio, + dim=(d.nworld, m.nu), + inputs=[ + m.actuator_biastype, + m.actuator_gainprm, + d.moment_rownnz, + d.moment_rowadr, + d.moment_colind, + d.actuator_moment, + dof_M0, + m.nv, + ], + outputs=[m.actuator_biasprm], ) wp.copy(d.qpos, qpos_saved) @@ -2255,16 +2413,12 @@ def set_const(m: types.Model, d: types.Data): Field | Notes ---------------------------------|---------------------------------------------- qpos0, qpos_spring | - body_mass, body_inertia, | Mass and inertia are usually scaled - together + body_mass, body_inertia, | Mass and inertia are usually scaled together body_ipos, body_iquat | since inertia is sum(m * r^2). - body_pos, body_quat | Unsafe for static bodies (invalidates - BVH). - body_gravcomp | If changing from 0 to >0 bodies, - required. + body_pos, body_quat | Unsafe for static bodies (invalidates BVH). + body_gravcomp | If changing from 0 to >0 bodies, required. dof_armature | - eq_data | For connect/weld, offsets computed if not - set. + eq_data | For connect/weld, offsets computed if not set. hfield_size | tendon_stiffness, tendon_damping | Only if changing from/to zero. actuator_gainprm, actuator_biasprm | For position actuators with dampratio. @@ -2319,19 +2473,19 @@ def set_length_range(m: types.Model, d: types.Data, index: int = -1): return wp.launch( - _set_length_range, - dim=(d.nworld, m.nu), - inputs=[ - m.actuator_trntype, - m.actuator_trnid, - m.actuator_gear, - m.jnt_limited, - m.jnt_range, - m.tendon_limited, - m.tendon_range, - m.ntendon, - ], - outputs=[m.actuator_lengthrange], + _set_length_range, + dim=(d.nworld, m.nu), + inputs=[ + m.actuator_trntype, + m.actuator_trnid, + m.actuator_gear, + m.jnt_limited, + m.jnt_range, + m.tendon_limited, + m.tendon_range, + m.ntendon, + ], + outputs=[m.actuator_lengthrange], ) @@ -2486,6 +2640,7 @@ def create_render_context( cam_res: list[tuple[int, int]] | tuple[int, int] | None = None, render_rgb: list[bool] | bool | None = None, render_depth: list[bool] | bool | None = None, + render_seg: list[bool] | bool | None = None, use_textures: bool = True, use_shadows: bool = False, enabled_geom_groups: list[int] = [0, 1, 2], @@ -2502,6 +2657,8 @@ def create_render_context( MuJoCo model values. render_rgb: Whether to render RGB images. If None, uses the MuJoCo model values. render_depth: Whether to render depth images. If None, uses the MuJoCo model values. + render_seg: Whether to render segmentation (per-pixel geom IDs). If None, + uses the MuJoCo model values. use_textures: Whether to use textures. use_shadows: Whether to use shadows. enabled_geom_groups: The geom groups to render. @@ -2517,10 +2674,13 @@ def create_render_context( mjd = mujoco.MjData(mjm) mujoco.mj_forward(mjm, mjd) - # TODO(team): remove after mjwarp depends on warp-lang >= 1.12 in pyproject.toml - if use_textures and not hasattr(wp, "Texture2D"): - warnings.warn("Textures require warp >= 1.12. Disabling textures.") - use_textures = False + constructor = "sah" + if check_version("warp>=1.13.0.dev20260325"): + # TODO: The cubql constructor and is_cubql_available exist only in + # recent Warp 1.13+ builds, modify this after warp is updated to 1.13+. + _cubql_avail = getattr(wp, "is_cubql_available", None) + if callable(_cubql_avail) and _cubql_avail(): + constructor = "cubql" # Mesh BVHs nmesh = mjm.nmesh @@ -2534,7 +2694,7 @@ def create_render_context( mesh_bounds_size = [wp.vec3(0.0, 0.0, 0.0) for _ in range(nmesh)] for mid in used_mesh_id: - mesh, half = bvh.build_mesh_bvh(mjm, mid) + mesh, half = bvh.build_mesh_bvh(mjm, mid, constructor=constructor) mesh_registry[mesh.id] = mesh mesh_bvh_id[mid] = mesh.id mesh_bounds_size[mid] = half @@ -2551,7 +2711,7 @@ def create_render_context( hfield_bounds_size = [wp.vec3(0.0, 0.0, 0.0) for _ in range(nhfield)] for hid in used_hfield_id: - hmesh, hhalf = bvh.build_hfield_bvh(mjm, hid) + hmesh, hhalf = bvh.build_hfield_bvh(mjm, hid, constructor=constructor) hfield_registry[hmesh.id] = hmesh hfield_bvh_id[hid] = hmesh.id hfield_bounds_size[hid] = hhalf @@ -2560,65 +2720,33 @@ def create_render_context( hfield_bounds_size_arr = wp.array(hfield_bounds_size, dtype=wp.vec3) # Flex BVHs - flex_bvh_id = wp.uint64(0) - flex_group_root = wp.zeros(nworld, dtype=int) - flex_mesh = None - flex_face_point = None - flex_elemdataadr = None - flex_shell = None - flex_shelldataadr = None - flex_faceadr = None - flex_nface = 0 - flex_radius = None - flex_workadr = None - flex_worknum = None - flex_nwork = 0 + nflex = mjm.nflex + flex_registry = {} - if mjm.nflex > 0: - ( - fmesh, - face_point, - flex_group_roots, - flex_shell_data, - flex_faceadr_data, - flex_nface, - ) = bvh.build_flex_bvh(mjm, mjd, nworld) + # Scene BVH flex primitives: 1D → one capsule per edge, 2D/3D → one box per flex + flex_geom_flexid = [] + flex_geom_edgeid = [] + flex_bvh_id = np.full(nflex, 0, dtype=wp.uint64) + flex_group_root = np.zeros((nflex, nworld), dtype=int) - flex_mesh = fmesh - flex_bvh_id = fmesh.id - flex_face_point = face_point - flex_group_root = flex_group_roots - flex_elemdataadr = wp.array(mjm.flex_elemdataadr, dtype=int) - flex_shell = flex_shell_data - flex_shelldataadr = wp.array(mjm.flex_shelldataadr, dtype=int) - flex_faceadr = wp.array(flex_faceadr_data, dtype=int) - flex_radius = wp.array(mjm.flex_radius, dtype=float) - - # precompute work item layout for unified refit kernel - nflex = mjm.nflex - workadr = np.zeros(nflex, dtype=np.int32) - worknum = np.zeros(nflex, dtype=np.int32) - cumsum = 0 - for f in range(nflex): - workadr[f] = cumsum - if mjm.flex_dim[f] == 2: - worknum[f] = mjm.flex_elemnum[f] + mjm.flex_shellnum[f] - else: - worknum[f] = mjm.flex_shellnum[f] - cumsum += worknum[f] - flex_workadr = wp.array(workadr, dtype=int) - flex_worknum = wp.array(worknum, dtype=int) - flex_nwork = int(cumsum) + for f in range(nflex): + if mjm.flex_dim[f] == 1: + edge_adr = mjm.flex_edgeadr[f] + flex_geom_flexid.extend([f] * mjm.flex_edgenum[f]) + flex_geom_edgeid.extend([edge_adr + e for e in range(mjm.flex_edgenum[f])]) + flex_group_root[f] = np.zeros(nworld, dtype=int) + else: + flex_geom_flexid.append(f) + flex_geom_edgeid.append(-1) + fmesh, group_root = bvh.build_flex_bvh(mjm, mjd, nworld, f) + flex_registry[f] = fmesh + flex_bvh_id[f] = fmesh.id + flex_group_root[f] = group_root.numpy() textures_registry = [] - # TODO: remove after mjwarp depends on warp-lang >= 1.12 in pyproject.toml - if hasattr(wp, "Texture2D"): - for i in range(mjm.ntex): - textures_registry.append(render_util.create_warp_texture(mjm, i)) - textures = wp.array(textures_registry, dtype=wp.Texture2D) - else: - # Dummy array when texture support isn't available (warp < 1.12) - textures = wp.zeros(1, dtype=int) + for i in range(mjm.ntex): + textures_registry.append(render_util.create_warp_texture(mjm, i)) + textures = wp.array(textures_registry, dtype=wp.Texture2D) # Filter active cameras if cam_active is not None: @@ -2642,31 +2770,31 @@ def create_render_context( cam_res_arr = wp.array(active_cam_res, dtype=wp.vec2i) if render_rgb is None: - render_rgb = [ - mjm.cam_output[i] & mujoco.mjtCamOutBit.mjCAMOUT_RGB - for i in active_cam_indices - ] + render_rgb = [mjm.cam_output[i] & mujoco.mjtCamOutBit.mjCAMOUT_RGB for i in active_cam_indices] elif isinstance(render_rgb, bool): render_rgb = [render_rgb] * ncam if render_depth is None: - render_depth = [ - mjm.cam_output[i] & mujoco.mjtCamOutBit.mjCAMOUT_DEPTH - for i in active_cam_indices - ] + render_depth = [mjm.cam_output[i] & mujoco.mjtCamOutBit.mjCAMOUT_DEPTH for i in active_cam_indices] if isinstance(render_depth, bool): render_depth = [render_depth] * ncam - assert len(render_rgb) == ncam and len(render_depth) == ncam, ( - "render_rgb and render_depth must be a bool or a list of bools with" - f" length {ncam}" + if render_seg is None: + render_seg = [mjm.cam_output[i] & mujoco.mjtCamOutBit.mjCAMOUT_SEG for i in active_cam_indices] + elif isinstance(render_seg, bool): + render_seg = [render_seg] * ncam + + assert len(render_rgb) == ncam and len(render_depth) == ncam and len(render_seg) == ncam, ( + f"render_rgb, render_depth, and render_seg must be a bool or a list of bools with length {ncam}" ) rgb_adr = -1 * np.ones(ncam, dtype=int) depth_adr = -1 * np.ones(ncam, dtype=int) + seg_adr = -1 * np.ones(ncam, dtype=int) cam_res_np = cam_res_arr.numpy() ri = 0 di = 0 + si = 0 total = 0 for idx in range(ncam): @@ -2676,6 +2804,9 @@ def create_render_context( if render_depth[idx]: depth_adr[idx] = di di += cam_res_np[idx][0] * cam_res_np[idx][1] + if render_seg[idx]: + seg_adr[idx] = si + si += cam_res_np[idx][0] * cam_res_np[idx][1] total += cam_res_np[idx][0] * cam_res_np[idx][1] @@ -2729,26 +2860,20 @@ def create_render_context( hfield_registry=hfield_registry, hfield_bvh_id=hfield_bvh_id_arr, hfield_bounds_size=hfield_bounds_size_arr, - flex_mesh=flex_mesh, + flex_mesh_registry=flex_registry, flex_rgba=wp.array(mjm.flex_rgba, dtype=wp.vec4), - flex_bvh_id=flex_bvh_id, - flex_face_point=flex_face_point, - flex_faceadr=flex_faceadr, - flex_nface=flex_nface, - flex_nwork=flex_nwork, - flex_group_root=flex_group_root, - flex_elemdataadr=flex_elemdataadr, - flex_shell=flex_shell, - flex_shelldataadr=flex_shelldataadr, - flex_radius=flex_radius, - flex_workadr=flex_workadr, - flex_worknum=flex_worknum, + flex_bvh_id=wp.array(flex_bvh_id, dtype=wp.uint64), + flex_group_root=wp.array(flex_group_root, dtype=int), flex_render_smooth=flex_render_smooth, + bvh_nflexgeom=len(flex_geom_flexid), + flex_dim_np=mjm.flex_dim, + flex_geom_flexid=wp.array(flex_geom_flexid, dtype=int), + flex_geom_edgeid=wp.array(flex_geom_edgeid, dtype=int), bvh=None, bvh_id=None, - lower=wp.zeros(nworld * bvh_ngeom, dtype=wp.vec3), - upper=wp.zeros(nworld * bvh_ngeom, dtype=wp.vec3), - group=wp.zeros(nworld * bvh_ngeom, dtype=int), + lower=wp.zeros(nworld * (bvh_ngeom + len(flex_geom_flexid)), dtype=wp.vec3), + upper=wp.zeros(nworld * (bvh_ngeom + len(flex_geom_flexid)), dtype=wp.vec3), + group=wp.zeros(nworld * (bvh_ngeom + len(flex_geom_flexid)), dtype=int), group_root=wp.zeros(nworld, dtype=int), ray=ray, rgb_data=wp.zeros((nworld, ri), dtype=wp.uint32), @@ -2757,6 +2882,9 @@ def create_render_context( depth_adr=wp.array(depth_adr, dtype=int), render_rgb=wp.array(render_rgb, dtype=bool), render_depth=wp.array(render_depth, dtype=bool), + seg_data=wp.zeros((nworld, max(si, 1)), dtype=int), + seg_adr=wp.array(seg_adr, dtype=int), + render_seg=wp.array(render_seg, dtype=bool), znear=znear, total_rays=int(total), ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py index b5ae2846..c021db3f 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py @@ -13,12 +13,13 @@ # limitations under the License. # ============================================================================== +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src import types from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType from mujoco.mjx.third_party.mujoco_warp._src.types import EqType from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp @wp.kernel @@ -180,17 +181,17 @@ def tree_edges(m: types.Model, d: types.Data, tree_tree: wp.array3d(dtype=int)): @wp.kernel def _flood_fill( - # Model: - ntree: int, - # In: - tree_tree_in: wp.array3d(dtype=int), - labels_in: wp.array2d(dtype=int), - stack_in: wp.array2d(dtype=int), - # Data out: - nisland_out: wp.array(dtype=int), - tree_island_out: wp.array2d(dtype=int), - # Out: - stack_out: wp.array2d(dtype=int), + # Model: + ntree: int, + # In: + tree_tree_in: wp.array3d(dtype=int), + labels_in: wp.array2d(dtype=int), + stack_in: wp.array2d(dtype=int), + # Data out: + nisland_out: wp.array(dtype=int), + tree_island_out: wp.array2d(dtype=int), + # Out: + stack_out: wp.array2d(dtype=int), ): """DFS flood fill to discover islands using tree_tree matrix.""" worldid = wp.tid() @@ -257,8 +258,8 @@ def island(m: types.Model, d: types.Data): stack_scratch = wp.empty((d.nworld, m.ntree * m.ntree), dtype=int) wp.launch( - _flood_fill, - dim=d.nworld, - inputs=[m.ntree, tree_tree, d.tree_island, stack_scratch], - outputs=[d.nisland, d.tree_island, stack_scratch], + _flood_fill, + dim=d.nworld, + inputs=[m.ntree, tree_tree, d.tree_island, stack_scratch], + outputs=[d.nisland, d.tree_island, stack_scratch], ) 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 ec49041e..3ce2d1c5 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py @@ -83,6 +83,35 @@ def quat_to_mat(quat: wp.quat) -> wp.mat33: ) +@wp.func +def quat_z2vec(vec: wp.vec3) -> wp.quat: + """Compute quaternion performing rotation from z-axis to given vector.""" + quat = wp.quat(0.0, 0.0, 0.0, 1.0) + + # normalize vector; if too small, no rotation + norm = wp.length(vec) + if norm < types.MJ_MINVAL: + return quat + vec = vec / norm + + axis = wp.vec3(-vec[1], vec[0], 0.0) + a = wp.length(axis) + + # almost parallel + if a < types.MJ_MINVAL: + # opposite: 180 deg rotation around x axis + if vec[2] < 0.0: + quat = wp.quat(1.0, 0.0, 0.0, 0.0) + return quat + + # make quaternion from angle and axis + axis = axis / a + angle = wp.atan2(a, vec[2]) + quat = axis_angle_to_quat(axis, angle) + + return quat + + @wp.func def quat_inv(quat: wp.quat) -> wp.quat: return wp.quat(quat[0], -quat[1], -quat[2], -quat[3]) 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 0bcd13af..4abcff26 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -89,8 +89,8 @@ def _spring_damper_dof_passive( stiffness = jnt_stiffness[worldid % jnt_stiffness.shape[0], jntid] damping = dof_damping[worldid % dof_damping.shape[0], dofid] - has_stiffness = stiffness != 0.0 and not opt_disableflags & DisableBit.SPRING - has_damping = damping != 0.0 and not opt_disableflags & DisableBit.DAMPER + has_stiffness = stiffness != 0.0 and not (opt_disableflags & DisableBit.SPRING) + has_damping = damping != 0.0 and not (opt_disableflags & DisableBit.DAMPER) if not has_stiffness: qfrc_spring_out[worldid, dofid] = 0.0 @@ -182,11 +182,14 @@ def _spring_damper_dof_passive( @wp.kernel def _spring_damper_tendon_passive( # Model: + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), tendon_stiffness: wp.array2d(dtype=float), tendon_damping: wp.array2d(dtype=float), tendon_lengthspring: wp.array2d(dtype=wp.vec2), # Data in: - ten_J_in: wp.array3d(dtype=float), + ten_J_in: wp.array2d(dtype=float), ten_length_in: wp.array2d(dtype=float), ten_velocity_in: wp.array2d(dtype=float), # In: @@ -196,7 +199,7 @@ def _spring_damper_tendon_passive( qfrc_spring_out: wp.array2d(dtype=float), qfrc_damper_out: wp.array2d(dtype=float), ): - worldid, tenid, dofid = wp.tid() + worldid, tenid, dofid_sparse = wp.tid() stiffness = tendon_stiffness[worldid % tendon_stiffness.shape[0], tenid] damping = tendon_damping[worldid % tendon_damping.shape[0], tenid] @@ -207,7 +210,13 @@ def _spring_damper_tendon_passive( if not has_stiffness and not has_damping: return - J = ten_J_in[worldid, tenid, dofid] + rownnz = ten_J_rownnz[tenid] + if dofid_sparse >= rownnz: + return + rowadr = ten_J_rowadr[tenid] + sparseid = rowadr + dofid_sparse + J = ten_J_in[worldid, sparseid] + dofid = ten_J_colind[sparseid] if has_stiffness: # compute spring force along tendon @@ -265,28 +274,28 @@ def _gravity_force( @wp.kernel def _fluid_force( - # Model: - opt_wind: wp.array(dtype=wp.vec3), - opt_density: wp.array(dtype=float), - opt_viscosity: wp.array(dtype=float), - body_rootid: wp.array(dtype=int), - body_geomnum: wp.array(dtype=int), - body_geomadr: wp.array(dtype=int), - body_mass: wp.array2d(dtype=float), - body_inertia: wp.array2d(dtype=wp.vec3), - geom_type: wp.array(dtype=int), - geom_size: wp.array2d(dtype=wp.vec3), - geom_fluid: wp.array2d(dtype=float), - body_fluid_ellipsoid: wp.array(dtype=bool), - # Data in: - xipos_in: wp.array2d(dtype=wp.vec3), - ximat_in: wp.array2d(dtype=wp.mat33), - geom_xpos_in: wp.array2d(dtype=wp.vec3), - geom_xmat_in: wp.array2d(dtype=wp.mat33), - subtree_com_in: wp.array2d(dtype=wp.vec3), - cvel_in: wp.array2d(dtype=wp.spatial_vector), - # Out: - fluid_applied_out: wp.array2d(dtype=wp.spatial_vector), + # Model: + opt_wind: wp.array(dtype=wp.vec3), + opt_density: wp.array(dtype=float), + opt_viscosity: wp.array(dtype=float), + body_rootid: wp.array(dtype=int), + body_geomnum: wp.array(dtype=int), + body_geomadr: wp.array(dtype=int), + body_mass: wp.array2d(dtype=float), + body_inertia: wp.array2d(dtype=wp.vec3), + geom_type: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + geom_fluid: wp.array2d(dtype=float), + body_fluid_ellipsoid: wp.array(dtype=bool), + # Data in: + xipos_in: wp.array2d(dtype=wp.vec3), + ximat_in: wp.array2d(dtype=wp.mat33), + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # Out: + fluid_applied_out: wp.array2d(dtype=wp.spatial_vector), ): """Computes body-space fluid forces for both inertia-box and ellipsoid models.""" worldid, bodyid = wp.tid() @@ -495,29 +504,29 @@ def _fluid(m: Model, d: Data): fluid_applied = wp.empty((d.nworld, m.nbody), dtype=wp.spatial_vector) wp.launch( - _fluid_force, - dim=(d.nworld, m.nbody), - inputs=[ - m.opt.wind, - m.opt.density, - m.opt.viscosity, - m.body_rootid, - m.body_geomnum, - m.body_geomadr, - m.body_mass, - m.body_inertia, - m.geom_type, - m.geom_size, - m.geom_fluid, - m.body_fluid_ellipsoid, - d.xipos, - d.ximat, - d.geom_xpos, - d.geom_xmat, - d.subtree_com, - d.cvel, - ], - outputs=[fluid_applied], + _fluid_force, + dim=(d.nworld, m.nbody), + inputs=[ + m.opt.wind, + m.opt.density, + m.opt.viscosity, + m.body_rootid, + m.body_geomnum, + m.body_geomadr, + m.body_mass, + m.body_inertia, + m.geom_type, + m.geom_size, + m.geom_fluid, + m.body_fluid_ellipsoid, + d.xipos, + d.ximat, + d.geom_xpos, + d.geom_xmat, + d.subtree_com, + d.cvel, + ], + outputs=[fluid_applied], ) support.apply_ft(m, d, fluid_applied, d.qfrc_fluid, False) @@ -565,6 +574,7 @@ def _flex_elasticity( flex_edgeadr: wp.array(dtype=int), flex_elemadr: wp.array(dtype=int), flex_elemnum: wp.array(dtype=int), + flex_elemdataadr: wp.array(dtype=int), flex_elemedgeadr: wp.array(dtype=int), flex_vertbodyid: wp.array(dtype=int), flex_elem: wp.array(dtype=int), @@ -590,32 +600,39 @@ def _flex_elasticity( f = i break + local_elemid = elemid - flex_elemadr[f] dim = flex_dim[f] nvert = dim + 1 nedge = nvert * (nvert - 1) / 2 edges = wp.where( - dim == 3, - wp.matrix(0, 1, 1, 2, 2, 0, 2, 3, 0, 3, 1, 3, shape=(6, 2), dtype=int), - wp.matrix(1, 2, 2, 0, 0, 1, 0, 0, 0, 0, 0, 0, shape=(6, 2), dtype=int), + dim == 1, + wp.matrix(0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, shape=(6, 2), dtype=int), + wp.where( + dim == 3, + wp.matrix(0, 1, 1, 2, 2, 0, 2, 3, 0, 3, 1, 3, shape=(6, 2), dtype=int), + wp.matrix(1, 2, 2, 0, 0, 1, 0, 0, 0, 0, 0, 0, shape=(6, 2), dtype=int), + ), ) if timestep > 0.0 and not dsbl_damper: kD = flex_damping[f] / timestep else: kD = 0.0 + elem_data_adr = flex_elemdataadr[f] + local_elemid * (dim + 1) + vbase = flex_vertadr[f] gradient = wp.matrix(0.0, shape=(6, 6)) for e in range(nedge): - vert0 = flex_elem[(dim + 1) * elemid + edges[e, 0]] - vert1 = flex_elem[(dim + 1) * elemid + edges[e, 1]] - xpos0 = flexvert_xpos_in[worldid, vert0] - xpos1 = flexvert_xpos_in[worldid, vert1] + vert0 = flex_elem[elem_data_adr + edges[e, 0]] + vert1 = flex_elem[elem_data_adr + edges[e, 1]] + xpos0 = flexvert_xpos_in[worldid, vbase + vert0] + xpos1 = flexvert_xpos_in[worldid, vbase + vert1] for i in range(3): gradient[e, 0 + i] = xpos0[i] - xpos1[i] gradient[e, 3 + i] = xpos1[i] - xpos0[i] elongation = wp.spatial_vectorf(0.0) for e in range(nedge): - idx = flex_elemedge[elemid * nedge + e] + idx = flex_elemedge[flex_elemedgeadr[f] + local_elemid * nedge + e] vel = flexedge_velocity_in[worldid, flex_edgeadr[f] + idx] deformed = flexedge_length_in[worldid, flex_edgeadr[f] + idx] reference = flexedge_length0[flex_edgeadr[f] + idx] @@ -638,7 +655,7 @@ def _flex_elasticity( force[edges[ed2, i], x] -= elongation[ed1] * gradient[ed2, 3 * i + x] * metric[ed1, ed2] for v in range(nvert): - vert = flex_elem[(dim + 1) * elemid + v] + vert = flex_elem[elem_data_adr + v] bodyid = flex_vertbodyid[flex_vertadr[f] + vert] for x in range(3): wp.atomic_add(qfrc_spring_out, worldid, body_dofadr[bodyid] + x, force[v, x]) @@ -742,8 +759,11 @@ def passive(m: Model, d: Data): if m.ntendon: wp.launch( _spring_damper_tendon_passive, - dim=(d.nworld, m.ntendon, m.nv), + dim=(d.nworld, m.ntendon, m.max_ten_J_rownnz), inputs=[ + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, m.tendon_stiffness, m.tendon_damping, m.tendon_lengthspring, @@ -772,6 +792,7 @@ def passive(m: Model, d: Data): m.flex_edgeadr, m.flex_elemadr, m.flex_elemnum, + m.flex_elemdataadr, m.flex_elemedgeadr, m.flex_vertbodyid, m.flex_elem, 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 5eaea821..57d6a5ba 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py @@ -752,7 +752,8 @@ def ray_mesh_with_bvh_anyhit( @wp.func def ray_flex_with_bvh( # In: - bvh_id: wp.uint64, + flex_bvh_id: wp.array(dtype=wp.uint64), + flexid: int, group_root: int, pnt: wp.vec3, vec: wp.vec3, @@ -769,7 +770,7 @@ def ray_flex_with_bvh( n = wp.vec3(0.0, 0.0, 0.0) f = int(-1) - hit = wp.mesh_query_ray(bvh_id, pnt, vec, max_t, t, u, v, sign, n, f, group_root) + hit = wp.mesh_query_ray(flex_bvh_id[flexid], pnt, vec, max_t, t, u, v, sign, n, f, group_root) if hit: return t, n, u, v, f @@ -777,6 +778,23 @@ def ray_flex_with_bvh( return -1.0, wp.vec3(0.0, 0.0, 0.0), 0.0, 0.0, -1 +@wp.func +def ray_flex_with_bvh_anyhit( + # In: + flex_bvh_id: wp.array(dtype=wp.uint64), + flexid: int, + group_root: int, + pnt: wp.vec3, + vec: wp.vec3, + max_t: float, +) -> bool: + """Returns True if there is any hit for ray flex intersections. + + Requires wp.Mesh be constructed and their ids to be passed. Flex are already in world space. + """ + return wp.mesh_query_ray_anyhit(flex_bvh_id[flexid], pnt, vec, max_t, group_root) + + @wp.func 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. diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py index 28a4284f..bc8d16c3 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py @@ -23,6 +23,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_capsule from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_cylinder from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_ellipsoid from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_flex_with_bvh +from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_flex_with_bvh_anyhit from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_mesh_with_bvh from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_mesh_with_bvh_anyhit from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_plane @@ -39,10 +40,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope wp.set_module_options({"enable_backward": False}) -# TODO(team): remove after mjwarp depends on warp-lang >= 1.12 in pyproject.toml -from mujoco.mjx.third_party.mujoco_warp._src.types import TEXTURE_DTYPE - - @wp.func def sample_texture( # Model: @@ -51,7 +48,7 @@ def sample_texture( # In: geom_id: int, tex_repeat: wp.vec2, - tex: TEXTURE_DTYPE, + tex: wp.Texture2D, pos: wp.vec3, rot: wp.mat33, mesh_facetexcoord: wp.array(dtype=wp.vec3i), @@ -94,17 +91,26 @@ def cast_ray( geom_type: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), + flex_vertadr: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flex_radius: wp.array(dtype=float), # Data in: geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), # In: bvh_id: wp.uint64, group_root: int, - world_id: int, + worldid: int, bvh_ngeom: int, + flex_bvh_ngeom: int, enabled_geom_ids: wp.array(dtype=int), mesh_bvh_id: wp.array(dtype=wp.uint64), hfield_bvh_id: wp.array(dtype=wp.uint64), + flex_geom_flexid: wp.array(dtype=int), + flex_geom_edgeid: wp.array(dtype=int), + flex_bvh_id: wp.array(dtype=wp.uint64), + flex_group_root: wp.array2d(dtype=int), ray_origin_world: wp.vec3, ray_dir_world: wp.vec3, ) -> Tuple[int, float, wp.vec3, float, float, int, int]: @@ -118,91 +124,127 @@ def cast_ray( query = wp.bvh_query_ray(bvh_id, ray_origin_world, ray_dir_world, group_root) bounds_nr = int(0) + ngeom = bvh_ngeom + flex_bvh_ngeom while wp.bvh_query_next(query, bounds_nr, dist): gi_global = bounds_nr - gi_bvh_local = gi_global - (world_id * bvh_ngeom) - gi = enabled_geom_ids[gi_bvh_local] + local_id = gi_global - (worldid * ngeom) + d = float(-1.0) hit_mesh_id = int(-1) u = float(0.0) v = float(0.0) f = int(-1) n = wp.vec3(0.0, 0.0, 0.0) + hit_geom_id = int(-1) + + if local_id < bvh_ngeom: + gi = enabled_geom_ids[local_id] + gtype = geom_type[gi] + else: + gi = local_id - bvh_ngeom + gtype = GeomType.FLEX + + hit_geom_id = gi # TODO: Investigate branch elimination with static loop unrolling - if geom_type[gi] == GeomType.PLANE: + if gtype == GeomType.PLANE: d, n = ray_plane( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.HFIELD: + if gtype == GeomType.HFIELD: d, n, u, v, f, geom_hfield_id = ray_mesh_with_bvh( hfield_bvh_id, geom_dataid[gi], - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], ray_origin_world, ray_dir_world, dist, ) - if geom_type[gi] == GeomType.SPHERE: + if gtype == GeomType.SPHERE: d, n = ray_sphere( - geom_xpos_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi][0] * geom_size[world_id % geom_size.shape[0], gi][0], + geom_xpos_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi][0] * geom_size[worldid % geom_size.shape[0], gi][0], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.ELLIPSOID: + if gtype == GeomType.ELLIPSOID: d, n = ray_ellipsoid( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.CAPSULE: + if gtype == GeomType.CAPSULE: d, n = ray_capsule( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.CYLINDER: + if gtype == GeomType.CYLINDER: d, n = ray_cylinder( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.BOX: + if gtype == GeomType.BOX: d, all, n = ray_box( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.MESH: + if gtype == GeomType.MESH: d, n, u, v, f, hit_mesh_id = ray_mesh_with_bvh( mesh_bvh_id, geom_dataid[gi], - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], ray_origin_world, ray_dir_world, dist, ) + if gtype == GeomType.FLEX: + hit_geom_id = -2 + flexid = flex_geom_flexid[gi] + edge_id = flex_geom_edgeid[gi] + + if edge_id >= 0: + edge = flex_edge[edge_id] + vert_adr = flex_vertadr[flexid] + v0 = flexvert_xpos_in[worldid, vert_adr + edge[0]] + v1 = flexvert_xpos_in[worldid, vert_adr + edge[1]] + pos = 0.5 * (v0 + v1) + vec = v1 - v0 + + length = wp.length(vec) + edgeq = math.quat_z2vec(vec) + mat = math.quat_to_mat(edgeq) + size = wp.vec3(flex_radius[flexid], 0.5 * length, 0.0) + + d, n = ray_capsule(pos, mat, size, ray_origin_world, ray_dir_world) + hit_mesh_id = flexid + else: + flex_gr = flex_group_root[worldid, flexid] + d, n, u, v, f = ray_flex_with_bvh(flex_bvh_id, flexid, flex_gr, ray_origin_world, ray_dir_world, dist) + if d >= 0.0: + hit_mesh_id = flexid if d >= 0.0 and d < dist: dist = d normal = n - geom_id = gi + geom_id = hit_geom_id bary_u = u bary_v = v face_idx = f @@ -217,17 +259,26 @@ def cast_ray_first_hit( geom_type: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), + flex_vertadr: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flex_radius: wp.array(dtype=float), # Data in: geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), # In: bvh_id: wp.uint64, group_root: int, - world_id: int, + worldid: int, bvh_ngeom: int, + bvh_nflexgeom: int, enabled_geom_ids: wp.array(dtype=int), mesh_bvh_id: wp.array(dtype=wp.uint64), hfield_bvh_id: wp.array(dtype=wp.uint64), + flex_geom_flexid: wp.array(dtype=int), + flex_geom_edgeid: wp.array(dtype=int), + flex_bvh_id: wp.array(dtype=wp.uint64), + flex_group_root: wp.array2d(dtype=int), ray_origin_world: wp.vec3, ray_dir_world: wp.vec3, max_dist: float, @@ -235,81 +286,119 @@ def cast_ray_first_hit( """A simpler version of casting rays that only checks for the first hit.""" query = wp.bvh_query_ray(bvh_id, ray_origin_world, ray_dir_world, group_root) bounds_nr = int(0) + ngeom = bvh_ngeom + bvh_nflexgeom while wp.bvh_query_next(query, bounds_nr, max_dist): gi_global = bounds_nr - gi_bvh_local = gi_global - (world_id * bvh_ngeom) - gi = enabled_geom_ids[gi_bvh_local] + local_id = gi_global - (worldid * ngeom) + + d = float(-1.0) + n = wp.vec3(0.0, 0.0, 0.0) + + if local_id < bvh_ngeom: + gi = enabled_geom_ids[local_id] + gtype = geom_type[gi] + else: + gi = local_id - bvh_ngeom + gtype = GeomType.FLEX # TODO: Investigate branch elimination with static loop unrolling - if geom_type[gi] == GeomType.PLANE: + if gtype == GeomType.PLANE: d, n = ray_plane( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.HFIELD: + if gtype == GeomType.HFIELD: d, n, u, v, f, geom_hfield_id = ray_mesh_with_bvh( hfield_bvh_id, geom_dataid[gi], - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], ray_origin_world, ray_dir_world, max_dist, ) - if geom_type[gi] == GeomType.SPHERE: + if gtype == GeomType.SPHERE: d, n = ray_sphere( - geom_xpos_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi][0] * geom_size[world_id % geom_size.shape[0], gi][0], + geom_xpos_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi][0] * geom_size[worldid % geom_size.shape[0], gi][0], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.ELLIPSOID: + if gtype == GeomType.ELLIPSOID: d, n = ray_ellipsoid( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.CAPSULE: + if gtype == GeomType.CAPSULE: d, n = ray_capsule( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.CYLINDER: + if gtype == GeomType.CYLINDER: d, n = ray_cylinder( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.BOX: + if gtype == GeomType.BOX: d, all, n = ray_box( - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], - geom_size[world_id % geom_size.shape[0], gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], + geom_size[worldid % geom_size.shape[0], gi], ray_origin_world, ray_dir_world, ) - if geom_type[gi] == GeomType.MESH: + if gtype == GeomType.MESH: hit = ray_mesh_with_bvh_anyhit( mesh_bvh_id, geom_dataid[gi], - geom_xpos_in[world_id, gi], - geom_xmat_in[world_id, gi], + geom_xpos_in[worldid, gi], + geom_xmat_in[worldid, gi], ray_origin_world, ray_dir_world, max_dist, ) d = 0.0 if hit else -1.0 + if gtype == GeomType.FLEX: + flexid = flex_geom_flexid[gi] + edge_id = flex_geom_edgeid[gi] + + if edge_id >= 0: + edge = flex_edge[edge_id] + vert_adr = flex_vertadr[flexid] + v0 = flexvert_xpos_in[worldid, vert_adr + edge[0]] + v1 = flexvert_xpos_in[worldid, vert_adr + edge[1]] + pos = 0.5 * (v0 + v1) + vec = v1 - v0 + + length = wp.length(vec) + edgeq = math.quat_z2vec(vec) + mat = math.quat_to_mat(edgeq) + size = wp.vec3(flex_radius[flexid], 0.5 * length, 0.0) + + d, n = ray_capsule(pos, mat, size, ray_origin_world, ray_dir_world) + else: + hit = ray_flex_with_bvh_anyhit( + flex_bvh_id, + flexid, + flex_group_root[worldid, flexid], + ray_origin_world, + ray_dir_world, + max_dist, + ) + d = 0.0 if hit else -1.0 if d >= 0.0 and d < max_dist: return True @@ -323,18 +412,27 @@ def compute_lighting( geom_type: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), + flex_vertadr: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flex_radius: wp.array(dtype=float), # Data in: geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), # In: use_shadows: bool, bvh_id: wp.uint64, group_root: int, bvh_ngeom: int, + bvh_nflexgeom: int, enabled_geom_ids: wp.array(dtype=int), - world_id: int, + worldid: int, mesh_bvh_id: wp.array(dtype=wp.uint64), hfield_bvh_id: wp.array(dtype=wp.uint64), + flex_geom_flexid: wp.array(dtype=int), + flex_geom_edgeid: wp.array(dtype=int), + flex_bvh_id: wp.array(dtype=wp.uint64), + flex_group_root: wp.array2d(dtype=int), lightactive: bool, lighttype: int, lightcastshadow: bool, @@ -385,15 +483,24 @@ def compute_lighting( geom_type, geom_dataid, geom_size, + flex_vertadr, + flex_edge, + flex_radius, geom_xpos_in, geom_xmat_in, + flexvert_xpos_in, bvh_id, group_root, - world_id, + worldid, bvh_ngeom, + bvh_nflexgeom, enabled_geom_ids, mesh_bvh_id, hfield_bvh_id, + flex_geom_flexid, + flex_geom_edgeid, + flex_bvh_id, + flex_group_root, shadow_origin, L, max_t, @@ -418,6 +525,7 @@ def render(m: Model, d: Data, rc: RenderContext): """ rc.rgb_data.fill_(rc.background_color) rc.depth_data.fill_(0.0) + rc.seg_data.fill_(-1) @wp.kernel(module="unique", enable_backward=False) def _render_megakernel( @@ -434,6 +542,9 @@ def render(m: Model, d: Data, rc: RenderContext): light_type: wp.array2d(dtype=int), light_castshadow: wp.array2d(dtype=bool), light_active: wp.array2d(dtype=bool), + flex_vertadr: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flex_radius: wp.array(dtype=float), mesh_faceadr: wp.array(dtype=int), mat_texid: wp.array3d(dtype=int), mat_texrepeat: wp.array2d(dtype=wp.vec2), @@ -445,21 +556,25 @@ def render(m: Model, d: Data, rc: RenderContext): cam_xmat_in: wp.array2d(dtype=wp.mat33), light_xpos_in: wp.array2d(dtype=wp.vec3), light_xdir_in: wp.array2d(dtype=wp.vec3), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), # In: nrender: int, use_shadows: bool, bvh_ngeom: int, + bvh_nflexgeom: int, cam_res: wp.array(dtype=wp.vec2i), cam_id_map: wp.array(dtype=int), ray: wp.array(dtype=wp.vec3), rgb_adr: wp.array(dtype=int), depth_adr: wp.array(dtype=int), + seg_adr: wp.array(dtype=int), render_rgb: wp.array(dtype=bool), render_depth: wp.array(dtype=bool), + render_seg: wp.array(dtype=bool), bvh_id: wp.uint64, group_root: wp.array(dtype=int), - flex_bvh_id: wp.uint64, - flex_group_root: wp.array(dtype=int), + flex_bvh_id: wp.array(dtype=wp.uint64), + flex_group_root: wp.array2d(dtype=int), enabled_geom_ids: wp.array(dtype=int), mesh_bvh_id: wp.array(dtype=wp.uint64), mesh_facetexcoord: wp.array(dtype=wp.vec3i), @@ -467,46 +582,48 @@ def render(m: Model, d: Data, rc: RenderContext): mesh_texcoord_offsets: wp.array(dtype=int), hfield_bvh_id: wp.array(dtype=wp.uint64), flex_rgba: wp.array(dtype=wp.vec4), - # TODO: remove after mjwarp depends on warp-lang >= 1.12 in pyproject.toml - textures: wp.array(dtype=TEXTURE_DTYPE), + flex_geom_flexid: wp.array(dtype=int), + flex_geom_edgeid: wp.array(dtype=int), + textures: wp.array(dtype=wp.Texture2D), # Out: rgb_out: wp.array2d(dtype=wp.uint32), depth_out: wp.array2d(dtype=float), + seg_out: wp.array2d(dtype=int), ): - world_idx, ray_idx = wp.tid() + worldid, rayid = wp.tid() - # Map global ray_idx -> (cam_idx, ray_idx_local) using cumulative sizes + # Map global rayid -> (cam_idx, rayid_local) using cumulative sizes cam_idx = int(-1) - ray_idx_local = int(-1) + rayid_local = int(-1) accum = int(0) for i in range(nrender): num_i = cam_res[i][0] * cam_res[i][1] - if ray_idx < accum + num_i: + if rayid < accum + num_i: cam_idx = i - ray_idx_local = ray_idx - accum + rayid_local = rayid - accum break accum += num_i - if cam_idx == -1 or ray_idx_local < 0: + if cam_idx == -1 or rayid_local < 0: return - if not render_rgb[cam_idx] and not render_depth[cam_idx]: + if not render_rgb[cam_idx] and not render_depth[cam_idx] and not render_seg[cam_idx]: return # Map active camera index to MuJoCo camera ID mujoco_cam_id = cam_id_map[cam_idx] if wp.static(rc.use_precomputed_rays): - ray_dir_local_cam = ray[ray_idx] + ray_dir_local_cam = ray[rayid] else: img_w = cam_res[cam_idx][0] img_h = cam_res[cam_idx][1] - px = ray_idx_local % img_w - py = ray_idx_local // img_w + px = rayid_local % img_w + py = rayid_local // img_w ray_dir_local_cam = compute_ray( cam_projection[mujoco_cam_id], - cam_fovy[world_idx % cam_fovy.shape[0], mujoco_cam_id], + cam_fovy[worldid % cam_fovy.shape[0], mujoco_cam_id], cam_sensorsize[mujoco_cam_id], - cam_intrinsic[world_idx % cam_intrinsic.shape[0], mujoco_cam_id], + cam_intrinsic[worldid % cam_intrinsic.shape[0], mujoco_cam_id], img_w, img_h, px, @@ -514,38 +631,37 @@ def render(m: Model, d: Data, rc: RenderContext): wp.static(rc.znear), ) - ray_dir_world = cam_xmat_in[world_idx, mujoco_cam_id] @ ray_dir_local_cam - ray_origin_world = cam_xpos_in[world_idx, mujoco_cam_id] + ray_dir_world = cam_xmat_in[worldid, mujoco_cam_id] @ ray_dir_local_cam + ray_origin_world = cam_xpos_in[worldid, mujoco_cam_id] geom_id, dist, normal, u, v, f, mesh_id = cast_ray( geom_type, geom_dataid, geom_size, + flex_vertadr, + flex_edge, + flex_radius, geom_xpos_in, geom_xmat_in, + flexvert_xpos_in, bvh_id, - group_root[world_idx], - world_idx, + group_root[worldid], + worldid, bvh_ngeom, + bvh_nflexgeom, enabled_geom_ids, mesh_bvh_id, hfield_bvh_id, + flex_geom_flexid, + flex_geom_edgeid, + flex_bvh_id, + flex_group_root, ray_origin_world, ray_dir_world, ) - if wp.static(m.nflex > 0): - d, n, u, v, f = ray_flex_with_bvh( - flex_bvh_id, - flex_group_root[world_idx], - ray_origin_world, - ray_dir_world, - dist, - ) - if d >= 0.0 and d < dist: - dist = d - normal = n - geom_id = -2 + if render_seg[cam_idx] and geom_id != -1: + seg_out[worldid, seg_adr[cam_idx] + rayid_local] = geom_id # Early Out if geom_id == -1: @@ -556,9 +672,7 @@ def render(m: Model, d: Data, rc: RenderContext): # In camera-local coordinates, the optical axis is -Z. The Z-component of the # normalized ray direction is negative, so -ray_dir_local_cam[2] gives cos(θ) # between the ray and the optical axis. - depth_out[world_idx, depth_adr[cam_idx] + ray_idx_local] = dist * ( - -ray_dir_local_cam[2] - ) + depth_out[worldid, depth_adr[cam_idx] + rayid_local] = dist * (-ray_dir_local_cam[2]) if not render_rgb[cam_idx]: return @@ -567,31 +681,30 @@ def render(m: Model, d: Data, rc: RenderContext): hit_point = ray_origin_world + ray_dir_world * dist if geom_id == -2: - # TODO: Currently flex textures are not supported, and only the first rgba value - # is used until further flex support is added. - color = flex_rgba[0] - elif geom_matid[world_idx % geom_matid.shape[0], geom_id] == -1: - color = geom_rgba[world_idx % geom_rgba.shape[0], geom_id] + # We encode flex_id in mesh_id for flex ray hits during cast_ray + color = flex_rgba[mesh_id] + elif geom_matid[worldid % geom_matid.shape[0], geom_id] == -1: + color = geom_rgba[worldid % geom_rgba.shape[0], geom_id] else: - color = mat_rgba[world_idx % mat_rgba.shape[0], geom_matid[world_idx % geom_matid.shape[0], geom_id]] + color = mat_rgba[worldid % mat_rgba.shape[0], geom_matid[worldid % geom_matid.shape[0], geom_id]] base_color = wp.vec3(color[0], color[1], color[2]) hit_color = base_color if wp.static(rc.use_textures): if geom_id != -2: - mat_id = geom_matid[world_idx % geom_matid.shape[0], geom_id] + mat_id = geom_matid[worldid % geom_matid.shape[0], geom_id] if mat_id >= 0: - tex_id = mat_texid[world_idx % mat_texid.shape[0], mat_id, 1] + tex_id = mat_texid[worldid % mat_texid.shape[0], mat_id, 1] if tex_id >= 0: tex_color = sample_texture( geom_type, mesh_faceadr, geom_id, - mat_texrepeat[world_idx % mat_texrepeat.shape[0], mat_id], + mat_texrepeat[worldid % mat_texrepeat.shape[0], mat_id], textures[tex_id], - geom_xpos_in[world_idx, geom_id], - geom_xmat_in[world_idx, geom_id], + geom_xpos_in[worldid, geom_id], + geom_xmat_in[worldid, geom_id], mesh_facetexcoord, mesh_texcoord, mesh_texcoord_offsets, @@ -616,21 +729,30 @@ def render(m: Model, d: Data, rc: RenderContext): geom_type, geom_dataid, geom_size, + flex_vertadr, + flex_edge, + flex_radius, geom_xpos_in, geom_xmat_in, + flexvert_xpos_in, use_shadows, bvh_id, - group_root[world_idx], + group_root[worldid], bvh_ngeom, + bvh_nflexgeom, enabled_geom_ids, - world_idx, + worldid, mesh_bvh_id, hfield_bvh_id, - light_active[world_idx % light_active.shape[0], l], - light_type[world_idx % light_type.shape[0], l], - light_castshadow[world_idx % light_castshadow.shape[0], l], - light_xpos_in[world_idx, l], - light_xdir_in[world_idx, l], + flex_geom_flexid, + flex_geom_edgeid, + flex_bvh_id, + flex_group_root, + light_active[worldid % light_active.shape[0], l], + light_type[worldid % light_type.shape[0], l], + light_castshadow[worldid % light_castshadow.shape[0], l], + light_xpos_in[worldid, l], + light_xdir_in[worldid, l], normal, hit_point, ) @@ -639,7 +761,7 @@ def render(m: Model, d: Data, rc: RenderContext): hit_color = wp.min(result, wp.vec3(1.0, 1.0, 1.0)) hit_color = wp.max(hit_color, wp.vec3(0.0, 0.0, 0.0)) - rgb_out[world_idx, rgb_adr[cam_idx] + ray_idx_local] = pack_rgba_to_uint32( + rgb_out[worldid, rgb_adr[cam_idx] + rayid_local] = pack_rgba_to_uint32( hit_color[0] * 255.0, hit_color[1] * 255.0, hit_color[2] * 255.0, @@ -662,6 +784,9 @@ def render(m: Model, d: Data, rc: RenderContext): m.light_type, m.light_castshadow, m.light_active, + m.flex_vertadr, + m.flex_edge, + m.flex_radius, m.mesh_faceadr, m.mat_texid, m.mat_texrepeat, @@ -672,16 +797,20 @@ def render(m: Model, d: Data, rc: RenderContext): d.cam_xmat, d.light_xpos, d.light_xdir, + d.flexvert_xpos, rc.nrender, rc.use_shadows, rc.bvh_ngeom, + rc.bvh_nflexgeom, rc.cam_res, rc.cam_id_map, rc.ray, rc.rgb_adr, rc.depth_adr, + rc.seg_adr, rc.render_rgb, rc.render_depth, + rc.render_seg, rc.bvh_id, rc.group_root, rc.flex_bvh_id, @@ -693,10 +822,13 @@ def render(m: Model, d: Data, rc: RenderContext): rc.mesh_texcoord_offsets, rc.hfield_bvh_id, rc.flex_rgba, + rc.flex_geom_flexid, + rc.flex_geom_edgeid, rc.textures, ], outputs=[ rc.rgb_data, rc.depth_data, + rc.seg_data, ], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py index ccb808ff..36958f8e 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py @@ -206,3 +206,41 @@ def get_depth(rc: RenderContext, camera_index: int, depth_scale: float, depth_ou inputs=[rc.depth_data, rc.depth_adr, camera_index, depth_scale], outputs=[depth_out], ) + + +@wp.kernel +def _extract_seg_kernel( + # In: + seg_data: wp.array2d(dtype=int), + seg_adr: wp.array(dtype=int), + camera_index: int, + # Out: + seg_out: wp.array3d(dtype=int), +): + """Extract per-pixel geom IDs from the render context buffers for a given camera index.""" + worldid, pixelid = wp.tid() + xid = pixelid % seg_out.shape[2] + yid = pixelid // seg_out.shape[2] + + seg_adr_offset = seg_adr[camera_index] + seg_out[worldid, yid, xid] = seg_data[worldid, seg_adr_offset + pixelid] + + +def get_segmentation(rc: RenderContext, camera_index: int, seg_out: wp.array3d(dtype=int)): + """Get the segmentation data from the render context buffers for a given camera index. + + Each pixel contains the MuJoCo geom ID of the geometry hit by the ray, -1 for + background, or -2 for flex bodies. + + Args: + rc: The render context on device. + camera_index: The index of the camera to get the segmentation data for. + seg_out: The output array to store the geom IDs in, with shape + (nworld, height, width). + """ + wp.launch( + _extract_seg_kernel, + dim=(seg_out.shape[0], seg_out.shape[1] * seg_out.shape[2]), + inputs=[rc.seg_data, rc.seg_adr, camera_index], + outputs=[seg_out], + ) 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 859ddf8a..2c8177b8 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -15,12 +15,17 @@ from typing import Any, Tuple +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src import math from mujoco.mjx.third_party.mujoco_warp._src import ray from mujoco.mjx.third_party.mujoco_warp._src import smooth from mujoco.mjx.third_party.mujoco_warp._src import support from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import get_sdf_params from mujoco.mjx.third_party.mujoco_warp._src.collision_sdf import sdf +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType @@ -28,9 +33,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import DataType from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit from mujoco.mjx.third_party.mujoco_warp._src.types import JointType -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import Model from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType from mujoco.mjx.third_party.mujoco_warp._src.types import SensorType @@ -40,10 +42,10 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 from mujoco.mjx.third_party.mujoco_warp._src.types import vec6 from mujoco.mjx.third_party.mujoco_warp._src.types import vec8 from mujoco.mjx.third_party.mujoco_warp._src.types import vec8i +from mujoco.mjx.third_party.mujoco_warp._src.types import vec_pluginattr from mujoco.mjx.third_party.mujoco_warp._src.util_misc import inside_geom from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -2081,16 +2083,16 @@ def _transform_spatial(vec: wp.spatial_vector, dif: wp.vec3) -> wp.vec3: @wp.kernel def _preprocess_tactile_contacts( - # Model: - body_weldid: wp.array(dtype=int), - geom_bodyid: wp.array(dtype=int), - # Data in: - contact_geom_in: wp.array(dtype=wp.vec2i), - contact_worldid_in: wp.array(dtype=int), - nacon_in: wp.array(dtype=int), - # Out: - weld_geom_count_out: wp.array2d(dtype=int), - weld_geom_list_out: wp.array3d(dtype=int), + # Model: + body_weldid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + # Data in: + contact_geom_in: wp.array(dtype=wp.vec2i), + contact_worldid_in: wp.array(dtype=int), + nacon_in: wp.array(dtype=int), + # Out: + weld_geom_count_out: wp.array2d(dtype=int), + weld_geom_list_out: wp.array3d(dtype=int), ): conid = wp.tid() ncon = nacon_in[0] @@ -2118,42 +2120,43 @@ def _preprocess_tactile_contacts( @wp.kernel def _sensor_tactile( - # Model: - body_rootid: wp.array(dtype=int), - body_weldid: wp.array(dtype=int), - oct_child: wp.array(dtype=vec8i), - oct_aabb: wp.array2d(dtype=wp.vec3), - oct_coeff: wp.array(dtype=vec8), - geom_type: wp.array(dtype=int), - geom_bodyid: wp.array(dtype=int), - geom_size: wp.array2d(dtype=wp.vec3), - mesh_vertadr: wp.array(dtype=int), - mesh_vertnum: wp.array(dtype=int), - mesh_octadr: wp.array(dtype=int), - mesh_normaladr: wp.array(dtype=int), - mesh_normalnum: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), - mesh_normal: wp.array(dtype=wp.vec3), - mesh_quat: wp.array(dtype=wp.quat), - sensor_objid: wp.array(dtype=int), - sensor_refid: wp.array(dtype=int), - sensor_dim: wp.array(dtype=int), - sensor_adr: wp.array(dtype=int), - plugin: wp.array(dtype=int), - plugin_attr: wp.array(dtype=wp.vec3f), - geom_plugin_index: wp.array(dtype=int), - taxel_vertadr: wp.array(dtype=int), - taxel_sensorid: wp.array(dtype=int), - # Data in: - geom_xpos_in: wp.array2d(dtype=wp.vec3), - geom_xmat_in: wp.array2d(dtype=wp.mat33), - subtree_com_in: wp.array2d(dtype=wp.vec3), - cvel_in: wp.array2d(dtype=wp.spatial_vector), - # In: - weld_geom_count_in: wp.array2d(dtype=int), - weld_geom_list_in: wp.array3d(dtype=int), - # Data out: - sensordata_out: wp.array2d(dtype=float), + # Model: + body_rootid: wp.array(dtype=int), + body_weldid: wp.array(dtype=int), + oct_child: wp.array(dtype=vec8i), + oct_aabb: wp.array2d(dtype=wp.vec3), + oct_coeff: wp.array(dtype=vec8), + geom_type: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + mesh_vertadr: wp.array(dtype=int), + mesh_vertnum: wp.array(dtype=int), + mesh_octadr: wp.array(dtype=int), + mesh_normaladr: wp.array(dtype=int), + mesh_normalnum: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_normal: wp.array(dtype=wp.vec3), + mesh_quat: wp.array(dtype=wp.quat), + sensor_objid: wp.array(dtype=int), + sensor_refid: wp.array(dtype=int), + sensor_dim: wp.array(dtype=int), + sensor_adr: wp.array(dtype=int), + plugin: wp.array(dtype=int), + plugin_attr: wp.array(dtype=vec_pluginattr), + geom_plugin_index: wp.array(dtype=int), + taxel_vertadr: wp.array(dtype=int), + taxel_sensorid: wp.array(dtype=int), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + # In: + weld_geom_count_in: wp.array2d(dtype=int), + weld_geom_list_in: wp.array3d(dtype=int), + # Data out: + sensordata_out: wp.array2d(dtype=float), ): worldid, taxelid = wp.tid() @@ -2211,40 +2214,25 @@ def _sensor_tactile( contact_type = geom_type[geom] plugin_attributes, plugin_index, volume_data, mesh_data = get_sdf_params( - oct_child, - oct_aabb, - oct_coeff, - mesh_octadr, - plugin, - plugin_attr, - contact_type, - geom_size[worldid % geom_size.shape[0], geom], - plugin_id, - mesh_id, + oct_child, + oct_aabb, + oct_coeff, + mesh_octadr, + plugin, + plugin_attr, + contact_type, + geom_size[worldid % geom_size.shape[0], geom], + plugin_id, + geom_dataid[geom], ) - depth = wp.min( - sdf( - contact_type, - lpos, - plugin_attributes, - plugin_index, - volume_data, - mesh_data, - ), - 0.0, - ) + depth = wp.min(sdf(contact_type, lpos, plugin_attributes, plugin_index, volume_data, mesh_data), 0.0) if depth >= 0.0: continue - vel_sensor = _transform_spatial( - cvel_in[worldid, parent_weld], - xpos - subtree_com_in[worldid, body_rootid[parent_weld]], - ) + vel_sensor = _transform_spatial(cvel_in[worldid, parent_weld], xpos - subtree_com_in[worldid, body_rootid[parent_weld]]) vel_other = _transform_spatial( - cvel_in[worldid, body], - geom_xpos_in[worldid, geom] - - subtree_com_in[worldid, body_rootid[body]], + cvel_in[worldid, body], geom_xpos_in[worldid, geom] - subtree_com_in[worldid, body_rootid[body]] ) vel_rel = vel_sensor - vel_other @@ -2259,24 +2247,9 @@ def _sensor_tactile( forceT[2] = wp.abs(wp.dot(vel_rel, tang2)) dim = sensor_dim[sensor_id] // 3 - wp.atomic_add( - sensordata_out, - worldid, - sensor_adr[sensor_id] + 0 * dim + vertid, - forceT[0], - ) - wp.atomic_add( - sensordata_out, - worldid, - sensor_adr[sensor_id] + 1 * dim + vertid, - forceT[1], - ) - wp.atomic_add( - sensordata_out, - worldid, - sensor_adr[sensor_id] + 2 * dim + vertid, - forceT[2], - ) + wp.atomic_add(sensordata_out, worldid, sensor_adr[sensor_id] + 0 * dim + vertid, forceT[0]) + wp.atomic_add(sensordata_out, worldid, sensor_adr[sensor_id] + 1 * dim + vertid, forceT[1]) + wp.atomic_add(sensordata_out, worldid, sensor_adr[sensor_id] + 2 * dim + vertid, forceT[2]) @wp.func @@ -2507,60 +2480,61 @@ def sensor_acc(m: Model, d: Data): weld_geom_count = wp.zeros((d.nworld, m.nbody), dtype=int) weld_geom_list = wp.full((d.nworld, m.nbody, MJ_MAXCONPAIR), -1, dtype=int) wp.launch( - _preprocess_tactile_contacts, - dim=d.naconmax, - inputs=[ - m.body_weldid, - m.geom_bodyid, - d.contact.geom, - d.contact.worldid, - d.nacon, - ], - outputs=[ - weld_geom_count, - weld_geom_list, - ], + _preprocess_tactile_contacts, + dim=d.naconmax, + inputs=[ + m.body_weldid, + m.geom_bodyid, + d.contact.geom, + d.contact.worldid, + d.nacon, + ], + outputs=[ + weld_geom_count, + weld_geom_list, + ], ) wp.launch( - _sensor_tactile, - dim=(d.nworld, m.nsensortaxel), - inputs=[ - m.body_rootid, - m.body_weldid, - m.oct_child, - m.oct_aabb, - m.oct_coeff, - m.geom_type, - m.geom_bodyid, - m.geom_size, - m.mesh_vertadr, - m.mesh_vertnum, - m.mesh_octadr, - m.mesh_normaladr, - m.mesh_normalnum, - m.mesh_vert, - m.mesh_normal, - m.mesh_quat, - m.sensor_objid, - m.sensor_refid, - m.sensor_dim, - m.sensor_adr, - m.plugin, - m.plugin_attr, - m.geom_plugin_index, - m.taxel_vertadr, - m.taxel_sensorid, - d.geom_xpos, - d.geom_xmat, - d.subtree_com, - d.cvel, - weld_geom_count, - weld_geom_list, - ], - outputs=[ - d.sensordata, - ], + _sensor_tactile, + dim=(d.nworld, m.nsensortaxel), + inputs=[ + m.body_rootid, + m.body_weldid, + m.oct_child, + m.oct_aabb, + m.oct_coeff, + m.geom_type, + m.geom_bodyid, + m.geom_dataid, + m.geom_size, + m.mesh_vertadr, + m.mesh_vertnum, + m.mesh_octadr, + m.mesh_normaladr, + m.mesh_normalnum, + m.mesh_vert, + m.mesh_normal, + m.mesh_quat, + m.sensor_objid, + m.sensor_refid, + m.sensor_dim, + m.sensor_adr, + m.plugin, + m.plugin_attr, + m.geom_plugin_index, + m.taxel_vertadr, + m.taxel_sensorid, + d.geom_xpos, + d.geom_xmat, + d.subtree_com, + d.cvel, + weld_geom_count, + weld_geom_list, + ], + outputs=[ + d.sensordata, + ], ) sensor_contact_nmatch = wp.empty((d.nworld, m.nsensorcontact), dtype=int) @@ -2882,12 +2856,12 @@ def energy_pos(m: Model, d: Data): wp.launch(_energy_pos_zero, dim=d.nworld, outputs=[d.energy]) # init potential energy: -sum_i(body_i.mass * dot(gravity, body_i.pos)) - if not m.opt.disableflags & DisableBit.GRAVITY: + if not (m.opt.disableflags & DisableBit.GRAVITY): wp.launch( _energy_pos_gravity, dim=(d.nworld, m.nbody - 1), inputs=[m.opt.gravity, m.body_mass, d.xipos], outputs=[d.energy] ) - if not m.opt.disableflags & DisableBit.SPRING: + if not (m.opt.disableflags & DisableBit.SPRING): # add joint-level springs wp.launch( _energy_pos_passive_joint, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py index 51b36403..51dfadb4 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py @@ -14,29 +14,29 @@ # ============================================================================== +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src import math from mujoco.mjx.third_party.mujoco_warp._src import support from mujoco.mjx.third_party.mujoco_warp._src import util_misc +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import CamLightType from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit from mujoco.mjx.third_party.mujoco_warp._src.types import EqType from mujoco.mjx.third_party.mujoco_warp._src.types import JointType -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL -from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import Model from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType -from mujoco.mjx.third_party.mujoco_warp._src.types import SPARSE_CONSTRAINT_JACOBIAN from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType +from mujoco.mjx.third_party.mujoco_warp._src.types import WrapType +from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 from mujoco.mjx.third_party.mujoco_warp._src.types import vec10 from mujoco.mjx.third_party.mujoco_warp._src.types import vec11 -from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 -from mujoco.mjx.third_party.mujoco_warp._src.types import WrapType from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -227,38 +227,59 @@ def _site_local_to_global( @wp.kernel def _flex_vertices( # Model: + nflex: int, + flex_vertadr: wp.array(dtype=int), + flex_vertnum: wp.array(dtype=int), flex_vertbodyid: wp.array(dtype=int), + flex_vert: wp.array(dtype=wp.vec3), + flex_centered: wp.array(dtype=bool), # Data in: xpos_in: wp.array2d(dtype=wp.vec3), + xmat_in: wp.array2d(dtype=wp.mat33), # Data out: flexvert_xpos_out: wp.array2d(dtype=wp.vec3), ): worldid, vertid = wp.tid() - flexvert_xpos_out[worldid, vertid] = xpos_in[worldid, flex_vertbodyid[vertid]] + + for f in range(nflex): + locid = vertid - flex_vertadr[f] + if locid >= 0 and locid < flex_vertnum[f]: + break + + bodyid = flex_vertbodyid[vertid] + xpos = xpos_in[worldid, bodyid] + + if flex_centered[f]: + flexvert_xpos_out[worldid, vertid] = xpos + else: + xmat = xmat_in[worldid, bodyid] + local_pos = flex_vert[vertid] + flexvert_xpos_out[worldid, vertid] = xmat @ local_pos + xpos @wp.kernel def _flex_edges( - # Model: - nflex: int, - body_rootid: wp.array(dtype=int), - body_dofadr: wp.array(dtype=int), - flex_vertadr: wp.array(dtype=int), - flex_edgeadr: wp.array(dtype=int), - flex_edgenum: wp.array(dtype=int), - flex_vertbodyid: wp.array(dtype=int), - flex_edge: wp.array(dtype=wp.vec2i), - flexedge_J_rowadr: wp.array(dtype=int), - flexedge_J_colind: wp.array(dtype=int), - # Data in: - qvel_in: wp.array2d(dtype=float), - subtree_com_in: wp.array2d(dtype=wp.vec3), - cdof_in: wp.array2d(dtype=wp.spatial_vector), - flexvert_xpos_in: wp.array2d(dtype=wp.vec3), - # Data out: - flexedge_J_out: wp.array2d(dtype=float), - flexedge_length_out: wp.array2d(dtype=float), - flexedge_velocity_out: wp.array2d(dtype=float), + # Model: + nflex: int, + body_rootid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_edgeadr: wp.array(dtype=int), + flex_edgenum: wp.array(dtype=int), + flex_vertbodyid: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), + flexedge_J_rowadr: wp.array(dtype=int), + flexedge_J_colind: wp.array(dtype=int), + # Data in: + qvel_in: wp.array2d(dtype=float), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + # Data out: + flexedge_J_out: wp.array2d(dtype=float), + flexedge_length_out: wp.array2d(dtype=float), + flexedge_velocity_out: wp.array2d(dtype=float), ): worldid, edgeid = wp.tid() for i in range(nflex): @@ -281,42 +302,56 @@ def _flex_edges( b1 = flex_vertbodyid[vbase0] b2 = flex_vertbodyid[vbase1] - dofi = body_dofadr[b1] - dofj = body_dofadr[b2] + dofnum1 = body_dofnum[b1] + dofnum2 = body_dofnum[b2] - vel1 = wp.vec3( - qvel_in[worldid, dofi], - qvel_in[worldid, dofi + 1], - qvel_in[worldid, dofi + 2], - ) - vel2 = wp.vec3( - qvel_in[worldid, dofj], - qvel_in[worldid, dofj + 1], - qvel_in[worldid, dofj + 2], - ) - flexedge_velocity_out[worldid, edgeid] = wp.dot(vel2 - vel1, edge) + # velocity via Jacobian: sum_k J_k * qvel_k for each body + vel = float(0.0) + if dofnum1 > 0: + dofi = body_dofadr[b1] + offset1 = pos1 - wp.vec3(subtree_com_in[worldid, body_rootid[b1]]) + for k in range(dofnum1): + cdof = cdof_in[worldid, dofi + k] + cdof_ang = wp.spatial_top(cdof) + cdof_lin = wp.spatial_bottom(cdof) + jacp1 = cdof_lin + wp.cross(cdof_ang, offset1) + vel -= wp.dot(jacp1, edge) * qvel_in[worldid, dofi + k] + if dofnum2 > 0: + dofj = body_dofadr[b2] + offset2 = pos2 - wp.vec3(subtree_com_in[worldid, body_rootid[b2]]) + for k in range(dofnum2): + cdof = cdof_in[worldid, dofj + k] + cdof_ang = wp.spatial_top(cdof) + cdof_lin = wp.spatial_bottom(cdof) + jacp2 = cdof_lin + wp.cross(cdof_ang, offset2) + vel += wp.dot(jacp2, edge) * qvel_in[worldid, dofj + k] + flexedge_velocity_out[worldid, edgeid] = vel rowadr = flexedge_J_rowadr[edgeid] - - # compute offsets once per body (avoids 12 redundant tree-ancestry walks in jac_dof) - offset1 = pos1 - wp.vec3(subtree_com_in[worldid, body_rootid[b1]]) - offset2 = pos2 - wp.vec3(subtree_com_in[worldid, body_rootid[b2]]) + nnz_offset = 0 # body1 DOFs: b1 is in subtree, b2 is not -> jacdif = 0 - jacp1 = -jacp1 - for k in range(3): - cdof = cdof_in[worldid, dofi + k] - cdof_ang = wp.spatial_top(cdof) - cdof_lin = wp.spatial_bottom(cdof) - jacp1 = cdof_lin + wp.cross(cdof_ang, offset1) - flexedge_J_out[worldid, rowadr + k] = wp.dot(-jacp1, edge) + if dofnum1 > 0: + dofi = body_dofadr[b1] + offset1 = pos1 - wp.vec3(subtree_com_in[worldid, body_rootid[b1]]) + for k in range(dofnum1): + cdof = cdof_in[worldid, dofi + k] + cdof_ang = wp.spatial_top(cdof) + cdof_lin = wp.spatial_bottom(cdof) + jacp1 = cdof_lin + wp.cross(cdof_ang, offset1) + flexedge_J_out[worldid, rowadr + nnz_offset + k] = wp.dot(-jacp1, edge) + nnz_offset += dofnum1 # body2 DOFs: b2 is in subtree, b1 is not -> jacdif = jacp2 - 0 = jacp2 - for k in range(3): - cdof = cdof_in[worldid, dofj + k] - cdof_ang = wp.spatial_top(cdof) - cdof_lin = wp.spatial_bottom(cdof) - jacp2 = cdof_lin + wp.cross(cdof_ang, offset2) - flexedge_J_out[worldid, rowadr + 3 + k] = wp.dot(jacp2, edge) + if dofnum2 > 0: + dofj = body_dofadr[b2] + offset2 = pos2 - wp.vec3(subtree_com_in[worldid, body_rootid[b2]]) + for k in range(dofnum2): + cdof = cdof_in[worldid, dofj + k] + cdof_ang = wp.spatial_top(cdof) + cdof_lin = wp.spatial_bottom(cdof) + jacp2 = cdof_lin + wp.cross(cdof_ang, offset2) + flexedge_J_out[worldid, rowadr + nnz_offset + k] = wp.dot(jacp2, edge) @event_scope @@ -382,13 +417,28 @@ def kinematics(m: Model, d: Data): @event_scope def flex(m: Model, d: Data): - wp.launch(_flex_vertices, dim=(d.nworld, m.nflexvert), inputs=[m.flex_vertbodyid, d.xpos], outputs=[d.flexvert_xpos]) + wp.launch( + _flex_vertices, + dim=(d.nworld, m.nflexvert), + inputs=[ + m.nflex, + m.flex_vertadr, + m.flex_vertnum, + m.flex_vertbodyid, + m.flex_vert, + m.flex_centered, + d.xpos, + d.xmat, + ], + outputs=[d.flexvert_xpos], + ) wp.launch( _flex_edges, dim=(d.nworld, m.nflexedge), inputs=[ m.nflex, m.body_rootid, + m.body_dofnum, m.body_dofadr, m.flex_vertadr, m.flex_edgeadr, @@ -790,9 +840,7 @@ def _qM_sparse( bodyid = dof_bodyid[dofid] # init M(i,i) with armature inertia - qM_out[worldid, 0, madr_ij] = dof_armature[ - worldid % dof_armature.shape[0], dofid - ] + qM_out[worldid, 0, madr_ij] = dof_armature[worldid % dof_armature.shape[0], dofid] # precompute buf = crb_body_i * cdof_i buf = math.inert_vec(crb_in[worldid, bodyid], cdof_in[worldid, dofid]) @@ -869,35 +917,55 @@ def _tendon_armature( # Model: dof_parentid: wp.array(dtype=int), dof_Madr: wp.array(dtype=int), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), tendon_armature: wp.array2d(dtype=float), is_sparse: bool, # Data in: - ten_J_in: wp.array3d(dtype=float), + ten_J_in: wp.array2d(dtype=float), # Data out: qM_out: wp.array3d(dtype=float), ): worldid, tenid, dofid = wp.tid() - if is_sparse: # is_sparse is not batched - madr_ij = dof_Madr[dofid] - armature = tendon_armature[worldid % tendon_armature.shape[0], tenid] if armature == 0.0: return - ten_Ji = ten_J_in[worldid, tenid, dofid] + rownnz = ten_J_rownnz[tenid] + if dofid >= rownnz: + return + rowadr = ten_J_rowadr[tenid] + dofid_sparse = dofid + sparseid = rowadr + dofid_sparse + dofid = ten_J_colind[sparseid] + ten_Ji = ten_J_in[worldid, sparseid] if ten_Ji == 0.0: return + if is_sparse: + madr_ij = dof_Madr[dofid] + # sparse backward pass over ancestors dofidi = dofid + ptr = dofid_sparse while dofid >= 0: - if dofid != dofidi: - ten_Jj = ten_J_in[worldid, tenid, dofid] - else: + if dofid == dofidi: ten_Jj = ten_Ji + else: + # scan pointer backward to find matching colind entry + while ptr >= 0: + sparseid = rowadr + ptr + if ten_J_colind[sparseid] <= dofid: + break + ptr -= 1 + if ptr >= 0 and ten_J_colind[sparseid] == dofid: + ten_Jj = ten_J_in[worldid, sparseid] + else: + ten_Jj = float(0.0) qMij = armature * ten_Jj * ten_Ji @@ -917,8 +985,17 @@ def tendon_armature(m: Model, d: Data): """Add tendon armature to qM.""" wp.launch( _tendon_armature, - dim=(d.nworld, m.ntendon, m.nv), - inputs=[m.dof_parentid, m.dof_Madr, m.tendon_armature, m.is_sparse, d.ten_J], + dim=(d.nworld, m.ntendon, m.max_ten_J_rownnz), + inputs=[ + m.dof_parentid, + m.dof_Madr, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, + m.tendon_armature, + m.is_sparse, + d.ten_J, + ], outputs=[d.qM], ) @@ -1504,19 +1581,93 @@ def rne_postconstraint(m: Model, d: Data): _rne_cfrc_backward(m, d) +@wp.func +def _accumulate_jac_dot_chain( + # Model: + body_parentid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_jntid: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + # Data in: + cdof_in: wp.array2d(dtype=wp.spatial_vector), + cvel_in: wp.array2d(dtype=wp.spatial_vector), + cdof_dot_in: wp.array2d(dtype=wp.spatial_vector), + # In: + offset: wp.vec3, + pvel_lin: wp.vec3, + dpnt: wp.vec3, + dvel: wp.vec3, + bodyid: int, + rowadr: int, + rownnz: int, + scale: float, + worldid: int, + # Out: + ten_Jdot_out: wp.array2d(dtype=float), +): + """Walk body chain from bodyid to root, accumulate Jdot contributions.""" + ptr = rownnz - 1 + bid = bodyid + while bid > 0: + bdofadr = body_dofadr[bid] + bdofnum = body_dofnum[bid] + # iterate DOFs in this body in descending order + for k_rev in range(bdofnum): + dof = bdofadr + bdofnum - 1 - k_rev + # scan pointer backward to find matching colind entry + while ptr >= 0: + sparseid = rowadr + ptr + if ten_J_colind[sparseid] <= dof: + break + ptr -= 1 + if ptr >= 0 and ten_J_colind[sparseid] == dof: + cdof = cdof_in[worldid, dof] + cdof_ang = wp.spatial_top(cdof) + cdof_lin = wp.spatial_bottom(cdof) + cdof_dot = cdof_dot_in[worldid, dof] + + # quaternion override: use cvel of DOF's body (which is bid) + dofjntid = dof_jntid[dof] + jnttype = jnt_type[dofjntid] + jntdofadr = jnt_dofadr[dofjntid] + if (jnttype == JointType.BALL) or ((jnttype == JointType.FREE) and dof >= jntdofadr + 3): + cdof_dot = math.motion_cross(cvel_in[worldid, bid], cdof) + + cdof_dot_ang = wp.spatial_top(cdof_dot) + cdof_dot_lin = wp.spatial_bottom(cdof_dot) + + # jacp_dot (from jac_dot_dof) + jacp_dot = cdof_dot_lin + wp.cross(cdof_dot_ang, offset) + wp.cross(cdof_ang, pvel_lin) + + # jacp (from jac_dof) + jacp = cdof_lin + wp.cross(cdof_ang, offset) + + # combined: dot(jacdot, dpnt) + dot(jac, dvel) + Jdot = (wp.dot(jacp_dot, dpnt) + wp.dot(jacp, dvel)) * scale + if Jdot != 0.0: + wp.atomic_add(ten_Jdot_out[worldid], sparseid, Jdot) + bid = body_parentid[bid] + + @wp.kernel def _tendon_dot( # Model: - nv: int, body_parentid: wp.array(dtype=int), body_rootid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), jnt_type: wp.array(dtype=int), jnt_dofadr: wp.array(dtype=int), - dof_bodyid: wp.array(dtype=int), dof_jntid: wp.array(dtype=int), site_bodyid: wp.array(dtype=int), tendon_adr: wp.array(dtype=int), tendon_num: wp.array(dtype=int), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), tendon_armature: wp.array2d(dtype=float), wrap_type: wp.array(dtype=int), wrap_objid: wp.array(dtype=int), @@ -1528,7 +1679,7 @@ def _tendon_dot( cvel_in: wp.array2d(dtype=wp.spatial_vector), cdof_dot_in: wp.array2d(dtype=wp.spatial_vector), # Out: - ten_Jdot_out: wp.array3d(dtype=float), + ten_Jdot_out: wp.array2d(dtype=float), ): worldid, tenid = wp.tid() @@ -1565,13 +1716,11 @@ def _tendon_dot( # init sequence; assume it start with site wpnt0 = site_xpos_in[worldid, id0] - bodyid0 = site_bodyid[id0] - pos0 = site_xpos_in[worldid, id0] - cvel0 = cvel_in[worldid, bodyid0] - subtree_com0 = subtree_com_in[worldid, body_rootid[bodyid0]] - dif0 = pos0 - subtree_com0 - wvel0 = wp.spatial_bottom(cvel0) - wp.cross(dif0, wp.spatial_top(cvel0)) wbody0 = site_bodyid[id0] + cvel0 = cvel_in[worldid, wbody0] + subtree_com0 = subtree_com_in[worldid, body_rootid[wbody0]] + offset0 = wpnt0 - subtree_com0 + pvel_lin0 = wp.spatial_bottom(cvel0) - wp.cross(offset0, wp.spatial_top(cvel0)) # second object is geom: process site-geom-site if (type1 == WrapType.SPHERE) or (type1 == WrapType.CYLINDER): @@ -1582,12 +1731,10 @@ def _tendon_dot( wbody1 = site_bodyid[id1] wpnt1 = site_xpos_in[worldid, id1] - bodyid1 = site_bodyid[id1] - pos1 = site_xpos_in[worldid, id1] - cvel1 = cvel_in[worldid, bodyid1] - subtree_com1 = subtree_com_in[worldid, body_rootid[bodyid1]] - dif1 = pos1 - subtree_com1 - wvel1 = wp.spatial_bottom(cvel1) - wp.cross(dif1, wp.spatial_top(cvel1)) + cvel1 = cvel_in[worldid, wbody1] + subtree_com1 = subtree_com_in[worldid, body_rootid[wbody1]] + offset1 = wpnt1 - subtree_com1 + pvel_lin1 = wp.spatial_bottom(cvel1) - wp.cross(offset1, wp.spatial_top(cvel1)) # accumulate moments if consecutive points are in different bodies if wbody0 != wbody1: @@ -1595,6 +1742,8 @@ def _tendon_dot( dpnt, norm = math.normalize_with_norm(wpnt1 - wpnt0) # dvel = d / dt (dpnt) + wvel0 = wp.spatial_bottom(cvel0) - wp.cross(wpnt0 - subtree_com0, wp.spatial_top(cvel0)) + wvel1 = wp.spatial_bottom(cvel1) - wp.cross(wpnt1 - subtree_com1, wp.spatial_top(cvel1)) dvel = wvel1 - wvel0 dot = wp.dot(dpnt, dvel) dvel += dpnt * (-dot) @@ -1603,75 +1752,55 @@ def _tendon_dot( else: dvel = wp.vec3(0.0) - # get endpoint Jacobian time derivatives, subtract - # TODO(team): parallelize? - for i in range(nv): - jac1, _ = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - wpnt0, - wbody0, - i, - worldid, - ) - jac2, _ = support.jac_dot_dof( - body_parentid, - body_rootid, - jnt_type, - jnt_dofadr, - dof_bodyid, - dof_jntid, - subtree_com_in, - cdof_in, - cvel_in, - cdof_dot_in, - wpnt1, - wbody1, - i, - worldid, - ) - jacdif = jac2 - jac1 + rownnz = ten_J_rownnz[tenid] + rowadr = ten_J_rowadr[tenid] + inv_divisor = math.safe_div(float(1.0), divisor) - # chain rule, first term: Jdot += d / dt (jac2 - jac1) * dpnt - Jdot = wp.dot(jacdif, dpnt) - - # get endpoint Jacobians, subtract - jac1, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - wpnt0, - wbody0, - i, - worldid, - ) - jac2, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - wpnt1, - wbody1, - i, - worldid, - ) - jacdif = jac2 - jac1 - - # chain rule, second term: Jdot += (jac2 - jac1) * d / dt (dpnt) - Jdot += wp.dot(jacdif, dvel) - - ten_Jdot_out[worldid, tenid, i] += math.safe_div(Jdot, divisor) + # body0 contributes with negative sign, body1 with positive + _accumulate_jac_dot_chain( + body_parentid, + body_dofnum, + body_dofadr, + jnt_type, + jnt_dofadr, + dof_jntid, + ten_J_colind, + cdof_in, + cvel_in, + cdof_dot_in, + offset0, + pvel_lin0, + dpnt, + dvel, + wbody0, + rowadr, + rownnz, + -inv_divisor, + worldid, + ten_Jdot_out, + ) + _accumulate_jac_dot_chain( + body_parentid, + body_dofnum, + body_dofadr, + jnt_type, + jnt_dofadr, + dof_jntid, + ten_J_colind, + cdof_in, + cvel_in, + cdof_dot_in, + offset1, + pvel_lin1, + dpnt, + dvel, + wbody1, + rowadr, + rownnz, + inv_divisor, + worldid, + ten_Jdot_out, + ) # TODO(team): j += 2 if geom wrapping j += 1 @@ -1680,33 +1809,45 @@ def _tendon_dot( @wp.kernel def _tendon_bias_coef( # Model: + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), tendon_armature: wp.array2d(dtype=float), # Data in: qvel_in: wp.array2d(dtype=float), # In: - ten_Jdot_in: wp.array3d(dtype=float), + ten_Jdot_in: wp.array2d(dtype=float), # Out: ten_bias_coef_out: wp.array2d(dtype=float), ): - worldid, tenid, dofid = wp.tid() + worldid, tenid, dofid_sparse = wp.tid() armature = tendon_armature[worldid % tendon_armature.shape[0], tenid] if armature == 0.0: return - ten_Jdot = ten_Jdot_in[worldid, tenid, dofid] + rownnz = ten_J_rownnz[tenid] + if dofid_sparse >= rownnz: + return + rowadr = ten_J_rowadr[tenid] + sparseid = rowadr + dofid_sparse + ten_Jdot = ten_Jdot_in[worldid, sparseid] if ten_Jdot == 0.0: return + dofid = ten_J_colind[sparseid] wp.atomic_add(ten_bias_coef_out[worldid], tenid, ten_Jdot * qvel_in[worldid, dofid]) @wp.kernel def _tendon_bias_qfrc( # Model: + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), tendon_armature: wp.array2d(dtype=float), # Data in: - ten_J_in: wp.array3d(dtype=float), + ten_J_in: wp.array2d(dtype=float), # In: ten_bias_coef_in: wp.array2d(dtype=float), # Out: @@ -1718,10 +1859,18 @@ def _tendon_bias_qfrc( if armature == 0.0: return - ten_J = ten_J_in[worldid, tenid, dofid] + rownnz = ten_J_rownnz[tenid] + if dofid >= rownnz: + return + rowadr = ten_J_rowadr[tenid] + sparseid = rowadr + dofid + ten_J = ten_J_in[worldid, sparseid] + if ten_J == 0.0: return + dofid = ten_J_colind[sparseid] + wp.atomic_add(qfrc_out[worldid], dofid, ten_J * armature * ten_bias_coef_in[worldid, tenid]) @@ -1735,21 +1884,24 @@ def tendon_bias(m: Model, d: Data, qfrc: wp.array2d(dtype=float)): qfrc: Force. """ # time derivative of tendon Jacobian - ten_Jdot = wp.zeros((d.nworld, m.ntendon, m.nv), dtype=float) + ten_Jdot = wp.zeros((d.nworld, m.nJten), dtype=float) wp.launch( _tendon_dot, dim=(d.nworld, m.ntendon), inputs=[ - m.nv, m.body_parentid, m.body_rootid, + m.body_dofnum, + m.body_dofadr, m.jnt_type, m.jnt_dofadr, - m.dof_bodyid, m.dof_jntid, m.site_bodyid, m.tendon_adr, m.tendon_num, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, m.tendon_armature, m.wrap_type, m.wrap_objid, @@ -1767,15 +1919,15 @@ def tendon_bias(m: Model, d: Data, qfrc: wp.array2d(dtype=float)): ten_bias_coef = wp.zeros((d.nworld, m.ntendon), dtype=float) wp.launch( _tendon_bias_coef, - dim=(d.nworld, m.ntendon, m.nv), - inputs=[m.tendon_armature, d.qvel, ten_Jdot], + dim=(d.nworld, m.ntendon, m.max_ten_J_rownnz), + inputs=[m.ten_J_rownnz, m.ten_J_rowadr, m.ten_J_colind, m.tendon_armature, d.qvel, ten_Jdot], outputs=[ten_bias_coef], ) wp.launch( _tendon_bias_qfrc, - dim=(d.nworld, m.ntendon, m.nv), - inputs=[m.tendon_armature, d.ten_J, ten_bias_coef], + dim=(d.nworld, m.ntendon, m.max_ten_J_rownnz), + inputs=[m.ten_J_rownnz, m.ten_J_rowadr, m.ten_J_colind, m.tendon_armature, d.ten_J, ten_bias_coef], outputs=[qfrc], ) @@ -1888,45 +2040,44 @@ def com_vel(m: Model, d: Data): @wp.kernel def _transmission( - # Model: - nv: int, - body_parentid: wp.array(dtype=int), - body_rootid: wp.array(dtype=int), - body_weldid: wp.array(dtype=int), - body_dofnum: wp.array(dtype=int), - body_dofadr: wp.array(dtype=int), - jnt_type: wp.array(dtype=int), - jnt_qposadr: wp.array(dtype=int), - jnt_dofadr: wp.array(dtype=int), - dof_bodyid: wp.array(dtype=int), - dof_parentid: wp.array(dtype=int), - site_bodyid: wp.array(dtype=int), - site_quat: wp.array2d(dtype=wp.quat), - tendon_adr: wp.array(dtype=int), - tendon_num: wp.array(dtype=int), - wrap_type: wp.array(dtype=int), - wrap_objid: wp.array(dtype=int), - actuator_trntype: wp.array(dtype=int), - actuator_trnid: wp.array(dtype=wp.vec2i), - actuator_gear: wp.array2d(dtype=wp.spatial_vector), - actuator_cranklength: wp.array2d(dtype=float), - # Data in: - qpos_in: wp.array2d(dtype=float), - xquat_in: wp.array2d(dtype=wp.quat), - site_xpos_in: wp.array2d(dtype=wp.vec3), - site_xmat_in: wp.array2d(dtype=wp.mat33), - subtree_com_in: wp.array2d(dtype=wp.vec3), - cdof_in: wp.array2d(dtype=wp.spatial_vector), - ten_J_in: wp.array3d(dtype=float), - ten_length_in: wp.array2d(dtype=float), - # In: - moment_nnz: wp.array(dtype=int), - # Data out: - actuator_length_out: wp.array2d(dtype=float), - moment_rownnz_out: wp.array2d(dtype=int), - moment_rowadr_out: wp.array2d(dtype=int), - moment_colind_out: wp.array2d(dtype=int), - actuator_moment_out: wp.array2d(dtype=float), + # Model: + nv: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_weldid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + jnt_qposadr: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + site_quat: wp.array2d(dtype=wp.quat), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + actuator_trntype: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + actuator_gear: wp.array2d(dtype=wp.spatial_vector), + actuator_cranklength: wp.array2d(dtype=float), + # Data in: + qpos_in: wp.array2d(dtype=float), + xquat_in: wp.array2d(dtype=wp.quat), + site_xpos_in: wp.array2d(dtype=wp.vec3), + site_xmat_in: wp.array2d(dtype=wp.mat33), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + ten_J_in: wp.array2d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + # In: + moment_nnz: wp.array(dtype=int), + # Data out: + actuator_length_out: wp.array2d(dtype=float), + moment_rownnz_out: wp.array2d(dtype=int), + moment_rowadr_out: wp.array2d(dtype=int), + moment_colind_out: wp.array2d(dtype=int), + actuator_moment_out: wp.array2d(dtype=float), ): worldid, actid = wp.tid() trntype = actuator_trntype[actid] @@ -2068,28 +2219,12 @@ def _transmission( # get Jacobians of axis(jacA) and vec(jac) jacp, jacr = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - site_xpos_idslider, - site_bodyid[idslider], - da, - worldid, + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_idslider, site_bodyid[idslider], da, worldid ) jacS = jacp jacA = wp.cross(jacr, axis) jac, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - site_xpos_id, - site_bodyid[id], - da, - worldid, + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_id, site_bodyid[id], da, worldid ) jac -= jacS @@ -2110,38 +2245,18 @@ def _transmission( gear0 = gear[0] actuator_length_out[worldid, actid] = ten_length_in[worldid, tenid] * gear0 - # fixed - adr = tendon_adr[tenid] - if wrap_type[adr] == WrapType.JOINT: - ten_num = tendon_num[tenid] - rowadr = wp.atomic_add(moment_nnz, worldid, ten_num) - moment_rownnz_out[worldid, actid] = ten_num - moment_rowadr_out[worldid, actid] = rowadr + rownnz_ten = ten_J_rownnz[tenid] + rowadr_ten = ten_J_rowadr[tenid] - for i in range(ten_num): - dofadr = jnt_dofadr[wrap_objid[adr + i]] - sparseid = rowadr + i - moment_colind_out[worldid, sparseid] = dofadr - actuator_moment_out[worldid, sparseid] = ( - ten_J_in[worldid, tenid, dofadr] * gear0 - ) - else: # spatial - # TODO(team): sparse tendon jacobian - ten_nnz = int(0) - for dofadr in range(nv): - if ten_J_in[worldid, tenid, dofadr] != 0.0: - ten_nnz += 1 - rowadr = wp.atomic_add(moment_nnz, worldid, ten_nnz) - moment_rownnz_out[worldid, actid] = ten_nnz - moment_rowadr_out[worldid, actid] = rowadr - ptr = int(0) - for dofadr in range(nv): - J = ten_J_in[worldid, tenid, dofadr] - if J != 0.0: - sparseid = rowadr + ptr - moment_colind_out[worldid, sparseid] = dofadr - actuator_moment_out[worldid, sparseid] = J * gear0 - ptr += 1 + rowadr_mom = wp.atomic_add(moment_nnz, worldid, rownnz_ten) + moment_rownnz_out[worldid, actid] = rownnz_ten + moment_rowadr_out[worldid, actid] = rowadr_mom + + for k in range(rownnz_ten): + sparseid_ten = rowadr_ten + k + sparseid_mom = rowadr_mom + k + moment_colind_out[worldid, sparseid_mom] = ten_J_colind[sparseid_ten] + actuator_moment_out[worldid, sparseid_mom] = ten_J_in[worldid, sparseid_ten] * gear0 elif trntype == TrnType.BODY: # cannot compute meaningful length, set to zero actuator_length_out[worldid, actid] = 0.0 @@ -2195,19 +2310,17 @@ def _transmission( ptr = ndof - 1 while da >= 0: jacp, jacr = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - site_xpos_in[worldid, siteid], - site_bodyid[siteid], - da, - worldid, - ) - moment = wp.dot(jacp, wrench_translation) + wp.dot( - jacr, wrench_rotation + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + site_xpos_in[worldid, siteid], + site_bodyid[siteid], + da, + worldid, ) + moment = wp.dot(jacp, wrench_translation) + wp.dot(jacr, wrench_rotation) sparseid = rowadr + ptr moment_colind_out[worldid, sparseid] = da actuator_moment_out[worldid, sparseid] = moment @@ -2306,26 +2419,10 @@ def _transmission( break jacp, jacr = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - site_xpos, - site_bodyid[siteid], - da, - worldid, + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos, site_bodyid[siteid], da, worldid ) jacpref, jacrref = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - ref_xpos, - site_bodyid[refid], - da, - worldid, + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, ref_xpos, site_bodyid[refid], da, worldid ) moment = float(0.0) @@ -2349,37 +2446,37 @@ def _transmission( @wp.kernel def _transmission_body_moment( - # Model: - opt_cone: int, - body_parentid: wp.array(dtype=int), - body_rootid: wp.array(dtype=int), - dof_bodyid: wp.array(dtype=int), - geom_bodyid: wp.array(dtype=int), - actuator_trnid: wp.array(dtype=wp.vec2i), - actuator_trntype_body_adr: wp.array(dtype=int), - # Data in: - subtree_com_in: wp.array2d(dtype=wp.vec3), - cdof_in: wp.array2d(dtype=wp.spatial_vector), - moment_rowadr_in: wp.array2d(dtype=int), - contact_dist_in: wp.array(dtype=float), - contact_pos_in: wp.array(dtype=wp.vec3), - contact_frame_in: wp.array(dtype=wp.mat33), - contact_includemargin_in: wp.array(dtype=float), - contact_dim_in: wp.array(dtype=int), - contact_geom_in: wp.array(dtype=wp.vec2i), - contact_efc_address_in: wp.array2d(dtype=int), - contact_worldid_in: wp.array(dtype=int), - efc_J_rownnz_in: wp.array2d(dtype=int), - efc_J_rowadr_in: wp.array2d(dtype=int), - efc_J_colind_in: wp.array3d(dtype=int), - efc_J_in: wp.array3d(dtype=float), - nacon_in: wp.array(dtype=int), - # In: - efc_is_sparse: bool, - # Data out: - actuator_moment_out: wp.array2d(dtype=float), - # Out: - actuator_trntype_body_ncon_out: wp.array2d(dtype=int), + # Model: + opt_cone: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + actuator_trntype_body_adr: wp.array(dtype=int), + # Data in: + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + moment_rowadr_in: wp.array2d(dtype=int), + contact_dist_in: wp.array(dtype=float), + contact_pos_in: wp.array(dtype=wp.vec3), + contact_frame_in: wp.array(dtype=wp.mat33), + contact_includemargin_in: wp.array(dtype=float), + contact_dim_in: wp.array(dtype=int), + contact_geom_in: wp.array(dtype=wp.vec2i), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), + efc_J_colind_in: wp.array3d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + nacon_in: wp.array(dtype=int), + # In: + efc_is_sparse: bool, + # Data out: + actuator_moment_out: wp.array2d(dtype=float), + # Out: + actuator_trntype_body_ncon_out: wp.array2d(dtype=int), ): trnbodyid, conid, dofid = wp.tid() actid = actuator_trntype_body_adr[trnbodyid] @@ -2427,20 +2524,12 @@ def _transmission_body_moment( efc_rowadr = efc_J_rowadr_in[worldid, efcid0] efc_sparseid = efc_rowadr + dofid colind = efc_J_colind_in[worldid, 0, efc_sparseid] - wp.atomic_add( - actuator_moment_out[worldid], - rowadr + colind, - efc_J_in[worldid, 0, efc_sparseid], - ) + wp.atomic_add(actuator_moment_out[worldid], rowadr + colind, efc_J_in[worldid, 0, efc_sparseid]) else: return else: colind = dofid - wp.atomic_add( - actuator_moment_out[worldid], - rowadr + colind, - efc_J_in[worldid, efcid0, dofid], - ) + wp.atomic_add(actuator_moment_out[worldid], rowadr + colind, efc_J_in[worldid, efcid0, dofid]) else: npyramid = contact_dim - 1 # number of frictional directions efc_force = 0.5 / float(npyramid) @@ -2453,20 +2542,12 @@ def _transmission_body_moment( efc_rowadr = efc_J_rowadr_in[worldid, efcid] efc_sparseid = efc_rowadr + dofid colind = efc_J_colind_in[worldid, 0, efc_sparseid] - wp.atomic_add( - actuator_moment_out[worldid], - rowadr + colind, - efc_J_in[worldid, 0, efc_sparseid] * efc_force, - ) + wp.atomic_add(actuator_moment_out[worldid], rowadr + colind, efc_J_in[worldid, 0, efc_sparseid] * efc_force) else: return else: colind = dofid - wp.atomic_add( - actuator_moment_out[worldid], - rowadr + colind, - efc_J_in[worldid, efcid, dofid] * efc_force, - ) + wp.atomic_add(actuator_moment_out[worldid], rowadr + colind, efc_J_in[worldid, efcid, dofid] * efc_force) # excluded contact in gap: get Jacobian, accumulate elif contact_exclude == 1: @@ -2487,46 +2568,28 @@ def _transmission_body_moment( colind = dofid jacp1, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - contact_pos, - b1, - colind, - worldid, + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b1, colind, worldid ) jacp2, _ = support.jac_dof( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - contact_pos, - b2, - colind, - worldid, + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b2, colind, worldid ) jacdif = jacp2 - jacp1 # project Jacobian along the normal of the contact frame - wp.atomic_add( - actuator_moment_out[worldid], rowadr + colind, wp.dot(normal, jacdif) - ) + wp.atomic_add(actuator_moment_out[worldid], rowadr + colind, wp.dot(normal, jacdif)) @wp.kernel def _transmission_body_moment_scale( - # Model: - actuator_trntype_body_adr: wp.array(dtype=int), - # Data in: - moment_rowadr_in: wp.array2d(dtype=int), - # In: - actuator_trntype_body_ncon_in: wp.array2d(dtype=int), - # Data out: - actuator_moment_out: wp.array2d(dtype=float), + # Model: + actuator_trntype_body_adr: wp.array(dtype=int), + # Data in: + moment_rowadr_in: wp.array2d(dtype=int), + # In: + actuator_trntype_body_ncon_in: wp.array2d(dtype=int), + # Data out: + actuator_moment_out: wp.array2d(dtype=float), ): worldid, trnbodyid, dofid = wp.tid() @@ -2549,47 +2612,40 @@ def transmission(m: Model, d: Data): moment_nnz = wp.zeros((d.nworld,), dtype=int) wp.launch( - _transmission, - dim=(d.nworld, m.nu), - inputs=[ - m.nv, - m.body_parentid, - m.body_rootid, - m.body_weldid, - m.body_dofnum, - m.body_dofadr, - m.jnt_type, - m.jnt_qposadr, - m.jnt_dofadr, - m.dof_bodyid, - m.dof_parentid, - m.site_bodyid, - m.site_quat, - m.tendon_adr, - m.tendon_num, - m.wrap_type, - m.wrap_objid, - m.actuator_trntype, - m.actuator_trnid, - m.actuator_gear, - m.actuator_cranklength, - d.qpos, - d.xquat, - d.site_xpos, - d.site_xmat, - d.subtree_com, - d.cdof, - d.ten_J, - d.ten_length, - moment_nnz, - ], - outputs=[ - d.actuator_length, - d.moment_rownnz, - d.moment_rowadr, - d.moment_colind, - d.actuator_moment, - ], + _transmission, + dim=(d.nworld, m.nu), + inputs=[ + m.nv, + m.body_parentid, + m.body_rootid, + m.body_weldid, + m.body_dofnum, + m.body_dofadr, + m.jnt_type, + m.jnt_qposadr, + m.jnt_dofadr, + m.dof_bodyid, + m.dof_parentid, + m.site_bodyid, + m.site_quat, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, + m.actuator_trntype, + m.actuator_trnid, + m.actuator_gear, + m.actuator_cranklength, + d.qpos, + d.xquat, + d.site_xpos, + d.site_xmat, + d.subtree_com, + d.cdof, + d.ten_J, + d.ten_length, + moment_nnz, + ], + outputs=[d.actuator_length, d.moment_rownnz, d.moment_rowadr, d.moment_colind, d.actuator_moment], ) if m.nacttrnbody: @@ -2597,83 +2653,105 @@ def transmission(m: Model, d: Data): ncon = wp.zeros((d.nworld, m.nacttrnbody), dtype=int) wp.launch( - _transmission_body_moment, - dim=(m.nacttrnbody, d.naconmax, m.nv), - inputs=[ - m.opt.cone, - m.body_parentid, - m.body_rootid, - m.dof_bodyid, - m.geom_bodyid, - m.actuator_trnid, - m.actuator_trntype_body_adr, - d.subtree_com, - d.cdof, - d.moment_rowadr, - d.contact.dist, - d.contact.pos, - d.contact.frame, - d.contact.includemargin, - d.contact.dim, - d.contact.geom, - d.contact.efc_address, - d.contact.worldid, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.nacon, - SPARSE_CONSTRAINT_JACOBIAN, - ], - outputs=[d.actuator_moment, ncon], + _transmission_body_moment, + dim=(m.nacttrnbody, d.naconmax, m.nv), + inputs=[ + m.opt.cone, + m.body_parentid, + m.body_rootid, + m.dof_bodyid, + m.geom_bodyid, + m.actuator_trnid, + m.actuator_trntype_body_adr, + d.subtree_com, + d.cdof, + d.moment_rowadr, + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.includemargin, + d.contact.dim, + d.contact.geom, + d.contact.efc_address, + d.contact.worldid, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.nacon, + m.is_sparse, + ], + outputs=[d.actuator_moment, ncon], ) # scale moments wp.launch( - _transmission_body_moment_scale, - dim=(d.nworld, m.nacttrnbody, m.nv), - inputs=[m.actuator_trntype_body_adr, d.moment_rowadr, ncon], - outputs=[d.actuator_moment], + _transmission_body_moment_scale, + dim=(d.nworld, m.nacttrnbody, m.nv), + inputs=[m.actuator_trntype_body_adr, d.moment_rowadr, ncon], + outputs=[d.actuator_moment], ) -@wp.kernel -def _solve_LD_sparse_x_acc_up( - # In: - L: wp.array3d(dtype=float), - qLD_updates_: wp.array(dtype=wp.vec3i), - # Out: - x: wp.array2d(dtype=float), -): - worldid, nodeid = wp.tid() - update = qLD_updates_[nodeid] - i, k, Madr_ki = update[0], update[1], update[2] - wp.atomic_sub(x[worldid], i, L[worldid, 0, Madr_ki] * x[worldid, k]) +@cache_kernel +def _solve_LD_sparse_fused(nv: int, nlevels: int): + """Fused sparse backsubstitution: UP + diag + DOWN in one kernel.""" + @wp.func_native(snippet="WP_TILE_SYNC();") + def _syncthreads(): + pass -@wp.kernel -def _solve_LD_sparse_qLDiag_mul( - # In: - D: wp.array2d(dtype=float), - # Out: - out: wp.array2d(dtype=float), -): - worldid, dofid = wp.tid() - out[worldid, dofid] *= D[worldid, dofid] + @wp.kernel(module="unique", enable_backward=False) + def kernel( + # In: + L: wp.array3d(dtype=float), + D: wp.array2d(dtype=float), + all_updates: wp.array(dtype=wp.vec3i), + level_offsets: wp.array(dtype=int), + y: wp.array2d(dtype=float), + # Out: + x_out: wp.array2d(dtype=float), + ): + worldid, tid = wp.tid() + NV = wp.static(nv) + NLEVELS = wp.static(nlevels) + BLOCK_DIM = wp.block_dim() + # Copy y to x_out + for dofid in range(tid, NV, BLOCK_DIM): + x_out[worldid, dofid] = y[worldid, dofid] + _syncthreads() -@wp.kernel -def _solve_LD_sparse_x_acc_down( - # In: - L: wp.array3d(dtype=float), - qLD_updates_: wp.array(dtype=wp.vec3i), - # Out: - x: wp.array2d(dtype=float), -): - worldid, nodeid = wp.tid() - update = qLD_updates_[nodeid] - i, k, Madr_ki = update[0], update[1], update[2] - wp.atomic_sub(x[worldid], k, L[worldid, 0, Madr_ki] * x[worldid, i]) + # Forward substitution + for level in range(NLEVELS): + level_idx = NLEVELS - 1 - level + level_offset = level_offsets[level_idx] + level_size = level_offsets[level_idx + 1] - level_offset + + for u in range(tid, level_size, BLOCK_DIM): + update = all_updates[level_offset + u] + i, k, Madr_ki = update[0], update[1], update[2] + wp.atomic_sub(x_out[worldid], i, L[worldid, 0, Madr_ki] * x_out[worldid, k]) + _syncthreads() + + # Diagonal multiply + for dofid in range(tid, NV, BLOCK_DIM): + x_out[worldid, dofid] *= D[worldid, dofid] + _syncthreads() + + # Backward substitution + for level in range(NLEVELS): + level_idx = level + level_offset = level_offsets[level_idx] + level_size = level_offsets[level_idx + 1] - level_offset + + for u in range(tid, level_size, BLOCK_DIM): + update = all_updates[level_offset + u] + i, k, Madr_ki = update[0], update[1], update[2] + wp.atomic_sub(x_out[worldid], k, L[worldid, 0, Madr_ki] * x_out[worldid, i]) + _syncthreads() + + return kernel def _solve_LD_sparse( @@ -2685,14 +2763,20 @@ def _solve_LD_sparse( y: wp.array2d(dtype=float), ): """Computes sparse backsubstitution: x = inv(L'*D*L)*y.""" - wp.copy(x, y) - for qLD_updates in reversed(m.qLD_updates): - wp.launch(_solve_LD_sparse_x_acc_up, dim=(d.nworld, qLD_updates.size), inputs=[L, qLD_updates], outputs=[x]) + nlevels = len(m.qLD_updates) + if wp.get_device().is_cuda: + dim_block = m.block_dim.solve_LD_sparse_fused + else: + # Fallback for CPU + dim_block = 1 - wp.launch(_solve_LD_sparse_qLDiag_mul, dim=(d.nworld, m.nv), inputs=[D], outputs=[x]) - - for qLD_updates in m.qLD_updates: - wp.launch(_solve_LD_sparse_x_acc_down, dim=(d.nworld, qLD_updates.size), inputs=[L, qLD_updates], outputs=[x]) + wp.launch( + _solve_LD_sparse_fused(m.nv, nlevels), + dim=(d.nworld, dim_block), + inputs=[L, D, m.qLD_all_updates, m.qLD_level_offsets, y], + outputs=[x], + block_dim=dim_block, + ) @cache_kernel @@ -3005,6 +3089,9 @@ def _joint_tendon( # Model: jnt_qposadr: wp.array(dtype=int), jnt_dofadr: wp.array(dtype=int), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), wrap_objid: wp.array(dtype=int), wrap_prm: wp.array(dtype=float), tendon_jnt_adr: wp.array(dtype=int), @@ -3012,34 +3099,87 @@ def _joint_tendon( # Data in: qpos_in: wp.array2d(dtype=float), # Data out: - ten_J_out: wp.array3d(dtype=float), + ten_J_out: wp.array2d(dtype=float), ten_length_out: wp.array2d(dtype=float), ): worldid, wrapid = wp.tid() - tendon_jnt_adr_ = tendon_jnt_adr[wrapid] - wrap_jnt_adr_ = wrap_jnt_adr[wrapid] - - wrap_objid_ = wrap_objid[wrap_jnt_adr_] - prm = wrap_prm[wrap_jnt_adr_] + tenid = tendon_jnt_adr[wrapid] + wrapjntid = wrap_jnt_adr[wrapid] + wrapobjid = wrap_objid[wrapjntid] + prm = wrap_prm[wrapjntid] # add to length - L = prm * qpos_in[worldid, jnt_qposadr[wrap_objid_]] - # TODO(team): compare atomic_add and for loop - wp.atomic_add(ten_length_out[worldid], tendon_jnt_adr_, L) + L = prm * qpos_in[worldid, jnt_qposadr[wrapobjid]] + wp.atomic_add(ten_length_out[worldid], tenid, L) # add to moment - ten_J_out[worldid, tendon_jnt_adr_, jnt_dofadr[wrap_objid_]] = prm + dofadr = jnt_dofadr[wrapobjid] + rowadr = ten_J_rowadr[tenid] + rownnz = ten_J_rownnz[tenid] + for k in range(rownnz): + if ten_J_colind[rowadr + k] == dofadr: + ten_J_out[worldid, rowadr + k] = prm + break + + +@wp.func +def _accumulate_jac_chain( + # Model: + body_parentid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + # Data in: + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # In: + offset: wp.vec3, + vec: wp.vec3, + bodyid: int, + rowadr: int, + rownnz: int, + scale: float, + worldid: int, + # Data out: + ten_J_out: wp.array2d(dtype=float), +): + """Walk body chain from bodyid to root, accumulate Jacobian contributions.""" + ptr = rownnz - 1 + bid = bodyid + while bid > 0: + bdofadr = body_dofadr[bid] + bdofnum = body_dofnum[bid] + # iterate DOFs in this body in descending order + for k_rev in range(bdofnum): + dof = bdofadr + bdofnum - 1 - k_rev + # scan pointer backward to find matching colind entry + while ptr >= 0: + sparseid = rowadr + ptr + if ten_J_colind[sparseid] <= dof: + break + ptr -= 1 + if ptr >= 0 and ten_J_colind[sparseid] == dof: + cdof = cdof_in[worldid, dof] + cdof_ang = wp.spatial_top(cdof) + cdof_lin = wp.spatial_bottom(cdof) + jacp = cdof_lin + wp.cross(cdof_ang, offset) + J = wp.dot(jacp, vec) * scale + if J != 0.0: + wp.atomic_add(ten_J_out[worldid], sparseid, J) + bid = body_parentid[bid] @wp.kernel def _spatial_site_tendon( # Model: - nv: int, body_parentid: wp.array(dtype=int), body_rootid: wp.array(dtype=int), - dof_bodyid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), site_bodyid: wp.array(dtype=int), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), wrap_objid: wp.array(dtype=int), tendon_site_pair_adr: wp.array(dtype=int), wrap_site_pair_adr: wp.array(dtype=int), @@ -3049,14 +3189,14 @@ def _spatial_site_tendon( subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), # Data out: - ten_J_out: wp.array3d(dtype=float), + ten_J_out: wp.array2d(dtype=float), ten_length_out: wp.array2d(dtype=float), ): worldid, elementid = wp.tid() # site pairs site_pair_adr = wrap_site_pair_adr[elementid] - ten_adr = tendon_site_pair_adr[elementid] + tenid = tendon_site_pair_adr[elementid] # pulley scaling pulley_scale = wrap_pulley_scale[site_pair_adr] @@ -3068,7 +3208,7 @@ def _spatial_site_tendon( pnt1 = site_xpos_in[worldid, id1] dif = pnt1 - pnt0 vec, length = math.normalize_with_norm(dif) - wp.atomic_add(ten_length_out[worldid], ten_adr, length * pulley_scale) + wp.atomic_add(ten_length_out[worldid], tenid, length * pulley_scale) if length < MJ_MINVAL: vec = wp.vec3(1.0, 0.0, 0.0) @@ -3076,26 +3216,55 @@ def _spatial_site_tendon( body0 = site_bodyid[id0] body1 = site_bodyid[id1] if body0 != body1: - # TODO(team): parallelize - for i in range(nv): - jacp1, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pnt0, body0, i, worldid) - jacp2, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pnt1, body1, i, worldid) - - J = wp.dot(jacp2 - jacp1, vec) - if J: - wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale) + rownnz = ten_J_rownnz[tenid] + rowadr = ten_J_rowadr[tenid] + offset0 = pnt0 - subtree_com_in[worldid, body_rootid[body0]] + offset1 = pnt1 - subtree_com_in[worldid, body_rootid[body1]] + _accumulate_jac_chain( + body_parentid, + body_dofnum, + body_dofadr, + ten_J_colind, + cdof_in, + offset0, + vec, + body0, + rowadr, + rownnz, + -pulley_scale, + worldid, + ten_J_out, + ) + _accumulate_jac_chain( + body_parentid, + body_dofnum, + body_dofadr, + ten_J_colind, + cdof_in, + offset1, + vec, + body1, + rowadr, + rownnz, + pulley_scale, + worldid, + ten_J_out, + ) @wp.kernel def _spatial_geom_tendon( # Model: - nv: int, body_parentid: wp.array(dtype=int), body_rootid: wp.array(dtype=int), - dof_bodyid: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), site_bodyid: wp.array(dtype=int), + ten_J_rownnz: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), wrap_type: wp.array(dtype=int), wrap_objid: wp.array(dtype=int), wrap_prm: wp.array(dtype=float), @@ -3109,14 +3278,14 @@ def _spatial_geom_tendon( subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), # Data out: - ten_J_out: wp.array3d(dtype=float), + ten_J_out: wp.array2d(dtype=float), ten_length_out: wp.array2d(dtype=float), # Out: wrap_geom_xpos_out: wp.array2d(dtype=wp.spatial_vector), ): worldid, elementid = wp.tid() wrap_adr = wrap_geom_adr[elementid] - ten_adr = tendon_geom_adr[elementid] + tenid = tendon_geom_adr[elementid] # pulley scaling pulley_scale = wrap_pulley_scale[wrap_adr] @@ -3154,6 +3323,9 @@ def _spatial_geom_tendon( # store geom points wrap_geom_xpos_out[worldid, elementid] = wp.spatial_vector(geom_pnt0, geom_pnt1) + rownnz = ten_J_rownnz[tenid] + rowadr = ten_J_rowadr[tenid] + if length_geomgeom >= 0.0: dif_sitegeom = geom_pnt0 - site_pnt0 dif_geomsite = site_pnt1 - geom_pnt1 @@ -3164,7 +3336,7 @@ def _spatial_geom_tendon( length_sitegeomsite = length_sitegeom + length_geomgeom + length_geomsite if length_sitegeomsite: - wp.atomic_add(ten_length_out[worldid], ten_adr, length_sitegeomsite * pulley_scale) + wp.atomic_add(ten_length_out[worldid], tenid, length_sitegeomsite * pulley_scale) # moment if length_sitegeom < MJ_MINVAL: @@ -3176,61 +3348,120 @@ def _spatial_geom_tendon( dif_body_sitegeom = bodyid_site0 != bodyid_geom dif_body_geomsite = bodyid_geom != bodyid_site1 - # TODO(team): parallelize - for i in range(nv): - J = float(0.0) - # site-geom - if dif_body_sitegeom: - jacp_site0, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt0, bodyid_site0, i, worldid - ) + # site-geom segment + if dif_body_sitegeom: + offset_site0 = site_pnt0 - subtree_com_in[worldid, body_rootid[bodyid_site0]] + offset_geom0 = geom_pnt0 - subtree_com_in[worldid, body_rootid[bodyid_geom]] + _accumulate_jac_chain( + body_parentid, + body_dofnum, + body_dofadr, + ten_J_colind, + cdof_in, + offset_site0, + vec_sitegeom, + bodyid_site0, + rowadr, + rownnz, + -pulley_scale, + worldid, + ten_J_out, + ) + _accumulate_jac_chain( + body_parentid, + body_dofnum, + body_dofadr, + ten_J_colind, + cdof_in, + offset_geom0, + vec_sitegeom, + bodyid_geom, + rowadr, + rownnz, + pulley_scale, + worldid, + ten_J_out, + ) - jacp_geom0, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, geom_pnt0, bodyid_geom, i, worldid - ) - - J += wp.dot(jacp_geom0 - jacp_site0, vec_sitegeom) - - # geom-site - if dif_body_geomsite: - jacp_geom1, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, geom_pnt1, bodyid_geom, i, worldid - ) - - jacp_site1, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt1, bodyid_site1, i, worldid - ) - - J += wp.dot(jacp_site1 - jacp_geom1, vec_geomsite) - - if J: - wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale) + # geom-site segment + if dif_body_geomsite: + offset_geom1 = geom_pnt1 - subtree_com_in[worldid, body_rootid[bodyid_geom]] + offset_site1 = site_pnt1 - subtree_com_in[worldid, body_rootid[bodyid_site1]] + _accumulate_jac_chain( + body_parentid, + body_dofnum, + body_dofadr, + ten_J_colind, + cdof_in, + offset_geom1, + vec_geomsite, + bodyid_geom, + rowadr, + rownnz, + -pulley_scale, + worldid, + ten_J_out, + ) + _accumulate_jac_chain( + body_parentid, + body_dofnum, + body_dofadr, + ten_J_colind, + cdof_in, + offset_site1, + vec_geomsite, + bodyid_site1, + rowadr, + rownnz, + pulley_scale, + worldid, + ten_J_out, + ) else: dif_sitesite = site_pnt1 - site_pnt0 vec_sitesite, length_sitesite = math.normalize_with_norm(dif_sitesite) # length if length_sitesite: - wp.atomic_add(ten_length_out[worldid], ten_adr, length_sitesite * pulley_scale) + wp.atomic_add(ten_length_out[worldid], tenid, length_sitesite * pulley_scale) # moment if length_sitesite < MJ_MINVAL: vec_sitesite = wp.vec3(1.0, 0.0, 0.0) if bodyid_site0 != bodyid_site1: - # TODO(team): parallelize - for i in range(nv): - jacp1, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt0, bodyid_site0, i, worldid - ) - jacp2, _ = support.jac_dof( - body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt1, bodyid_site1, i, worldid - ) - - J = wp.dot(jacp2 - jacp1, vec_sitesite) - - if J: - wp.atomic_add(ten_J_out[worldid, ten_adr], i, J * pulley_scale) + offset_site0 = site_pnt0 - subtree_com_in[worldid, body_rootid[bodyid_site0]] + offset_site1 = site_pnt1 - subtree_com_in[worldid, body_rootid[bodyid_site1]] + _accumulate_jac_chain( + body_parentid, + body_dofnum, + body_dofadr, + ten_J_colind, + cdof_in, + offset_site0, + vec_sitesite, + bodyid_site0, + rowadr, + rownnz, + -pulley_scale, + worldid, + ten_J_out, + ) + _accumulate_jac_chain( + body_parentid, + body_dofnum, + body_dofadr, + ten_J_colind, + cdof_in, + offset_site1, + vec_sitesite, + bodyid_site1, + rowadr, + rownnz, + pulley_scale, + worldid, + ten_J_out, + ) @wp.kernel @@ -3412,7 +3643,18 @@ def tendon(m: Model, d: Data): wp.launch( _joint_tendon, dim=(d.nworld, m.wrap_jnt_adr.size), - inputs=[m.jnt_qposadr, m.jnt_dofadr, m.wrap_objid, m.wrap_prm, m.tendon_jnt_adr, m.wrap_jnt_adr, d.qpos], + inputs=[ + m.jnt_qposadr, + m.jnt_dofadr, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, + m.wrap_objid, + m.wrap_prm, + m.tendon_jnt_adr, + m.wrap_jnt_adr, + d.qpos, + ], outputs=[d.ten_J, d.ten_length], ) @@ -3428,11 +3670,14 @@ def tendon(m: Model, d: Data): _spatial_site_tendon, dim=(d.nworld, m.wrap_site_pair_adr.size), inputs=[ - m.nv, m.body_parentid, m.body_rootid, - m.dof_bodyid, + m.body_dofnum, + m.body_dofadr, m.site_bodyid, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, m.wrap_objid, m.tendon_site_pair_adr, m.wrap_site_pair_adr, @@ -3449,13 +3694,16 @@ def tendon(m: Model, d: Data): _spatial_geom_tendon, dim=(d.nworld, m.wrap_geom_adr.size), inputs=[ - m.nv, m.body_parentid, m.body_rootid, - m.dof_bodyid, + m.body_dofnum, + m.body_dofadr, m.geom_bodyid, m.geom_size, m.site_bodyid, + m.ten_J_rownnz, + m.ten_J_rowadr, + m.ten_J_colind, m.wrap_type, m.wrap_objid, m.wrap_prm, 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 82ac23c7..2fabebe0 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py @@ -17,17 +17,17 @@ import dataclasses from math import ceil from math import sqrt +import warp as wp + from mujoco.mjx.third_party.mujoco_warp._src import math from mujoco.mjx.third_party.mujoco_warp._src import smooth from mujoco.mjx.third_party.mujoco_warp._src import support from mujoco.mjx.third_party.mujoco_warp._src import types from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_func from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_solve_func -from mujoco.mjx.third_party.mujoco_warp._src.types import SPARSE_CONSTRAINT_JACOBIAN from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope from mujoco.mjx.third_party.mujoco_warp._src.warp_util import scoped_mathdx_gemm_disabled -import warp as wp wp.set_module_options({"enable_backward": False}) @@ -91,14 +91,14 @@ def create_inverse_context(m: types.Model, d: types.Data) -> InverseContext: njmax = d.njmax return InverseContext( - Jaref=wp.empty((nworld, njmax), dtype=float), - search_dot=wp.empty((nworld,), dtype=float), - gauss=wp.empty((nworld,), dtype=float), - cost=wp.empty((nworld,), dtype=float), - prev_cost=wp.empty((nworld,), dtype=float), - done=wp.empty((nworld,), dtype=bool), - changed_efc_ids=wp.empty((nworld, 0), dtype=int), - changed_efc_count=wp.empty((0,), dtype=int), + Jaref=wp.empty((nworld, njmax), dtype=float), + search_dot=wp.empty((nworld,), dtype=float), + gauss=wp.empty((nworld,), dtype=float), + cost=wp.empty((nworld,), dtype=float), + prev_cost=wp.empty((nworld,), dtype=float), + done=wp.empty((nworld,), dtype=bool), + changed_efc_ids=wp.empty((nworld, 0), dtype=int), + changed_efc_count=wp.empty((0,), dtype=int), ) @@ -121,36 +121,28 @@ def create_solver_context(m: types.Model, d: types.Data) -> SolverContext: alloc_hfactor = alloc_h and nv > _BLOCK_CHOLESKY_DIM return SolverContext( - Jaref=wp.empty((nworld, njmax), dtype=float), - search_dot=wp.empty((nworld,), dtype=float), - gauss=wp.empty((nworld,), dtype=float), - cost=wp.empty((nworld,), dtype=float), - prev_cost=wp.empty((nworld,), dtype=float), - done=wp.empty((nworld,), dtype=bool), - grad=wp.zeros((nworld, nv_pad), dtype=float), - grad_dot=wp.empty((nworld,), dtype=float), - Mgrad=wp.zeros((nworld, nv_pad), dtype=float), - search=wp.empty((nworld, nv), dtype=float), - mv=wp.empty((nworld, nv), dtype=float), - jv=wp.empty((nworld, njmax), dtype=float), - quad=wp.empty((nworld, njmax), dtype=wp.vec3), - quad_gauss=wp.empty((nworld,), dtype=wp.vec3), - alpha=wp.empty((nworld,), dtype=float), - prev_grad=wp.empty((nworld, nv), dtype=float), - prev_Mgrad=wp.empty((nworld, nv), dtype=float), - beta=wp.empty((nworld,), dtype=float), - h=wp.zeros((nworld, nv_pad, nv_pad), dtype=float) - if alloc_h - else wp.empty((nworld, 0, 0), dtype=float), - hfactor=wp.zeros((nworld, nv_pad, nv_pad), dtype=float) - if alloc_hfactor - else wp.empty((nworld, 0, 0), dtype=float), - changed_efc_ids=wp.empty((nworld, njmax), dtype=int) - if alloc_h - else wp.empty((nworld, 0), dtype=int), - changed_efc_count=wp.empty((nworld,), dtype=int) - if alloc_h - else wp.empty((0,), dtype=int), + Jaref=wp.empty((nworld, njmax), dtype=float), + search_dot=wp.empty((nworld,), dtype=float), + gauss=wp.empty((nworld,), dtype=float), + cost=wp.empty((nworld,), dtype=float), + prev_cost=wp.empty((nworld,), dtype=float), + done=wp.empty((nworld,), dtype=bool), + grad=wp.zeros((nworld, nv_pad), dtype=float), + grad_dot=wp.empty((nworld,), dtype=float), + Mgrad=wp.zeros((nworld, nv_pad), dtype=float), + search=wp.empty((nworld, nv), dtype=float), + mv=wp.empty((nworld, nv), dtype=float), + jv=wp.empty((nworld, njmax), dtype=float), + quad=wp.empty((nworld, njmax), dtype=wp.vec3), + quad_gauss=wp.empty((nworld,), dtype=wp.vec3), + alpha=wp.empty((nworld,), dtype=float), + prev_grad=wp.empty((nworld, nv), dtype=float), + prev_Mgrad=wp.empty((nworld, nv), dtype=float), + beta=wp.empty((nworld,), dtype=float), + h=wp.zeros((nworld, nv_pad, nv_pad), dtype=float) if alloc_h else wp.empty((nworld, 0, 0), dtype=float), + hfactor=wp.zeros((nworld, nv_pad, nv_pad), dtype=float) if alloc_hfactor else wp.empty((nworld, 0, 0), dtype=float), + changed_efc_ids=wp.empty((nworld, njmax), dtype=int) if alloc_h else wp.empty((nworld, 0), dtype=int), + changed_efc_count=wp.empty((nworld,), dtype=int) if alloc_h else wp.empty((0,), dtype=int), ) @@ -892,18 +884,19 @@ def _compute_efc_eval_pt_3alphas_elliptic( @cache_kernel -def linesearch_iterative(ls_iterations: int, cone_type: types.ConeType, fuse_jv: bool): +def linesearch_iterative(ls_iterations: int, cone_type: types.ConeType, fuse_jv: bool, is_sparse: bool): """Factory for iterative linesearch kernel. Args: - block_dim: Number of threads per block for tile reductions. ls_iterations: Max linesearch iterations (compile-time constant for loop optimization). cone_type: Friction cone type (PYRAMIDAL or ELLIPTIC) for compile-time optimization. fuse_jv: Whether to compute jv = J @ search in-kernel (efficient for small nv). + is_sparse: Use sparse matrix representation for constraint Jacobian. """ LS_ITERATIONS = ls_iterations IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC FUSE_JV = fuse_jv + IS_SPARSE = is_sparse # Native snippet for CUDA __syncthreads() @wp.func_native(snippet="WP_TILE_SYNC();") @@ -922,46 +915,46 @@ def linesearch_iterative(ls_iterations: int, cone_type: types.ConeType, fuse_jv: @wp.kernel(module="unique", enable_backward=False) def kernel( - # Model: - nv: int, - opt_tolerance: wp.array(dtype=float), - opt_ls_tolerance: wp.array(dtype=float), - opt_impratio_invsqrt: wp.array(dtype=float), - stat_meaninertia: wp.array(dtype=float), - # Data in: - ne_in: wp.array(dtype=int), - nf_in: wp.array(dtype=int), - nefc_in: wp.array(dtype=int), - qfrc_smooth_in: wp.array2d(dtype=float), - contact_friction_in: wp.array(dtype=types.vec5), - contact_dim_in: wp.array(dtype=int), - contact_efc_address_in: wp.array2d(dtype=int), - efc_type_in: wp.array2d(dtype=int), - efc_id_in: wp.array2d(dtype=int), - efc_J_rownnz_in: wp.array2d(dtype=int), - efc_J_rowadr_in: wp.array2d(dtype=int), - efc_J_colind_in: wp.array3d(dtype=int), - efc_J_in: wp.array3d(dtype=float), - efc_D_in: wp.array2d(dtype=float), - efc_frictionloss_in: wp.array2d(dtype=float), - njmax_in: int, - nacon_in: wp.array(dtype=int), - # In: - ctx_Jaref_in: wp.array2d(dtype=float), - ctx_search_in: wp.array2d(dtype=float), - ctx_search_dot_in: wp.array(dtype=float), - ctx_gauss_in: wp.array(dtype=float), - ctx_mv_in: wp.array2d(dtype=float), - ctx_jv_in: wp.array2d(dtype=float), - ctx_quad_in: wp.array2d(dtype=wp.vec3), - ctx_done_in: wp.array(dtype=bool), - # Data out: - qacc_out: wp.array2d(dtype=float), - efc_Ma_out: wp.array2d(dtype=float), - # Out: - ctx_Jaref_out: wp.array2d(dtype=float), - ctx_jv_out: wp.array2d(dtype=float), - ctx_quad_out: wp.array2d(dtype=wp.vec3), + # Model: + nv: int, + opt_tolerance: wp.array(dtype=float), + opt_ls_tolerance: wp.array(dtype=float), + opt_impratio_invsqrt: wp.array(dtype=float), + stat_meaninertia: wp.array(dtype=float), + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + qfrc_smooth_in: wp.array2d(dtype=float), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + efc_type_in: wp.array2d(dtype=int), + efc_id_in: wp.array2d(dtype=int), + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), + efc_J_colind_in: wp.array3d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_D_in: wp.array2d(dtype=float), + efc_frictionloss_in: wp.array2d(dtype=float), + njmax_in: int, + nacon_in: wp.array(dtype=int), + # In: + ctx_Jaref_in: wp.array2d(dtype=float), + ctx_search_in: wp.array2d(dtype=float), + ctx_search_dot_in: wp.array(dtype=float), + ctx_gauss_in: wp.array(dtype=float), + ctx_mv_in: wp.array2d(dtype=float), + ctx_jv_in: wp.array2d(dtype=float), + ctx_quad_in: wp.array2d(dtype=wp.vec3), + ctx_done_in: wp.array(dtype=bool), + # Data out: + qacc_out: wp.array2d(dtype=float), + efc_Ma_out: wp.array2d(dtype=float), + # Out: + ctx_Jaref_out: wp.array2d(dtype=float), + ctx_jv_out: wp.array2d(dtype=float), + ctx_quad_out: wp.array2d(dtype=wp.vec3), ): worldid, tid = wp.tid() @@ -976,15 +969,13 @@ def linesearch_iterative(ls_iterations: int, cone_type: types.ConeType, fuse_jv: if wp.static(FUSE_JV): for efcid in range(tid, nefc, wp.block_dim()): jv = float(0.0) - if wp.static(SPARSE_CONSTRAINT_JACOBIAN): + if wp.static(IS_SPARSE): rownnz = efc_J_rownnz_in[worldid, efcid] rowadr = efc_J_rowadr_in[worldid, efcid] for k in range(rownnz): sparseid = rowadr + k colind = efc_J_colind_in[worldid, 0, sparseid] - jv += ( - efc_J_in[worldid, 0, sparseid] * ctx_search_in[worldid, colind] - ) + jv += efc_J_in[worldid, 0, sparseid] * ctx_search_in[worldid, colind] else: for i in range(nv): jv += efc_J_in[worldid, efcid, i] * ctx_search_in[worldid, i] @@ -1360,42 +1351,42 @@ def _linesearch_iterative(m: types.Model, d: types.Data, ctx: SolverContext, fus fuse_jv: Whether jv is computed in-kernel (True) or pre-computed (False). """ wp.launch_tiled( - linesearch_iterative(m.opt.ls_iterations, m.opt.cone, fuse_jv), - dim=d.nworld, - inputs=[ - m.nv, - m.opt.tolerance, - m.opt.ls_tolerance, - m.opt.impratio_invsqrt, - m.stat.meaninertia, - d.ne, - d.nf, - d.nefc, - d.qfrc_smooth, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.efc.type, - d.efc.id, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.D, - d.efc.frictionloss, - d.njmax, - d.nacon, - ctx.Jaref, - ctx.search, - ctx.search_dot, - ctx.gauss, - ctx.mv, - ctx.jv, - ctx.quad, - ctx.done, - ], - outputs=[d.qacc, d.efc.Ma, ctx.Jaref, ctx.jv, ctx.quad], - block_dim=m.block_dim.linesearch_iterative, + linesearch_iterative(m.opt.ls_iterations, m.opt.cone, fuse_jv, m.is_sparse), + dim=d.nworld, + inputs=[ + m.nv, + m.opt.tolerance, + m.opt.ls_tolerance, + m.opt.impratio_invsqrt, + m.stat.meaninertia, + d.ne, + d.nf, + d.nefc, + d.qfrc_smooth, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.efc.type, + d.efc.id, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.D, + d.efc.frictionloss, + d.njmax, + d.nacon, + ctx.Jaref, + ctx.search, + ctx.search_dot, + ctx.gauss, + ctx.mv, + ctx.jv, + ctx.quad, + ctx.done, + ], + outputs=[d.qacc, d.efc.Ma, ctx.Jaref, ctx.jv, ctx.quad], + block_dim=m.block_dim.linesearch_iterative, ) @@ -1420,20 +1411,20 @@ def linesearch_zero_jv( @cache_kernel -def linesearch_jv_fused(opt_is_sparse: bool, nv: int, dofs_per_thread: int): +def linesearch_jv_fused(is_sparse: bool, nv: int, dofs_per_thread: int): @wp.kernel(module="unique", enable_backward=False) def kernel( - # Data in: - nefc_in: wp.array(dtype=int), - efc_J_rownnz_in: wp.array2d(dtype=int), - efc_J_rowadr_in: wp.array2d(dtype=int), - efc_J_colind_in: wp.array3d(dtype=int), - efc_J_in: wp.array3d(dtype=float), - # In: - ctx_search_in: wp.array2d(dtype=float), - ctx_done_in: wp.array(dtype=bool), - # Out: - ctx_jv_out: wp.array2d(dtype=float), + # Data in: + nefc_in: wp.array(dtype=int), + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), + efc_J_colind_in: wp.array3d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + # In: + ctx_search_in: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_jv_out: wp.array2d(dtype=float), ): worldid, efcid, dofstart = wp.tid() @@ -1446,23 +1437,21 @@ def linesearch_jv_fused(opt_is_sparse: bool, nv: int, dofs_per_thread: int): jv_out = float(0.0) if wp.static(dofs_per_thread >= nv): - if wp.static(SPARSE_CONSTRAINT_JACOBIAN): + if wp.static(is_sparse): # Sparse: iterate over non-zero entries in the row rownnz = efc_J_rownnz_in[worldid, efcid] rowadr = efc_J_rowadr_in[worldid, efcid] for k in range(rownnz): sparseid = rowadr + k colind = efc_J_colind_in[worldid, 0, sparseid] - jv_out += ( - efc_J_in[worldid, 0, sparseid] * ctx_search_in[worldid, colind] - ) + jv_out += efc_J_in[worldid, 0, sparseid] * ctx_search_in[worldid, colind] else: for i in range(wp.static(min(dofs_per_thread, nv))): jv_out += efc_J_in[worldid, efcid, i] * ctx_search_in[worldid, i] ctx_jv_out[worldid, efcid] = jv_out else: - if wp.static(SPARSE_CONSTRAINT_JACOBIAN): + if wp.static(is_sparse): # Sparse: thread 0 handles entire row (sparse entries << nv typically) if dofstart == 0: rownnz = efc_J_rownnz_in[worldid, efcid] @@ -1470,9 +1459,7 @@ def linesearch_jv_fused(opt_is_sparse: bool, nv: int, dofs_per_thread: int): for k in range(rownnz): sparseid = rowadr + k colind = efc_J_colind_in[worldid, 0, sparseid] - jv_out += ( - efc_J_in[worldid, 0, sparseid] * ctx_search_in[worldid, colind] - ) + jv_out += efc_J_in[worldid, 0, sparseid] * ctx_search_in[worldid, colind] ctx_jv_out[worldid, efcid] = jv_out else: for i in range(wp.static(dofs_per_thread)): @@ -1583,10 +1570,7 @@ def linesearch_prepare_quad( dim = contact_dim_in[conid] friction = contact_friction_in[conid] - mu = ( - friction[0] - * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] - ) + mu = friction[0] * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] u0 = Jaref * mu v0 = jv * mu @@ -1707,18 +1691,10 @@ def _linesearch(m: types.Model, d: types.Data, ctx: SolverContext, cost: wp.arra ) wp.launch( - linesearch_jv_fused(m.is_sparse, m.nv, dofs_per_thread), - dim=(d.nworld, d.njmax, threads_per_efc), - inputs=[ - d.nefc, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - ctx.search, - ctx.done, - ], - outputs=[ctx.jv], + linesearch_jv_fused(m.is_sparse, m.nv, dofs_per_thread), + dim=(d.nworld, d.njmax, threads_per_efc), + inputs=[d.nefc, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, d.efc.J, ctx.search, ctx.done], + outputs=[ctx.jv], ) if m.opt.ls_parallel: @@ -1744,19 +1720,19 @@ def solve_init_efc( @cache_kernel -def solve_init_jaref(opt_is_sparse: bool, nv: int, dofs_per_thread: int): +def solve_init_jaref(is_sparse: bool, nv: int, dofs_per_thread: int): @wp.kernel(module="unique", enable_backward=False) def kernel( - # Data in: - nefc_in: wp.array(dtype=int), - qacc_in: wp.array2d(dtype=float), - efc_J_rownnz_in: wp.array2d(dtype=int), - efc_J_rowadr_in: wp.array2d(dtype=int), - efc_J_colind_in: wp.array3d(dtype=int), - efc_J_in: wp.array3d(dtype=float), - efc_aref_in: wp.array2d(dtype=float), - # Out: - ctx_Jaref_out: wp.array2d(dtype=float), + # Data in: + nefc_in: wp.array(dtype=int), + qacc_in: wp.array2d(dtype=float), + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), + efc_J_colind_in: wp.array3d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_aref_in: wp.array2d(dtype=float), + # Out: + ctx_Jaref_out: wp.array2d(dtype=float), ): worldid, efcid, dofstart = wp.tid() @@ -1764,7 +1740,7 @@ def solve_init_jaref(opt_is_sparse: bool, nv: int, dofs_per_thread: int): return jaref = float(0.0) - if wp.static(SPARSE_CONSTRAINT_JACOBIAN): + if wp.static(is_sparse): rownnz = efc_J_rownnz_in[worldid, efcid] rowadr = efc_J_rowadr_in[worldid, efcid] for i in range(rownnz): @@ -1785,9 +1761,7 @@ def solve_init_jaref(opt_is_sparse: bool, nv: int, dofs_per_thread: int): jaref += efc_J_in[worldid, efcid, ii] * qacc_in[worldid, ii] if dofstart == 0: - wp.atomic_add( - ctx_Jaref_out, worldid, efcid, jaref - efc_aref_in[worldid, efcid] - ) + wp.atomic_add(ctx_Jaref_out, worldid, efcid, jaref - efc_aref_in[worldid, efcid]) else: wp.atomic_add(ctx_Jaref_out, worldid, efcid, jaref) @@ -1834,30 +1808,30 @@ def update_constraint_efc(track_changes: bool): @wp.kernel(module="unique", enable_backward=False) def kernel( - # Model: - opt_impratio_invsqrt: wp.array(dtype=float), - # Data in: - ne_in: wp.array(dtype=int), - nf_in: wp.array(dtype=int), - nefc_in: wp.array(dtype=int), - contact_friction_in: wp.array(dtype=types.vec5), - contact_dim_in: wp.array(dtype=int), - contact_efc_address_in: wp.array2d(dtype=int), - efc_type_in: wp.array2d(dtype=int), - efc_id_in: wp.array2d(dtype=int), - efc_D_in: wp.array2d(dtype=float), - efc_frictionloss_in: wp.array2d(dtype=float), - nacon_in: wp.array(dtype=int), - # In: - ctx_Jaref_in: wp.array2d(dtype=float), - ctx_done_in: wp.array(dtype=bool), - # Data out: - efc_force_out: wp.array2d(dtype=float), - efc_state_out: wp.array2d(dtype=int), - # Out: - ctx_cost_out: wp.array(dtype=float), - changed_ids_out: wp.array2d(dtype=int), - changed_count_out: wp.array(dtype=int), + # Model: + opt_impratio_invsqrt: wp.array(dtype=float), + # Data in: + ne_in: wp.array(dtype=int), + nf_in: wp.array(dtype=int), + nefc_in: wp.array(dtype=int), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + efc_type_in: wp.array2d(dtype=int), + efc_id_in: wp.array2d(dtype=int), + efc_D_in: wp.array2d(dtype=float), + efc_frictionloss_in: wp.array2d(dtype=float), + nacon_in: wp.array(dtype=int), + # In: + ctx_Jaref_in: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Data out: + efc_force_out: wp.array2d(dtype=float), + efc_state_out: wp.array2d(dtype=int), + # Out: + ctx_cost_out: wp.array(dtype=float), + changed_ids_out: wp.array2d(dtype=int), + changed_count_out: wp.array(dtype=int), ): worldid, efcid = wp.tid() @@ -1869,9 +1843,7 @@ def update_constraint_efc(track_changes: bool): # Read old QUADRATIC status before overwriting if wp.static(TRACK_CHANGES): - old_quad = ( - efc_state_out[worldid, efcid] == types.ConstraintState.QUADRATIC.value - ) + old_quad = efc_state_out[worldid, efcid] == types.ConstraintState.QUADRATIC.value efc_D = efc_D_in[worldid, efcid] Jaref = ctx_Jaref_in[worldid, efcid] @@ -1919,10 +1891,7 @@ def update_constraint_efc(track_changes: bool): dim = contact_dim_in[conid] friction = contact_friction_in[conid] - mu = ( - friction[0] - * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] - ) + mu = friction[0] * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] efcid0 = contact_efc_address_in[conid, 0] if efcid0 < 0: @@ -1984,17 +1953,17 @@ def update_constraint_efc(track_changes: bool): @wp.kernel def update_constraint_init_qfrc_constraint_sparse( - # Data in: - nefc_in: wp.array(dtype=int), - efc_J_rownnz_in: wp.array2d(dtype=int), - efc_J_rowadr_in: wp.array2d(dtype=int), - efc_J_colind_in: wp.array3d(dtype=int), - efc_J_in: wp.array3d(dtype=float), - efc_force_in: wp.array2d(dtype=float), - # In: - ctx_done_in: wp.array(dtype=bool), - # Data out: - qfrc_constraint_out: wp.array2d(dtype=float), + # Data in: + nefc_in: wp.array(dtype=int), + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), + efc_J_colind_in: wp.array3d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_force_in: wp.array2d(dtype=float), + # In: + ctx_done_in: wp.array(dtype=bool), + # Data out: + qfrc_constraint_out: wp.array2d(dtype=float), ): worldid, efcid = wp.tid() @@ -2017,15 +1986,15 @@ def update_constraint_init_qfrc_constraint_sparse( @wp.kernel def update_constraint_init_qfrc_constraint_dense( - # Data in: - nefc_in: wp.array(dtype=int), - efc_J_in: wp.array3d(dtype=float), - efc_force_in: wp.array2d(dtype=float), - njmax_in: int, - # In: - ctx_done_in: wp.array(dtype=bool), - # Data out: - qfrc_constraint_out: wp.array2d(dtype=float), + # Data in: + nefc_in: wp.array(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_force_in: wp.array2d(dtype=float), + njmax_in: int, + # In: + ctx_done_in: wp.array(dtype=bool), + # Data out: + qfrc_constraint_out: wp.array2d(dtype=float), ): worldid, dofid = wp.tid() @@ -2076,23 +2045,23 @@ def update_constraint_gauss_cost(nv: int, dofs_per_thread: int): gauss_cost += (efc_Ma_in[worldid, ii] - qfrc_smooth_in[worldid, ii]) * ( qacc_in[worldid, ii] - qacc_smooth_in[worldid, ii] ) - wp.atomic_add(ctx_gauss_out, worldid, gauss_cost) - wp.atomic_add(ctx_cost_out, worldid, gauss_cost) + wp.atomic_add(ctx_gauss_out, worldid, 0.5 * gauss_cost) + wp.atomic_add(ctx_cost_out, worldid, 0.5 * gauss_cost) return kernel @wp.kernel def update_gradient_h_incremental( - # Data in: - efc_J_in: wp.array3d(dtype=float), - efc_D_in: wp.array2d(dtype=float), - efc_state_in: wp.array2d(dtype=int), - # In: - changed_ids_in: wp.array2d(dtype=int), - changed_count_in: wp.array(dtype=int), - # Out: - ctx_h_out: wp.array3d(dtype=float), + # Data in: + efc_J_in: wp.array3d(dtype=float), + efc_D_in: wp.array2d(dtype=float), + efc_state_in: wp.array2d(dtype=int), + # In: + changed_ids_in: wp.array2d(dtype=int), + changed_count_in: wp.array(dtype=int), + # Out: + ctx_h_out: wp.array3d(dtype=float), ): """Incrementally update lower triangle of H for changed constraints. @@ -2131,18 +2100,18 @@ def update_gradient_h_incremental( @wp.kernel def update_gradient_h_incremental_sparse( - # Data in: - efc_J_rownnz_in: wp.array2d(dtype=int), - efc_J_rowadr_in: wp.array2d(dtype=int), - efc_J_colind_in: wp.array3d(dtype=int), - efc_J_in: wp.array3d(dtype=float), - efc_D_in: wp.array2d(dtype=float), - efc_state_in: wp.array2d(dtype=int), - # In: - changed_ids_in: wp.array2d(dtype=int), - changed_count_in: wp.array(dtype=int), - # Out: - ctx_h_out: wp.array3d(dtype=float), + # Data in: + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), + efc_J_colind_in: wp.array3d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_D_in: wp.array2d(dtype=float), + efc_state_in: wp.array2d(dtype=int), + # In: + changed_ids_in: wp.array2d(dtype=int), + changed_count_in: wp.array(dtype=int), + # Out: + ctx_h_out: wp.array3d(dtype=float), ): """Incrementally update lower triangle of H for changed constraints (sparse J).""" worldid, change_idx = wp.tid() @@ -2182,12 +2151,7 @@ def update_gradient_h_incremental_sparse( wp.atomic_add(ctx_h_out[worldid, colindj], colindi, h) -def _update_constraint( - m: types.Model, - d: types.Data, - ctx: SolverContext | InverseContext, - track_changes: bool = False, -): +def _update_constraint(m: types.Model, d: types.Data, ctx: SolverContext | InverseContext, track_changes: bool = False): """Update constraint arrays after each solve iteration.""" wp.launch( update_constraint_init_cost, @@ -2197,58 +2161,44 @@ def _update_constraint( ) efc_inputs = [ - m.opt.impratio_invsqrt, - d.ne, - d.nf, - d.nefc, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.efc.type, - d.efc.id, - d.efc.D, - d.efc.frictionloss, - d.nacon, - ctx.Jaref, - ctx.done, + m.opt.impratio_invsqrt, + d.ne, + d.nf, + d.nefc, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.efc.type, + d.efc.id, + d.efc.D, + d.efc.frictionloss, + d.nacon, + ctx.Jaref, + ctx.done, ] wp.launch( - update_constraint_efc(track_changes), - dim=(d.nworld, d.njmax), - inputs=efc_inputs, - outputs=[ - d.efc.force, - d.efc.state, - ctx.cost, - ctx.changed_efc_ids, - ctx.changed_efc_count, - ], + update_constraint_efc(track_changes), + dim=(d.nworld, d.njmax), + inputs=efc_inputs, + outputs=[d.efc.force, d.efc.state, ctx.cost, ctx.changed_efc_ids, ctx.changed_efc_count], ) # qfrc_constraint = efc_J.T @ efc_force - if SPARSE_CONSTRAINT_JACOBIAN: + if m.is_sparse: d.qfrc_constraint.zero_() wp.launch( - update_constraint_init_qfrc_constraint_sparse, - dim=(d.nworld, d.njmax), - inputs=[ - d.nefc, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.force, - ctx.done, - ], - outputs=[d.qfrc_constraint], + update_constraint_init_qfrc_constraint_sparse, + dim=(d.nworld, d.njmax), + inputs=[d.nefc, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, d.efc.J, d.efc.force, ctx.done], + outputs=[d.qfrc_constraint], ) else: wp.launch( - update_constraint_init_qfrc_constraint_dense, - dim=(d.nworld, m.nv), - inputs=[d.nefc, d.efc.J, d.efc.force, d.njmax, ctx.done], - outputs=[d.qfrc_constraint], + update_constraint_init_qfrc_constraint_dense, + dim=(d.nworld, m.nv), + inputs=[d.nefc, d.efc.J, d.efc.force, d.njmax, ctx.done], + outputs=[d.qfrc_constraint], ) # if we are only using 1 thread, it makes sense to do more dofs and skip the atomics. @@ -2441,9 +2391,7 @@ def update_gradient_JTDAJ_dense_tiled(nv_pad: int, tile_size: int, njmax: int): nefc = nefc_in[worldid] - sum_val = wp.tile_load( - qM_in[worldid], shape=(nv_pad, nv_pad), bounds_check=True - ) + sum_val = wp.tile_load(qM_in[worldid], shape=(nv_pad, nv_pad), bounds_check=True) # Each tile processes one output tile by looping over all constraints for k in range(0, njmax, TILE_SIZE_K): @@ -2453,12 +2401,7 @@ def update_gradient_JTDAJ_dense_tiled(nv_pad: int, tile_size: int, njmax: int): # AD: leaving bounds-check disabled here because I'm not entirely sure that # everything always hits the fast path. The padding takes care of any # potential OOB accesses. - J_kj = wp.tile_load( - efc_J_in[worldid], - shape=(TILE_SIZE_K, nv_pad), - offset=(k, 0), - bounds_check=False, - ) + J_kj = wp.tile_load(efc_J_in[worldid], shape=(TILE_SIZE_K, nv_pad), offset=(k, 0), bounds_check=False) # state check D_k = wp.tile_load(efc_D_in[worldid], shape=TILE_SIZE_K, offset=k, bounds_check=False) @@ -2473,11 +2416,7 @@ def update_gradient_JTDAJ_dense_tiled(nv_pad: int, tile_size: int, njmax: int): active_tile = wp.tile_map(active_check, tid_tile, threshold_tile) D_k = wp.tile_map(wp.mul, active_tile, D_k) - J_ki = wp.tile_map( - wp.mul, - wp.tile_transpose(J_kj), - wp.tile_broadcast(D_k, shape=(nv_pad, TILE_SIZE_K)), - ) + J_ki = wp.tile_map(wp.mul, wp.tile_transpose(J_kj), wp.tile_broadcast(D_k, shape=(nv_pad, TILE_SIZE_K))) sum_val += wp.tile_matmul(J_ki, J_kj) @@ -2489,32 +2428,32 @@ def update_gradient_JTDAJ_dense_tiled(nv_pad: int, tile_size: int, njmax: int): # TODO(thowell): combine with JTDAJ ? @wp.kernel def update_gradient_JTCJ_sparse( - # Model: - opt_impratio_invsqrt: wp.array(dtype=float), - dof_tri_row: wp.array(dtype=int), - dof_tri_col: wp.array(dtype=int), - # Data in: - contact_dist_in: wp.array(dtype=float), - contact_includemargin_in: wp.array(dtype=float), - contact_friction_in: wp.array(dtype=types.vec5), - contact_dim_in: wp.array(dtype=int), - contact_efc_address_in: wp.array2d(dtype=int), - contact_worldid_in: wp.array(dtype=int), - efc_J_rownnz_in: wp.array2d(dtype=int), - efc_J_rowadr_in: wp.array2d(dtype=int), - efc_J_colind_in: wp.array3d(dtype=int), - efc_J_in: wp.array3d(dtype=float), - efc_D_in: wp.array2d(dtype=float), - efc_state_in: wp.array2d(dtype=int), - naconmax_in: int, - nacon_in: wp.array(dtype=int), - # In: - ctx_Jaref_in: wp.array2d(dtype=float), - ctx_done_in: wp.array(dtype=bool), - nblocks_perblock: int, - dim_block: int, - # Out: - h_out: wp.array3d(dtype=float), + # Model: + opt_impratio_invsqrt: wp.array(dtype=float), + dof_tri_row: wp.array(dtype=int), + dof_tri_col: wp.array(dtype=int), + # Data in: + contact_dist_in: wp.array(dtype=float), + contact_includemargin_in: wp.array(dtype=float), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), + efc_J_colind_in: wp.array3d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_D_in: wp.array2d(dtype=float), + efc_state_in: wp.array2d(dtype=int), + naconmax_in: int, + nacon_in: wp.array(dtype=int), + # In: + ctx_Jaref_in: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + nblocks_perblock: int, + dim_block: int, + # Out: + ctx_h_out: wp.array3d(dtype=float), ): conid_start, elementid = wp.tid() @@ -2529,20 +2468,37 @@ def update_gradient_JTCJ_sparse( worldid = contact_worldid_in[conid] if ctx_done_in[worldid]: - return + continue condim = contact_dim_in[conid] if condim == 1: - return + continue # check contact status if contact_dist_in[conid] - contact_includemargin_in[conid] >= 0.0: - return + continue efcid0 = contact_efc_address_in[conid, 0] if efc_state_in[worldid, efcid0] != types.ConstraintState.CONE: - return + continue + + # All dims share the same sparsity pattern. Scan colind once to find + # the sparse positions of dof1id and dof2id. Skip if either is absent. + rownnz = efc_J_rownnz_in[worldid, efcid0] + rowadr0 = efc_J_rowadr_in[worldid, efcid0] + pos1 = int(-1) + pos2 = int(-1) + for k in range(rownnz): + col = efc_J_colind_in[worldid, 0, rowadr0 + k] + if col == dof1id: + pos1 = k + if col == dof2id: + pos2 = k + if pos1 >= 0 and pos2 >= 0: + break + if pos1 < 0 or pos2 < 0: + continue fri = contact_friction_in[conid] mu = fri[0] * opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] @@ -2551,7 +2507,7 @@ def update_gradient_JTCJ_sparse( dm = math.safe_div(efc_D_in[worldid, efcid0], mu2 * (1.0 + mu2)) if dm == 0.0: - return + continue n = ctx_Jaref_in[worldid, efcid0] * mu u = types.vec6(n, 0.0, 0.0, 0.0, 0.0, 0.0) @@ -2570,52 +2526,40 @@ def update_gradient_JTCJ_sparse( t = wp.max(t, types.MJ_MINVAL) ttt = wp.max(t * t * t, types.MJ_MINVAL) + # Precompute common subexpressions. + mu_over_t = math.safe_div(mu, t) + mu_n_over_ttt = mu * math.safe_div(n, ttt) + mu2_minus_mu_n_over_t = mu2 - mu * math.safe_div(n, t) + h = float(0.0) for dim1id in range(condim): if dim1id == 0: - efcid1 = efcid0 + rowadr1 = rowadr0 + dm_fri1 = dm * mu else: efcid1 = contact_efc_address_in[conid, dim1id] + rowadr1 = efc_J_rowadr_in[worldid, efcid1] + dm_fri1 = dm * fri[dim1id - 1] - # TODO(team): improve performance for sparse code path - rownnz1 = efc_J_rownnz_in[worldid, efcid1] - rowadr1 = efc_J_rowadr_in[worldid, efcid1] - - efc_J11 = float(0.0) - efc_J12 = float(0.0) - for i1 in range(rownnz1): - sparseid1 = rowadr1 + i1 - colind1 = efc_J_colind_in[worldid, 0, sparseid1] - if dof1id == colind1: - efc_J11 = efc_J_in[worldid, 0, sparseid1] - if dof2id == colind1: - efc_J12 = efc_J_in[worldid, 0, sparseid1] - if efc_J11 != 0.0 and efc_J12 != 0.0: - break + # Direct J reads using cached sparse positions. + efc_J11 = efc_J_in[worldid, 0, rowadr1 + pos1] + efc_J12 = efc_J_in[worldid, 0, rowadr1 + pos2] ui = u[dim1id] for dim2id in range(0, dim1id + 1): if dim2id == 0: - efcid2 = efcid0 + rowadr2 = rowadr0 + dm_fri12 = dm_fri1 * mu else: efcid2 = contact_efc_address_in[conid, dim2id] + rowadr2 = efc_J_rowadr_in[worldid, efcid2] + dm_fri12 = dm_fri1 * fri[dim2id - 1] - rownnz2 = efc_J_rownnz_in[worldid, efcid2] - rowadr2 = efc_J_rowadr_in[worldid, efcid2] - - efc_J21 = float(0.0) - efc_J22 = float(0.0) - for i2 in range(rownnz2): - sparseid2 = rowadr2 + i2 - colind2 = efc_J_colind_in[worldid, 0, sparseid2] - if dof1id == colind2: - efc_J21 = efc_J_in[worldid, 0, sparseid2] - if dof2id == colind2: - efc_J22 = efc_J_in[worldid, 0, sparseid2] - if efc_J21 != 0.0 and efc_J22 != 0.0: - break + # Direct J reads using cached sparse positions. + efc_J21 = efc_J_in[worldid, 0, rowadr2 + pos1] + efc_J22 = efc_J_in[worldid, 0, rowadr2 + pos2] uj = u[dim2id] @@ -2623,28 +2567,17 @@ def update_gradient_JTCJ_sparse( if dim1id == 0 and dim2id == 0: hcone = 1.0 elif dim1id == 0: - hcone = -math.safe_div(mu, t) * uj + hcone = -mu_over_t * uj elif dim2id == 0: - hcone = -math.safe_div(mu, t) * ui + hcone = -mu_over_t * ui else: - hcone = mu * math.safe_div(n, ttt) * ui * uj + hcone = mu_n_over_ttt * ui * uj # add to diagonal: mu^2 - mu * n / t if dim1id == dim2id: - hcone += mu2 - mu * math.safe_div(n, t) + hcone += mu2_minus_mu_n_over_t - # pre and post multiply by diag(mu, friction) scale by dm - if dim1id == 0: - fri1 = mu - else: - fri1 = fri[dim1id - 1] - - if dim2id == 0: - fri2 = mu - else: - fri2 = fri[dim2id - 1] - - hcone *= dm * fri1 * fri2 + hcone *= dm_fri12 if hcone != 0.0: h += hcone * efc_J11 * efc_J22 @@ -2652,34 +2585,34 @@ def update_gradient_JTCJ_sparse( if dim1id != dim2id: h += hcone * efc_J12 * efc_J21 - h_out[worldid, dof1id, dof2id] += h + ctx_h_out[worldid, dof1id, dof2id] += h @wp.kernel def update_gradient_JTCJ_dense( - # Model: - opt_impratio_invsqrt: wp.array(dtype=float), - dof_tri_row: wp.array(dtype=int), - dof_tri_col: wp.array(dtype=int), - # Data in: - contact_dist_in: wp.array(dtype=float), - contact_includemargin_in: wp.array(dtype=float), - contact_friction_in: wp.array(dtype=types.vec5), - contact_dim_in: wp.array(dtype=int), - contact_efc_address_in: wp.array2d(dtype=int), - contact_worldid_in: wp.array(dtype=int), - efc_J_in: wp.array3d(dtype=float), - efc_D_in: wp.array2d(dtype=float), - efc_state_in: wp.array2d(dtype=int), - naconmax_in: int, - nacon_in: wp.array(dtype=int), - # In: - ctx_Jaref_in: wp.array2d(dtype=float), - ctx_done_in: wp.array(dtype=bool), - nblocks_perblock: int, - dim_block: int, - # Out: - ctx_h_out: wp.array3d(dtype=float), + # Model: + opt_impratio_invsqrt: wp.array(dtype=float), + dof_tri_row: wp.array(dtype=int), + dof_tri_col: wp.array(dtype=int), + # Data in: + contact_dist_in: wp.array(dtype=float), + contact_includemargin_in: wp.array(dtype=float), + contact_friction_in: wp.array(dtype=types.vec5), + contact_dim_in: wp.array(dtype=int), + contact_efc_address_in: wp.array2d(dtype=int), + contact_worldid_in: wp.array(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_D_in: wp.array2d(dtype=float), + efc_state_in: wp.array2d(dtype=int), + naconmax_in: int, + nacon_in: wp.array(dtype=int), + # In: + ctx_Jaref_in: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + nblocks_perblock: int, + dim_block: int, + # Out: + ctx_h_out: wp.array3d(dtype=float), ): conid_start, elementid = wp.tid() @@ -2863,40 +2796,86 @@ def padding_h(nv: int, ctx_done_in: wp.array(dtype=bool), ctx_h_out: wp.array3d( ctx_h_out[worldid, dofid, dofid] = 1.0 -def _cholesky_factorize_solve( - m: types.Model, d: types.Data, ctx: SolverContext -): +def _cholesky_factorize_solve(m: types.Model, d: types.Data, ctx: SolverContext): """Cholesky factorize ctx.h and solve for Mgrad.""" if m.nv <= _BLOCK_CHOLESKY_DIM: wp.launch_tiled( - update_gradient_cholesky(m.nv), - dim=d.nworld, - inputs=[ctx.grad, ctx.h, ctx.done], - outputs=[ctx.Mgrad], - block_dim=m.block_dim.update_gradient_cholesky, + update_gradient_cholesky(m.nv), + dim=d.nworld, + inputs=[ctx.grad, ctx.h, ctx.done], + outputs=[ctx.Mgrad], + block_dim=m.block_dim.update_gradient_cholesky, ) else: wp.launch( - padding_h, - dim=(d.nworld, m.nv_pad - m.nv), - inputs=[m.nv, ctx.done], - outputs=[ctx.h], + padding_h, + dim=(d.nworld, m.nv_pad - m.nv), + inputs=[m.nv, ctx.done], + outputs=[ctx.h], ) wp.launch_tiled( - update_gradient_cholesky_blocked(types.TILE_SIZE_JTDAJ_DENSE, m.nv_pad), - dim=d.nworld, - inputs=[ - ctx.done, - ctx.grad.reshape(shape=(d.nworld, ctx.grad.shape[1], 1)), - ctx.h, - ctx.hfactor, - ], - outputs=[ctx.Mgrad.reshape(shape=(d.nworld, ctx.Mgrad.shape[1], 1))], - block_dim=m.block_dim.update_gradient_cholesky_blocked, + update_gradient_cholesky_blocked(types.TILE_SIZE_JTDAJ_DENSE, m.nv_pad), + dim=d.nworld, + inputs=[ctx.done, ctx.grad.reshape(shape=(d.nworld, ctx.grad.shape[1], 1)), ctx.h, ctx.hfactor], + outputs=[ctx.Mgrad.reshape(shape=(d.nworld, ctx.Mgrad.shape[1], 1))], + block_dim=m.block_dim.update_gradient_cholesky_blocked, ) +@wp.kernel +def _JTDAJ_sparse( + # Data in: + nefc_in: wp.array(dtype=int), + efc_J_rownnz_in: wp.array2d(dtype=int), + efc_J_rowadr_in: wp.array2d(dtype=int), + efc_J_colind_in: wp.array3d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + efc_D_in: wp.array2d(dtype=float), + efc_state_in: wp.array2d(dtype=int), + # In: + ctx_done_in: wp.array(dtype=bool), + # Out: + h_out: wp.array3d(dtype=float), +): + worldid, efcid = wp.tid() + + if ctx_done_in[worldid]: + return + + if efcid >= nefc_in[worldid]: + return + + efc_D = efc_D_in[worldid, efcid] + efc_state = efc_state_in[worldid, efcid] + + if state_check(efc_D, efc_state) == 0.0: + return + + rownnz = efc_J_rownnz_in[worldid, efcid] + rowadr = efc_J_rowadr_in[worldid, efcid] + + for i in range(rownnz): + sparseidi = rowadr + i + Ji = efc_J_in[worldid, 0, sparseidi] + colindi = efc_J_colind_in[worldid, 0, sparseidi] + for j in range(i, rownnz): + if j == i: + sparseidj = sparseidi + Jj = Ji + colindj = colindi + else: + sparseidj = rowadr + j + Jj = efc_J_in[worldid, 0, sparseidj] + colindj = efc_J_colind_in[worldid, 0, sparseidj] + + h = Ji * Jj * efc_D + wp.atomic_add(h_out[worldid, colindi], colindj, h) + + if i != j: + wp.atomic_add(h_out[worldid, colindj], colindi, h) + + def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): # grad = Ma - qfrc_smooth - qfrc_constraint wp.launch(update_gradient_zero_grad_dot, dim=(d.nworld), inputs=[ctx.done], outputs=[ctx.grad_dot]) @@ -2912,126 +2891,15 @@ def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): smooth.solve_m(m, d, ctx.Mgrad, ctx.grad) elif m.opt.solver == types.SolverType.NEWTON: # h = qM + (efc_J.T * efc_D * active) @ efc_J - if SPARSE_CONSTRAINT_JACOBIAN: - # TODO(team): improve performance for sparse code path - @wp.kernel(module="unique", enable_backward=False) - def _JTDAJ_sparse( - # Data in: - nefc_in: wp.array(dtype=int), - efc_J_rownnz_in: wp.array2d(dtype=int), - efc_J_rowadr_in: wp.array2d(dtype=int), - efc_J_colind_in: wp.array3d(dtype=int), - efc_J_in: wp.array3d(dtype=float), - efc_D_in: wp.array2d(dtype=float), - efc_state_in: wp.array2d(dtype=int), - # In: - ctx_done_in: wp.array(dtype=bool), - # Out: - h_out: wp.array3d(dtype=float), - ): - worldid, efcid = wp.tid() - - if ctx_done_in[worldid]: - return - - if efcid >= nefc_in[worldid]: - return - - efc_D = efc_D_in[worldid, efcid] - efc_state = efc_state_in[worldid, efcid] - - if state_check(efc_D, efc_state) == 0.0: - return - - rownnz = efc_J_rownnz_in[worldid, efcid] - rowadr = efc_J_rowadr_in[worldid, efcid] - - for i in range(rownnz): - sparseidi = rowadr + i - Ji = efc_J_in[worldid, 0, sparseidi] - colindi = efc_J_colind_in[worldid, 0, sparseidi] - for j in range(i, rownnz): - if j == i: - sparseidj = sparseidi - Jj = Ji - colindj = colindi - else: - sparseidj = rowadr + j - Jj = efc_J_in[worldid, 0, sparseidj] - colindj = efc_J_colind_in[worldid, 0, sparseidj] - - h = Ji * Jj * efc_D - wp.atomic_add(h_out[worldid, colindi], colindj, h) - - if i != j: - wp.atomic_add(h_out[worldid, colindj], colindi, h) - + if m.is_sparse: + ctx.h.zero_() wp.launch( - _JTDAJ_sparse, - dim=(d.nworld, d.njmax), - inputs=[ - d.nefc, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.D, - d.efc.state, - ctx.done, - ], - outputs=[ctx.h], + _JTDAJ_sparse, + dim=(d.nworld, d.njmax), + inputs=[d.nefc, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, d.efc.J, d.efc.D, d.efc.state, ctx.done], + outputs=[ctx.h], ) - if m.is_sparse: - wp.launch( - update_gradient_set_h_qM_lower_sparse, - dim=(d.nworld, m.qM_fullm_i.size), - inputs=[m.qM_fullm_i, m.qM_fullm_j, d.qM, ctx.done], - outputs=[ctx.h], - ) - else: - # dense M: copy qM directly into h - @wp.kernel(module="unique", enable_backward=False) - def _set_h_qM_dense( - nv: int, - qM_in: wp.array3d(dtype=float), - ctx_done_in: wp.array(dtype=bool), - ctx_h_out: wp.array3d(dtype=float), - ): - worldid, i, j = wp.tid() - if ctx_done_in[worldid]: - return - if i >= nv or j >= nv: - return - if i >= j: - ctx_h_out[worldid, i, j] += qM_in[worldid, i, j] - - wp.launch( - _set_h_qM_dense, - dim=(d.nworld, m.nv, m.nv), - inputs=[m.nv, d.qM, ctx.done], - outputs=[ctx.h], - ) - elif m.is_sparse: - num_blocks_ceil = ceil(m.nv / types.TILE_SIZE_JTDAJ_SPARSE) - lower_triangle_dim = int(num_blocks_ceil * (num_blocks_ceil + 1) / 2) - with scoped_mathdx_gemm_disabled(): - wp.launch_tiled( - update_gradient_JTDAJ_sparse_tiled( - types.TILE_SIZE_JTDAJ_SPARSE, d.njmax - ), - dim=(d.nworld, lower_triangle_dim), - inputs=[ - d.nefc, - d.efc.J, - d.efc.D, - d.efc.state, - ctx.done, - ], - outputs=[ctx.h], - block_dim=m.block_dim.update_gradient_JTDAJ_sparse, - ) - wp.launch( update_gradient_set_h_qM_lower_sparse, dim=(d.nworld, m.qM_fullm_i.size), @@ -3041,20 +2909,18 @@ def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): else: with scoped_mathdx_gemm_disabled(): wp.launch_tiled( - update_gradient_JTDAJ_dense_tiled( - m.nv_pad, types.TILE_SIZE_JTDAJ_DENSE, d.njmax - ), - dim=d.nworld, - inputs=[ - d.nefc, - d.qM, - d.efc.J, - d.efc.D, - d.efc.state, - ctx.done, - ], - outputs=[ctx.h], - block_dim=m.block_dim.update_gradient_JTDAJ_dense, + update_gradient_JTDAJ_dense_tiled(m.nv_pad, types.TILE_SIZE_JTDAJ_DENSE, d.njmax), + dim=d.nworld, + inputs=[ + d.nefc, + d.qM, + d.efc.J, + d.efc.D, + d.efc.state, + ctx.done, + ], + outputs=[ctx.h], + block_dim=m.block_dim.update_gradient_JTDAJ_dense, ) if m.opt.cone == types.ConeType.ELLIPTIC: @@ -3081,60 +2947,60 @@ def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): nblocks_perblock = int((d.naconmax + dim_block - 1) / dim_block) - if SPARSE_CONSTRAINT_JACOBIAN: + if m.is_sparse: wp.launch( - update_gradient_JTCJ_sparse, - dim=(d.naconmax, m.dof_tri_row.size), - inputs=[ - m.opt.impratio_invsqrt, - m.dof_tri_row, - m.dof_tri_col, - d.contact.dist, - d.contact.includemargin, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.contact.worldid, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.D, - d.efc.state, - d.naconmax, - d.nacon, - ctx.Jaref, - ctx.done, - nblocks_perblock, - dim_block, - ], - outputs=[ctx.h], + update_gradient_JTCJ_sparse, + dim=(dim_block, m.dof_tri_row.size), + inputs=[ + m.opt.impratio_invsqrt, + m.dof_tri_row, + m.dof_tri_col, + d.contact.dist, + d.contact.includemargin, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.D, + d.efc.state, + d.naconmax, + d.nacon, + ctx.Jaref, + ctx.done, + nblocks_perblock, + dim_block, + ], + outputs=[ctx.h], ) else: wp.launch( - update_gradient_JTCJ_dense, - dim=(dim_block, m.dof_tri_row.size), - inputs=[ - m.opt.impratio_invsqrt, - m.dof_tri_row, - m.dof_tri_col, - d.contact.dist, - d.contact.includemargin, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.contact.worldid, - d.efc.J, - d.efc.D, - d.efc.state, - d.naconmax, - d.nacon, - ctx.Jaref, - ctx.done, - nblocks_perblock, - dim_block, - ], - outputs=[ctx.h], + update_gradient_JTCJ_dense, + dim=(dim_block, m.dof_tri_row.size), + inputs=[ + m.opt.impratio_invsqrt, + m.dof_tri_row, + m.dof_tri_col, + d.contact.dist, + d.contact.includemargin, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.contact.worldid, + d.efc.J, + d.efc.D, + d.efc.state, + d.naconmax, + d.nacon, + ctx.Jaref, + ctx.done, + nblocks_perblock, + dim_block, + ], + outputs=[ctx.h], ) _cholesky_factorize_solve(m, d, ctx) @@ -3142,58 +3008,51 @@ def _update_gradient(m: types.Model, d: types.Data, ctx: SolverContext): raise ValueError(f"Unknown solver type: {m.opt.solver}") -def _update_gradient_incremental( - m: types.Model, d: types.Data, ctx: SolverContext -): +def _update_gradient_incremental(m: types.Model, d: types.Data, ctx: SolverContext): """Incremental gradient update: update H for changed constraints + re-factorize. Skips the full J^T*D*J rebuild by applying only the delta from constraints that changed QUADRATIC state, then re-factorizes and solves. """ - wp.launch( - update_gradient_zero_grad_dot, - dim=(d.nworld), - inputs=[ctx.done], - outputs=[ctx.grad_dot], - ) + wp.launch(update_gradient_zero_grad_dot, dim=(d.nworld), inputs=[ctx.done], outputs=[ctx.grad_dot]) wp.launch( - update_gradient_grad, - dim=(d.nworld, m.nv), - inputs=[d.qfrc_smooth, d.qfrc_constraint, d.efc.Ma, ctx.done], - outputs=[ctx.grad, ctx.grad_dot], + update_gradient_grad, + dim=(d.nworld, m.nv), + inputs=[d.qfrc_smooth, d.qfrc_constraint, d.efc.Ma, ctx.done], + outputs=[ctx.grad, ctx.grad_dot], ) # Update lower triangle of H with delta from changed constraints - if SPARSE_CONSTRAINT_JACOBIAN: + if m.is_sparse: wp.launch( - update_gradient_h_incremental_sparse, - dim=(d.nworld, ctx.changed_efc_ids.shape[1]), - inputs=[ - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.D, - d.efc.state, - ctx.changed_efc_ids, - ctx.changed_efc_count, - ], - outputs=[ctx.h], + update_gradient_h_incremental_sparse, + dim=(d.nworld, ctx.changed_efc_ids.shape[1]), + inputs=[ + d.efc.J_rownnz, + d.efc.J_rowadr, + d.efc.J_colind, + d.efc.J, + d.efc.D, + d.efc.state, + ctx.changed_efc_ids, + ctx.changed_efc_count, + ], + outputs=[ctx.h], ) else: lower_tri_dim = m.nv * (m.nv + 1) // 2 wp.launch( - update_gradient_h_incremental, - dim=(d.nworld, lower_tri_dim), - inputs=[ - d.efc.J, - d.efc.D, - d.efc.state, - ctx.changed_efc_ids, - ctx.changed_efc_count, - ], - outputs=[ctx.h], + update_gradient_h_incremental, + dim=(d.nworld, lower_tri_dim), + inputs=[ + d.efc.J, + d.efc.D, + d.efc.state, + ctx.changed_efc_ids, + ctx.changed_efc_count, + ], + outputs=[ctx.h], ) _cholesky_factorize_solve(m, d, ctx) @@ -3347,10 +3206,7 @@ def _solver_iteration( # path in update_constraint_efc has early returns that skip state change # tracking, and the additional JTCJ Hessian term depends on Jaref which # changes every iteration. - incremental = ( - m.opt.solver == types.SolverType.NEWTON - and m.opt.cone != types.ConeType.ELLIPTIC - ) + incremental = m.opt.solver == types.SolverType.NEWTON and m.opt.cone != types.ConeType.ELLIPTIC if incremental: # Must complete before update_constraint_efc which atomically increments. @@ -3422,18 +3278,10 @@ def init_context(m: types.Model, d: types.Data, ctx: SolverContext | InverseCont ctx.Jaref.zero_() wp.launch( - solve_init_jaref(m.is_sparse, m.nv, dofs_per_thread), - dim=(d.nworld, d.njmax, threads_per_efc), - inputs=[ - d.nefc, - d.qacc, - d.efc.J_rownnz, - d.efc.J_rowadr, - d.efc.J_colind, - d.efc.J, - d.efc.aref, - ], - outputs=[ctx.Jaref], + solve_init_jaref(m.is_sparse, m.nv, dofs_per_thread), + dim=(d.nworld, d.njmax, threads_per_efc), + inputs=[d.nefc, d.qacc, d.efc.J_rownnz, d.efc.J_rowadr, d.efc.J_colind, d.efc.J, d.efc.aref], + outputs=[ctx.Jaref], ) # Ma = qM @ qacc diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py index b45472c4..995b2a7c 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py @@ -18,18 +18,52 @@ from typing import Optional, Tuple import warp as wp from mujoco.mjx.third_party.mujoco_warp._src.math import motion_cross +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import DynType from mujoco.mjx.third_party.mujoco_warp._src.types import JointType from mujoco.mjx.third_party.mujoco_warp._src.types import Model from mujoco.mjx.third_party.mujoco_warp._src.types import State from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 +from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope wp.set_module_options({"enable_backward": False}) +# TODO(team): kernel analyzer array slice? +@wp.func +def next_act( + # Model: + opt_timestep: float, # kernel_analyzer: ignore + actuator_dyntype: int, # kernel_analyzer: ignore + actuator_dynprm: vec10f, # kernel_analyzer: ignore + actuator_actrange: wp.vec2, # kernel_analyzer: ignore + # Data In: + act_in: float, # kernel_analyzer: ignore + act_dot_in: float, # kernel_analyzer: ignore + # In: + act_dot_scale: float, + clamp: bool, +) -> float: + # advance actuation + if actuator_dyntype == DynType.FILTEREXACT: + tau = wp.max(MJ_MINVAL, actuator_dynprm[0]) + act = act_in + act_dot_scale * act_dot_in * tau * (1.0 - wp.exp(-opt_timestep / tau)) + elif actuator_dyntype == DynType.USER: + return act_in + else: + act = act_in + act_dot_scale * act_dot_in * opt_timestep + + # clamp to actrange + if clamp: + act = wp.clamp(act, actuator_actrange[0], actuator_actrange[1]) + + return act + + @cache_kernel def mul_m_sparse(check_skip: bool): @wp.kernel(module="unique") diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py index 94434ff9..11b7a7c0 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py @@ -33,11 +33,8 @@ MJ_MAX_EPAFACES = 5 TILE_SIZE_JTDAJ_SPARSE = 16 TILE_SIZE_JTDAJ_DENSE = 16 -# TODO(team): remove after improving performance for sparse constraint jacobian -SPARSE_CONSTRAINT_JACOBIAN = False - -# TODO(team): remove after mjwarp depends on warp-lang >= 1.12 in pyproject.toml -TEXTURE_DTYPE = wp.Texture2D if hasattr(wp, "Texture2D") else int +# maximum number of plugin attributes +_NPLUGINATTR = 128 # TODO(team): add check that all wp.launch_tiled 'block_dim' settings are configurable @@ -53,7 +50,6 @@ class BlockDim: # forward euler_dense: int = 32 actuator_velocity: int = 32 - tendon_velocity: int = 32 # ray ray: int = 64 # sensor @@ -63,6 +59,7 @@ class BlockDim: cholesky_factorize: int = 32 cholesky_solve: int = 32 cholesky_factorize_solve: int = 32 + solve_LD_sparse_fused: int = 64 # solver update_gradient_cholesky: int = 64 update_gradient_cholesky_blocked: int = 32 @@ -351,6 +348,7 @@ class GeomType(enum.IntEnum): BOX: box MESH: mesh SDF: sdf + FLEX: flex """ PLANE = mujoco.mjtGeom.mjGEOM_PLANE @@ -362,6 +360,7 @@ class GeomType(enum.IntEnum): BOX = mujoco.mjtGeom.mjGEOM_BOX MESH = mujoco.mjtGeom.mjGEOM_MESH SDF = mujoco.mjtGeom.mjGEOM_SDF + FLEX = mujoco.mjtGeom.mjGEOM_FLEX # unsupported: NGEOMTYPES, ARROW*, LINE, SKIN, LABEL, NONE @@ -662,6 +661,10 @@ class vec11f(wp.types.vector(length=11, dtype=float)): pass +class vec_pluginattr(wp.types.vector(length=_NPLUGINATTR, dtype=float)): + pass + + class mat23f(wp.types.matrix(shape=(2, 3), dtype=float)): pass @@ -679,6 +682,7 @@ vec6 = vec6f vec8 = vec8f vec10 = vec10f vec11 = vec11f +vec128 = vec_pluginattr mat23 = mat23f mat43 = mat43f mat63 = mat63f @@ -841,6 +845,7 @@ class Model: nflexelem: number of elements in all flexes nflexelemdata: number of element vertex ids in all flexes nflexelemedge: number of element edge ids in all flexes + nflexshelldata: number of shell fragment vertex ids in all flexes nJfe: number of non-zeros in sparse flexedge Jacobian nmesh: number of meshes nmeshvert: number of vertices for all meshes @@ -857,6 +862,7 @@ class Model: nexclude: number of excluded geom pairs neq: number of equality constraints ntendon: number of tendons + nJten: number of non-zeros in sparse tendon Jacobian nwrap: number of wrap objects in all tendon paths nsensor: number of sensors nmocap: number of mocap bodies @@ -973,6 +979,11 @@ class Model: light_poscom0: global position rel. to sub-com in qpos0 (*, nlight, 3) light_pos0: global position rel. to body in qpos0 (*, nlight, 3) light_dir0: global direction in qpos0 (*, nlight, 3) + flex_contype: flex contact type (nflex,) + flex_conaffinity: flex contact affinity (nflex,) + flex_condim: contact dimensionality (1, 3, 4, 6) (nflex,) + flex_friction: friction for (slide, spin, roll) (nflex, 3) + flex_margin: detect contact if dist= 1.12 in pyproject.toml - textures: array("*", TEXTURE_DTYPE) - textures_registry: list[TEXTURE_DTYPE] + textures: array("*", wp.Texture2D) + textures_registry: list[wp.Texture2D] hfield_registry: dict hfield_bvh_id: array("nhfield", wp.uint64) hfield_bounds_size: array("nhfield", wp.vec3) - flex_mesh: wp.Mesh + flex_mesh_registry: dict flex_rgba: array("nflex", wp.vec4) - flex_bvh_id: wp.uint64 - flex_face_point: array("*", wp.vec3) - flex_faceadr: array("nflex", int) - flex_nface: int - flex_nwork: int - flex_group_root: array("nworld", int) - flex_elemdataadr: array("nflex", int) - flex_shell: array("*", int) - flex_shelldataadr: array("nflex", int) - flex_radius: array("nflex", float) - flex_workadr: array("nflex", int) - flex_worknum: array("nflex", int) + flex_bvh_id: array("*", wp.uint64) + flex_group_root: array("nworld", "*", int) flex_render_smooth: bool + bvh_nflexgeom: int + flex_dim_np: array("nflex", int) + flex_geom_flexid: array("*", int) + flex_geom_edgeid: array("*", int) bvh: wp.Bvh bvh_id: wp.uint64 lower: array("*", wp.vec3) @@ -1988,5 +1978,8 @@ class RenderContext: depth_adr: array("ncam", int) render_rgb: array("ncam", bool) render_depth: array("ncam", bool) + seg_data: array("*", int) + seg_adr: array("ncam", int) + render_seg: array("ncam", bool) znear: float total_rays: int diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_pkg.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_pkg.py index b7debea4..e2f8acf9 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_pkg.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_pkg.py @@ -37,9 +37,7 @@ def _parse_version(version_str: str) -> tuple[tuple[int, int | str], ...]: """ # Split on both '.' and '-' parts = re.split(r"[.\-]", version_str) - return tuple( - [(0, int(p)) if p.isdigit() else (-1, p) for p in parts] + [(0, 0)] - ) + return tuple([(0, int(p)) if p.isdigit() else (-1, p) for p in parts] + [(0, 0)]) def check_version(spec: str) -> bool: @@ -65,9 +63,7 @@ def check_version(spec: str) -> bool: """ match = re.match(r"^([a-zA-Z0-9_\-]+)(>=|<=|>|<|==|!=)(.+)$", spec) if not match: - raise ValueError( - f"Invalid version spec '{spec}'. Expected format: 'package>=version'" - ) + raise ValueError(f"Invalid version spec '{spec}'. Expected format: 'package>=version'") package_name, op, version_str = match.groups() required_version = _parse_version(version_str) @@ -87,11 +83,11 @@ def check_version(spec: str) -> bool: installed_version = _parse_version(installed_str) ops = { - ">=": operator.ge, - "<=": operator.le, - ">": operator.gt, - "<": operator.lt, - "==": operator.eq, - "!=": operator.ne, + ">=": operator.ge, + "<=": operator.le, + ">": operator.gt, + "<": operator.lt, + "==": operator.eq, + "!=": operator.ne, } return ops[op](installed_version, required_version) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py index 4ac1cb0d..960e4266 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/warp_util.py @@ -146,13 +146,13 @@ def check_toolkit_driver(): if wp.get_device().is_cuda: if not wp.is_conditional_graph_supported(): warnings.warn( - """ + """ CUDA version < 12.4 detected - graph capture may be unreliable for < 12.3 - conditional graph nodes are not available for < 12.4 Model.opt.graph_conditional should be set to False """, - stacklevel=2, + stacklevel=2, ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml index 37305248..e8d98eb4 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml @@ -28,7 +28,7 @@ requires-python = ">=3.10" dependencies = [ "absl-py", "etils[epath]", - "mujoco>=3.5.0", + "mujoco>=3.6.0", "numpy", "warp-lang>=1.12", ] @@ -55,7 +55,7 @@ dev = [ "ruff", "pygls>=1.0.0,<2.0.0", "lsprotocol>=2023.0.1,<2024.0.0", - "mujoco>=3.5.0.dev0", + "mujoco>=3.6.0.dev0", "warp-lang>=1.11.0.dev0", ] # TODO(team): cpu and cuda JAX optional dependencies are temporary, remove after we land MJX:Warp diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py index 5f1de2d4..cf65f02a 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py @@ -18,7 +18,7 @@ Usage: mjwarp-viewer [flags] Example: - mjwarp-viewer benchmark/humanoid/humanoid.xml -o "opt.solver=cg" + mjwarp-viewer benchmarks/humanoid/humanoid.xml -o "opt.solver=cg" """ import copy @@ -56,6 +56,7 @@ _CLEAR_WARP_CACHE = flags.DEFINE_bool("clear_warp_cache", False, "Clear warp cac _ENGINE = flags.DEFINE_enum_class("engine", EngineOptions.WARP, EngineOptions, "Simulation engine") _NCONMAX = flags.DEFINE_integer("nconmax", None, "Maximum number of contacts.") _NJMAX = flags.DEFINE_integer("njmax", None, "Maximum number of constraints per world.") +_NJMAX_NNZ = flags.DEFINE_integer("njmax_nnz", None, "Maximum number of non-zeros in constraint Jacobian.") _NCCDMAX = flags.DEFINE_integer("nccdmax", None, "Maximum number of CCD contacts per world.") _OVERRIDE = flags.DEFINE_multi_string("override", [], "Model overrides (notation: foo.bar = baz)", short_name="o") _KEYFRAME = flags.DEFINE_integer("keyframe", 0, "keyframe to initialize simulation.") @@ -149,7 +150,7 @@ def _main(argv: Sequence[str]) -> None: override_model(mjm, _OVERRIDE.value) m = mjw.put_model(mjm) override_model(m, _OVERRIDE.value) - d = mjw.put_data(mjm, mjd, nconmax=_NCONMAX.value, njmax=_NJMAX.value, nccdmax=_NCCDMAX.value) + d = mjw.put_data(mjm, mjd, nconmax=_NCONMAX.value, njmax=_NJMAX.value, njmax_nnz=_NJMAX_NNZ.value, nccdmax=_NCCDMAX.value) graph = _compile_step(m, d) if wp.get_device().is_cuda else None if graph is None: mjw.step(m, d) # warmup step diff --git a/mjx/mujoco/mjx/warp/bvh.py b/mjx/mujoco/mjx/warp/bvh.py index 2bb522f2..2c1fee49 100644 --- a/mjx/mujoco/mjx/warp/bvh.py +++ b/mjx/mujoco/mjx/warp/bvh.py @@ -48,20 +48,27 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) + @ffi.format_args_for_warp def _refit_bvh_shim( # Model nworld: int, flex_dim: wp.array(dtype=int), + flex_edge: wp.array(dtype=wp.vec2i), flex_elem: wp.array(dtype=int), + flex_elemadr: wp.array(dtype=int), + flex_elemdataadr: wp.array(dtype=int), flex_elemnum: wp.array(dtype=int), + flex_radius: wp.array(dtype=float), + flex_shell: wp.array(dtype=int), + flex_shelldataadr: wp.array(dtype=int), flex_vertadr: wp.array(dtype=int), + flex_vertnum: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), geom_type: wp.array(dtype=int), nflex: int, - nflexelemdata: int, - nflexvert: int, + nflexelem: int, # Data flexvert_xpos: wp.array2d(dtype=wp.vec3), geom_xmat: wp.array2d(dtype=wp.mat33), @@ -77,15 +84,21 @@ def _refit_bvh_shim( _d.efc = _e _d.contact = _c _m.flex_dim = flex_dim + _m.flex_edge = flex_edge _m.flex_elem = flex_elem + _m.flex_elemadr = flex_elemadr + _m.flex_elemdataadr = flex_elemdataadr _m.flex_elemnum = flex_elemnum + _m.flex_radius = flex_radius + _m.flex_shell = flex_shell + _m.flex_shelldataadr = flex_shelldataadr _m.flex_vertadr = flex_vertadr + _m.flex_vertnum = flex_vertnum _m.geom_dataid = geom_dataid _m.geom_size = geom_size _m.geom_type = geom_type _m.nflex = nflex - _m.nflexelemdata = nflexelemdata - _m.nflexvert = nflexvert + _m.nflexelem = nflexelem _d.flexvert_xpos = flexvert_xpos _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos @@ -113,15 +126,21 @@ def _refit_bvh_jax_impl( out = jf( d.qpos.shape[0], m._impl.flex_dim, + m._impl.flex_edge, m._impl.flex_elem, + m._impl.flex_elemadr, + m._impl.flex_elemdataadr, m._impl.flex_elemnum, + m._impl.flex_radius, + m._impl.flex_shell, + m._impl.flex_shelldataadr, m._impl.flex_vertadr, + m._impl.flex_vertnum, m.geom_dataid, m.geom_size, m.geom_type, m._impl.nflex, - m._impl.nflexelemdata, - m._impl.nflexvert, + m._impl.nflexelem, d._impl.flexvert_xpos, d.geom_xmat, d.geom_xpos, diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index 265bfba3..9d5412d9 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -52,8 +52,26 @@ def _collision_shim( # Model nworld: int, block_dim: mjwp_types.BlockDim, + flex_conaffinity: wp.array(dtype=int), + flex_condim: wp.array(dtype=int), + flex_contype: wp.array(dtype=int), + flex_dim: wp.array(dtype=int), + flex_elem: wp.array(dtype=int), + flex_elemadr: wp.array(dtype=int), + flex_elemdataadr: wp.array(dtype=int), + flex_elemnum: wp.array(dtype=int), + flex_friction: wp.array(dtype=wp.vec3), + flex_margin: wp.array(dtype=float), + flex_radius: wp.array(dtype=float), + flex_shell: wp.array(dtype=int), + flex_shelldataadr: wp.array(dtype=int), + flex_shellnum: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_vertflexid: wp.array(dtype=int), geom_aabb: wp.array3d(dtype=wp.vec3), + geom_conaffinity: wp.array(dtype=int), geom_condim: wp.array(dtype=int), + geom_contype: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_friction: wp.array2d(dtype=wp.vec3), geom_gap: wp.array2d(dtype=float), @@ -90,6 +108,10 @@ def _collision_shim( mesh_vert: wp.array(dtype=wp.vec3), mesh_vertadr: wp.array(dtype=int), mesh_vertnum: wp.array(dtype=int), + nflex: int, + nflexelem: int, + nflexshelldata: int, + nflexvert: int, ngeom: int, nmaxmeshdeg: int, nmaxpolygon: int, @@ -108,7 +130,7 @@ def _collision_shim( pair_solref: wp.array2d(dtype=wp.vec2), pair_solreffriction: wp.array2d(dtype=wp.vec2), plugin: wp.array(dtype=int), - plugin_attr: wp.array(dtype=wp.vec3f), + plugin_attr: wp.array(dtype=mjwp_types.vec_pluginattr), opt__broadphase: int, opt__broadphase_filter: int, opt__ccd_iterations: int, @@ -120,12 +142,15 @@ def _collision_shim( # Data naccdmax: int, naconmax: int, + flexvert_xpos: wp.array2d(dtype=wp.vec3), geom_xmat: wp.array2d(dtype=wp.mat33), geom_xpos: wp.array2d(dtype=wp.vec3), nacon: wp.array(dtype=int), ncollision: wp.array(dtype=int), contact__dim: wp.array(dtype=int), contact__dist: wp.array(dtype=float), + contact__efc_address: wp.array2d(dtype=int), + contact__flex: wp.array(dtype=wp.vec2i), contact__frame: wp.array(dtype=wp.mat33), contact__friction: wp.array(dtype=mjwp_types.vec5), contact__geom: wp.array(dtype=wp.vec2i), @@ -136,6 +161,7 @@ def _collision_shim( contact__solref: wp.array(dtype=wp.vec2), contact__solreffriction: wp.array(dtype=wp.vec2), contact__type: wp.array(dtype=int), + contact__vert: wp.array(dtype=wp.vec2i), contact__worldid: wp.array(dtype=int), ): _m.stat = _s @@ -144,8 +170,26 @@ def _collision_shim( _d.efc = _e _d.contact = _c _m.block_dim = block_dim + _m.flex_conaffinity = flex_conaffinity + _m.flex_condim = flex_condim + _m.flex_contype = flex_contype + _m.flex_dim = flex_dim + _m.flex_elem = flex_elem + _m.flex_elemadr = flex_elemadr + _m.flex_elemdataadr = flex_elemdataadr + _m.flex_elemnum = flex_elemnum + _m.flex_friction = flex_friction + _m.flex_margin = flex_margin + _m.flex_radius = flex_radius + _m.flex_shell = flex_shell + _m.flex_shelldataadr = flex_shelldataadr + _m.flex_shellnum = flex_shellnum + _m.flex_vertadr = flex_vertadr + _m.flex_vertflexid = flex_vertflexid _m.geom_aabb = geom_aabb + _m.geom_conaffinity = geom_conaffinity _m.geom_condim = geom_condim + _m.geom_contype = geom_contype _m.geom_dataid = geom_dataid _m.geom_friction = geom_friction _m.geom_gap = geom_gap @@ -182,6 +226,10 @@ def _collision_shim( _m.mesh_vert = mesh_vert _m.mesh_vertadr = mesh_vertadr _m.mesh_vertnum = mesh_vertnum + _m.nflex = nflex + _m.nflexelem = nflexelem + _m.nflexshelldata = nflexshelldata + _m.nflexvert = nflexvert _m.ngeom = ngeom _m.nmaxmeshdeg = nmaxmeshdeg _m.nmaxpolygon = nmaxpolygon @@ -211,6 +259,8 @@ def _collision_shim( _m.plugin_attr = plugin_attr _d.contact.dim = contact__dim _d.contact.dist = contact__dist + _d.contact.efc_address = contact__efc_address + _d.contact.flex = contact__flex _d.contact.frame = contact__frame _d.contact.friction = contact__friction _d.contact.geom = contact__geom @@ -221,7 +271,9 @@ def _collision_shim( _d.contact.solref = contact__solref _d.contact.solreffriction = contact__solreffriction _d.contact.type = contact__type + _d.contact.vert = contact__vert _d.contact.worldid = contact__worldid + _d.flexvert_xpos = flexvert_xpos _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos _d.naccdmax = naccdmax @@ -238,6 +290,8 @@ def _collision_jax_impl(m: types.Model, d: types.Data): 'ncollision': d._impl.ncollision.shape, 'contact__dim': d._impl.contact__dim.shape, 'contact__dist': d._impl.contact__dist.shape, + 'contact__efc_address': d._impl.contact__efc_address.shape, + 'contact__flex': d._impl.contact__flex.shape, 'contact__frame': d._impl.contact__frame.shape, 'contact__friction': d._impl.contact__friction.shape, 'contact__geom': d._impl.contact__geom.shape, @@ -248,11 +302,12 @@ def _collision_jax_impl(m: types.Model, d: types.Data): 'contact__solref': d._impl.contact__solref.shape, 'contact__solreffriction': d._impl.contact__solreffriction.shape, 'contact__type': d._impl.contact__type.shape, + 'contact__vert': d._impl.contact__vert.shape, 'contact__worldid': d._impl.contact__worldid.shape, } jf = ffi.jax_callable_variadic_tuple( _collision_shim, - num_outputs=15, + num_outputs=18, output_dims=output_dims, vmap_method=None, in_out_argnames=set([ @@ -260,6 +315,8 @@ def _collision_jax_impl(m: types.Model, d: types.Data): 'ncollision', 'contact__dim', 'contact__dist', + 'contact__efc_address', + 'contact__flex', 'contact__frame', 'contact__friction', 'contact__geom', @@ -270,6 +327,7 @@ def _collision_jax_impl(m: types.Model, d: types.Data): 'contact__solref', 'contact__solreffriction', 'contact__type', + 'contact__vert', 'contact__worldid', ]), stage_in_argnames=set([ @@ -299,8 +357,26 @@ def _collision_jax_impl(m: types.Model, d: types.Data): out = jf( d.qpos.shape[0], m._impl.block_dim, + m._impl.flex_conaffinity, + m._impl.flex_condim, + m._impl.flex_contype, + m._impl.flex_dim, + m._impl.flex_elem, + m._impl.flex_elemadr, + m._impl.flex_elemdataadr, + m._impl.flex_elemnum, + m._impl.flex_friction, + m._impl.flex_margin, + m._impl.flex_radius, + m._impl.flex_shell, + m._impl.flex_shelldataadr, + m._impl.flex_shellnum, + m._impl.flex_vertadr, + m._impl.flex_vertflexid, m.geom_aabb, + m.geom_conaffinity, m.geom_condim, + m.geom_contype, m.geom_dataid, m.geom_friction, m.geom_gap, @@ -337,6 +413,10 @@ def _collision_jax_impl(m: types.Model, d: types.Data): m.mesh_vert, m.mesh_vertadr, m.mesh_vertnum, + m._impl.nflex, + m._impl.nflexelem, + m._impl.nflexshelldata, + m._impl.nflexvert, m.ngeom, m._impl.nmaxmeshdeg, m._impl.nmaxpolygon, @@ -366,12 +446,15 @@ def _collision_jax_impl(m: types.Model, d: types.Data): m.opt._impl.sdf_iterations, d._impl.naccdmax, d._impl.naconmax, + d._impl.flexvert_xpos, d.geom_xmat, d.geom_xpos, d._impl.nacon, d._impl.ncollision, d._impl.contact__dim, d._impl.contact__dist, + d._impl.contact__efc_address, + d._impl.contact__flex, d._impl.contact__frame, d._impl.contact__friction, d._impl.contact__geom, @@ -382,6 +465,7 @@ def _collision_jax_impl(m: types.Model, d: types.Data): d._impl.contact__solref, d._impl.contact__solreffriction, d._impl.contact__type, + d._impl.contact__vert, d._impl.contact__worldid, ) d = d.tree_replace({ @@ -389,17 +473,20 @@ def _collision_jax_impl(m: types.Model, d: types.Data): '_impl.ncollision': out[1], '_impl.contact__dim': out[2], '_impl.contact__dist': out[3], - '_impl.contact__frame': out[4], - '_impl.contact__friction': out[5], - '_impl.contact__geom': out[6], - '_impl.contact__geomcollisionid': out[7], - '_impl.contact__includemargin': out[8], - '_impl.contact__pos': out[9], - '_impl.contact__solimp': out[10], - '_impl.contact__solref': out[11], - '_impl.contact__solreffriction': out[12], - '_impl.contact__type': out[13], - '_impl.contact__worldid': out[14], + '_impl.contact__efc_address': out[4], + '_impl.contact__flex': out[5], + '_impl.contact__frame': out[6], + '_impl.contact__friction': out[7], + '_impl.contact__geom': out[8], + '_impl.contact__geomcollisionid': out[9], + '_impl.contact__includemargin': out[10], + '_impl.contact__pos': out[11], + '_impl.contact__solimp': out[12], + '_impl.contact__solref': out[13], + '_impl.contact__solreffriction': out[14], + '_impl.contact__type': out[15], + '_impl.contact__vert': out[16], + '_impl.contact__worldid': out[17], }) return d diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index e3833601..e0f5c053 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -138,6 +138,10 @@ def _forward_shim( eq_type: wp.array(dtype=int), eq_wld_adr: wp.array(dtype=int), flex_bending: wp.array2d(dtype=float), + flex_centered: wp.array(dtype=bool), + flex_conaffinity: wp.array(dtype=int), + flex_condim: wp.array(dtype=int), + flex_contype: wp.array(dtype=int), flex_damping: wp.array(dtype=float), flex_dim: wp.array(dtype=int), flex_edge: wp.array(dtype=wp.vec2i), @@ -146,12 +150,22 @@ def _forward_shim( flex_edgenum: wp.array(dtype=int), flex_elem: wp.array(dtype=int), flex_elemadr: wp.array(dtype=int), + flex_elemdataadr: wp.array(dtype=int), flex_elemedge: wp.array(dtype=int), flex_elemedgeadr: wp.array(dtype=int), flex_elemnum: wp.array(dtype=int), + flex_friction: wp.array(dtype=wp.vec3), + flex_margin: wp.array(dtype=float), + flex_radius: wp.array(dtype=float), + flex_shell: wp.array(dtype=int), + flex_shelldataadr: wp.array(dtype=int), + flex_shellnum: wp.array(dtype=int), flex_stiffness: wp.array2d(dtype=float), + flex_vert: wp.array(dtype=wp.vec3), flex_vertadr: wp.array(dtype=int), flex_vertbodyid: wp.array(dtype=int), + flex_vertflexid: wp.array(dtype=int), + flex_vertnum: wp.array(dtype=int), flexedge_J_colind: wp.array(dtype=int), flexedge_J_rowadr: wp.array(dtype=int), flexedge_J_rownnz: wp.array(dtype=int), @@ -159,7 +173,9 @@ def _forward_shim( flexedge_length0: wp.array(dtype=float), geom_aabb: wp.array3d(dtype=wp.vec3), geom_bodyid: wp.array(dtype=int), + geom_conaffinity: wp.array(dtype=int), geom_condim: wp.array(dtype=int), + geom_contype: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_fluid: wp.array2d(dtype=float), geom_friction: wp.array2d(dtype=wp.vec3), @@ -213,6 +229,7 @@ def _forward_shim( light_targetbodyid: wp.array(dtype=int), mapM2M: wp.array(dtype=int), mat_rgba: wp.array2d(dtype=wp.vec4), + max_ten_J_rownnz: int, mesh_face: wp.array(dtype=wp.vec3i), mesh_faceadr: wp.array(dtype=int), mesh_graph: wp.array(dtype=int), @@ -235,6 +252,7 @@ def _forward_shim( mesh_vertadr: wp.array(dtype=int), mesh_vertnum: wp.array(dtype=int), nC: int, + nJten: int, na: int, nacttrnbody: int, nbody: int, @@ -244,6 +262,7 @@ def _forward_shim( nflex: int, nflexedge: int, nflexelem: int, + nflexshelldata: int, nflexvert: int, ngeom: int, ngravcomp: int, @@ -279,7 +298,9 @@ def _forward_shim( pair_solref: wp.array2d(dtype=wp.vec2), pair_solreffriction: wp.array2d(dtype=wp.vec2), plugin: wp.array(dtype=int), - plugin_attr: wp.array(dtype=wp.vec3f), + plugin_attr: wp.array(dtype=mjwp_types.vec_pluginattr), + qLD_all_updates: wp.array(dtype=wp.vec3i), + qLD_level_offsets: wp.array(dtype=int), qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...], qM_fullm_i: wp.array(dtype=int), qM_fullm_j: wp.array(dtype=int), @@ -323,6 +344,9 @@ def _forward_shim( site_type: wp.array(dtype=int), taxel_sensorid: wp.array(dtype=int), taxel_vertadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_rownnz: wp.array(dtype=int), tendon_actfrclimited: wp.array(dtype=bool), tendon_actfrcrange: wp.array2d(dtype=wp.vec2), tendon_adr: wp.array(dtype=int), @@ -382,6 +406,7 @@ def _forward_shim( naccdmax: int, naconmax: int, njmax: int, + njmax_nnz: int, act: wp.array2d(dtype=float), act_dot: wp.array2d(dtype=float), actuator_force: wp.array2d(dtype=float), @@ -446,7 +471,7 @@ def _forward_shim( subtree_angmom: wp.array2d(dtype=wp.vec3), subtree_com: wp.array2d(dtype=wp.vec3), subtree_linvel: wp.array2d(dtype=wp.vec3), - ten_J: wp.array3d(dtype=float), + ten_J: wp.array2d(dtype=float), ten_length: wp.array2d(dtype=float), ten_velocity: wp.array2d(dtype=float), ten_wrapadr: wp.array2d(dtype=int), @@ -466,6 +491,7 @@ def _forward_shim( contact__dim: wp.array(dtype=int), contact__dist: wp.array(dtype=float), contact__efc_address: wp.array2d(dtype=int), + contact__flex: wp.array(dtype=wp.vec2i), contact__frame: wp.array(dtype=wp.mat33), contact__friction: wp.array(dtype=mjwp_types.vec5), contact__geom: wp.array(dtype=wp.vec2i), @@ -476,6 +502,7 @@ def _forward_shim( contact__solref: wp.array(dtype=wp.vec2), contact__solreffriction: wp.array(dtype=wp.vec2), contact__type: wp.array(dtype=int), + contact__vert: wp.array(dtype=wp.vec2i), contact__worldid: wp.array(dtype=int), efc__D: wp.array2d(dtype=float), efc__J: wp.array3d(dtype=float), @@ -585,6 +612,10 @@ def _forward_shim( _m.eq_type = eq_type _m.eq_wld_adr = eq_wld_adr _m.flex_bending = flex_bending + _m.flex_centered = flex_centered + _m.flex_conaffinity = flex_conaffinity + _m.flex_condim = flex_condim + _m.flex_contype = flex_contype _m.flex_damping = flex_damping _m.flex_dim = flex_dim _m.flex_edge = flex_edge @@ -593,12 +624,22 @@ def _forward_shim( _m.flex_edgenum = flex_edgenum _m.flex_elem = flex_elem _m.flex_elemadr = flex_elemadr + _m.flex_elemdataadr = flex_elemdataadr _m.flex_elemedge = flex_elemedge _m.flex_elemedgeadr = flex_elemedgeadr _m.flex_elemnum = flex_elemnum + _m.flex_friction = flex_friction + _m.flex_margin = flex_margin + _m.flex_radius = flex_radius + _m.flex_shell = flex_shell + _m.flex_shelldataadr = flex_shelldataadr + _m.flex_shellnum = flex_shellnum _m.flex_stiffness = flex_stiffness + _m.flex_vert = flex_vert _m.flex_vertadr = flex_vertadr _m.flex_vertbodyid = flex_vertbodyid + _m.flex_vertflexid = flex_vertflexid + _m.flex_vertnum = flex_vertnum _m.flexedge_J_colind = flexedge_J_colind _m.flexedge_J_rowadr = flexedge_J_rowadr _m.flexedge_J_rownnz = flexedge_J_rownnz @@ -606,7 +647,9 @@ def _forward_shim( _m.flexedge_length0 = flexedge_length0 _m.geom_aabb = geom_aabb _m.geom_bodyid = geom_bodyid + _m.geom_conaffinity = geom_conaffinity _m.geom_condim = geom_condim + _m.geom_contype = geom_contype _m.geom_dataid = geom_dataid _m.geom_fluid = geom_fluid _m.geom_friction = geom_friction @@ -660,6 +703,7 @@ def _forward_shim( _m.light_targetbodyid = light_targetbodyid _m.mapM2M = mapM2M _m.mat_rgba = mat_rgba + _m.max_ten_J_rownnz = max_ten_J_rownnz _m.mesh_face = mesh_face _m.mesh_faceadr = mesh_faceadr _m.mesh_graph = mesh_graph @@ -682,6 +726,7 @@ def _forward_shim( _m.mesh_vertadr = mesh_vertadr _m.mesh_vertnum = mesh_vertnum _m.nC = nC + _m.nJten = nJten _m.na = na _m.nacttrnbody = nacttrnbody _m.nbody = nbody @@ -691,6 +736,7 @@ def _forward_shim( _m.nflex = nflex _m.nflexedge = nflexedge _m.nflexelem = nflexelem + _m.nflexshelldata = nflexshelldata _m.nflexvert = nflexvert _m.ngeom = ngeom _m.ngravcomp = ngravcomp @@ -753,6 +799,8 @@ def _forward_shim( _m.pair_solreffriction = pair_solreffriction _m.plugin = plugin _m.plugin_attr = plugin_attr + _m.qLD_all_updates = qLD_all_updates + _m.qLD_level_offsets = qLD_level_offsets _m.qLD_updates = qLD_updates _m.qM_fullm_i = qM_fullm_i _m.qM_fullm_j = qM_fullm_j @@ -797,6 +845,9 @@ def _forward_shim( _m.stat.meaninertia = stat__meaninertia _m.taxel_sensorid = taxel_sensorid _m.taxel_vertadr = taxel_vertadr + _m.ten_J_colind = ten_J_colind + _m.ten_J_rowadr = ten_J_rowadr + _m.ten_J_rownnz = ten_J_rownnz _m.tendon_actfrclimited = tendon_actfrclimited _m.tendon_actfrcrange = tendon_actfrcrange _m.tendon_adr = tendon_adr @@ -842,6 +893,7 @@ def _forward_shim( _d.contact.dim = contact__dim _d.contact.dist = contact__dist _d.contact.efc_address = contact__efc_address + _d.contact.flex = contact__flex _d.contact.frame = contact__frame _d.contact.friction = contact__friction _d.contact.geom = contact__geom @@ -852,6 +904,7 @@ def _forward_shim( _d.contact.solref = contact__solref _d.contact.solreffriction = contact__solreffriction _d.contact.type = contact__type + _d.contact.vert = contact__vert _d.contact.worldid = contact__worldid _d.crb = crb _d.ctrl = ctrl @@ -895,6 +948,7 @@ def _forward_shim( _d.nf = nf _d.nisland = nisland _d.njmax = njmax + _d.njmax_nnz = njmax_nnz _d.nl = nl _d.qLD = qLD _d.qLDiagInv = qLDiagInv @@ -1018,6 +1072,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__dim': d._impl.contact__dim.shape, 'contact__dist': d._impl.contact__dist.shape, 'contact__efc_address': d._impl.contact__efc_address.shape, + 'contact__flex': d._impl.contact__flex.shape, 'contact__frame': d._impl.contact__frame.shape, 'contact__friction': d._impl.contact__friction.shape, 'contact__geom': d._impl.contact__geom.shape, @@ -1028,6 +1083,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__solref': d._impl.contact__solref.shape, 'contact__solreffriction': d._impl.contact__solreffriction.shape, 'contact__type': d._impl.contact__type.shape, + 'contact__vert': d._impl.contact__vert.shape, 'contact__worldid': d._impl.contact__worldid.shape, 'efc__D': d._impl.efc__D.shape, 'efc__J': d._impl.efc__J.shape, @@ -1047,7 +1103,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _forward_shim, - num_outputs=100, + num_outputs=102, output_dims=output_dims, vmap_method=None, in_out_argnames=set([ @@ -1125,6 +1181,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__dim', 'contact__dist', 'contact__efc_address', + 'contact__flex', 'contact__frame', 'contact__friction', 'contact__geom', @@ -1135,6 +1192,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__solref', 'contact__solreffriction', 'contact__type', + 'contact__vert', 'contact__worldid', 'efc__D', 'efc__J', @@ -1417,6 +1475,10 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.eq_type, m._impl.eq_wld_adr, m._impl.flex_bending, + m._impl.flex_centered, + m._impl.flex_conaffinity, + m._impl.flex_condim, + m._impl.flex_contype, m._impl.flex_damping, m._impl.flex_dim, m._impl.flex_edge, @@ -1425,12 +1487,22 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.flex_edgenum, m._impl.flex_elem, m._impl.flex_elemadr, + m._impl.flex_elemdataadr, m._impl.flex_elemedge, m._impl.flex_elemedgeadr, m._impl.flex_elemnum, + m._impl.flex_friction, + m._impl.flex_margin, + m._impl.flex_radius, + m._impl.flex_shell, + m._impl.flex_shelldataadr, + m._impl.flex_shellnum, m._impl.flex_stiffness, + m._impl.flex_vert, m._impl.flex_vertadr, m._impl.flex_vertbodyid, + m._impl.flex_vertflexid, + m._impl.flex_vertnum, m._impl.flexedge_J_colind, m._impl.flexedge_J_rowadr, m._impl.flexedge_J_rownnz, @@ -1438,7 +1510,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.flexedge_length0, m.geom_aabb, m.geom_bodyid, + m.geom_conaffinity, m.geom_condim, + m.geom_contype, m.geom_dataid, m.geom_fluid, m.geom_friction, @@ -1492,6 +1566,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.light_targetbodyid, m._impl.mapM2M, m.mat_rgba, + m._impl.max_ten_J_rownnz, m.mesh_face, m.mesh_faceadr, m.mesh_graph, @@ -1514,6 +1589,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.mesh_vertadr, m.mesh_vertnum, m.nC, + m.nJten, m.na, m._impl.nacttrnbody, m.nbody, @@ -1523,6 +1599,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.nflex, m._impl.nflexedge, m._impl.nflexelem, + m._impl.nflexshelldata, m._impl.nflexvert, m.ngeom, m.ngravcomp, @@ -1559,6 +1636,8 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.pair_solreffriction, m._impl.plugin, m._impl.plugin_attr, + m._impl.qLD_all_updates, + m._impl.qLD_level_offsets, m._impl.qLD_updates, m._impl.qM_fullm_i, m._impl.qM_fullm_j, @@ -1602,6 +1681,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.site_type, m._impl.taxel_sensorid, m._impl.taxel_vertadr, + m._impl.ten_J_colind, + m._impl.ten_J_rowadr, + m._impl.ten_J_rownnz, m.tendon_actfrclimited, m.tendon_actfrcrange, m.tendon_adr, @@ -1660,6 +1742,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.naccdmax, d._impl.naconmax, d._impl.njmax, + d._impl.njmax_nnz, d.act, d.act_dot, d.actuator_force, @@ -1744,6 +1827,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.contact__dim, d._impl.contact__dist, d._impl.contact__efc_address, + d._impl.contact__flex, d._impl.contact__frame, d._impl.contact__friction, d._impl.contact__geom, @@ -1754,6 +1838,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.contact__solref, d._impl.contact__solreffriction, d._impl.contact__type, + d._impl.contact__vert, d._impl.contact__worldid, d._impl.efc__D, d._impl.efc__J, @@ -1846,32 +1931,34 @@ def _forward_jax_impl(m: types.Model, d: types.Data): '_impl.contact__dim': out[71], '_impl.contact__dist': out[72], '_impl.contact__efc_address': out[73], - '_impl.contact__frame': out[74], - '_impl.contact__friction': out[75], - '_impl.contact__geom': out[76], - '_impl.contact__geomcollisionid': out[77], - '_impl.contact__includemargin': out[78], - '_impl.contact__pos': out[79], - '_impl.contact__solimp': out[80], - '_impl.contact__solref': out[81], - '_impl.contact__solreffriction': out[82], - '_impl.contact__type': out[83], - '_impl.contact__worldid': out[84], - '_impl.efc__D': out[85], - '_impl.efc__J': out[86], - '_impl.efc__J_colind': out[87], - '_impl.efc__J_rowadr': out[88], - '_impl.efc__J_rownnz': out[89], - '_impl.efc__Ma': out[90], - '_impl.efc__aref': out[91], - '_impl.efc__force': out[92], - '_impl.efc__frictionloss': out[93], - '_impl.efc__id': out[94], - '_impl.efc__margin': out[95], - '_impl.efc__pos': out[96], - '_impl.efc__state': out[97], - '_impl.efc__type': out[98], - '_impl.efc__vel': out[99], + '_impl.contact__flex': out[74], + '_impl.contact__frame': out[75], + '_impl.contact__friction': out[76], + '_impl.contact__geom': out[77], + '_impl.contact__geomcollisionid': out[78], + '_impl.contact__includemargin': out[79], + '_impl.contact__pos': out[80], + '_impl.contact__solimp': out[81], + '_impl.contact__solref': out[82], + '_impl.contact__solreffriction': out[83], + '_impl.contact__type': out[84], + '_impl.contact__vert': out[85], + '_impl.contact__worldid': out[86], + '_impl.efc__D': out[87], + '_impl.efc__J': out[88], + '_impl.efc__J_colind': out[89], + '_impl.efc__J_rowadr': out[90], + '_impl.efc__J_rownnz': out[91], + '_impl.efc__Ma': out[92], + '_impl.efc__aref': out[93], + '_impl.efc__force': out[94], + '_impl.efc__frictionloss': out[95], + '_impl.efc__id': out[96], + '_impl.efc__margin': out[97], + '_impl.efc__pos': out[98], + '_impl.efc__state': out[99], + '_impl.efc__type': out[100], + '_impl.efc__vel': out[101], }) return d @@ -1980,6 +2067,10 @@ def _step_shim( eq_type: wp.array(dtype=int), eq_wld_adr: wp.array(dtype=int), flex_bending: wp.array2d(dtype=float), + flex_centered: wp.array(dtype=bool), + flex_conaffinity: wp.array(dtype=int), + flex_condim: wp.array(dtype=int), + flex_contype: wp.array(dtype=int), flex_damping: wp.array(dtype=float), flex_dim: wp.array(dtype=int), flex_edge: wp.array(dtype=wp.vec2i), @@ -1988,12 +2079,22 @@ def _step_shim( flex_edgenum: wp.array(dtype=int), flex_elem: wp.array(dtype=int), flex_elemadr: wp.array(dtype=int), + flex_elemdataadr: wp.array(dtype=int), flex_elemedge: wp.array(dtype=int), flex_elemedgeadr: wp.array(dtype=int), flex_elemnum: wp.array(dtype=int), + flex_friction: wp.array(dtype=wp.vec3), + flex_margin: wp.array(dtype=float), + flex_radius: wp.array(dtype=float), + flex_shell: wp.array(dtype=int), + flex_shelldataadr: wp.array(dtype=int), + flex_shellnum: wp.array(dtype=int), flex_stiffness: wp.array2d(dtype=float), + flex_vert: wp.array(dtype=wp.vec3), flex_vertadr: wp.array(dtype=int), flex_vertbodyid: wp.array(dtype=int), + flex_vertflexid: wp.array(dtype=int), + flex_vertnum: wp.array(dtype=int), flexedge_J_colind: wp.array(dtype=int), flexedge_J_rowadr: wp.array(dtype=int), flexedge_J_rownnz: wp.array(dtype=int), @@ -2001,7 +2102,9 @@ def _step_shim( flexedge_length0: wp.array(dtype=float), geom_aabb: wp.array3d(dtype=wp.vec3), geom_bodyid: wp.array(dtype=int), + geom_conaffinity: wp.array(dtype=int), geom_condim: wp.array(dtype=int), + geom_contype: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_fluid: wp.array2d(dtype=float), geom_friction: wp.array2d(dtype=wp.vec3), @@ -2055,6 +2158,7 @@ def _step_shim( light_targetbodyid: wp.array(dtype=int), mapM2M: wp.array(dtype=int), mat_rgba: wp.array2d(dtype=wp.vec4), + max_ten_J_rownnz: int, mesh_face: wp.array(dtype=wp.vec3i), mesh_faceadr: wp.array(dtype=int), mesh_graph: wp.array(dtype=int), @@ -2077,6 +2181,7 @@ def _step_shim( mesh_vertadr: wp.array(dtype=int), mesh_vertnum: wp.array(dtype=int), nC: int, + nJten: int, nM: int, na: int, nacttrnbody: int, @@ -2087,6 +2192,7 @@ def _step_shim( nflex: int, nflexedge: int, nflexelem: int, + nflexshelldata: int, nflexvert: int, ngeom: int, ngravcomp: int, @@ -2122,7 +2228,9 @@ def _step_shim( pair_solref: wp.array2d(dtype=wp.vec2), pair_solreffriction: wp.array2d(dtype=wp.vec2), plugin: wp.array(dtype=int), - plugin_attr: wp.array(dtype=wp.vec3f), + plugin_attr: wp.array(dtype=mjwp_types.vec_pluginattr), + qLD_all_updates: wp.array(dtype=wp.vec3i), + qLD_level_offsets: wp.array(dtype=int), qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...], qM_fullm_i: wp.array(dtype=int), qM_fullm_j: wp.array(dtype=int), @@ -2166,6 +2274,9 @@ def _step_shim( site_type: wp.array(dtype=int), taxel_sensorid: wp.array(dtype=int), taxel_vertadr: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_rownnz: wp.array(dtype=int), tendon_actfrclimited: wp.array(dtype=bool), tendon_actfrcrange: wp.array2d(dtype=wp.vec2), tendon_adr: wp.array(dtype=int), @@ -2226,6 +2337,7 @@ def _step_shim( naccdmax: int, naconmax: int, njmax: int, + njmax_nnz: int, act: wp.array2d(dtype=float), act_dot: wp.array2d(dtype=float), actuator_force: wp.array2d(dtype=float), @@ -2290,7 +2402,7 @@ def _step_shim( subtree_angmom: wp.array2d(dtype=wp.vec3), subtree_com: wp.array2d(dtype=wp.vec3), subtree_linvel: wp.array2d(dtype=wp.vec3), - ten_J: wp.array3d(dtype=float), + ten_J: wp.array2d(dtype=float), ten_length: wp.array2d(dtype=float), ten_velocity: wp.array2d(dtype=float), ten_wrapadr: wp.array2d(dtype=int), @@ -2310,6 +2422,7 @@ def _step_shim( contact__dim: wp.array(dtype=int), contact__dist: wp.array(dtype=float), contact__efc_address: wp.array2d(dtype=int), + contact__flex: wp.array(dtype=wp.vec2i), contact__frame: wp.array(dtype=wp.mat33), contact__friction: wp.array(dtype=mjwp_types.vec5), contact__geom: wp.array(dtype=wp.vec2i), @@ -2320,6 +2433,7 @@ def _step_shim( contact__solref: wp.array(dtype=wp.vec2), contact__solreffriction: wp.array(dtype=wp.vec2), contact__type: wp.array(dtype=int), + contact__vert: wp.array(dtype=wp.vec2i), contact__worldid: wp.array(dtype=int), efc__D: wp.array2d(dtype=float), efc__J: wp.array3d(dtype=float), @@ -2429,6 +2543,10 @@ def _step_shim( _m.eq_type = eq_type _m.eq_wld_adr = eq_wld_adr _m.flex_bending = flex_bending + _m.flex_centered = flex_centered + _m.flex_conaffinity = flex_conaffinity + _m.flex_condim = flex_condim + _m.flex_contype = flex_contype _m.flex_damping = flex_damping _m.flex_dim = flex_dim _m.flex_edge = flex_edge @@ -2437,12 +2555,22 @@ def _step_shim( _m.flex_edgenum = flex_edgenum _m.flex_elem = flex_elem _m.flex_elemadr = flex_elemadr + _m.flex_elemdataadr = flex_elemdataadr _m.flex_elemedge = flex_elemedge _m.flex_elemedgeadr = flex_elemedgeadr _m.flex_elemnum = flex_elemnum + _m.flex_friction = flex_friction + _m.flex_margin = flex_margin + _m.flex_radius = flex_radius + _m.flex_shell = flex_shell + _m.flex_shelldataadr = flex_shelldataadr + _m.flex_shellnum = flex_shellnum _m.flex_stiffness = flex_stiffness + _m.flex_vert = flex_vert _m.flex_vertadr = flex_vertadr _m.flex_vertbodyid = flex_vertbodyid + _m.flex_vertflexid = flex_vertflexid + _m.flex_vertnum = flex_vertnum _m.flexedge_J_colind = flexedge_J_colind _m.flexedge_J_rowadr = flexedge_J_rowadr _m.flexedge_J_rownnz = flexedge_J_rownnz @@ -2450,7 +2578,9 @@ def _step_shim( _m.flexedge_length0 = flexedge_length0 _m.geom_aabb = geom_aabb _m.geom_bodyid = geom_bodyid + _m.geom_conaffinity = geom_conaffinity _m.geom_condim = geom_condim + _m.geom_contype = geom_contype _m.geom_dataid = geom_dataid _m.geom_fluid = geom_fluid _m.geom_friction = geom_friction @@ -2504,6 +2634,7 @@ def _step_shim( _m.light_targetbodyid = light_targetbodyid _m.mapM2M = mapM2M _m.mat_rgba = mat_rgba + _m.max_ten_J_rownnz = max_ten_J_rownnz _m.mesh_face = mesh_face _m.mesh_faceadr = mesh_faceadr _m.mesh_graph = mesh_graph @@ -2526,6 +2657,7 @@ def _step_shim( _m.mesh_vertadr = mesh_vertadr _m.mesh_vertnum = mesh_vertnum _m.nC = nC + _m.nJten = nJten _m.nM = nM _m.na = na _m.nacttrnbody = nacttrnbody @@ -2536,6 +2668,7 @@ def _step_shim( _m.nflex = nflex _m.nflexedge = nflexedge _m.nflexelem = nflexelem + _m.nflexshelldata = nflexshelldata _m.nflexvert = nflexvert _m.ngeom = ngeom _m.ngravcomp = ngravcomp @@ -2599,6 +2732,8 @@ def _step_shim( _m.pair_solreffriction = pair_solreffriction _m.plugin = plugin _m.plugin_attr = plugin_attr + _m.qLD_all_updates = qLD_all_updates + _m.qLD_level_offsets = qLD_level_offsets _m.qLD_updates = qLD_updates _m.qM_fullm_i = qM_fullm_i _m.qM_fullm_j = qM_fullm_j @@ -2643,6 +2778,9 @@ def _step_shim( _m.stat.meaninertia = stat__meaninertia _m.taxel_sensorid = taxel_sensorid _m.taxel_vertadr = taxel_vertadr + _m.ten_J_colind = ten_J_colind + _m.ten_J_rowadr = ten_J_rowadr + _m.ten_J_rownnz = ten_J_rownnz _m.tendon_actfrclimited = tendon_actfrclimited _m.tendon_actfrcrange = tendon_actfrcrange _m.tendon_adr = tendon_adr @@ -2688,6 +2826,7 @@ def _step_shim( _d.contact.dim = contact__dim _d.contact.dist = contact__dist _d.contact.efc_address = contact__efc_address + _d.contact.flex = contact__flex _d.contact.frame = contact__frame _d.contact.friction = contact__friction _d.contact.geom = contact__geom @@ -2698,6 +2837,7 @@ def _step_shim( _d.contact.solref = contact__solref _d.contact.solreffriction = contact__solreffriction _d.contact.type = contact__type + _d.contact.vert = contact__vert _d.contact.worldid = contact__worldid _d.crb = crb _d.ctrl = ctrl @@ -2741,6 +2881,7 @@ def _step_shim( _d.nf = nf _d.nisland = nisland _d.njmax = njmax + _d.njmax_nnz = njmax_nnz _d.nl = nl _d.qLD = qLD _d.qLDiagInv = qLDiagInv @@ -2868,6 +3009,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__dim': d._impl.contact__dim.shape, 'contact__dist': d._impl.contact__dist.shape, 'contact__efc_address': d._impl.contact__efc_address.shape, + 'contact__flex': d._impl.contact__flex.shape, 'contact__frame': d._impl.contact__frame.shape, 'contact__friction': d._impl.contact__friction.shape, 'contact__geom': d._impl.contact__geom.shape, @@ -2878,6 +3020,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__solref': d._impl.contact__solref.shape, 'contact__solreffriction': d._impl.contact__solreffriction.shape, 'contact__type': d._impl.contact__type.shape, + 'contact__vert': d._impl.contact__vert.shape, 'contact__worldid': d._impl.contact__worldid.shape, 'efc__D': d._impl.efc__D.shape, 'efc__J': d._impl.efc__J.shape, @@ -2897,7 +3040,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _step_shim, - num_outputs=104, + num_outputs=106, output_dims=output_dims, vmap_method=None, in_out_argnames=set([ @@ -2979,6 +3122,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__dim', 'contact__dist', 'contact__efc_address', + 'contact__flex', 'contact__frame', 'contact__friction', 'contact__geom', @@ -2989,6 +3133,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__solref', 'contact__solreffriction', 'contact__type', + 'contact__vert', 'contact__worldid', 'efc__D', 'efc__J', @@ -3275,6 +3420,10 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.eq_type, m._impl.eq_wld_adr, m._impl.flex_bending, + m._impl.flex_centered, + m._impl.flex_conaffinity, + m._impl.flex_condim, + m._impl.flex_contype, m._impl.flex_damping, m._impl.flex_dim, m._impl.flex_edge, @@ -3283,12 +3432,22 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.flex_edgenum, m._impl.flex_elem, m._impl.flex_elemadr, + m._impl.flex_elemdataadr, m._impl.flex_elemedge, m._impl.flex_elemedgeadr, m._impl.flex_elemnum, + m._impl.flex_friction, + m._impl.flex_margin, + m._impl.flex_radius, + m._impl.flex_shell, + m._impl.flex_shelldataadr, + m._impl.flex_shellnum, m._impl.flex_stiffness, + m._impl.flex_vert, m._impl.flex_vertadr, m._impl.flex_vertbodyid, + m._impl.flex_vertflexid, + m._impl.flex_vertnum, m._impl.flexedge_J_colind, m._impl.flexedge_J_rowadr, m._impl.flexedge_J_rownnz, @@ -3296,7 +3455,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.flexedge_length0, m.geom_aabb, m.geom_bodyid, + m.geom_conaffinity, m.geom_condim, + m.geom_contype, m.geom_dataid, m.geom_fluid, m.geom_friction, @@ -3350,6 +3511,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.light_targetbodyid, m._impl.mapM2M, m.mat_rgba, + m._impl.max_ten_J_rownnz, m.mesh_face, m.mesh_faceadr, m.mesh_graph, @@ -3372,6 +3534,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.mesh_vertadr, m.mesh_vertnum, m.nC, + m.nJten, m.nM, m.na, m._impl.nacttrnbody, @@ -3382,6 +3545,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.nflex, m._impl.nflexedge, m._impl.nflexelem, + m._impl.nflexshelldata, m._impl.nflexvert, m.ngeom, m.ngravcomp, @@ -3418,6 +3582,8 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.pair_solreffriction, m._impl.plugin, m._impl.plugin_attr, + m._impl.qLD_all_updates, + m._impl.qLD_level_offsets, m._impl.qLD_updates, m._impl.qM_fullm_i, m._impl.qM_fullm_j, @@ -3461,6 +3627,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.site_type, m._impl.taxel_sensorid, m._impl.taxel_vertadr, + m._impl.ten_J_colind, + m._impl.ten_J_rowadr, + m._impl.ten_J_rownnz, m.tendon_actfrclimited, m.tendon_actfrcrange, m.tendon_adr, @@ -3520,6 +3689,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.naccdmax, d._impl.naconmax, d._impl.njmax, + d._impl.njmax_nnz, d.act, d.act_dot, d.actuator_force, @@ -3604,6 +3774,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.contact__dim, d._impl.contact__dist, d._impl.contact__efc_address, + d._impl.contact__flex, d._impl.contact__frame, d._impl.contact__friction, d._impl.contact__geom, @@ -3614,6 +3785,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.contact__solref, d._impl.contact__solreffriction, d._impl.contact__type, + d._impl.contact__vert, d._impl.contact__worldid, d._impl.efc__D, d._impl.efc__J, @@ -3710,32 +3882,34 @@ def _step_jax_impl(m: types.Model, d: types.Data): '_impl.contact__dim': out[75], '_impl.contact__dist': out[76], '_impl.contact__efc_address': out[77], - '_impl.contact__frame': out[78], - '_impl.contact__friction': out[79], - '_impl.contact__geom': out[80], - '_impl.contact__geomcollisionid': out[81], - '_impl.contact__includemargin': out[82], - '_impl.contact__pos': out[83], - '_impl.contact__solimp': out[84], - '_impl.contact__solref': out[85], - '_impl.contact__solreffriction': out[86], - '_impl.contact__type': out[87], - '_impl.contact__worldid': out[88], - '_impl.efc__D': out[89], - '_impl.efc__J': out[90], - '_impl.efc__J_colind': out[91], - '_impl.efc__J_rowadr': out[92], - '_impl.efc__J_rownnz': out[93], - '_impl.efc__Ma': out[94], - '_impl.efc__aref': out[95], - '_impl.efc__force': out[96], - '_impl.efc__frictionloss': out[97], - '_impl.efc__id': out[98], - '_impl.efc__margin': out[99], - '_impl.efc__pos': out[100], - '_impl.efc__state': out[101], - '_impl.efc__type': out[102], - '_impl.efc__vel': out[103], + '_impl.contact__flex': out[78], + '_impl.contact__frame': out[79], + '_impl.contact__friction': out[80], + '_impl.contact__geom': out[81], + '_impl.contact__geomcollisionid': out[82], + '_impl.contact__includemargin': out[83], + '_impl.contact__pos': out[84], + '_impl.contact__solimp': out[85], + '_impl.contact__solref': out[86], + '_impl.contact__solreffriction': out[87], + '_impl.contact__type': out[88], + '_impl.contact__vert': out[89], + '_impl.contact__worldid': out[90], + '_impl.efc__D': out[91], + '_impl.efc__J': out[92], + '_impl.efc__J_colind': out[93], + '_impl.efc__J_rowadr': out[94], + '_impl.efc__J_rownnz': out[95], + '_impl.efc__Ma': out[96], + '_impl.efc__aref': out[97], + '_impl.efc__force': out[98], + '_impl.efc__frictionloss': out[99], + '_impl.efc__id': out[100], + '_impl.efc__margin': out[101], + '_impl.efc__pos': out[102], + '_impl.efc__state': out[103], + '_impl.efc__type': out[104], + '_impl.efc__vel': out[105], }) return d diff --git a/mjx/mujoco/mjx/warp/forward_test.py b/mjx/mujoco/mjx/warp/forward_test.py index 15e874e3..b80c921d 100644 --- a/mjx/mujoco/mjx/warp/forward_test.py +++ b/mjx/mujoco/mjx/warp/forward_test.py @@ -157,7 +157,16 @@ class ForwardTest(parameterized.TestCase): m.ten_J_rowadr, m.ten_J_colind, ) - tu.assert_eq(dx._impl.ten_J, ten_J, 'ten_J') + # convert sparse warp ten_J to dense representation + warp_ten_J = np.zeros((m.ntendon, m.nv)) + mujoco.mju_sparse2dense( + warp_ten_J, + np.asarray(dx._impl.ten_J), + mx._impl.ten_J_rownnz, + mx._impl.ten_J_rowadr, + mx._impl.ten_J_colind, + ) + tu.assert_eq(warp_ten_J, ten_J, 'ten_J') tu.assert_attr_eq(dx._impl, d, 'ten_wrapadr') tu.assert_attr_eq(dx._impl, d, 'ten_wrapnum') tu.assert_attr_eq(dx._impl, d, 'wrap_xpos') diff --git a/mjx/mujoco/mjx/warp/render.py b/mjx/mujoco/mjx/warp/render.py index 0bab51a3..e723e6b0 100644 --- a/mjx/mujoco/mjx/warp/render.py +++ b/mjx/mujoco/mjx/warp/render.py @@ -48,6 +48,7 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) + @ffi.format_args_for_warp def _render_shim( # Model @@ -56,6 +57,9 @@ def _render_shim( cam_intrinsic: wp.array2d(dtype=wp.vec4), cam_projection: wp.array(dtype=int), cam_sensorsize: wp.array(dtype=wp.vec2), + flex_edge: wp.array(dtype=wp.vec2i), + flex_radius: wp.array(dtype=float), + flex_vertadr: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_matid: wp.array2d(dtype=int), geom_rgba: wp.array2d(dtype=wp.vec4), @@ -68,11 +72,11 @@ def _render_shim( mat_texid: wp.array3d(dtype=int), mat_texrepeat: wp.array2d(dtype=wp.vec2), mesh_faceadr: wp.array(dtype=int), - nflex: int, nlight: int, # Data cam_xmat: wp.array2d(dtype=wp.mat33), cam_xpos: wp.array2d(dtype=wp.vec3), + flexvert_xpos: wp.array2d(dtype=wp.vec3), geom_xmat: wp.array2d(dtype=wp.mat33), geom_xpos: wp.array2d(dtype=wp.vec3), light_xdir: wp.array2d(dtype=wp.vec3), @@ -91,6 +95,9 @@ def _render_shim( _m.cam_intrinsic = cam_intrinsic _m.cam_projection = cam_projection _m.cam_sensorsize = cam_sensorsize + _m.flex_edge = flex_edge + _m.flex_radius = flex_radius + _m.flex_vertadr = flex_vertadr _m.geom_dataid = geom_dataid _m.geom_matid = geom_matid _m.geom_rgba = geom_rgba @@ -103,10 +110,10 @@ def _render_shim( _m.mat_texid = mat_texid _m.mat_texrepeat = mat_texrepeat _m.mesh_faceadr = mesh_faceadr - _m.nflex = nflex _m.nlight = nlight _d.cam_xmat = cam_xmat _d.cam_xpos = cam_xpos + _d.flexvert_xpos = flexvert_xpos _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos _d.light_xdir = light_xdir @@ -155,6 +162,9 @@ def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContextPytree): m.cam_intrinsic, m._impl.cam_projection, m.cam_sensorsize, + m._impl.flex_edge, + m._impl.flex_radius, + m._impl.flex_vertadr, m.geom_dataid, m.geom_matid, m.geom_rgba, @@ -167,10 +177,10 @@ def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContextPytree): m.mat_texid, m._impl.mat_texrepeat, m.mesh_faceadr, - m._impl.nflex, m.nlight, d.cam_xmat, d.cam_xpos, + d._impl.flexvert_xpos, d.geom_xmat, d.geom_xpos, d._impl.light_xdir, diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index eeef2d92..3fe3a712 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -297,17 +297,20 @@ def kinematics_vmap( def _tendon_shim( # Model nworld: int, + body_dofadr: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), body_parentid: wp.array(dtype=int), body_rootid: wp.array(dtype=int), - dof_bodyid: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), jnt_dofadr: wp.array(dtype=int), jnt_qposadr: wp.array(dtype=int), ntendon: int, - nv: int, nwrap: int, site_bodyid: wp.array(dtype=int), + ten_J_colind: wp.array(dtype=int), + ten_J_rowadr: wp.array(dtype=int), + ten_J_rownnz: wp.array(dtype=int), tendon_adr: wp.array(dtype=int), tendon_geom_adr: wp.array(dtype=int), tendon_jnt_adr: wp.array(dtype=int), @@ -327,7 +330,7 @@ def _tendon_shim( qpos: wp.array2d(dtype=float), site_xpos: wp.array2d(dtype=wp.vec3), subtree_com: wp.array2d(dtype=wp.vec3), - ten_J: wp.array3d(dtype=float), + ten_J: wp.array2d(dtype=float), ten_length: wp.array2d(dtype=float), ten_wrapadr: wp.array2d(dtype=int), ten_wrapnum: wp.array2d(dtype=int), @@ -339,17 +342,20 @@ def _tendon_shim( _m.callback = _cb _d.efc = _e _d.contact = _c + _m.body_dofadr = body_dofadr + _m.body_dofnum = body_dofnum _m.body_parentid = body_parentid _m.body_rootid = body_rootid - _m.dof_bodyid = dof_bodyid _m.geom_bodyid = geom_bodyid _m.geom_size = geom_size _m.jnt_dofadr = jnt_dofadr _m.jnt_qposadr = jnt_qposadr _m.ntendon = ntendon - _m.nv = nv _m.nwrap = nwrap _m.site_bodyid = site_bodyid + _m.ten_J_colind = ten_J_colind + _m.ten_J_rowadr = ten_J_rowadr + _m.ten_J_rownnz = ten_J_rownnz _m.tendon_adr = tendon_adr _m.tendon_geom_adr = tendon_geom_adr _m.tendon_jnt_adr = tendon_jnt_adr @@ -416,17 +422,20 @@ def _tendon_jax_impl(m: types.Model, d: types.Data): ) out = jf( d.qpos.shape[0], + m.body_dofadr, + m.body_dofnum, m.body_parentid, m.body_rootid, - m.dof_bodyid, m.geom_bodyid, m.geom_size, m.jnt_dofadr, m.jnt_qposadr, m.ntendon, - m.nv, m.nwrap, m.site_bodyid, + m._impl.ten_J_colind, + m._impl.ten_J_rowadr, + m._impl.ten_J_rownnz, m.tendon_adr, m._impl.tendon_geom_adr, m._impl.tendon_jnt_adr, diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index ff85189f..ef75edc4 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -83,7 +83,7 @@ class BlockDim: qderiv_actuator_dense: int ray: int segmented_sort: int - tendon_velocity: int + solve_LD_sparse_fused: int update_gradient_JTDAJ_dense: int update_gradient_JTDAJ_sparse: int update_gradient_cholesky: int @@ -141,6 +141,10 @@ class ModelWarp(PyTreeNode): eq_ten_adr: np.ndarray eq_wld_adr: np.ndarray flex_bending: np.ndarray + flex_centered: np.ndarray + flex_conaffinity: np.ndarray + flex_condim: np.ndarray + flex_contype: np.ndarray flex_damping: np.ndarray flex_dim: np.ndarray flex_edge: np.ndarray @@ -149,12 +153,21 @@ class ModelWarp(PyTreeNode): flex_edgenum: np.ndarray flex_elem: np.ndarray flex_elemadr: np.ndarray + flex_elemdataadr: np.ndarray flex_elemedge: np.ndarray flex_elemedgeadr: np.ndarray flex_elemnum: np.ndarray + flex_friction: np.ndarray + flex_margin: np.ndarray + flex_radius: np.ndarray + flex_shell: np.ndarray + flex_shelldataadr: np.ndarray + flex_shellnum: np.ndarray flex_stiffness: np.ndarray + flex_vert: np.ndarray flex_vertadr: np.ndarray flex_vertbodyid: np.ndarray + flex_vertflexid: np.ndarray flex_vertnum: np.ndarray flexedge_J_colind: np.ndarray flexedge_J_rowadr: np.ndarray @@ -173,6 +186,7 @@ class ModelWarp(PyTreeNode): light_targetbodyid: np.ndarray mapM2M: np.ndarray mat_texrepeat: jax.Array + max_ten_J_rownnz: int mesh_polyadr: np.ndarray mesh_polymap: np.ndarray mesh_polymapadr: np.ndarray @@ -191,6 +205,7 @@ class ModelWarp(PyTreeNode): nflexelem: int nflexelemdata: int nflexelemedge: int + nflexshelldata: int nflexvert: int nmaxcondim: int nmaxmeshdeg: int @@ -213,6 +228,8 @@ class ModelWarp(PyTreeNode): oct_coeff: np.ndarray plugin: np.ndarray plugin_attr: np.ndarray + qLD_all_updates: np.ndarray + qLD_level_offsets: np.ndarray qLD_updates: Tuple[np.ndarray, ...] qM_fullm_i: np.ndarray qM_fullm_j: np.ndarray @@ -240,6 +257,9 @@ class ModelWarp(PyTreeNode): sensor_vel_adr: np.ndarray taxel_sensorid: np.ndarray taxel_vertadr: np.ndarray + ten_J_colind: np.ndarray + ten_J_rowadr: np.ndarray + ten_J_rownnz: np.ndarray ten_wrapadr_site: np.ndarray ten_wrapnum_site: np.ndarray tendon_geom_adr: np.ndarray @@ -266,6 +286,7 @@ class DataWarp(PyTreeNode): contact__dim: jax.Array contact__dist: jax.Array contact__efc_address: jax.Array + contact__flex: jax.Array contact__frame: jax.Array contact__friction: jax.Array contact__geom: jax.Array @@ -276,6 +297,7 @@ class DataWarp(PyTreeNode): contact__solref: jax.Array contact__solreffriction: jax.Array contact__type: jax.Array + contact__vert: jax.Array contact__worldid: jax.Array crb: jax.Array efc__D: jax.Array @@ -312,6 +334,7 @@ class DataWarp(PyTreeNode): nf: jax.Array nisland: jax.Array njmax: int + njmax_nnz: int njmax_pad: int nl: jax.Array nworld: int @@ -335,6 +358,7 @@ DATA_NON_VMAP = { 'contact__dim', 'contact__dist', 'contact__efc_address', + 'contact__flex', 'contact__frame', 'contact__friction', 'contact__geom', @@ -345,12 +369,14 @@ DATA_NON_VMAP = { 'contact__solref', 'contact__solreffriction', 'contact__type', + 'contact__vert', 'contact__worldid', 'naccdmax', 'nacon', 'naconmax', 'ncollision', 'njmax', + 'njmax_nnz', 'njmax_pad', 'nworld', } @@ -398,6 +424,7 @@ _NDIM = { 'contact__dim': 1, 'contact__dist': 1, 'contact__efc_address': 2, + 'contact__flex': 2, 'contact__frame': 3, 'contact__friction': 2, 'contact__geom': 2, @@ -408,6 +435,7 @@ _NDIM = { 'contact__solref': 2, 'contact__solreffriction': 2, 'contact__type': 1, + 'contact__vert': 2, 'contact__worldid': 1, 'crb': 3, 'ctrl': 2, @@ -451,6 +479,7 @@ _NDIM = { 'nf': 1, 'nisland': 1, 'njmax': 0, + 'njmax_nnz': 0, 'njmax_pad': 0, 'nl': 1, 'nworld': 0, @@ -480,7 +509,7 @@ _NDIM = { 'subtree_angmom': 3, 'subtree_com': 3, 'subtree_linvel': 3, - 'ten_J': 3, + 'ten_J': 2, 'ten_length': 2, 'ten_velocity': 2, 'ten_wrapadr': 2, @@ -535,7 +564,7 @@ _NDIM = { 'block_dim__qderiv_actuator_dense': 0, 'block_dim__ray': 0, 'block_dim__segmented_sort': 0, - 'block_dim__tendon_velocity': 0, + 'block_dim__solve_LD_sparse_fused': 0, 'block_dim__update_gradient_JTDAJ_dense': 0, 'block_dim__update_gradient_JTDAJ_sparse': 0, 'block_dim__update_gradient_cholesky': 0, @@ -608,6 +637,10 @@ _NDIM = { 'eq_wld_adr': 1, 'exclude_signature': 1, 'flex_bending': 2, + 'flex_centered': 1, + 'flex_conaffinity': 1, + 'flex_condim': 1, + 'flex_contype': 1, 'flex_damping': 1, 'flex_dim': 1, 'flex_edge': 2, @@ -616,12 +649,21 @@ _NDIM = { 'flex_edgenum': 1, 'flex_elem': 1, 'flex_elemadr': 1, + 'flex_elemdataadr': 1, 'flex_elemedge': 1, 'flex_elemedgeadr': 1, 'flex_elemnum': 1, + 'flex_friction': 2, + 'flex_margin': 1, + 'flex_radius': 1, + 'flex_shell': 1, + 'flex_shelldataadr': 1, + 'flex_shellnum': 1, 'flex_stiffness': 2, + 'flex_vert': 2, 'flex_vertadr': 1, 'flex_vertbodyid': 1, + 'flex_vertflexid': 1, 'flex_vertnum': 1, 'flexedge_J_colind': 1, 'flexedge_J_rowadr': 1, @@ -692,6 +734,7 @@ _NDIM = { 'mat_rgba': 3, 'mat_texid': 3, 'mat_texrepeat': 3, + 'max_ten_J_rownnz': 0, 'mesh_face': 2, 'mesh_faceadr': 1, 'mesh_graph': 1, @@ -717,6 +760,7 @@ _NDIM = { 'nC': 0, 'nJfe': 0, 'nJmom': 0, + 'nJten': 0, 'nM': 0, 'na': 0, 'nacttrnbody': 0, @@ -730,6 +774,7 @@ _NDIM = { 'nflexelem': 0, 'nflexelemdata': 0, 'nflexelemedge': 0, + 'nflexshelldata': 0, 'nflexvert': 0, 'ngeom': 0, 'ngravcomp': 0, @@ -811,6 +856,8 @@ _NDIM = { 'pair_solreffriction': 3, 'plugin': 1, 'plugin_attr': 2, + 'qLD_all_updates': 2, + 'qLD_level_offsets': 1, 'qLD_updates': -1, 'qM_fullm_i': 1, 'qM_fullm_j': 1, @@ -856,6 +903,9 @@ _NDIM = { 'stat__meaninertia': 1, 'taxel_sensorid': 1, 'taxel_vertadr': 1, + 'ten_J_colind': 1, + 'ten_J_rowadr': 1, + 'ten_J_rownnz': 1, 'ten_wrapadr_site': 1, 'ten_wrapnum_site': 1, 'tendon_actfrclimited': 1, @@ -940,6 +990,7 @@ _BATCH_DIM = { 'contact__dim': False, 'contact__dist': False, 'contact__efc_address': False, + 'contact__flex': False, 'contact__frame': False, 'contact__friction': False, 'contact__geom': False, @@ -950,6 +1001,7 @@ _BATCH_DIM = { 'contact__solref': False, 'contact__solreffriction': False, 'contact__type': False, + 'contact__vert': False, 'contact__worldid': False, 'crb': True, 'ctrl': True, @@ -993,6 +1045,7 @@ _BATCH_DIM = { 'nf': True, 'nisland': True, 'njmax': False, + 'njmax_nnz': False, 'njmax_pad': False, 'nl': True, 'nworld': False, @@ -1077,7 +1130,7 @@ _BATCH_DIM = { 'block_dim__qderiv_actuator_dense': False, 'block_dim__ray': False, 'block_dim__segmented_sort': False, - 'block_dim__tendon_velocity': False, + 'block_dim__solve_LD_sparse_fused': False, 'block_dim__update_gradient_JTDAJ_dense': False, 'block_dim__update_gradient_JTDAJ_sparse': False, 'block_dim__update_gradient_cholesky': False, @@ -1150,6 +1203,10 @@ _BATCH_DIM = { 'eq_wld_adr': False, 'exclude_signature': False, 'flex_bending': False, + 'flex_centered': False, + 'flex_conaffinity': False, + 'flex_condim': False, + 'flex_contype': False, 'flex_damping': False, 'flex_dim': False, 'flex_edge': False, @@ -1158,12 +1215,21 @@ _BATCH_DIM = { 'flex_edgenum': False, 'flex_elem': False, 'flex_elemadr': False, + 'flex_elemdataadr': False, 'flex_elemedge': False, 'flex_elemedgeadr': False, 'flex_elemnum': False, + 'flex_friction': False, + 'flex_margin': False, + 'flex_radius': False, + 'flex_shell': False, + 'flex_shelldataadr': False, + 'flex_shellnum': False, 'flex_stiffness': False, + 'flex_vert': False, 'flex_vertadr': False, 'flex_vertbodyid': False, + 'flex_vertflexid': False, 'flex_vertnum': False, 'flexedge_J_colind': False, 'flexedge_J_rowadr': False, @@ -1234,6 +1300,7 @@ _BATCH_DIM = { 'mat_rgba': True, 'mat_texid': True, 'mat_texrepeat': True, + 'max_ten_J_rownnz': False, 'mesh_face': False, 'mesh_faceadr': False, 'mesh_graph': False, @@ -1259,6 +1326,7 @@ _BATCH_DIM = { 'nC': False, 'nJfe': False, 'nJmom': False, + 'nJten': False, 'nM': False, 'na': False, 'nacttrnbody': False, @@ -1272,6 +1340,7 @@ _BATCH_DIM = { 'nflexelem': False, 'nflexelemdata': False, 'nflexelemedge': False, + 'nflexshelldata': False, 'nflexvert': False, 'ngeom': False, 'ngravcomp': False, @@ -1353,6 +1422,8 @@ _BATCH_DIM = { 'pair_solreffriction': True, 'plugin': False, 'plugin_attr': False, + 'qLD_all_updates': False, + 'qLD_level_offsets': False, 'qLD_updates': False, 'qM_fullm_i': False, 'qM_fullm_j': False, @@ -1398,6 +1469,9 @@ _BATCH_DIM = { 'stat__meaninertia': True, 'taxel_sensorid': False, 'taxel_vertadr': False, + 'ten_J_colind': False, + 'ten_J_rowadr': False, + 'ten_J_rownnz': False, 'ten_wrapadr_site': False, 'ten_wrapnum_site': False, 'tendon_actfrclimited': False,