diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 3965a010..81c5ae2c 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -871,10 +871,10 @@ class Model(PyTreeNode): cam_poscom0: jax.Array cam_pos0: jax.Array cam_mat0: jax.Array - cam_fovy: np.ndarray + cam_fovy: jax.Array cam_resolution: np.ndarray cam_sensorsize: np.ndarray - cam_intrinsic: np.ndarray + cam_intrinsic: jax.Array light_mode: np.ndarray light_type: jax.Array light_castshadow: jax.Array diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py index 24a01166..99e3c032 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py @@ -28,6 +28,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import Model as Model from mujoco.mjx.third_party.mujoco_warp._src.types import Data as Data # isort: on +from mujoco.mjx.third_party.mujoco_warp._src.bvh import refit_bvh as refit_bvh from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import collision as collision from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import nxn_broadphase as nxn_broadphase from mujoco.mjx.third_party.mujoco_warp._src.collision_driver import sap_broadphase as sap_broadphase @@ -46,14 +47,19 @@ from mujoco.mjx.third_party.mujoco_warp._src.forward import rungekutta4 as runge from mujoco.mjx.third_party.mujoco_warp._src.forward import step1 as step1 from mujoco.mjx.third_party.mujoco_warp._src.forward import step2 as step2 from mujoco.mjx.third_party.mujoco_warp._src.inverse import inverse as inverse +from mujoco.mjx.third_party.mujoco_warp._src.io import create_render_context as create_render_context from mujoco.mjx.third_party.mujoco_warp._src.io import get_data_into as get_data_into from mujoco.mjx.third_party.mujoco_warp._src.io import make_data as make_data from mujoco.mjx.third_party.mujoco_warp._src.io import put_data as put_data from mujoco.mjx.third_party.mujoco_warp._src.io import put_model as put_model from mujoco.mjx.third_party.mujoco_warp._src.io import reset_data as reset_data +from mujoco.mjx.third_party.mujoco_warp._src.io import set_const as set_const +from mujoco.mjx.third_party.mujoco_warp._src.io import set_const_0 as set_const_0 +from mujoco.mjx.third_party.mujoco_warp._src.io import set_const_fixed as set_const_fixed from mujoco.mjx.third_party.mujoco_warp._src.passive import passive as passive from mujoco.mjx.third_party.mujoco_warp._src.ray import ray as ray from mujoco.mjx.third_party.mujoco_warp._src.ray import rays as rays +from mujoco.mjx.third_party.mujoco_warp._src.render import render as render 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 @@ -75,6 +81,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.smooth import transmission as trans from mujoco.mjx.third_party.mujoco_warp._src.solver import solve as solve from mujoco.mjx.third_party.mujoco_warp._src.support import contact_force as contact_force from mujoco.mjx.third_party.mujoco_warp._src.support import get_state as get_state +from mujoco.mjx.third_party.mujoco_warp._src.support import jac as jac from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m as mul_m from mujoco.mjx.third_party.mujoco_warp._src.support import set_state as set_state from mujoco.mjx.third_party.mujoco_warp._src.support import xfrc_accumulate as xfrc_accumulate @@ -92,6 +99,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType as GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType as IntegratorType from mujoco.mjx.third_party.mujoco_warp._src.types import JointType as JointType from mujoco.mjx.third_party.mujoco_warp._src.types import Option as Option +from mujoco.mjx.third_party.mujoco_warp._src.types import RenderContext as RenderContext from mujoco.mjx.third_party.mujoco_warp._src.types import SolverType as SolverType from mujoco.mjx.third_party.mujoco_warp._src.types import State as State from mujoco.mjx.third_party.mujoco_warp._src.types import Statistic as Statistic diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py index 0263928f..42912470 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py @@ -16,7 +16,7 @@ """Utilities for benchmarking MuJoCo Warp.""" import time -from typing import Callable, Optional, Tuple +from typing import Callable, Tuple import numpy as np import warp as wp @@ -24,6 +24,7 @@ import warp as wp from mujoco.mjx.third_party.mujoco_warp._src import warp_util from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import RenderContext from mujoco.mjx.third_party.mujoco_warp._src.util_misc import halton @@ -87,10 +88,11 @@ def benchmark( m: Model, d: Data, nstep: int, - ctrls: Optional[np.ndarray] = None, + ctrls: np.ndarray | None = None, event_trace: bool = False, measure_alloc: bool = False, measure_solver_niter: bool = False, + render_context: RenderContext | None = None, ) -> Tuple[float, float, dict, list, list, list, int]: """Benchmark a function of Model and Data. @@ -103,6 +105,7 @@ def benchmark( event_trace: If True, time routines decorated with @event_scope. measure_alloc: If True, record number of contacts and constraints. measure_solver_niter: If True, record the number of solver iterations. + render_context: The render context to use for rendering. Returns: - Time to JIT fn. @@ -120,8 +123,14 @@ def benchmark( with warp_util.EventTracer(enabled=event_trace) as tracer: # capture the whole function as a CUDA graph jit_beg = time.perf_counter() - with wp.ScopedCapture() as capture: - fn(m, d) + + if render_context is not None: + with wp.ScopedCapture() as capture: + fn(m, d, render_context) + else: + with wp.ScopedCapture() as capture: + fn(m, d) + jit_end = time.perf_counter() jit_duration = jit_end - jit_beg diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py new file mode 100644 index 00000000..65e7fad5 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/bvh.py @@ -0,0 +1,1013 @@ +# 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. +# ============================================================================== + +from __future__ import annotations + +from typing import Tuple + +import mujoco +import numpy as np +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import RenderContext +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope + +wp.set_module_options({"enable_backward": False}) + + +@event_scope +def refit_bvh(m: Model, d: Data, rc: RenderContext): + """Refit the dynamic BVH structures in the render context.""" + refit_scene_bvh(m, d, rc) + if m.nflex: + refit_flex_bvh(m, d, rc) + + +@wp.func +def _compute_box_bounds( + # In: + pos: wp.vec3, + rot: wp.mat33, + size: wp.vec3, +) -> Tuple[wp.vec3, wp.vec3]: + min_bound = wp.vec3(MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL) + max_bound = wp.vec3(-MJ_MAXVAL, -MJ_MAXVAL, -MJ_MAXVAL) + + for i in range(2): + for j in range(2): + for k in range(2): + local_corner = wp.vec3( + size[0] * (2.0 * float(i) - 1.0), + size[1] * (2.0 * float(j) - 1.0), + size[2] * (2.0 * float(k) - 1.0), + ) + world_corner = pos + rot @ local_corner + min_bound = wp.min(min_bound, world_corner) + max_bound = wp.max(max_bound, world_corner) + + return min_bound, max_bound + + +@wp.func +def _compute_sphere_bounds( + # In: + pos: wp.vec3, + rot: wp.mat33, + size: wp.vec3, +) -> Tuple[wp.vec3, wp.vec3]: + radius = size[0] + return pos - wp.vec3(radius, radius, radius), pos + wp.vec3(radius, radius, radius) + + +@wp.func +def _compute_capsule_bounds( + # In: + pos: wp.vec3, + rot: wp.mat33, + size: wp.vec3, +) -> Tuple[wp.vec3, wp.vec3]: + radius = size[0] + half_length = size[1] + z = wp.vec3(rot[0, 2], rot[1, 2], rot[2, 2]) + world_end1 = pos - z * half_length + world_end2 = pos + z * half_length + + seg_min = wp.min(world_end1, world_end2) + seg_max = wp.max(world_end1, world_end2) + + inflate = wp.vec3(radius, radius, radius) + return seg_min - inflate, seg_max + inflate + + +@wp.func +def _compute_plane_bounds( + # In: + pos: wp.vec3, + rot: wp.mat33, + size: wp.vec3, +) -> Tuple[wp.vec3, wp.vec3]: + # If plane size is non-positive, treat as infinite plane and use a large default extent + size_scale = wp.max(size[0], size[1]) * 2.0 + if size[0] <= 0.0 or size[1] <= 0.0: + size_scale = 1000.0 + min_bound = wp.vec3(MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL) + max_bound = wp.vec3(-MJ_MAXVAL, -MJ_MAXVAL, -MJ_MAXVAL) + + for i in range(2): + for j in range(2): + local_corner = wp.vec3( + size_scale * (2.0 * float(i) - 1.0), + size_scale * (2.0 * float(j) - 1.0), + 0.0, + ) + world_corner = pos + rot @ local_corner + min_bound = wp.min(min_bound, world_corner) + max_bound = wp.max(max_bound, world_corner) + + min_bound = min_bound - wp.vec3(0.01, 0.01, 0.01) + max_bound = max_bound + wp.vec3(0.01, 0.01, 0.01) + + return min_bound, max_bound + + +@wp.func +def _compute_ellipsoid_bounds( + # In: + pos: wp.vec3, + rot: wp.mat33, + size: wp.vec3, +) -> Tuple[wp.vec3, wp.vec3]: + # Half-extent along each world axis equals the norm of the corresponding row of rot*diag(size) + row0 = wp.vec3(rot[0, 0] * size[0], rot[0, 1] * size[1], rot[0, 2] * size[2]) + row1 = wp.vec3(rot[1, 0] * size[0], rot[1, 1] * size[1], rot[1, 2] * size[2]) + row2 = wp.vec3(rot[2, 0] * size[0], rot[2, 1] * size[1], rot[2, 2] * size[2]) + extent = wp.vec3(wp.length(row0), wp.length(row1), wp.length(row2)) + return pos - extent, pos + extent + + +@wp.func +def _compute_cylinder_bounds( + # In: + pos: wp.vec3, + rot: wp.mat33, + size: wp.vec3, +) -> Tuple[wp.vec3, wp.vec3]: + radius = size[0] + half_height = size[1] + + axis = wp.vec3(rot[0, 2], rot[1, 2], rot[2, 2]) + axis_abs = wp.vec3(wp.abs(axis[0]), wp.abs(axis[1]), wp.abs(axis[2])) + + basis_x = wp.vec3(rot[0, 0], rot[1, 0], rot[2, 0]) + basis_y = wp.vec3(rot[0, 1], rot[1, 1], rot[2, 1]) + + radial_x = radius * wp.sqrt(basis_x[0] * basis_x[0] + basis_y[0] * basis_y[0]) + radial_y = radius * wp.sqrt(basis_x[1] * basis_x[1] + basis_y[1] * basis_y[1]) + radial_z = radius * wp.sqrt(basis_x[2] * basis_x[2] + basis_y[2] * basis_y[2]) + + extent = wp.vec3( + radial_x + half_height * axis_abs[0], + radial_y + half_height * axis_abs[1], + radial_z + half_height * axis_abs[2], + ) + + return pos - extent, pos + extent + + +@wp.kernel +def _compute_bvh_bounds( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + nworld_in: int, + # In: + bvh_ngeom: int, + enabled_geom_ids: wp.array(dtype=int), + mesh_bounds_size: wp.array(dtype=wp.vec3), + hfield_bounds_size: wp.array(dtype=wp.vec3), + # Out: + lower_out: wp.array(dtype=wp.vec3), + upper_out: wp.array(dtype=wp.vec3), + group_out: wp.array(dtype=int), +): + world_id, 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] + type = geom_type[geom_id] + + # TODO: Investigate branch elimination with static loop unrolling + if type == GeomType.SPHERE: + lower_bound, upper_bound = _compute_sphere_bounds(pos, rot, size) + elif type == GeomType.CAPSULE: + lower_bound, upper_bound = _compute_capsule_bounds(pos, rot, size) + elif type == GeomType.PLANE: + lower_bound, upper_bound = _compute_plane_bounds(pos, rot, size) + elif type == GeomType.MESH: + size = mesh_bounds_size[geom_dataid[geom_id]] + lower_bound, upper_bound = _compute_box_bounds(pos, rot, size) + elif type == GeomType.ELLIPSOID: + lower_bound, upper_bound = _compute_ellipsoid_bounds(pos, rot, size) + elif type == GeomType.CYLINDER: + lower_bound, upper_bound = _compute_cylinder_bounds(pos, rot, size) + elif type == GeomType.BOX: + lower_bound, upper_bound = _compute_box_bounds(pos, rot, size) + elif type == GeomType.HFIELD: + size = hfield_bounds_size[geom_dataid[geom_id]] + lower_bound, upper_bound = _compute_box_bounds(pos, 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 + + +@wp.kernel +def compute_bvh_group_roots( + # In: + bvh_id: wp.uint64, + # Out: + group_root_out: wp.array(dtype=int), +): + tid = wp.tid() + root = wp.bvh_get_group_root(bvh_id, tid) + group_root_out[tid] = root + + +def build_scene_bvh(m: Model, d: Data, rc: RenderContext): + """Build a global BVH for all geometries in all worlds.""" + wp.launch( + kernel=_compute_bvh_bounds, + dim=(d.nworld, rc.bvh_ngeom), + inputs=[ + m.geom_type, + m.geom_dataid, + m.geom_size, + d.geom_xpos, + d.geom_xmat, + d.nworld, + rc.bvh_ngeom, + rc.enabled_geom_ids, + rc.mesh_bounds_size, + rc.hfield_bounds_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 + rc.bvh = bvh + rc.bvh_id = bvh.id + + wp.launch( + kernel=compute_bvh_group_roots, + dim=d.nworld, + inputs=[bvh.id], + outputs=[rc.group_root], + ) + + +def refit_scene_bvh(m: Model, d: Data, rc: RenderContext): + wp.launch( + kernel=_compute_bvh_bounds, + dim=(d.nworld, rc.bvh_ngeom), + inputs=[ + m.geom_type, + m.geom_dataid, + m.geom_size, + d.geom_xpos, + d.geom_xmat, + d.nworld, + rc.bvh_ngeom, + rc.enabled_geom_ids, + rc.mesh_bounds_size, + rc.hfield_bounds_size, + rc.lower, + rc.upper, + rc.group, + ], + ) + + rc.bvh.refit() + + +def build_mesh_bvh( + mjm: mujoco.MjModel, + meshid: int, + constructor: str = "sah", + leaf_size: int = 2, +) -> tuple[wp.Mesh, wp.vec3]: + """Create a Warp mesh BVH from mesh data.""" + v_start = mjm.mesh_vertadr[meshid] + v_end = v_start + mjm.mesh_vertnum[meshid] + points = mjm.mesh_vert[v_start:v_end] + + f_start = mjm.mesh_faceadr[meshid] + f_end = mjm.mesh_face.shape[0] if (meshid + 1) >= mjm.mesh_faceadr.shape[0] else mjm.mesh_faceadr[meshid + 1] + indices = mjm.mesh_face[f_start:f_end] + indices = indices.flatten() + pmin = np.min(points, axis=0) + pmax = np.max(points, axis=0) + half = 0.5 * (pmax - pmin) + + points = wp.array(points, dtype=wp.vec3) + indices = wp.array(indices, dtype=wp.int32) + mesh = wp.Mesh(points=points, indices=indices, bvh_constructor=constructor, bvh_leaf_size=leaf_size) + + return mesh, half + + +def _optimize_hfield_mesh( + data: np.ndarray, + nr: int, + nc: int, + sx: float, + sy: float, + sz_scale: float, + width: float, + height: float, +) -> tuple[np.ndarray, np.ndarray]: + """Greedy meshing for heightfield optimization. + + Merges coplanar adjacent cells into larger rectangles to + reduce triangle and vertex count. + """ + points_map = {} + points_list = [] + indices_list = [] + + def get_point_index(r, c): + if (r, c) in points_map: + return points_map[(r, c)] + + # Compute vertex position + x = sx * (float(c) / width - 1.0) + y = sy * (float(r) / height - 1.0) + z = float(data[r, c]) * sz_scale + + idx = len(points_list) + points_list.append([x, y, z]) + points_map[(r, c)] = idx + return idx + + visited = np.zeros((nr - 1, nc - 1), dtype=bool) + + for r in range(nr - 1): + for c in range(nc - 1): + if visited[r, c]: + continue + + # Check if current cell is planar + z00 = data[r, c] + z01 = data[r, c + 1] + z10 = data[r + 1, c] + z11 = data[r + 1, c + 1] + + # Approx check for planarity: z00 + z11 == z01 + z10 + is_planar = abs((z00 + z11) - (z01 + z10)) < 1e-5 + + if not is_planar: + # Must emit single cell (2 triangles) + idx00 = get_point_index(r, c) + idx01 = get_point_index(r, c + 1) + idx10 = get_point_index(r + 1, c) + idx11 = get_point_index(r + 1, c + 1) + + # Tri 1: TL, TR, BR + indices_list.extend([idx00, idx01, idx11]) + # Tri 2: TL, BR, BL + indices_list.extend([idx00, idx11, idx10]) + visited[r, c] = True + continue + + # If planar, try to expand + slope_x = z01 - z00 + slope_y = z10 - z00 + w = 1 + h = 1 + + def fits_plane(rr, cc): + if rr >= nr - 1 or cc >= nc - 1: + return False + # Check planarity of the cell itself + cz00 = data[rr, cc] + cz01 = data[rr, cc + 1] + cz10 = data[rr + 1, cc] + cz11 = data[rr + 1, cc + 1] + if abs((cz00 + cz11) - (cz01 + cz10)) >= 1e-5: + return False + + # Check if it lies on the SAME plane as start cell + # Expected z at (rr, cc) + z_pred = z00 + (rr - r) * slope_y + (cc - c) * slope_x + if abs(cz00 - z_pred) >= 1e-5: + return False + + # Since cell is planar and one corner matches, slopes must match if connected + cslope_x = cz01 - cz00 + cslope_y = cz10 - cz00 + if abs(cslope_x - slope_x) >= 1e-5 or abs(cslope_y - slope_y) >= 1e-5: + return False + + return True + + # Expand width + while c + w < nc - 1 and not visited[r, c + w] and fits_plane(r, c + w): + w += 1 + + # Expand height + while r + h < nr - 1: + # Check entire row + row_ok = True + for k in range(w): + if visited[r + h, c + k] or not fits_plane(r + h, c + k): + row_ok = False + break + if row_ok: + h += 1 + else: + break + + # Mark visited + visited[r : r + h, c : c + w] = True + + # Emit large quad + idx_tl = get_point_index(r, c) + idx_tr = get_point_index(r, c + w) + idx_bl = get_point_index(r + h, c) + idx_br = get_point_index(r + h, c + w) + + # Tri 1: TL, TR, BR + indices_list.extend([idx_tl, idx_tr, idx_br]) + # Tri 2: TL, BR, BL + indices_list.extend([idx_tl, idx_br, idx_bl]) + + return np.array(points_list, dtype=np.float32), np.array(indices_list, dtype=np.int32) + + +def build_hfield_bvh( + mjm: mujoco.MjModel, + hfieldid: int, + constructor: str = "sah", + leaf_size: int = 2, +) -> tuple[wp.Mesh, wp.vec3]: + """Create a Warp mesh BVH from heightfield data.""" + nr = mjm.hfield_nrow[hfieldid] + nc = mjm.hfield_ncol[hfieldid] + sz = np.asarray(mjm.hfield_size[hfieldid], dtype=np.float32) + + adr = mjm.hfield_adr[hfieldid] + data = mjm.hfield_data[adr : adr + nr * nc].reshape((nr, nc)) + + width = 0.5 * max(nc - 1, 1) + height = 0.5 * max(nr - 1, 1) + + points, indices = _optimize_hfield_mesh( + data, + nr, + nc, + sz[0], + sz[1], + sz[2], + width, + height, + ) + pmin = np.min(points, axis=0) + pmax = np.max(points, axis=0) + half = 0.5 * (pmax - pmin) + + points = wp.array(points, dtype=wp.vec3) + indices = wp.array(indices, dtype=wp.int32) + + mesh = wp.Mesh( + points=points, + indices=indices, + bvh_constructor=constructor, + bvh_leaf_size=leaf_size, + ) + + return mesh, half + + +@wp.kernel +def accumulate_flex_vertex_normals( + # Model: + flex_elem: wp.array(dtype=int), + # Data in: + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + # Out: + flexvert_norm_out: wp.array2d(dtype=wp.vec3), +): + """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] + + v0 = flexvert_xpos_in[worldid, i0] + v1 = flexvert_xpos_in[worldid, i1] + v2 = flexvert_xpos_in[worldid, i2] + + face_nrm = wp.cross(v1 - v0, v2 - v0) + face_nrm = wp.normalize(face_nrm) + flexvert_norm_out[worldid, i0] += face_nrm + flexvert_norm_out[worldid, i1] += face_nrm + flexvert_norm_out[worldid, i2] += face_nrm + + +@wp.kernel +def normalize_vertex_normals( + # Out: + flexvert_norm_out: wp.array2d(dtype=wp.vec3), +): + """Normalize accumulated vertex normals.""" + worldid, vertid = wp.tid() + flexvert_norm_out[worldid, vertid] = wp.normalize(flexvert_norm_out[worldid, vertid]) + + +@wp.kernel +def _build_flex_2d_elements( + # Model: + flex_elem: wp.array(dtype=int), + # Data in: + flexvert_xpos_in: wp.array2d(dtype=wp.vec3), + # In: + flexvert_norm_in: wp.array2d(dtype=wp.vec3), + elem_adr: int, + vert_adr: int, + face_offset: int, + radius: float, + nfaces: int, + # Out: + face_point_out: wp.array(dtype=wp.vec3), + face_index_out: wp.array(dtype=int), + group_out: wp.array(dtype=int), +): + """Create faces from 2D flex elements (triangles). + + Two faces (top/bottom) per element, separated by the radius of the flex element. + """ + worldid, elemid = wp.tid() + + base = elem_adr + elemid * 3 + i0 = vert_adr + flex_elem[base + 0] + i1 = vert_adr + flex_elem[base + 1] + i2 = vert_adr + flex_elem[base + 2] + + v0 = flexvert_xpos_in[worldid, i0] + v1 = flexvert_xpos_in[worldid, i1] + v2 = flexvert_xpos_in[worldid, i2] + + n0 = flexvert_norm_in[worldid, i0] + n1 = flexvert_norm_in[worldid, i1] + n2 = flexvert_norm_in[worldid, i2] + + 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 + + world_face_offset = worldid * nfaces + + # First face (top): i0, i1, i2 + 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_index_out[base0 + 0] = base0 + 0 + face_index_out[base0 + 1] = base0 + 1 + face_index_out[base0 + 2] = base0 + 2 + + group_out[face_id0] = worldid + + # Second face (bottom): i0, i2, i1 (opposite winding) + 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 + + face_index_out[base1 + 0] = base1 + 0 + face_index_out[base1 + 1] = base1 + 2 + face_index_out[base1 + 2] = base1 + 1 + + group_out[face_id1] = worldid + + +@wp.kernel +def _build_flex_2d_sides( + # 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, + radius: float, + nface: int, + # Out: + face_point_out: wp.array(dtype=wp.vec3), + face_index_out: wp.array(dtype=int), + group_out: wp.array(dtype=int), +): + """Create side faces from 2D flex shell fragments. + + For each shell fragment (edge i0 -> i1), we emit two triangles: + - one using +radius + - one using -radius (i0/i1 swapped) + """ + 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] + + v0 = flexvert_xpos_in[worldid, i0] + v1 = flexvert_xpos_in[worldid, i1] + + n0 = flexvert_norm_in[worldid, i0] + n1 = flexvert_norm_in[worldid, i1] + + neg_radius = -radius + + # First side i0, i1 with +radius + face_id0 = worldid * nface + face_offset + 2 * shellid + base0 = face_id0 * 3 + face_point_out[base0 + 0] = v0 + n0 * radius + face_point_out[base0 + 1] = v1 + n1 * neg_radius + face_point_out[base0 + 2] = v1 + n1 * radius + face_index_out[base0 + 0] = base0 + 0 + face_index_out[base0 + 1] = base0 + 1 + face_index_out[base0 + 2] = base0 + 2 + + # Second side i1, i0 with -radius + face_id1 = worldid * nface + face_offset + 2 * shellid + 1 + base1 = face_id1 * 3 + face_point_out[base1 + 0] = v1 + n1 * neg_radius + face_point_out[base1 + 1] = v0 + n0 * neg_radius + face_point_out[base1 + 2] = v0 + n0 * radius + face_index_out[base1 + 0] = base1 + 0 + face_index_out[base1 + 1] = base1 + 1 + face_index_out[base1 + 2] = base1 + 2 + + group_out[face_id0] = worldid + group_out[face_id1] = worldid + + +@wp.kernel +def _build_flex_3d_shells( + # 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, + nface: int, + # Out: + face_point_out: wp.array(dtype=wp.vec3), + face_index_out: wp.array(dtype=int), + group_out: wp.array(dtype=int), +): + """Create faces from 3D flex shell fragments (triangles). + + Each shell fragment contributes a single triangle whose vertices are taken + directly from the flex vertex positions (one-sided surface). + """ + 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] + + face_id = worldid * nface + face_offset + shellid + base = face_id * 3 + + v0 = flexvert_xpos_in[worldid, i0] + v1 = flexvert_xpos_in[worldid, i1] + v2 = flexvert_xpos_in[worldid, i2] + + face_point_out[base + 0] = v0 + face_point_out[base + 1] = v1 + face_point_out[base + 2] = v2 + + face_index_out[base + 0] = base + 0 + face_index_out[base + 1] = base + 1 + face_index_out[base + 2] = base + 2 + + group_out[face_id] = worldid + + +@wp.kernel +def _update_flex_face_points( + # Model: + nflex: int, + flex_dim: wp.array(dtype=int), + flex_vertadr: wp.array(dtype=int), + flex_elemnum: wp.array(dtype=int), + flex_elem: wp.array(dtype=int), + # 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, + 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 + + 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] + + 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 + + face_point_out[fbase + 0] = v0 + face_point_out[fbase + 1] = v1 + face_point_out[fbase + 2] = v2 + + +def build_flex_bvh( + mjm: mujoco.MjModel, m: Model, d: Data, 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 + + 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]) + + nface = int(flex_faceadr[-1]) + flex_faceadr = flex_faceadr[:-1] + + face_point = wp.zeros(nface * 3 * d.nworld, dtype=wp.vec3) + face_index = wp.zeros(nface * 3 * d.nworld, dtype=wp.int32) + group = wp.zeros(nface * d.nworld, dtype=int) + + flexvert_norm = wp.zeros(d.flexvert_xpos.shape, dtype=wp.vec3) + flex_shell = wp.array(mjm.flex_shell, dtype=int) + + wp.launch( + kernel=accumulate_flex_vertex_normals, + dim=(d.nworld, m.nflexelemdata // 3), + inputs=[m.flex_elem, d.flexvert_xpos], + outputs=[flexvert_norm], + ) + + wp.launch( + kernel=normalize_vertex_normals, + dim=(d.nworld, m.nflexvert), + 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] + + if dim == 2: + wp.launch( + kernel=_build_flex_2d_elements, + dim=(d.nworld, nelem), + inputs=[ + m.flex_elem, + d.flexvert_xpos, + flexvert_norm, + elem_adr, + vert_adr, + flex_faceadr[f], + mjm.flex_radius[f], + nface, + ], + outputs=[face_point, face_index, group], + ) + + wp.launch( + kernel=_build_flex_2d_sides, + dim=(d.nworld, nshell), + inputs=[ + d.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=(d.nworld, nshell), + inputs=[ + d.flexvert_xpos, + flex_shell, + shell_adr, + vert_adr, + flex_faceadr[f], + nface, + ], + outputs=[face_point, face_index, group], + ) + + flex_mesh = wp.Mesh( + points=face_point, + indices=face_index, + groups=group, + bvh_constructor=constructor, + bvh_leaf_size=leaf_size, + ) + + group_root = wp.zeros(d.nworld, dtype=int) + wp.launch( + kernel=compute_bvh_group_roots, + dim=d.nworld, + inputs=[flex_mesh.id], + outputs=[group_root], + ) + + return ( + flex_mesh, + face_point, + group_root, + flex_shell, + flex_faceadr, + nface, + ) + + +def refit_flex_bvh(m: Model, d: Data, rc: RenderContext): + """Refit the flex BVH.""" + flexvert_norm = wp.zeros(d.flexvert_xpos.shape, dtype=wp.vec3) + + wp.launch( + kernel=accumulate_flex_vertex_normals, + dim=(d.nworld, m.nflexelemdata // 3), + inputs=[ + m.flex_elem, + d.flexvert_xpos, + ], + outputs=[flexvert_norm], + ) + + wp.launch( + kernel=normalize_vertex_normals, + dim=(d.nworld, m.nflexvert), + 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], + ) + + rc.flex_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 e022df89..cd6dd6e6 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 @@ -29,6 +29,8 @@ from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index 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 CollisionContext 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 @@ -38,7 +40,6 @@ 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 -from mujoco.mjx.third_party.mujoco_warp._src.warp_util import nested_kernel # TODO(team): improve compile time to enable backward pass wp.set_module_options({"enable_backward": False}) @@ -46,42 +47,6 @@ wp.set_module_options({"enable_backward": False}) vec_maxconpair = wp.types.vector(length=MJ_MAXCONPAIR, dtype=float) mat_maxconpair = wp.types.matrix(shape=(MJ_MAXCONPAIR, 3), dtype=float) -_CONVEX_COLLISION_PAIRS = [ - (GeomType.HFIELD, GeomType.SPHERE), - (GeomType.HFIELD, GeomType.CAPSULE), - (GeomType.HFIELD, GeomType.ELLIPSOID), - (GeomType.HFIELD, GeomType.CYLINDER), - (GeomType.HFIELD, GeomType.BOX), - (GeomType.HFIELD, GeomType.MESH), - (GeomType.SPHERE, GeomType.ELLIPSOID), - (GeomType.SPHERE, GeomType.MESH), - (GeomType.CAPSULE, GeomType.ELLIPSOID), - (GeomType.CAPSULE, GeomType.CYLINDER), - (GeomType.CAPSULE, GeomType.MESH), - (GeomType.ELLIPSOID, GeomType.ELLIPSOID), - (GeomType.ELLIPSOID, GeomType.CYLINDER), - (GeomType.ELLIPSOID, GeomType.BOX), - (GeomType.ELLIPSOID, GeomType.MESH), - (GeomType.CYLINDER, GeomType.CYLINDER), - (GeomType.CYLINDER, GeomType.BOX), - (GeomType.CYLINDER, GeomType.MESH), - (GeomType.BOX, GeomType.MESH), - (GeomType.MESH, GeomType.MESH), -] - - -def _check_convex_collision_pairs(): - prev_idx = -1 - for pair in _CONVEX_COLLISION_PAIRS: - idx = upper_trid_index(len(GeomType), pair[0].value, pair[1].value) - if pair[1] < pair[0] or idx <= prev_idx: - return False - prev_idx = idx - return True - - -assert _check_convex_collision_pairs(), "_CONVEX_COLLISION_PAIRS is in invalid order." - @wp.func def _hfield_filter( @@ -91,6 +56,11 @@ def _hfield_filter( geom_size: wp.array2d(dtype=wp.vec3), geom_rbound: wp.array2d(dtype=float), geom_margin: wp.array2d(dtype=float), + 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), hfield_size: wp.array(dtype=wp.vec4), # Data in: geom_xpos_in: wp.array2d(dtype=wp.vec3), @@ -126,14 +96,14 @@ def _hfield_filter( # box-sphere test: horizontal plane for i in range(2): if (size1[i] < pos[i] - r2 - margin) or (-size1[i] > pos[i] + r2 + margin): - return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf + return True, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL # box-sphere test: vertical direction if size1[2] < pos[2] - r2 - margin: # up - return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf + return True, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL if -size1[3] > pos[2] + r2 + margin: # down - return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf + return True, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL mat2 = geom_xmat_in[worldid, g2] mat = mat1T @ mat2 @@ -148,6 +118,15 @@ def _hfield_filter( geomtype2 = geom_type[g2] + # load mesh vertex data for support function queries + if geomtype2 == GeomType.MESH: + dataid = geom_dataid[g2] + geom2.vertadr = wp.where(dataid >= 0, mesh_vertadr[dataid], -1) + geom2.vertnum = wp.where(dataid >= 0, mesh_vertnum[dataid], -1) + geom2.graphadr = wp.where(dataid >= 0, mesh_graphadr[dataid], -1) + geom2.vert = mesh_vert + geom2.graph = mesh_graph + # use support functions for tight AABB bounds xmax = support(geom2, geomtype2, wp.vec3(1.0, 0.0, 0.0)).point[0] xmin = support(geom2, geomtype2, wp.vec3(-1.0, 0.0, 0.0)).point[0] @@ -165,7 +144,7 @@ def _hfield_filter( or (zmin - margin > size1[2]) or (zmax + margin < -size1[3]) ): - return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf + return True, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL else: return False, xmin, xmax, ymin, ymax, zmin, zmax @@ -176,11 +155,12 @@ def ccd_hfield_kernel_builder( geomtype2: int, gjk_iterations: int, epa_iterations: int, + geomgeomid: int, ): """Kernel builder for heightfield CCD collisions (no multiccd args).""" # runs convex collision on a set of geom pairs to recover contact info - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def ccd_hfield_kernel( # Model: opt_ccd_tolerance: wp.array(dtype=float), @@ -196,15 +176,10 @@ def ccd_hfield_kernel_builder( geom_friction: wp.array2d(dtype=wp.vec3), geom_margin: wp.array2d(dtype=float), geom_gap: wp.array2d(dtype=float), - hfield_adr: wp.array(dtype=int), - hfield_nrow: wp.array(dtype=int), - hfield_ncol: wp.array(dtype=int), - hfield_size: wp.array(dtype=wp.vec4), - hfield_data: wp.array(dtype=float), mesh_vertadr: wp.array(dtype=int), mesh_vertnum: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), mesh_graphadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), mesh_graph: wp.array(dtype=int), mesh_polynum: wp.array(dtype=int), mesh_polyadr: wp.array(dtype=int), @@ -215,6 +190,11 @@ def ccd_hfield_kernel_builder( mesh_polymapadr: wp.array(dtype=int), mesh_polymapnum: wp.array(dtype=int), mesh_polymap: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), + hfield_nrow: wp.array(dtype=int), + hfield_ncol: wp.array(dtype=int), + hfield_adr: wp.array(dtype=int), + hfield_data: wp.array(dtype=float), pair_dim: wp.array(dtype=int), pair_solref: wp.array2d(dtype=wp.vec2), pair_solreffriction: wp.array2d(dtype=wp.vec2), @@ -223,24 +203,23 @@ def ccd_hfield_kernel_builder( pair_gap: wp.array2d(dtype=float), pair_friction: wp.array2d(dtype=vec5), # Data in: - naconmax_in: int, geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), + naconmax_in: int, + naccdmax_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), - ncollision_in: wp.array(dtype=int), - # In: - epa_vert1_in: wp.array2d(dtype=wp.vec3), - epa_vert2_in: wp.array2d(dtype=wp.vec3), - epa_vert_index1_in: wp.array2d(dtype=int), - epa_vert_index2_in: wp.array2d(dtype=int), + epa_vert_in: wp.array2d(dtype=wp.vec3), + epa_vert_index_in: wp.array2d(dtype=int), epa_face_in: wp.array2d(dtype=int), epa_pr_in: wp.array2d(dtype=wp.vec3), epa_norm2_in: wp.array2d(dtype=float), epa_horizon_in: wp.array2d(dtype=int), + nccd_in: wp.array(dtype=int), # Data out: - nacon_out: wp.array(dtype=int), contact_dist_out: wp.array(dtype=float), contact_pos_out: wp.array(dtype=wp.vec3), contact_frame_out: wp.array(dtype=wp.mat33), @@ -254,27 +233,48 @@ def ccd_hfield_kernel_builder( 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), ): - tid = wp.tid() - if tid >= ncollision_in[0]: + collisionid = wp.tid() + if collisionid >= ncollision_in[0]: return - geoms = collision_pair_in[tid] + geoms = collision_pair_in[collisionid] g1 = geoms[0] g2 = geoms[1] if geom_type[g1] != geomtype1 or geom_type[g2] != geomtype2: return - worldid = collision_worldid_in[tid] + worldid = collision_worldid_in[collisionid] # height field filter no_hf_collision, xmin, xmax, ymin, ymax, zmin, zmax = _hfield_filter( - geom_type, geom_dataid, geom_size, geom_rbound, geom_margin, hfield_size, geom_xpos_in, geom_xmat_in, worldid, g1, g2 + geom_type, + geom_dataid, + geom_size, + geom_rbound, + geom_margin, + mesh_vertadr, + mesh_vertnum, + mesh_graphadr, + mesh_vert, + mesh_graph, + hfield_size, + geom_xpos_in, + geom_xmat_in, + worldid, + g1, + g2, ) if no_hf_collision: return + ccdid = wp.atomic_add(nccd_in, wp.static(geomgeomid), 1) + if ccdid >= naccdmax_in: + wp.printf("CCD overflow - please increase naccdmax to %u\n", ccdid) + return + _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params( geom_condim, geom_priority, @@ -293,7 +293,7 @@ def ccd_hfield_kernel_builder( pair_friction, collision_pair_in, collision_pairid_in, - tid, + collisionid, worldid, ) @@ -365,26 +365,23 @@ def ccd_hfield_kernel_builder( hfield_contact_dist = vec_maxconpair() hfield_contact_pos = mat_maxconpair() hfield_contact_normal = mat_maxconpair() - min_dist = float(wp.inf) - min_normal = wp.vec3(wp.inf, wp.inf, wp.inf) - min_pos = wp.vec3(wp.inf, wp.inf, wp.inf) + min_dist = float(MJ_MAXVAL) + min_normal = wp.vec3(MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL) + min_pos = wp.vec3(MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL) min_id = int(-1) - # TODO(team): height field margin? - geom1.margin = margin + # geom1 margin added to z-axis of hfield prism top points below geom2.margin = margin # EPA memory - epa_vert1 = epa_vert1_in[tid] - epa_vert2 = epa_vert2_in[tid] - epa_vert_index1 = epa_vert_index1_in[tid] - epa_vert_index2 = epa_vert_index2_in[tid] - epa_face = epa_face_in[tid] - epa_pr = epa_pr_in[tid] - epa_norm2 = epa_norm2_in[tid] - epa_horizon = epa_horizon_in[tid] + epa_vert = epa_vert_in[ccdid] + epa_vert_index = epa_vert_index_in[ccdid] + epa_face = epa_face_in[ccdid] + epa_pr = epa_pr_in[ccdid] + epa_norm2 = epa_norm2_in[ccdid] + epa_horizon = epa_horizon_in[ccdid] - collision_pairid = collision_pairid_in[tid] + collision_pairid = collision_pairid_in[collisionid] # process all prisms in subgrid count = int(0) @@ -439,11 +436,7 @@ def ccd_hfield_kernel_builder( geom1.hfprism = prism # prism center - x1 = geom1.pos - x1_ = wp.vec3(0.0, 0.0, 0.0) - for i in range(6): - x1_ += prism[i] - x1 += geom1.rot @ (x1_ / 6.0) + x1 = geom1.pos + geom1.rot @ (prism[0] + prism[1] + prism[2] + prism[3] + prism[4] + prism[5]) * wp.static(1.0 / 6.0) dist, ncontact, w1, w2, idx = ccd( opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]], @@ -456,10 +449,8 @@ def ccd_hfield_kernel_builder( geomtype2, x1, geom2.pos, - epa_vert1, - epa_vert2, - epa_vert_index1, - epa_vert_index2, + epa_vert, + epa_vert_index, epa_face, epa_pr, epa_norm2, @@ -529,13 +520,12 @@ def ccd_hfield_kernel_builder( ) # TODO(team): routine for select subset of contacts - # TODO(team): if use_multiccd? if wp.static(True): MIN_DIST_TO_NEXT_CONTACT = 1.0e-3 # contact 1: furthest from minimum distance contact id1 = int(-1) - dist1 = float(-wp.inf) + dist1 = float(-MJ_MAXVAL) for i in range(count): if i == min_id: continue @@ -589,7 +579,7 @@ def ccd_hfield_kernel_builder( dist_min1 = wp.cross(min_normal, min_pos - pos1) id2 = int(-1) - dist_12 = float(-wp.inf) + dist_12 = float(-MJ_MAXVAL) for i in range(count): if i == min_id or i == id1: continue @@ -644,7 +634,7 @@ def ccd_hfield_kernel_builder( vec_12 = wp.cross(min_normal, pos1 - pos2) id3 = int(-1) - dist3 = float(-wp.inf) + dist3 = float(-MJ_MAXVAL) for i in range(count): if i == min_id or i == id1 or i == id2: continue @@ -704,6 +694,7 @@ def ccd_kernel_builder( gjk_iterations: int, epa_iterations: int, use_multiccd: bool, + geomgeomid: int, ): """Kernel builder for non-heightfield CCD collisions (no hfield args).""" @@ -711,14 +702,11 @@ def ccd_kernel_builder( def eval_ccd_write_contact( # Model: opt_ccd_tolerance: wp.array(dtype=float), - geom_type: wp.array(dtype=int), # Data in: naconmax_in: int, # In: - epa_vert1_in: wp.array2d(dtype=wp.vec3), - epa_vert2_in: wp.array2d(dtype=wp.vec3), - epa_vert_index1_in: wp.array2d(dtype=int), - epa_vert_index2_in: wp.array2d(dtype=int), + epa_vert_in: wp.array2d(dtype=wp.vec3), + epa_vert_index_in: wp.array2d(dtype=int), epa_face_in: wp.array2d(dtype=int), epa_pr_in: wp.array2d(dtype=wp.vec3), epa_norm2_in: wp.array2d(dtype=float), @@ -738,7 +726,7 @@ def ccd_kernel_builder( geom2: Geom, geoms: wp.vec2i, worldid: int, - tid: int, + ccdid: int, margin: float, gap: float, condim: int, @@ -748,7 +736,6 @@ def ccd_kernel_builder( solimp: vec5, x1: wp.vec3, x2: wp.vec3, - count: int, pairid: wp.vec2i, # Data out: contact_dist_out: wp.array(dtype=float), @@ -771,12 +758,12 @@ def ccd_kernel_builder( witness2 = mat43() geom1.margin = margin geom2.margin = margin - if pairid[1] >= 0: - # if collision sensor, set large cutoff to work with various sensor cutoff values + is_collision_sensor = pairid[1] >= 0 + if is_collision_sensor: cutoff = 1.0e32 else: cutoff = 0.0 - dist, ncontact, w1, w2, idx = ccd( + dist, ncollision, w1, w2, multiccd_idx = ccd( opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]], cutoff, gjk_iterations, @@ -787,14 +774,12 @@ def ccd_kernel_builder( geomtype2, x1, x2, - epa_vert1_in[tid], - epa_vert2_in[tid], - epa_vert_index1_in[tid], - epa_vert_index2_in[tid], - epa_face_in[tid], - epa_pr_in[tid], - epa_norm2_in[tid], - epa_horizon_in[tid], + epa_vert_in[ccdid], + epa_vert_index_in[ccdid], + epa_face_in[ccdid], + epa_pr_in[ccdid], + epa_norm2_in[ccdid], + epa_horizon_in[ccdid], ) if dist >= 0.0 and pairid[1] == -1: @@ -803,30 +788,33 @@ def ccd_kernel_builder( witness1[0] = w1 witness2[0] = w2 - if wp.static(use_multiccd): - if ( - geom1.margin == 0.0 - and geom2.margin == 0.0 - and (geomtype1 == GeomType.BOX or (geomtype1 == GeomType.MESH and geom1.mesh_polyadr > -1)) - and (geomtype2 == GeomType.BOX or (geomtype2 == GeomType.MESH and geom2.mesh_polyadr > -1)) - ): - ncontact, witness1, witness2 = multicontact( - multiccd_polygon_in[tid], - multiccd_clipped_in[tid], - multiccd_pnormal_in[tid], - multiccd_pdist_in[tid], - multiccd_idx1_in[tid], - multiccd_idx2_in[tid], - multiccd_n1_in[tid], - multiccd_n2_in[tid], - multiccd_endvert_in[tid], - multiccd_face1_in[tid], - multiccd_face2_in[tid], - epa_vert1_in[tid], - epa_vert2_in[tid], - epa_vert_index1_in[tid], - epa_vert_index2_in[tid], - epa_face_in[tid, idx], + if wp.static(use_multiccd or (geomtype1 == GeomType.BOX and geomtype2 == GeomType.BOX)): + if wp.static(geomtype1 == GeomType.MESH): + # verify that geom1 mesh data is present for multicontact + if geom1.mesh_polyadr < 0: + multiccd_idx = -1 + + if wp.static(geomtype2 == GeomType.MESH): + # verify that geom2 mesh data is present for multicontact + if geom2.mesh_polyadr < 0: + multiccd_idx = -1 + + if multiccd_idx > -1: + ncollision, witness1, witness2 = multicontact( + multiccd_polygon_in[ccdid], + multiccd_clipped_in[ccdid], + multiccd_pnormal_in[ccdid], + multiccd_pdist_in[ccdid], + multiccd_idx1_in[ccdid], + multiccd_idx2_in[ccdid], + multiccd_n1_in[ccdid], + multiccd_n2_in[ccdid], + multiccd_endvert_in[ccdid], + multiccd_face1_in[ccdid], + multiccd_face2_in[ccdid], + epa_vert_in[ccdid], + epa_vert_index_in[ccdid], + epa_face_in[ccdid, multiccd_idx], w1, w2, geom1, @@ -835,7 +823,7 @@ def ccd_kernel_builder( geomtype2, ) - for i in range(ncontact): + for i in range(ncollision): points[i] = 0.5 * (witness1[i] + witness2[i]) normal = witness1[0] - witness2[0] frame = make_frame(normal) @@ -845,8 +833,9 @@ def ccd_kernel_builder( frame *= -1.0 geoms = wp.vec2i(geoms[1], geoms[0]) - for i in range(ncontact): - write_contact( + nactive = int(0) # number of contacts contributing to the physics + for i in range(ncollision): + active = write_contact( naconmax_in, i, dist, @@ -877,13 +866,12 @@ def ccd_kernel_builder( contact_geomcollisionid_out, nacon_out, ) - if count + (i + 1) >= MJ_MAXCONPAIR: - return i + 1 + nactive += active - return ncontact + return nactive # runs convex collision on a set of geom pairs to recover contact info (non-heightfield) - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def ccd_kernel( # Model: opt_ccd_tolerance: wp.array(dtype=float), @@ -900,8 +888,8 @@ def ccd_kernel_builder( geom_gap: wp.array2d(dtype=float), mesh_vertadr: wp.array(dtype=int), mesh_vertnum: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), mesh_graphadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), mesh_graph: wp.array(dtype=int), mesh_polynum: wp.array(dtype=int), mesh_polyadr: wp.array(dtype=int), @@ -920,18 +908,17 @@ def ccd_kernel_builder( pair_gap: wp.array2d(dtype=float), pair_friction: wp.array2d(dtype=vec5), # Data in: - naconmax_in: int, geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), + naconmax_in: int, + naccdmax_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), - ncollision_in: wp.array(dtype=int), - # In: - epa_vert1_in: wp.array2d(dtype=wp.vec3), - epa_vert2_in: wp.array2d(dtype=wp.vec3), - epa_vert_index1_in: wp.array2d(dtype=int), - epa_vert_index2_in: wp.array2d(dtype=int), + epa_vert_in: wp.array2d(dtype=wp.vec3), + epa_vert_index_in: wp.array2d(dtype=int), epa_face_in: wp.array2d(dtype=int), epa_pr_in: wp.array2d(dtype=wp.vec3), epa_norm2_in: wp.array2d(dtype=float), @@ -947,8 +934,8 @@ def ccd_kernel_builder( multiccd_endvert_in: wp.array2d(dtype=wp.vec3), multiccd_face1_in: wp.array2d(dtype=wp.vec3), multiccd_face2_in: wp.array2d(dtype=wp.vec3), + nccd_in: wp.array(dtype=int), # Data out: - nacon_out: wp.array(dtype=int), contact_dist_out: wp.array(dtype=float), contact_pos_out: wp.array(dtype=wp.vec3), contact_frame_out: wp.array(dtype=wp.mat33), @@ -962,19 +949,25 @@ def ccd_kernel_builder( 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), ): - tid = wp.tid() - if tid >= ncollision_in[0]: + collisionid = wp.tid() + if collisionid >= ncollision_in[0]: return - geoms = collision_pair_in[tid] + geoms = collision_pair_in[collisionid] g1 = geoms[0] g2 = geoms[1] if geom_type[g1] != geomtype1 or geom_type[g2] != geomtype2: return - worldid = collision_worldid_in[tid] + ccdid = wp.atomic_add(nccd_in, wp.static(geomgeomid), 1) + if ccdid >= naccdmax_in: + wp.printf("CCD overflow - please increase naccdmax to %u\n", ccdid) + return + + worldid = collision_worldid_in[collisionid] _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params( geom_condim, @@ -994,7 +987,7 @@ def ccd_kernel_builder( pair_friction, collision_pair_in, collision_pairid_in, - tid, + collisionid, worldid, ) @@ -1024,12 +1017,9 @@ def ccd_kernel_builder( eval_ccd_write_contact( opt_ccd_tolerance, - geom_type, naconmax_in, - epa_vert1_in, - epa_vert2_in, - epa_vert_index1_in, - epa_vert_index2_in, + epa_vert_in, + epa_vert_index_in, epa_face_in, epa_pr_in, epa_norm2_in, @@ -1049,7 +1039,7 @@ def ccd_kernel_builder( geom2, geoms, worldid, - tid, + ccdid, margin, gap, condim, @@ -1059,8 +1049,7 @@ def ccd_kernel_builder( solimp, geom1.pos, geom2.pos, - 0, - collision_pairid_in[tid], + collision_pairid_in[collisionid], contact_dist_out, contact_pos_out, contact_frame_out, @@ -1080,37 +1069,8 @@ def ccd_kernel_builder( return ccd_kernel -# Heightfield collision pairs handled by ccd_hfield_kernel_builder -_HFIELD_COLLISION_PAIRS = [ - (GeomType.HFIELD, GeomType.SPHERE), - (GeomType.HFIELD, GeomType.CAPSULE), - (GeomType.HFIELD, GeomType.ELLIPSOID), - (GeomType.HFIELD, GeomType.CYLINDER), - (GeomType.HFIELD, GeomType.BOX), - (GeomType.HFIELD, GeomType.MESH), -] - -# Non-heightfield collision pairs handled by ccd_kernel_builder -_NON_HFIELD_COLLISION_PAIRS = [ - (GeomType.SPHERE, GeomType.ELLIPSOID), - (GeomType.SPHERE, GeomType.MESH), - (GeomType.CAPSULE, GeomType.ELLIPSOID), - (GeomType.CAPSULE, GeomType.CYLINDER), - (GeomType.CAPSULE, GeomType.MESH), - (GeomType.ELLIPSOID, GeomType.ELLIPSOID), - (GeomType.ELLIPSOID, GeomType.CYLINDER), - (GeomType.ELLIPSOID, GeomType.BOX), - (GeomType.ELLIPSOID, GeomType.MESH), - (GeomType.CYLINDER, GeomType.CYLINDER), - (GeomType.CYLINDER, GeomType.BOX), - (GeomType.CYLINDER, GeomType.MESH), - (GeomType.BOX, GeomType.MESH), - (GeomType.MESH, GeomType.MESH), -] - - @event_scope -def convex_narrowphase(m: Model, d: Data): +def convex_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table: list[tuple[GeomType, GeomType]]): """Runs narrowphase collision detection for convex geom pairs. This function handles collision detection for pairs of convex geometries that were @@ -1126,40 +1086,56 @@ def convex_narrowphase(m: Model, d: Data): computations for non-existent pair types. """ - def _pair_count(p1: int, p2: int) -> int: - return m.geom_pair_type_count[upper_trid_index(len(GeomType), p1, p2)] + def _pair_count(p1: int, p2: int) -> Tuple[int, int]: + idx = upper_trid_index(len(GeomType), p1, p2) + return m.geom_pair_type_count[idx], idx + ncollision = sum(_pair_count(g[0].value, g[1].value)[0] for g in collision_table) # no convex collisions, early return - if not any(_pair_count(g[0].value, g[1].value) for g in _CONVEX_COLLISION_PAIRS): + + if ncollision == 0: return - epa_iterations = m.opt.ccd_iterations + # compute nmaxpolygon and nmaxmeshdeg given the geom pairs for the model + nboxbox, _ = _pair_count(GeomType.BOX.value, GeomType.BOX.value) + if (GeomType.BOX, GeomType.BOX) not in collision_table: + nboxbox = 0 + nboxmesh, _ = _pair_count(GeomType.BOX.value, GeomType.MESH.value) + nmeshmesh, _ = _pair_count(GeomType.MESH.value, GeomType.MESH.value) + + epa_iterations = 16 if nboxbox == ncollision else m.opt.ccd_iterations # set to true to enable multiccd use_multiccd = m.opt.enableflags & EnableBit.MULTICCD - nmaxpolygon = m.nmaxpolygon if use_multiccd else 0 - nmaxmeshdeg = m.nmaxmeshdeg if use_multiccd else 0 - # epa_vert1: vertices in EPA polytope in geom 1 space - epa_vert1 = wp.empty(shape=(d.naconmax, 5 + epa_iterations), dtype=wp.vec3) - # epa_vert2: vertices in EPA polytope in geom 2 space - epa_vert2 = wp.empty(shape=(d.naconmax, 5 + epa_iterations), dtype=wp.vec3) - # epa_vert_index1: vertex indices in EPA polytope for geom 1 - epa_vert_index1 = wp.empty(shape=(d.naconmax, 5 + epa_iterations), dtype=int) - # epa_vert_index2: vertex indices in EPA polytope for geom 2 (naconmax, 5 + CCDiter) - epa_vert_index2 = wp.empty(shape=(d.naconmax, 5 + epa_iterations), dtype=int) + # need at least 4 (square sides) if there's a box collision needing multiccd + nmaxpolygon = 4 if nboxbox > 0 else 0 + nmaxmeshdeg = 3 if nboxbox > 0 else 0 + + # need to allocate more memory if there's meshes + if use_multiccd and nmeshmesh + nboxmesh > 0: + minval = 4 if nboxmesh else nmaxpolygon + nmaxpolygon = max(m.nmaxpolygon, minval) + nmaxmeshdeg = max(m.nmaxmeshdeg, 3) + + # ccd collider count + nccd = wp.zeros(len(GeomType) * (len(GeomType) + 1) // 2, dtype=int) + + # epa_vert: vertices in EPA polytope + epa_vert = wp.empty(shape=(d.naccdmax, 10 + 2 * epa_iterations), dtype=wp.vec3) + # epa_vert_index: vertex indices in EPA polytope + epa_vert_index = wp.empty(shape=(d.naccdmax, 10 + 2 * epa_iterations), dtype=int) # epa_face: faces of polytope represented by three indices - epa_face = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) + epa_face = wp.empty(shape=(d.naccdmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=int) # epa_pr: projection of origin on polytope faces - epa_pr = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=wp.vec3) + epa_pr = wp.empty(shape=(d.naccdmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=wp.vec3) # epa_norm2: epa_pr * epa_pr - epa_norm2 = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=float) + epa_norm2 = wp.empty(shape=(d.naccdmax, 6 + MJ_MAX_EPAFACES * epa_iterations), dtype=float) # epa_horizon: index pair (i j) of edges on horizon - epa_horizon = wp.empty(shape=(d.naconmax, MJ_MAX_EPAHORIZON), dtype=int) + epa_horizon = wp.empty(shape=(d.naccdmax, MJ_MAX_EPAHORIZON), dtype=int) # Contact outputs contact_outputs = [ - d.nacon, d.contact.dist, d.contact.pos, d.contact.frame, @@ -1173,15 +1149,17 @@ def convex_narrowphase(m: Model, d: Data): d.contact.worldid, d.contact.type, d.contact.geomcollisionid, + d.nacon, ] # Launch heightfield collision kernels (no multiccd args, 72 args total) - for geom_pair in _HFIELD_COLLISION_PAIRS: + for geom_pair in collision_table: g1 = geom_pair[0].value g2 = geom_pair[1].value - if _pair_count(g1, g2): + count, geomgeomid = _pair_count(g1, g2) + if (g1 == GeomType.HFIELD or g2 == GeomType.HFIELD) and count: wp.launch( - ccd_hfield_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations), + ccd_hfield_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations, geomgeomid), dim=d.naconmax, inputs=[ m.opt.ccd_tolerance, @@ -1197,15 +1175,10 @@ def convex_narrowphase(m: Model, d: Data): m.geom_friction, m.geom_margin, m.geom_gap, - m.hfield_adr, - m.hfield_nrow, - m.hfield_ncol, - m.hfield_size, - m.hfield_data, m.mesh_vertadr, m.mesh_vertnum, - m.mesh_vert, m.mesh_graphadr, + m.mesh_vert, m.mesh_graph, m.mesh_polynum, m.mesh_polyadr, @@ -1216,6 +1189,11 @@ def convex_narrowphase(m: Model, d: Data): m.mesh_polymapadr, m.mesh_polymapnum, m.mesh_polymap, + m.hfield_size, + m.hfield_nrow, + m.hfield_ncol, + m.hfield_adr, + m.hfield_data, m.pair_dim, m.pair_solref, m.pair_solreffriction, @@ -1223,56 +1201,57 @@ def convex_narrowphase(m: Model, d: Data): m.pair_margin, m.pair_gap, m.pair_friction, - d.naconmax, d.geom_xpos, d.geom_xmat, - d.collision_pair, - d.collision_pairid, - d.collision_worldid, + d.naconmax, + d.naccdmax, d.ncollision, - epa_vert1, - epa_vert2, - epa_vert_index1, - epa_vert_index2, + ctx.collision_pair, + ctx.collision_pairid, + ctx.collision_worldid, + epa_vert, + epa_vert_index, epa_face, epa_pr, epa_norm2, epa_horizon, + nccd, ], outputs=contact_outputs, ) # Allocate multiccd arrays only for non-heightfield collisions # multiccd_polygon: clipped contact surface - multiccd_polygon = wp.empty(shape=(d.naconmax, 2 * nmaxpolygon), dtype=wp.vec3) + multiccd_polygon = wp.empty(shape=(d.naccdmax, 2 * nmaxpolygon), dtype=wp.vec3) # multiccd_clipped: clipped contact surface (intermediate) - multiccd_clipped = wp.empty(shape=(d.naconmax, 2 * nmaxpolygon), dtype=wp.vec3) + multiccd_clipped = wp.empty(shape=(d.naccdmax, 2 * nmaxpolygon), dtype=wp.vec3) # multiccd_pnormal: plane normal of clipping polygon - multiccd_pnormal = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + multiccd_pnormal = wp.empty(shape=(d.naccdmax, nmaxpolygon), dtype=wp.vec3) # multiccd_pdist: plane distance of clipping polygon - multiccd_pdist = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=float) + multiccd_pdist = wp.empty(shape=(d.naccdmax, nmaxpolygon), dtype=float) # multiccd_idx1: list of normal index candidates for Geom 1 - multiccd_idx1 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=int) + multiccd_idx1 = wp.empty(shape=(d.naccdmax, nmaxmeshdeg), dtype=int) # multiccd_idx2: list of normal index candidates for Geom 2 - multiccd_idx2 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=int) + multiccd_idx2 = wp.empty(shape=(d.naccdmax, nmaxmeshdeg), dtype=int) # multiccd_n1: list of normal candidates for Geom 1 - multiccd_n1 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + multiccd_n1 = wp.empty(shape=(d.naccdmax, nmaxmeshdeg), dtype=wp.vec3) # multiccd_n2: list of normal candidates for Geom 1 - multiccd_n2 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + multiccd_n2 = wp.empty(shape=(d.naccdmax, nmaxmeshdeg), dtype=wp.vec3) # multiccd_endvert: list of edge vertices candidates - multiccd_endvert = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + multiccd_endvert = wp.empty(shape=(d.naccdmax, nmaxmeshdeg), dtype=wp.vec3) # multiccd_face1: contact face - multiccd_face1 = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + multiccd_face1 = wp.empty(shape=(d.naccdmax, nmaxpolygon), dtype=wp.vec3) # multiccd_face2: contact face - multiccd_face2 = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + multiccd_face2 = wp.empty(shape=(d.naccdmax, nmaxpolygon), dtype=wp.vec3) # Launch non-heightfield collision kernels (no hfield args, 78 args total) - for geom_pair in _NON_HFIELD_COLLISION_PAIRS: + for geom_pair in collision_table: g1 = geom_pair[0].value g2 = geom_pair[1].value - if _pair_count(g1, g2): + count, geomgeomid = _pair_count(g1, g2) + if g1 != GeomType.HFIELD and g2 != GeomType.HFIELD and count: wp.launch( - ccd_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations, use_multiccd), + ccd_kernel_builder(g1, g2, m.opt.ccd_iterations, epa_iterations, use_multiccd, geomgeomid), dim=d.naconmax, inputs=[ m.opt.ccd_tolerance, @@ -1289,8 +1268,8 @@ def convex_narrowphase(m: Model, d: Data): m.geom_gap, m.mesh_vertadr, m.mesh_vertnum, - m.mesh_vert, m.mesh_graphadr, + m.mesh_vert, m.mesh_graph, m.mesh_polynum, m.mesh_polyadr, @@ -1308,17 +1287,16 @@ def convex_narrowphase(m: Model, d: Data): m.pair_margin, m.pair_gap, m.pair_friction, - d.naconmax, d.geom_xpos, d.geom_xmat, - d.collision_pair, - d.collision_pairid, - d.collision_worldid, + d.naconmax, + d.naccdmax, d.ncollision, - epa_vert1, - epa_vert2, - epa_vert_index1, - epa_vert_index2, + ctx.collision_pair, + ctx.collision_pairid, + ctx.collision_worldid, + epa_vert, + epa_vert_index, epa_face, epa_pr, epa_norm2, @@ -1334,6 +1312,7 @@ def convex_narrowphase(m: Model, d: Data): multiccd_endvert, multiccd_face1, multiccd_face2, + nccd, ], outputs=contact_outputs, ) 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 60df7d3a..97fc1305 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 @@ -24,17 +24,65 @@ 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 CollisionContext +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.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 nested_kernel wp.set_module_options({"enable_backward": False}) +# Corresponding table to MuJoCo's mjCOLLISIONFUNC table in engine_collision_driver.c +MJ_COLLISION_TABLE = { + (GeomType.PLANE, GeomType.SPHERE): CollisionType.PRIMITIVE, + (GeomType.PLANE, GeomType.CAPSULE): CollisionType.PRIMITIVE, + (GeomType.PLANE, GeomType.ELLIPSOID): CollisionType.PRIMITIVE, + (GeomType.PLANE, GeomType.CYLINDER): CollisionType.PRIMITIVE, + (GeomType.PLANE, GeomType.BOX): CollisionType.PRIMITIVE, + (GeomType.PLANE, GeomType.MESH): CollisionType.PRIMITIVE, + (GeomType.HFIELD, GeomType.SPHERE): CollisionType.CONVEX, + (GeomType.HFIELD, GeomType.CAPSULE): CollisionType.CONVEX, + (GeomType.HFIELD, GeomType.ELLIPSOID): CollisionType.CONVEX, + (GeomType.HFIELD, GeomType.CYLINDER): CollisionType.CONVEX, + (GeomType.HFIELD, GeomType.BOX): CollisionType.CONVEX, + (GeomType.HFIELD, GeomType.MESH): CollisionType.CONVEX, + (GeomType.SPHERE, GeomType.SPHERE): CollisionType.PRIMITIVE, + (GeomType.SPHERE, GeomType.CAPSULE): CollisionType.PRIMITIVE, + (GeomType.SPHERE, GeomType.ELLIPSOID): CollisionType.CONVEX, + (GeomType.SPHERE, GeomType.CYLINDER): CollisionType.PRIMITIVE, + (GeomType.SPHERE, GeomType.BOX): CollisionType.PRIMITIVE, + (GeomType.SPHERE, GeomType.MESH): CollisionType.CONVEX, + (GeomType.CAPSULE, GeomType.CAPSULE): CollisionType.PRIMITIVE, + (GeomType.CAPSULE, GeomType.ELLIPSOID): CollisionType.CONVEX, + (GeomType.CAPSULE, GeomType.CYLINDER): CollisionType.CONVEX, + (GeomType.CAPSULE, GeomType.BOX): CollisionType.PRIMITIVE, + (GeomType.CAPSULE, GeomType.MESH): CollisionType.CONVEX, + (GeomType.ELLIPSOID, GeomType.ELLIPSOID): CollisionType.CONVEX, + (GeomType.ELLIPSOID, GeomType.CYLINDER): CollisionType.CONVEX, + (GeomType.ELLIPSOID, GeomType.BOX): CollisionType.CONVEX, + (GeomType.ELLIPSOID, GeomType.MESH): CollisionType.CONVEX, + (GeomType.CYLINDER, GeomType.CYLINDER): CollisionType.CONVEX, + (GeomType.CYLINDER, GeomType.BOX): CollisionType.CONVEX, + (GeomType.CYLINDER, GeomType.MESH): CollisionType.CONVEX, + (GeomType.BOX, GeomType.BOX): CollisionType.CONVEX, # overwritten by NATIVECCD disable flag + (GeomType.BOX, GeomType.MESH): CollisionType.CONVEX, + (GeomType.MESH, GeomType.MESH): CollisionType.CONVEX, +} + + +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), + ) + @wp.kernel def _zero_nacon_ncollision( @@ -293,10 +341,11 @@ def _add_geom_pair( worldid: int, nxnid: int, # Data out: + ncollision_out: wp.array(dtype=int), + # Out: collision_pair_out: wp.array(dtype=wp.vec2i), collision_pairid_out: wp.array(dtype=wp.vec2i), collision_worldid_out: wp.array(dtype=int), - ncollision_out: wp.array(dtype=int), ): pairid = wp.atomic_add(ncollision_out, 0, 1) @@ -329,15 +378,15 @@ def _binary_search(values: wp.array(dtype=Any), value: Any, lower: int, upper: i def _sap_project(opt_broadphase: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def sap_project( # Model: ngeom: int, geom_rbound: wp.array2d(dtype=float), geom_margin: wp.array2d(dtype=float), # Data in: - nworld_in: int, geom_xpos_in: wp.array2d(dtype=wp.vec3), + nworld_in: int, # In: direction_in: wp.vec3, # Out: @@ -402,7 +451,7 @@ def _sap_range( @cache_kernel def _sap_broadphase(opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound: int, ngeom_margin: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( # Model: ngeom: int, @@ -412,19 +461,20 @@ def _sap_broadphase(opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound: i geom_margin: wp.array2d(dtype=float), nxn_pairid: wp.array(dtype=wp.vec2i), # Data in: - nworld_in: int, - naconmax_in: int, geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), + nworld_in: int, + naconmax_in: int, # In: sort_index_in: wp.array2d(dtype=int), cumulative_sum_in: wp.array(dtype=int), nsweep_in: int, # Data out: + ncollision_out: wp.array(dtype=int), + # Out: collision_pair_out: wp.array(dtype=wp.vec2i), collision_pairid_out: wp.array(dtype=wp.vec2i), collision_worldid_out: wp.array(dtype=int), - ncollision_out: wp.array(dtype=int), ): worldgeomid = wp.tid() @@ -472,17 +522,17 @@ def _sap_broadphase(opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound: i geom2, worldid, idx, + ncollision_out, collision_pair_out, collision_pairid_out, collision_worldid_out, - ncollision_out, ) return kernel def _segmented_sort(tile_size: int): - @wp.kernel + @wp.kernel(module="unique") def segmented_sort( # In: projection_lower_in: wp.array2d(dtype=float), @@ -508,7 +558,7 @@ def _segmented_sort(tile_size: int): @event_scope -def sap_broadphase(m: Model, d: Data): +def sap_broadphase(m: Model, d: Data, ctx: CollisionContext): """Runs broadphase collision detection using a sweep-and-prune (SAP) algorithm. This method is more efficient than the N-squared approach for large numbers of @@ -543,7 +593,7 @@ def sap_broadphase(m: Model, d: Data): wp.launch( kernel=_sap_project(m.opt.broadphase), dim=(d.nworld, m.ngeom), - inputs=[m.ngeom, m.geom_rbound, m.geom_margin, d.nworld, d.geom_xpos, direction], + inputs=[m.ngeom, m.geom_rbound, m.geom_margin, d.geom_xpos, d.nworld, direction], outputs=[ projection_lower.reshape((-1, m.ngeom)), projection_upper, @@ -588,21 +638,21 @@ def sap_broadphase(m: Model, d: Data): m.geom_rbound, m.geom_margin, m.nxn_pairid, - d.nworld, - d.naconmax, d.geom_xpos, d.geom_xmat, + d.nworld, + d.naconmax, sort_index.reshape((-1, m.ngeom)), cumulative_sum.reshape(-1), nsweep, ], - outputs=[d.collision_pair, d.collision_pairid, d.collision_worldid, d.ncollision], + outputs=[d.ncollision, ctx.collision_pair, ctx.collision_pairid, ctx.collision_worldid], ) @cache_kernel def _nxn_broadphase(opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound: int, ngeom_margin: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( # Model: geom_type: wp.array(dtype=int), @@ -612,14 +662,15 @@ def _nxn_broadphase(opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound: i nxn_geom_pair: wp.array(dtype=wp.vec2i), nxn_pairid: wp.array(dtype=wp.vec2i), # Data in: - naconmax_in: int, geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), + naconmax_in: int, # Data out: + ncollision_out: wp.array(dtype=int), + # Out: collision_pair_out: wp.array(dtype=wp.vec2i), collision_pairid_out: wp.array(dtype=wp.vec2i), collision_worldid_out: wp.array(dtype=int), - ncollision_out: wp.array(dtype=int), ): worldid, elementid = wp.tid() @@ -641,17 +692,17 @@ def _nxn_broadphase(opt_broadphase_filter: int, ngeom_aabb: int, ngeom_rbound: i geom2, worldid, elementid, + ncollision_out, collision_pair_out, collision_pairid_out, collision_worldid_out, - ncollision_out, ) return kernel @event_scope -def nxn_broadphase(m: Model, d: Data): +def nxn_broadphase(m: Model, d: Data, ctx: CollisionContext): """Runs broadphase collision detection using a brute-force N-squared approach. This function iterates through a pre-filtered list of all possible geometry pairs and @@ -674,27 +725,34 @@ def nxn_broadphase(m: Model, d: Data): m.geom_margin, m.nxn_geom_pair_filtered, m.nxn_pairid_filtered, - d.naconmax, d.geom_xpos, d.geom_xmat, + d.naconmax, ], outputs=[ - d.collision_pair, - d.collision_pairid, - d.collision_worldid, d.ncollision, + ctx.collision_pair, + ctx.collision_pairid, + ctx.collision_worldid, ], ) -def _narrowphase(m, d): +def _narrowphase(m: Model, d: Data, ctx: CollisionContext): + collision_table = MJ_COLLISION_TABLE + if m.opt.disableflags & DisableBit.NATIVECCD: + collision_table[(GeomType.BOX, GeomType.BOX)] = CollisionType.PRIMITIVE + + convex_pairs = [key for key, value in collision_table.items() if value == CollisionType.CONVEX] + primitive_pairs = [key for key, value in collision_table.items() if value == CollisionType.PRIMITIVE] + # TODO(team): we should reject far-away contacts in the narrowphase instead of constraint # partitioning because we can move some pressure of the atomics - convex_narrowphase(m, d) - primitive_narrowphase(m, d) + convex_narrowphase(m, d, ctx, convex_pairs) + primitive_narrowphase(m, d, ctx, primitive_pairs) if m.has_sdf_geom: - sdf_narrowphase(m, d) + sdf_narrowphase(m, d, ctx) @event_scope @@ -715,15 +773,18 @@ def collision(m: Model, d: Data): This function will do nothing except zero out arrays if collision detection is disabled via `m.opt.disableflags` or if `d.nacon` is 0. """ - # zero contact and collision counters - wp.launch(_zero_nacon_ncollision, dim=1, outputs=[d.nacon, d.ncollision]) - if d.naconmax == 0 or m.opt.disableflags & (DisableBit.CONSTRAINT | DisableBit.CONTACT): + d.nacon.zero_() return - if m.opt.broadphase == BroadphaseType.NXN: - nxn_broadphase(m, d) - else: - sap_broadphase(m, d) + ctx = create_collision_context(d.naconmax) - _narrowphase(m, d) + # zero counters + wp.launch(_zero_nacon_ncollision, dim=1, outputs=[d.nacon, d.ncollision]) + + if m.opt.broadphase == BroadphaseType.NXN: + nxn_broadphase(m, d, ctx) + else: + sap_broadphase(m, d, ctx) + + _narrowphase(m, d, ctx) 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 45dd006b..b44264cd 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py @@ -13,6 +13,7 @@ # limitations under the License. # ============================================================================== +import math from typing import Tuple import warp as wp @@ -30,9 +31,11 @@ FLOAT_MAX = 1e30 MINVAL = 1e-15 MIN_DIST = 1e-10 -# TODO(kbayes): write out formulas to derive these constants -FACE_TOL = 0.99999872 -EDGE_TOL = 0.00159999931 +FACE_TOL = wp.static(math.cos(0.0016)) +EDGE_TOL = wp.static(math.sin(0.0016)) + +# tolarance used by multicontact for intersecting a plane and a line segment +INTERSECT_TOL = 0.0000003 # Bit flags for face status in EPA polytope. # Defined at module scope to avoid Warp's intermediate type issues with literals. @@ -59,11 +62,9 @@ class GJKResult: class Polytope: status: int - # vertices in polytope - vert1: wp.array(dtype=wp.vec3) - vert2: wp.array(dtype=wp.vec3) - vert_index1: wp.array(dtype=int) - vert_index2: wp.array(dtype=int) + # vertices in polytope (packed geom1 followed by geom2) + vert: wp.array(dtype=wp.vec3) + vert_index: wp.array(dtype=int) nvert: int # faces in polytope @@ -107,9 +108,9 @@ def support(geom: Geom, geomtype: int, dir: wp.vec3) -> SupportPoint: tmp = wp.sign(local_dir) res = wp.cw_mul(tmp, geom.size) sp.point = geom.rot @ res + geom.pos - sp.vertex_index = wp.where(tmp[0] > 0, 1, 0) - sp.vertex_index += wp.where(tmp[1] > 0, 2, 0) - sp.vertex_index += wp.where(tmp[2] > 0, 4, 0) + sp.vertex_index = wp.where(tmp[0] > 0.0, 1, 0) + sp.vertex_index += wp.where(tmp[1] > 0.0, 2, 0) + sp.vertex_index += wp.where(tmp[2] > 0.0, 4, 0) elif geomtype == GeomType.CAPSULE: res = local_dir * geom.size[0] # add cylinder contribution @@ -177,7 +178,7 @@ def support(geom: Geom, geomtype: int, dir: wp.vec3) -> SupportPoint: elif geomtype == GeomType.HFIELD: max_dist = float(FLOAT_MIN) # TODO(kbayes): Support edge prisms - sp.vertex_index = wp.where(dir[2] < 0, -2, -3) + sp.vertex_index = wp.where(dir[2] < 0.0, -2, -3) for i in range(6): vert = geom.hfprism[i] dist = wp.dot(vert, dir) @@ -197,7 +198,10 @@ def _attach_face(pt: Polytope, idx: int, v1: int, v2: int, v3: int) -> float: return 0.0 # compute witness point v - r, ret = _project_origin_plane(pt.vert1[v3] - pt.vert2[v3], pt.vert1[v2] - pt.vert2[v2], pt.vert1[v1] - pt.vert2[v1]) + p1 = pt.vert[2 * v1] - pt.vert[2 * v1 + 1] + p2 = pt.vert[2 * v2] - pt.vert[2 * v2 + 1] + p3 = pt.vert[2 * v3] - pt.vert[2 * v3 + 1] + r, ret = _project_origin_plane(p3, p2, p1) if ret: return 0.0 @@ -214,13 +218,13 @@ def _epa_support( pt: Polytope, idx: int, geom1: Geom, geom2: Geom, geom1_type: int, geom2_type: int, dir: wp.vec3 ) -> Tuple[int, int]: sp = support(geom1, geom1_type, dir) - pt.vert1[idx] = sp.point - pt.vert_index1[idx] = sp.vertex_index + pt.vert[2 * idx] = sp.point + pt.vert_index[2 * idx] = sp.vertex_index index1 = sp.cached_index sp = support(geom2, geom2_type, -dir) - pt.vert2[idx] = sp.point - pt.vert_index2[idx] = sp.vertex_index + pt.vert[2 * idx + 1] = sp.point + pt.vert_index[2 * idx + 1] = sp.vertex_index index2 = sp.cached_index return index1, index2 @@ -265,9 +269,9 @@ def _det3(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3) -> float: @wp.func def _same_sign(a: float, b: float) -> int: - if a > 0 and b > 0: + if a > 0.0 and b > 0.0: return 1 - if a < 0 and b < 0: + if a < 0.0 and b < 0.0: return -1 return 0 @@ -290,28 +294,25 @@ def _project_origin_plane(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3) -> Tuple[wp.vec n = wp.cross(diff32, diff21) nv = wp.dot(n, v2) nn = wp.dot(n, n) - if nn == 0: + if nn == 0.0: return z, 1 - if nv != 0 and nn > MINVAL: - v = (nv / nn) * n - return v, 0 + if nv != 0.0 and nn > MINVAL: + return (nv / nn) * n, 0 # n = (v2 - v1) x (v3 - v1) n = wp.cross(diff21, diff31) nv = wp.dot(n, v1) nn = wp.dot(n, n) - if nn == 0: + if nn == 0.0: return z, 1 - if nv != 0 and nn > MINVAL: - v = (nv / nn) * n - return v, 0 + if nv != 0.0 and nn > MINVAL: + return (nv / nn) * n, 0 # n = (v1 - v3) x (v2 - v3) n = wp.cross(diff31, diff32) nv = wp.dot(n, v3) nn = wp.dot(n, n) - v = (nv / nn) * n - return v, 0 + return (nv / nn) * n, 0 @wp.func @@ -610,14 +611,14 @@ def gjk( simplex[n] = simplex1[n] - simplex2[n] if cutoff == 0.0: - if wp.dot(x_k, simplex[n]) > 0: + if wp.dot(x_k, simplex[n]) > 0.0: result = GJKResult() result.dim = 0 result.dist = FLOAT_MAX return result elif cutoff < FLOAT_MAX: vs = wp.dot(x_k, simplex[n]) - if wp.dot(x_k, simplex[n]) > 0 and (vs * vs / xnorm) >= cutoff2: + if wp.dot(x_k, simplex[n]) > 0.0 and (vs * vs / xnorm) >= cutoff2: result = GJKResult() result.dim = 0 result.dist = FLOAT_MAX @@ -635,7 +636,7 @@ def gjk( # remove vertices from the simplex no longer needed n = int(0) for i in range(4): - if coordinates[i] == 0: + if coordinates[i] == 0.0: continue simplex[n] = simplex[i] @@ -692,7 +693,7 @@ def _same_side(p0: wp.vec3, p1: wp.vec3, p2: wp.vec3, p3: wp.vec3) -> bool: n = wp.cross(p1 - p0, p2 - p0) dot1 = wp.dot(n, p3 - p0) dot2 = wp.dot(n, -p0) - return (dot1 > 0 and dot2 > 0) or (dot1 < 0 and dot2 < 0) + return (dot1 > 0.0 and dot2 > 0.0) or (dot1 < 0.0 and dot2 < 0.0) @wp.func @@ -750,7 +751,7 @@ def _tri_point_intersect(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, p: wp.vec3) -> b l2 = coordinates[1] l3 = coordinates[2] - if l1 < 0 or l2 < 0 or l3 < 0: + if l1 < 0.0 or l2 < 0.0 or l3 < 0.0: return False pr = wp.vec3() @@ -766,14 +767,14 @@ def _replace_simplex3(pt: Polytope, v1: int, v2: int, v3: int) -> GJKResult: # reset GJK simplex simplex1 = mat43() - simplex1[0] = pt.vert1[v1] - simplex1[1] = pt.vert1[v2] - simplex1[2] = pt.vert1[v3] + simplex1[0] = pt.vert[2 * v1] + simplex1[1] = pt.vert[2 * v2] + simplex1[2] = pt.vert[2 * v3] simplex2 = mat43() - simplex2[0] = pt.vert2[v1] - simplex2[1] = pt.vert2[v2] - simplex2[2] = pt.vert2[v3] + simplex2[0] = pt.vert[2 * v1 + 1] + simplex2[1] = pt.vert[2 * v2 + 1] + simplex2[2] = pt.vert[2 * v3 + 1] simplex = mat43() simplex[0] = simplex1[0] - simplex2[0] @@ -781,14 +782,14 @@ def _replace_simplex3(pt: Polytope, v1: int, v2: int, v3: int) -> GJKResult: simplex[2] = simplex1[2] - simplex2[2] simplex_index1 = wp.vec4i() - simplex_index1[0] = pt.vert_index1[v1] - simplex_index1[1] = pt.vert_index1[v2] - simplex_index1[2] = pt.vert_index1[v3] + simplex_index1[0] = pt.vert_index[2 * v1] + simplex_index1[1] = pt.vert_index[2 * v2] + simplex_index1[2] = pt.vert_index[2 * v3] simplex_index2 = wp.vec4i() - simplex_index2[0] = pt.vert_index2[v1] - simplex_index2[1] = pt.vert_index2[v2] - simplex_index2[2] = pt.vert_index2[v3] + simplex_index2[0] = pt.vert_index[2 * v1 + 1] + simplex_index2[1] = pt.vert_index[2 * v2 + 1] + simplex_index2[2] = pt.vert_index[2 * v3 + 1] result.simplex = simplex result.simplex1 = simplex1 @@ -827,9 +828,9 @@ def _ray_triangle(v1: wp.vec3, v2: wp.vec3, v3: wp.vec3, v4: wp.vec3, v5: wp.vec vol2 = _det3(v4 - v1, v5 - v1, v2 - v1) vol3 = _det3(v5 - v1, v3 - v1, v2 - v1) - if vol1 >= 0 and vol2 >= 0 and vol3 >= 0: + if vol1 >= 0.0 and vol2 >= 0.0 and vol3 >= 0.0: return 1 - if vol1 <= 0 and vol2 <= 0 and vol3 <= 0: + if vol1 <= 0.0 and vol2 <= 0.0 and vol3 <= 0.0: return -1 return 0 @@ -867,9 +868,9 @@ def _epa_witness( ) -> Tuple[wp.vec3, wp.vec3, float]: face = _get_face_verts(pt.face[face_idx]) # compute affine coordinates for witness points on plane defined by face - v1 = pt.vert1[face[0]] - pt.vert2[face[0]] - v2 = pt.vert1[face[1]] - pt.vert2[face[1]] - v3 = pt.vert1[face[2]] - pt.vert2[face[2]] + v1 = pt.vert[2 * face[0]] - pt.vert[2 * face[0] + 1] + v2 = pt.vert[2 * face[1]] - pt.vert[2 * face[1] + 1] + v3 = pt.vert[2 * face[2]] - pt.vert[2 * face[2] + 1] coordinates = _tri_affine_coord(v1, v2, v3, pt.face_pr[face_idx]) l1 = coordinates[0] @@ -877,18 +878,18 @@ def _epa_witness( l3 = coordinates[2] # face on geom 2 - v1 = pt.vert2[face[0]] - v2 = pt.vert2[face[1]] - v3 = pt.vert2[face[2]] + v1 = pt.vert[2 * face[0] + 1] + v2 = pt.vert[2 * face[1] + 1] + v3 = pt.vert[2 * face[2] + 1] x2 = wp.vec3() x2[0] = v1[0] * l1 + v2[0] * l2 + v3[0] * l3 x2[1] = v1[1] * l1 + v2[1] * l2 + v3[1] * l3 x2[2] = v1[2] * l1 + v2[2] * l2 + v3[2] * l3 # correct witness points for hfield geoms - i1 = pt.vert_index1[face[0]] - i2 = pt.vert_index1[face[1]] - i3 = pt.vert_index1[face[2]] + i1 = pt.vert_index[2 * face[0]] + i2 = pt.vert_index[2 * face[1]] + i3 = pt.vert_index[2 * face[2]] if geomtype1 == GeomType.HFIELD and (i1 != i2 or i1 != i3): # TODO(kbayes): Fix case where geom2 is near bottom of height field or "extreme" prism heights n = wp.vec3(0.0, 0.0, 1.0) @@ -914,7 +915,7 @@ def _epa_witness( x2 = sp.point coordinates2 = _tri_affine_coord(a, b, c, x2) - if coordinates2[0] > 0 and coordinates2[1] > 0 and coordinates2[2] > 0: + if coordinates2[0] > 0.0 and coordinates2[1] > 0.0 and coordinates2[2] > 0.0: x1 = coordinates2[0] * a + coordinates2[1] * b + coordinates2[2] * c else: p = c @@ -924,9 +925,9 @@ def _epa_witness( return x1, x2, -wp.norm_l2(x1 - x2) # face on geom 1 - v1 = pt.vert1[face[0]] - v2 = pt.vert1[face[1]] - v3 = pt.vert1[face[2]] + v1 = pt.vert[2 * face[0]] + v2 = pt.vert[2 * face[1]] + v3 = pt.vert[2 * face[2]] x1 = wp.vec3() x1[0] = v1[0] * l1 + v2[0] * l2 + v3[0] * l3 x1[1] = v1[1] * l1 + v2[1] * l2 + v3[1] * l3 @@ -971,17 +972,15 @@ def _polytope2( d3 = R @ d2 # save vertices and get indices for each one - pt.vert1[0] = simplex1[0] - pt.vert1[1] = simplex1[1] + pt.vert[0] = simplex1[0] + pt.vert[1] = simplex2[0] + pt.vert[2] = simplex1[1] + pt.vert[3] = simplex2[1] - pt.vert_index1[0] = simplex_index1[0] - pt.vert_index1[1] = simplex_index1[1] - - pt.vert2[0] = simplex2[0] - pt.vert2[1] = simplex2[1] - - pt.vert_index2[0] = simplex_index2[0] - pt.vert_index2[1] = simplex_index2[1] + pt.vert_index[0] = simplex_index1[0] + pt.vert_index[1] = simplex_index2[0] + pt.vert_index[2] = simplex_index1[1] + pt.vert_index[3] = simplex_index2[1] _epa_support(pt, 2, geom1, geom2, geomtype1, geomtype2, d1 / wp.norm_l2(d1)) _epa_support(pt, 3, geom1, geom2, geomtype1, geomtype2, d2 / wp.norm_l2(d2)) @@ -1013,9 +1012,9 @@ def _polytope2( return pt, _replace_simplex3(pt, 1, 4, 3) # check hexahedron is convex - v2 = pt.vert1[2] - pt.vert2[2] - v3 = pt.vert1[3] - pt.vert2[3] - v4 = pt.vert1[4] - pt.vert2[4] + v2 = pt.vert[4] - pt.vert[5] + v3 = pt.vert[6] - pt.vert[7] + v4 = pt.vert[8] - pt.vert[9] if not _ray_triangle(simplex[0], simplex[1], v2, v3, v4): pt.status = 1 return pt, GJKResult() @@ -1049,21 +1048,19 @@ def _polytope3( pt.status = 2 return pt - pt.vert1[0] = simplex1[0] - pt.vert1[1] = simplex1[1] - pt.vert1[2] = simplex1[2] + pt.vert[0] = simplex1[0] + pt.vert[1] = simplex2[0] + pt.vert[2] = simplex1[1] + pt.vert[3] = simplex2[1] + pt.vert[4] = simplex1[2] + pt.vert[5] = simplex2[2] - pt.vert_index1[0] = simplex_index1[0] - pt.vert_index1[1] = simplex_index1[1] - pt.vert_index1[2] = simplex_index1[2] - - pt.vert2[0] = simplex2[0] - pt.vert2[1] = simplex2[1] - pt.vert2[2] = simplex2[2] - - pt.vert_index2[0] = simplex_index2[0] - pt.vert_index2[1] = simplex_index2[1] - pt.vert_index2[2] = simplex_index2[2] + pt.vert_index[0] = simplex_index1[0] + pt.vert_index[1] = simplex_index2[0] + pt.vert_index[2] = simplex_index1[1] + pt.vert_index[3] = simplex_index2[1] + pt.vert_index[4] = simplex_index1[2] + pt.vert_index[5] = simplex_index2[2] _epa_support(pt, 3, geom1, geom2, geomtype1, geomtype2, -n) _epa_support(pt, 4, geom1, geom2, geomtype1, geomtype2, n) @@ -1071,8 +1068,8 @@ def _polytope3( v1 = simplex[0] v2 = simplex[1] v3 = simplex[2] - v4 = pt.vert1[3] - pt.vert2[3] - v5 = pt.vert1[4] - pt.vert2[4] + v4 = pt.vert[6] - pt.vert[7] + v5 = pt.vert[8] - pt.vert[9] # check that v4 is not contained in the 2-simplex if _tri_point_intersect(v1, v2, v3, v4): @@ -1128,25 +1125,23 @@ def _polytope4( simplex_index2: wp.vec4i, ) -> Tuple[Polytope, GJKResult]: """Create polytope for EPA given a 3-simplex from GJK.""" - pt.vert1[0] = simplex1[0] - pt.vert1[1] = simplex1[1] - pt.vert1[2] = simplex1[2] - pt.vert1[3] = simplex1[3] + pt.vert[0] = simplex1[0] + pt.vert[1] = simplex2[0] + pt.vert[2] = simplex1[1] + pt.vert[3] = simplex2[1] + pt.vert[4] = simplex1[2] + pt.vert[5] = simplex2[2] + pt.vert[6] = simplex1[3] + pt.vert[7] = simplex2[3] - pt.vert_index1[0] = simplex_index1[0] - pt.vert_index1[1] = simplex_index1[1] - pt.vert_index1[2] = simplex_index1[2] - pt.vert_index1[3] = simplex_index1[3] - - pt.vert2[0] = simplex2[0] - pt.vert2[1] = simplex2[1] - pt.vert2[2] = simplex2[2] - pt.vert2[3] = simplex2[3] - - pt.vert_index2[0] = simplex_index2[0] - pt.vert_index2[1] = simplex_index2[1] - pt.vert_index2[2] = simplex_index2[2] - pt.vert_index2[3] = simplex_index2[3] + pt.vert_index[0] = simplex_index1[0] + pt.vert_index[1] = simplex_index2[0] + pt.vert_index[2] = simplex_index1[1] + pt.vert_index[3] = simplex_index2[1] + pt.vert_index[4] = simplex_index1[2] + pt.vert_index[5] = simplex_index2[2] + pt.vert_index[6] = simplex_index1[3] + pt.vert_index[7] = simplex_index2[3] # if the origin is on a face, replace the 3-simplex with a 2-simplex if _attach_face(pt, 0, 0, 1, 2) < MIN_DIST: @@ -1249,7 +1244,7 @@ def _epa( break # check if lower bound is 0 - if lower2 <= 0: + if lower2 <= 0.0: break # compute support point w from the closest face's normal @@ -1257,7 +1252,7 @@ def _epa( wi = pt.nvert face_pr_normalized = pt.face_pr[idx] / lower i1, i2 = _epa_support(pt, wi, geom1, geom2, geomtype1, geomtype2, face_pr_normalized) - w = pt.vert1[wi] - pt.vert2[wi] + w = pt.vert[2 * wi] - pt.vert[2 * wi + 1] geom1.index = i1 geom2.index = i2 pt.nvert += 1 @@ -1275,7 +1270,7 @@ def _epa( if is_discrete: found_repeated = bool(False) for i in range(pt.nvert - 1): - if pt.vert_index1[i] == pt.vert_index1[wi] and pt.vert_index2[i] == pt.vert_index2[wi]: + if pt.vert_index[2 * i] == pt.vert_index[2 * wi] and pt.vert_index[2 * i + 1] == pt.vert_index[2 * wi + 1]: found_repeated = True break if found_repeated: @@ -1313,7 +1308,7 @@ def _epa( for i in range(pt.nhorizon): edge = _get_edge(pt.horizon[i]) dist2 = _attach_face(pt, pt.nface, wi, edge[0], edge[1]) - if dist2 == 0: + if dist2 == 0.0: idx = -1 break @@ -1348,72 +1343,66 @@ def _area4(a: wp.vec3, b: wp.vec3, c: wp.vec3, d: wp.vec3) -> float: return 0.5 * wp.norm_l2(wp.cross(a - d, d - b) + wp.cross(b - c, c - a)) -@wp.func -def _next(n: int, i: int) -> int: - """Returns (i + 1) mod n for 0 <= i <= n - 1.""" - return wp.where(i == n - 1, 0, i + 1) - - @wp.func def _polygon_quad(polygon: wp.array(dtype=wp.vec3), npolygon: int) -> wp.vec4i: - """Returns the indices of a quadrilateral of maximum area in a convex polygon.""" - b = _next(npolygon, 0) - c = _next(npolygon, b) - d = _next(npolygon, c) + """Returns the indices of a quadrilateral of maximum area in a convex polygon (npolygon > 4).""" + b = int(1) + c = int(2) + d = int(3) res = wp.vec4i(0, b, c, d) m = _area4(polygon[0], polygon[b], polygon[c], polygon[d]) for a in range(npolygon): while True: - m_next = _area4(polygon[a], polygon[b], polygon[c], polygon[_next(npolygon, d)]) + m_next = _area4(polygon[a], polygon[b], polygon[c], polygon[(d + 1) % npolygon]) if m_next <= m: break m = m_next - d = _next(npolygon, d) + d = (d + 1) % npolygon res = wp.vec4i(a, b, c, d) while True: - m_next = _area4(polygon[a], polygon[b], polygon[_next(npolygon, c)], polygon[d]) + m_next = _area4(polygon[a], polygon[b], polygon[(c + 1) % npolygon], polygon[d]) if m_next <= m: break m = m_next - c = _next(npolygon, c) + c = (c + 1) % npolygon res = wp.vec4i(a, b, c, d) while True: - m_next = _area4(polygon[a], polygon[_next(npolygon, b)], polygon[c], polygon[d]) + m_next = _area4(polygon[a], polygon[(b + 1) % npolygon], polygon[c], polygon[d]) if m_next <= m: break m = m_next - b = _next(npolygon, b) + b = (b + 1) % npolygon res = wp.vec4i(a, b, c, d) if b == a: - b = _next(npolygon, b) + b = (b + 1) % npolygon if c == b: - c = _next(npolygon, c) + c = (c + 1) % npolygon if d == c: - d = _next(npolygon, d) + d = (d + 1) % npolygon return res # return number (1, 2 or 3) of dimensions of a simplex; reorder vertices if necessary @wp.func def _feature_dim( - face: wp.vec3i, vert_index: wp.array(dtype=int), vert: wp.array(dtype=wp.vec3) + face: wp.vec3i, vert_index: wp.array(dtype=int), vert: wp.array(dtype=wp.vec3), offset: int ) -> Tuple[int, wp.vec3i, wp.mat33]: - v1i = vert_index[face[0]] - v2i = vert_index[face[1]] - v3i = vert_index[face[2]] + v1i = vert_index[2 * face[0] + offset] + v2i = vert_index[2 * face[1] + offset] + v3i = vert_index[2 * face[2] + offset] feature_index = wp.vec3i(v1i, v2i, v3i) feature_vert = wp.mat33() - feature_vert[0] = vert[face[0]] - feature_vert[1] = vert[face[1]] - feature_vert[2] = vert[face[2]] + feature_vert[0] = vert[2 * face[0] + offset] + feature_vert[1] = vert[2 * face[1] + offset] + feature_vert[2] = vert[2 * face[2] + offset] if v1i != v2i: dim = wp.where(v3i == v1i or v3i == v2i, 2, 3) return dim, feature_index, feature_vert feature_index[1] = v3i - feature_vert[1] = vert[face[2]] + feature_vert[1] = vert[2 * face[2] + offset] dim = wp.where(v1i != v3i, 2, 1) return dim, feature_index, feature_vert @@ -1675,7 +1664,7 @@ def _box_normals( y = float((v1 & 2) and (v2 & 2)) - float(not (v1 & 2) and not (v2 & 2)) z = float((v1 & 4) and (v2 & 4)) - float(not (v1 & 4) and not (v2 & 4)) if x != 0.0: - normal_out[c] = mat @ wp.vec3(float(x), 0.0, 0.0) + normal_out[c] = mat @ wp.vec3(x, 0.0, 0.0) index_out[c] = wp.where(x > 0.0, 0, 1) c += 1 if y != 0.0: @@ -1686,7 +1675,9 @@ def _box_normals( normal_out[c] = mat @ wp.vec3(0.0, 0.0, z) index_out[c] = wp.where(z > 0.0, 4, 5) c += 1 - if c == 2: + # c is 1 if edge is diagonal of a box face + # c is 2 if edge is an external edge of box + if c == 1 or c == 2: return 2 return _box_normals2(mat, dir, normal_out, index_out) @@ -1824,18 +1815,15 @@ def _halfspace(a: wp.vec3, n: wp.vec3, p: wp.vec3) -> bool: @wp.func -def _plane_intersect(pn: wp.vec3, pd: float, a: wp.vec3, b: wp.vec3) -> Tuple[float, wp.vec3]: - res = wp.vec3() - ab = b - a - temp = wp.dot(pn, ab) - if temp == 0.0: - return FLOAT_MAX, res # parallel; no intersection - t = (pd - wp.dot(pn, a)) / temp - if t >= 0.0 and t <= 1.0: - res[0] = a[0] + t * ab[0] - res[1] = a[1] + t * ab[1] - res[2] = a[2] + t * ab[2] - return t, res +def _plane_intersect(pn: wp.vec3, pd: float, a: wp.vec3, b: wp.vec3) -> float: + """Returns the parameter t where the line a + t(b - a) intersects the given plane.""" + dot = wp.dot(pn, b - a) + + # parallel; no intersection + if wp.abs(dot) < 1e-10: + return FLOAT_MAX + + return (pd - wp.dot(pn, a)) / dot # clip a polygon against another polygon @@ -1901,9 +1889,10 @@ def _polygon_clip( continue # add new vertex to clipped polygon where PQ intersects the clipping edge - t, res = _plane_intersect(pn[e], pd[e], P, Q) - if t >= 0.0 and t <= 1.0: - clipped_out[nclipped] = res + t = _plane_intersect(pn[e], pd[e], P, Q) + if t > -INTERSECT_TOL and t < 1.0 + INTERSECT_TOL: + t = wp.clamp(t, 0.0, 1.0) + clipped_out[nclipped] = P + t * (Q - P) nclipped += 1 # add Q as PQ is now back inside the clipping edge @@ -1937,9 +1926,16 @@ def _polygon_clip( @wp.func def _set_edge( - vert1: wp.array(dtype=wp.vec3), vert2: wp.array(dtype=wp.vec3), start: int, end: int, face_out: wp.array(dtype=wp.vec3) + # In: + vert1: wp.array(dtype=wp.vec3), + vert2: wp.array(dtype=wp.vec3), + start: int, + end: int, + offset: int, + # Out: + face_out: wp.array(dtype=wp.vec3), ) -> int: - face_out[0] = vert1[start] + face_out[0] = vert1[2 * start + offset] face_out[1] = vert2[end] return 2 @@ -1959,10 +1955,8 @@ def multicontact( endvert: wp.array(dtype=wp.vec3), face1: wp.array(dtype=wp.vec3), face2: wp.array(dtype=wp.vec3), - epa_vert1: wp.array(dtype=wp.vec3), - epa_vert2: wp.array(dtype=wp.vec3), - epa_vert_index1: wp.array(dtype=int), - epa_vert_index2: wp.array(dtype=int), + epa_vert: wp.array(dtype=wp.vec3), + epa_vert_index: wp.array(dtype=int), epa_face: int, x1: wp.vec3, x2: wp.vec3, @@ -1998,8 +1992,8 @@ def multicontact( face = _get_face_verts(epa_face) # get dimensions of features of geoms 1 and 2 - nface1, feature_index1, feature_vertex1 = _feature_dim(face, epa_vert_index1, epa_vert1) - nface2, feature_index2, feature_vertex2 = _feature_dim(face, epa_vert_index2, epa_vert2) + nface1, feature_index1, feature_vertex1 = _feature_dim(face, epa_vert_index, epa_vert, 0) + nface2, feature_index2, feature_vertex2 = _feature_dim(face, epa_vert_index, epa_vert, 1) dir = x2 - x1 dir_neg = -dir @@ -2115,7 +2109,7 @@ def multicontact( # recover geom1 matching edge or face if is_edge_contact_geom1: - nface1 = _set_edge(epa_vert1, endvert, face[0], i, face1) + nface1 = _set_edge(epa_vert, endvert, face[0], i, 0, face1) else: ind = wp.where(is_edge_contact_geom2, idx1[j], idx1[i]) if geomtype1 == GeomType.BOX: @@ -2136,7 +2130,7 @@ def multicontact( # recover geom2 matching edge or face if is_edge_contact_geom2: - nface2 = _set_edge(epa_vert2, endvert, face[0], i, face2) + nface2 = _set_edge(epa_vert, endvert, face[0], i, 1, face2) else: if geomtype2 == GeomType.BOX: nface2 = _box_face(geom2.rot, geom2.pos, geom2.size, idx2[j], face2) @@ -2197,12 +2191,12 @@ def _inflate( c = geom1.hfprism[5] coordinates = _tri_affine_coord(a, b, c, x2) - if coordinates[0] > 0 and coordinates[1] > 0 and coordinates[2] > 0: + if coordinates[0] > 0.0 and coordinates[1] > 0.0 and coordinates[2] > 0.0: x1 = coordinates[0] * a + coordinates[1] * b + coordinates[2] * c else: p = c - p = wp.where(coordinates[1] > 0, b, p) - p = wp.where(coordinates[0] > 0, a, p) + p = wp.where(coordinates[1] > 0.0, b, p) + p = wp.where(coordinates[0] > 0.0, a, p) x1 = x2 - wp.dot(x2 - p, n) * n dist = -wp.norm_l2(x1 - x2) return dist, x1, x2 @@ -2230,10 +2224,8 @@ def ccd( geomtype2: int, x_1: wp.vec3, x_2: wp.vec3, - vert1: wp.array(dtype=wp.vec3), - vert2: wp.array(dtype=wp.vec3), - vert_index1: wp.array(dtype=int), - vert_index2: wp.array(dtype=int), + vert: wp.array(dtype=wp.vec3), + vert_index: wp.array(dtype=int), face: wp.array(dtype=int), face_pr: wp.array(dtype=wp.vec3), face_norm2: wp.array(dtype=float), @@ -2289,10 +2281,8 @@ def ccd( pt.nface = 0 pt.nvert = 0 pt.nhorizon = 0 - pt.vert1 = vert1 - pt.vert2 = vert2 - pt.vert_index1 = vert_index1 - pt.vert_index2 = vert_index2 + pt.vert = vert + pt.vert_index = vert_index pt.face = face pt.face_pr = face_pr pt.face_norm2 = face_norm2 @@ -2358,4 +2348,13 @@ def ccd( dist, x1, x2, idx = _epa(tolerance, gjk_iterations, epa_iterations, pt, geom1, geom2, geomtype1, geomtype2, is_discrete) if idx == -1: return FLOAT_MAX, 0, wp.vec3(), wp.vec3(), -1 + + # multicontact not supported for margin + if geom1.margin != 0.0 or geom2.margin != 0.0: + idx = -1 + + # multicontact only supported for boxes and meshes + if (geomtype1 != GeomType.BOX and geomtype1 != GeomType.MESH) or (geomtype2 != GeomType.BOX and geomtype2 != GeomType.MESH): + idx = -1 + return dist, 1, x1, x2, idx 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 ff6d80a1..54f05709 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 @@ -32,8 +32,10 @@ from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive_core import sph from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame from mujoco.mjx.third_party.mujoco_warp._src.math import safe_div from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL 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 CollisionContext from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType @@ -43,7 +45,6 @@ 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 -from mujoco.mjx.third_party.mujoco_warp._src.warp_util import nested_kernel wp.set_module_options({"enable_backward": False}) @@ -173,13 +174,13 @@ def plane_convex(plane_normal: wp.vec3, plane_pos: wp.vec3, convex: Geom) -> Tup convex: Convex geometry object containing position, rotation, and mesh data. Returns: - - Vector of contact distances (wp.inf for unpopulated contacts). + - Vector of contact distances (MJ_MAXVAL for unpopulated contacts). - Matrix of contact positions (one per row). - Matrix of contact normal vectors (one per row). """ _HUGE_VAL = 1e6 - contact_dist = wp.vec4(wp.inf) + contact_dist = wp.vec4(MJ_MAXVAL) contact_pos = mat43() contact_count = int(0) @@ -426,12 +427,12 @@ def write_contact( contact_type_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int), nacon_out: wp.array(dtype=int), -): +) -> int: active = dist_in < margin_in # skip contact and no collision sensor if (pairid_in[0] == -2 or not active) and pairid_in[1] == -1: - return + return 0 contact_type = 0 @@ -457,6 +458,8 @@ def write_contact( contact_solimp_out[cid] = solimp_in contact_type_out[cid] = contact_type contact_geomcollisionid_out[cid] = id_ + return int(active) + return 0 @wp.func @@ -477,10 +480,9 @@ def contact_params( pair_margin: wp.array2d(dtype=float), pair_gap: wp.array2d(dtype=float), pair_friction: wp.array2d(dtype=vec5), - # Data in: + # In: collision_pair_in: wp.array(dtype=wp.vec2i), collision_pairid_in: wp.array(dtype=wp.vec2i), - # In: cid: int, worldid: int, ): @@ -1540,6 +1542,7 @@ def box_box_wrapper( ) +# Map of supported primitive collision functions _PRIMITIVE_COLLISIONS = { (GeomType.PLANE, GeomType.SPHERE): plane_sphere_wrapper, (GeomType.PLANE, GeomType.CAPSULE): plane_capsule_wrapper, @@ -1557,23 +1560,9 @@ _PRIMITIVE_COLLISIONS = { } -# TODO(team): _check_collisions shared utility -def _check_primitive_collisions(): - prev_idx = -1 - for types in _PRIMITIVE_COLLISIONS.keys(): - idx = upper_trid_index(len(GeomType), types[0].value, types[1].value) - if types[1] < types[0] or idx <= prev_idx: - return False - prev_idx = idx - return True - - -assert _check_primitive_collisions(), "_PRIMITIVE_COLLISIONS is in invalid order" - - @cache_kernel def _primitive_narrowphase(primitive_collisions_types, primitive_collisions_func): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def primitive_narrowphase( # Model: geom_type: wp.array(dtype=int), @@ -1612,10 +1601,11 @@ def _primitive_narrowphase(primitive_collisions_types, primitive_collisions_func 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), - ncollision_in: wp.array(dtype=int), # Data out: contact_dist_out: wp.array(dtype=float), contact_pos_out: wp.array(dtype=wp.vec3), @@ -1730,7 +1720,7 @@ _PRIMITIVE_COLLISION_FUNC = [] @event_scope -def primitive_narrowphase(m: Model, d: Data): +def primitive_narrowphase(m: Model, d: Data, ctx: CollisionContext, collision_table: list[tuple[GeomType, GeomType]]): """Runs collision detection on primitive geom pairs discovered during broadphase. This function processes collision pairs involving primitive shapes that were @@ -1749,6 +1739,8 @@ def primitive_narrowphase(m: Model, d: Data): # for pair types without collisions, as well as updating the launch dimensions. for types, func in _PRIMITIVE_COLLISIONS.items(): + if types not in collision_table: + continue idx = upper_trid_index(len(GeomType), types[0].value, types[1].value) if m.geom_pair_type_count[idx] and types not in _PRIMITIVE_COLLISION_TYPES: _PRIMITIVE_COLLISION_TYPES.append(types) @@ -1793,10 +1785,10 @@ def primitive_narrowphase(m: Model, d: Data): d.geom_xpos, d.geom_xmat, d.naconmax, - d.collision_pair, - d.collision_pairid, - d.collision_worldid, d.ncollision, + ctx.collision_pair, + ctx.collision_pairid, + ctx.collision_worldid, ], outputs=[ d.contact.dist, 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 a54565da..306a301e 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 @@ -18,6 +18,7 @@ from typing import Any, Tuple import warp as wp MJ_MINVAL = 1e-15 +MJ_MAXVAL = 1e10 wp.set_module_options({"enable_backward": False}) @@ -411,13 +412,13 @@ def plane_box( margin: Collision tolerance. Returns: - - Vector of contact distances (wp.inf for unpopulated contacts). + - Vector of contact distances (MJ_MAXVAL for unpopulated contacts). - Matrix of contact positions (one per row). - Contact normal vector. """ center_dist = wp.dot(box_pos - plane_pos, plane_normal) - dist = vec8f(wp.inf) + dist = vec8f(MJ_MAXVAL) pos = mat83f() # test all corners, pick bottom 4 @@ -540,7 +541,7 @@ def plane_cylinder( - Matrix of contact normal vectors (one per row). """ # Initialize output matrices - contact_dist = wp.vec4(wp.inf) + contact_dist = wp.vec4(MJ_MAXVAL) contact_pos = mat43f() contact_count = 0 @@ -666,14 +667,14 @@ def box_box( margin: Collision tolerance. Returns: - - Vector of contact distances (wp.inf for unpopulated contacts). + - Vector of contact distances (MJ_MAXVAL for unpopulated contacts). - Matrix of contact positions (one per row). - Matrix of contact normal vectors (one per row). """ # Initialize output matrices contact_dist = vec8f() for i in range(8): - contact_dist[i] = wp.inf + contact_dist[i] = MJ_MAXVAL contact_pos = mat83f() contact_normals = mat83f() contact_count = 0 @@ -1176,7 +1177,7 @@ def capsule_box( box_size: Half-extents of the box along each axis. Returns: - - Vector of contact distances (wp.inf for unpopulated contacts). + - Vector of contact distances (MJ_MAXVAL for unpopulated contacts). - Matrix of contact positions (one per row). - Matrix of contact normal vectors (one per row). """ @@ -1348,7 +1349,7 @@ def capsule_box( c1 = wp.where((ee2 > 0) == w_neg, 1, 2) if cltype == -4: # invalid type - return wp.vec2(wp.inf), mat23f(), mat23f() + return wp.vec2(MJ_MAXVAL), mat23f(), mat23f() if cltype >= 0 and cltype // 3 != 1: # closest to a corner of the box c1 = axisdir ^ clcorner @@ -1479,7 +1480,7 @@ def capsule_box( # collide with sphere using core function dist2, pos2, normal2 = sphere_box(s2_pos_g, capsule_radius, box_pos, box_rot, box_size) else: - dist2 = wp.inf + dist2 = MJ_MAXVAL pos2 = wp.vec3() normal2 = wp.vec3() 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 65acf580..15ea3a41 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 @@ -22,6 +22,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom_col from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_mesh +from mujoco.mjx.third_party.mujoco_warp._src.types import CollisionContext 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 @@ -68,6 +69,7 @@ class MeshData: data_id: int pos: wp.vec3 mat: wp.mat33 + size: wp.vec3 pnt: wp.vec3 vec: wp.vec3 valid: bool = False @@ -375,6 +377,7 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: Volume mesh_data.data_id, mesh_data.pos, mesh_data.mat, + mesh_data.size, mesh_data.pnt, mesh_data.vec, ) @@ -388,6 +391,7 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: Volume mesh_data.data_id, mesh_data.pos, mesh_data.mat, + mesh_data.size, mesh_data.pnt, -mesh_data.vec, ) @@ -425,6 +429,7 @@ def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: V mesh_data.data_id, mesh_data.pos, mesh_data.mat, + mesh_data.size, mesh_data.pnt, mesh_data.vec, ) @@ -664,11 +669,11 @@ def _sdf_narrowphase( 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), - ncollision_in: wp.array(dtype=int), - # In: sdf_initpoints: int, sdf_iterations: int, # Data out: @@ -785,6 +790,7 @@ def _sdf_narrowphase( mesh_data1.data_id = geom_dataid[g1] mesh_data1.pos = geom1.pos mesh_data1.mat = geom1.rot + mesh_data1.size = geom1.size mesh_data1.pnt = wp.vec3(-1.0) mesh_data1.vec = wp.vec3(0.0) mesh_data1.valid = True @@ -797,6 +803,7 @@ def _sdf_narrowphase( mesh_data2.data_id = geom_dataid[g2] mesh_data2.pos = geom2.pos mesh_data2.mat = geom2.rot + mesh_data2.size = geom2.size mesh_data2.pnt = wp.vec3(-1.0) mesh_data2.vec = wp.vec3(0.0) mesh_data2.valid = True @@ -859,7 +866,7 @@ def _sdf_narrowphase( @event_scope -def sdf_narrowphase(m: Model, d: Data): +def sdf_narrowphase(m: Model, d: Data, ctx: CollisionContext): wp.launch( _sdf_narrowphase, dim=(m.opt.sdf_initpoints, d.naconmax), @@ -909,10 +916,10 @@ def sdf_narrowphase(m: Model, d: Data): d.geom_xpos, d.geom_xmat, d.naconmax, - d.collision_pair, - d.collision_pairid, - d.collision_worldid, d.ncollision, + ctx.collision_pair, + ctx.collision_pairid, + ctx.collision_worldid, m.opt.sdf_initpoints, m.opt.sdf_iterations, ], 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 13468c8c..f786e36d 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py @@ -34,21 +34,11 @@ def _zero_constraint_counts( nf_out: wp.array(dtype=int), nl_out: wp.array(dtype=int), nefc_out: wp.array(dtype=int), - ne_connect_out: wp.array(dtype=int), - ne_weld_out: wp.array(dtype=int), - ne_jnt_out: wp.array(dtype=int), - ne_ten_out: wp.array(dtype=int), - ne_flex_out: wp.array(dtype=int), ): worldid = wp.tid() # Zero all constraint counters ne_out[worldid] = 0 - ne_connect_out[worldid] = 0 - ne_weld_out[worldid] = 0 - ne_jnt_out[worldid] = 0 - ne_ten_out[worldid] = 0 - ne_flex_out[worldid] = 0 nf_out[worldid] = 0 nl_out[worldid] = 0 nefc_out[worldid] = 0 @@ -155,6 +145,7 @@ def _efc_equality_connect( # In: refsafe_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), @@ -165,7 +156,6 @@ def _efc_equality_connect( efc_vel_out: wp.array2d(dtype=float), efc_aref_out: wp.array2d(dtype=float), efc_frictionloss_out: wp.array2d(dtype=float), - ne_connect_out: wp.array(dtype=int), ): """Calculates constraint rows for connect equality constraints.""" worldid, eqconnectid = wp.tid() @@ -174,7 +164,7 @@ def _efc_equality_connect( if not eq_active_in[worldid, eqid]: return - wp.atomic_add(ne_connect_out, worldid, 3) + wp.atomic_add(ne_out, worldid, 3) efcid = wp.atomic_add(nefc_out, worldid, 3) if efcid + 3 >= njmax_in: @@ -205,7 +195,7 @@ def _efc_equality_connect( # compute Jacobian difference (opposite of contact: 0 - 1) Jqvel = wp.vec3f(0.0, 0.0, 0.0) for dofid in range(nv): # TODO: parallelize - jacp1, _ = support.jac( + jacp1, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, @@ -216,7 +206,7 @@ def _efc_equality_connect( dofid, worldid, ) - jacp2, _ = support.jac( + jacp2, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, @@ -293,6 +283,7 @@ def _efc_equality_joint( # In: refsafe_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), @@ -303,7 +294,6 @@ def _efc_equality_joint( efc_vel_out: wp.array2d(dtype=float), efc_aref_out: wp.array2d(dtype=float), efc_frictionloss_out: wp.array2d(dtype=float), - ne_jnt_out: wp.array(dtype=int), ): worldid, eqjntid = wp.tid() eqid = eq_jnt_adr[eqjntid] @@ -311,7 +301,7 @@ def _efc_equality_joint( if not eq_active_in[worldid, eqid]: return - wp.atomic_add(ne_jnt_out, worldid, 1) + wp.atomic_add(ne_out, worldid, 1) efcid = wp.atomic_add(nefc_out, worldid, 1) if efcid >= njmax_in: @@ -399,6 +389,7 @@ def _efc_equality_tendon( # In: refsafe_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), @@ -409,7 +400,6 @@ def _efc_equality_tendon( efc_vel_out: wp.array2d(dtype=float), efc_aref_out: wp.array2d(dtype=float), efc_frictionloss_out: wp.array2d(dtype=float), - ne_ten_out: wp.array(dtype=int), ): worldid, eqtenid = wp.tid() eqid = eq_ten_adr[eqtenid] @@ -417,7 +407,7 @@ def _efc_equality_tendon( if not eq_active_in[worldid, eqid]: return - wp.atomic_add(ne_ten_out, worldid, 1) + wp.atomic_add(ne_out, worldid, 1) efcid = wp.atomic_add(nefc_out, worldid, 1) if efcid >= njmax_in: @@ -494,6 +484,9 @@ def _efc_equality_flex( opt_timestep: wp.array(dtype=float), 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), @@ -505,6 +498,7 @@ def _efc_equality_flex( # In: refsafe_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), @@ -515,12 +509,11 @@ def _efc_equality_flex( efc_vel_out: wp.array2d(dtype=float), efc_aref_out: wp.array2d(dtype=float), efc_frictionloss_out: wp.array2d(dtype=float), - ne_flex_out: wp.array(dtype=int), ): worldid, eqflexid, edgeid = wp.tid() eqid = eq_flex_adr[eqflexid] - wp.atomic_add(ne_flex_out, worldid, 1) + wp.atomic_add(ne_out, worldid, 1) efcid = wp.atomic_add(nefc_out, worldid, 1) if efcid >= njmax_in: @@ -531,10 +524,20 @@ def _efc_equality_flex( solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] Jqvel = float(0.0) + + # TODO(team): remove once efc.J is sparse for i in range(nv): - J = flexedge_J_in[worldid, edgeid, i] - efc_J_out[worldid, efcid, i] = J - Jqvel += J * qvel_in[worldid, i] + efc_J_out[worldid, efcid, i] = 0.0 + + rownnz = flexedge_J_rownnz[edgeid] + rowadr = flexedge_J_rowadr[edgeid] + for i in range(rownnz): + sparseid = rowadr + i + colind = flexedge_J_colind[sparseid] + J = flexedge_J_in[worldid, 0, sparseid] + # TODO(team): sparse efc.J + efc_J_out[worldid, efcid, colind] = J + Jqvel += J * qvel_in[worldid, colind] _update_efc_row( worldid, @@ -748,6 +751,7 @@ def _efc_equality_weld( # In: refsafe_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), @@ -758,7 +762,6 @@ def _efc_equality_weld( efc_vel_out: wp.array2d(dtype=float), efc_aref_out: wp.array2d(dtype=float), efc_frictionloss_out: wp.array2d(dtype=float), - ne_weld_out: wp.array(dtype=int), ): worldid, eqweldid = wp.tid() eqid = eq_wld_adr[eqweldid] @@ -766,7 +769,7 @@ def _efc_equality_weld( if not eq_active_in[worldid, eqid]: return - wp.atomic_add(ne_weld_out, worldid, 6) + wp.atomic_add(ne_out, worldid, 6) efcid = wp.atomic_add(nefc_out, worldid, 6) if efcid + 6 >= njmax_in: @@ -808,7 +811,7 @@ def _efc_equality_weld( Jqvelr = wp.vec3f(0.0, 0.0, 0.0) for dofid in range(nv): # TODO: parallelize - jacp1, jacr1 = support.jac( + jacp1, jacr1 = support.jac_dof( body_parentid, body_rootid, dof_bodyid, @@ -819,7 +822,7 @@ def _efc_equality_weld( dofid, worldid, ) - jacp2, jacr2 = support.jac( + jacp2, jacr2 = support.jac_dof( body_parentid, body_rootid, dof_bodyid, @@ -1215,8 +1218,12 @@ def _efc_contact_pyramidal( 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), # Data in: qvel_in: wp.array2d(dtype=float), @@ -1302,49 +1309,73 @@ def _efc_contact_pyramidal( invweight = invweight * 2.0 * fri0 * fri0 * impratio_invsqrt * impratio_invsqrt Jqvel = float(0.0) - for i in range(nv): - J = float(0.0) - Ji = float(0.0) - jac1p, jac1r = support.jac( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - con_pos, - body1, - i, - worldid, - ) - jac2p, jac2r = support.jac( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - con_pos, - body2, - i, - worldid, - ) - jacp_dif = jac2p - jac1p - for xyz in range(3): - J += frame[0, xyz] * jacp_dif[xyz] + + # 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 + da = wp.max(da1, da2) + + for dofid in range(nv - 1, -1, -1): + if dofid == da: + # TODO(team): contact_jacobian + jac1p, jac1r = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + con_pos, + body1, + dofid, + worldid, + ) + jac2p, jac2r = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + con_pos, + body2, + dofid, + worldid, + ) + + J = float(0.0) + Ji = float(0.0) + if condim > 1: + dimid2 = dimid / 2 + 1 + + for xyz in range(3): + jacp_dif = jac2p[xyz] - jac1p[xyz] + J += frame[0, xyz] * jacp_dif + + if condim > 1: + if dimid2 < 3: + Ji += frame[dimid2, xyz] * jacp_dif + else: + Ji += frame[dimid2 - 3, xyz] * (jac2r[xyz] - jac1r[xyz]) if condim > 1: - if dimid2 < 3: - Ji += frame[dimid2, xyz] * jacp_dif[xyz] + if dimid % 2 == 0: + J += Ji * frii else: - Ji += frame[dimid2 - 3, xyz] * (jac2r[xyz] - jac1r[xyz]) + J -= Ji * frii - if condim > 1: - if dimid % 2 == 0: - J += Ji * frii - else: - J -= Ji * frii + efc_J_out[worldid, efcid, dofid] = J + Jqvel += J * qvel_in[worldid, dofid] - efc_J_out[worldid, efcid, i] = J - Jqvel += J * qvel_in[worldid, i] + # Advance tree pointers and recompute da for next iteration + if da1 == da: + da1 = dof_parentid[da1] + if da2 == da: + da2 = dof_parentid[da2] + da = wp.max(da1, da2) + else: + efc_J_out[worldid, efcid, dofid] = 0.0 if condim == 1: efc_type = ConstraintType.CONTACT_FRICTIONLESS @@ -1385,8 +1416,12 @@ def _efc_contact_elliptic( 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), # Data in: qvel_in: wp.array2d(dtype=float), @@ -1450,49 +1485,69 @@ def _efc_contact_elliptic( impratio_invsqrt = opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] contact_efc_address_out[conid, dimid] = efcid + con_pos = pos_in[conid] + frame = frame_in[conid] + geom = geom_in[conid] body1 = geom_bodyid[geom[0]] body2 = geom_bodyid[geom[1]] - cpos = pos_in[conid] - frame = frame_in[conid] - - # TODO(team): parallelize J and Jqvel computation? Jqvel = float(0.0) - for i in range(nv): - J = float(0.0) - jac1p, jac1r = support.jac( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - cpos, - body1, - i, - worldid, - ) - jac2p, jac2r = support.jac( - body_parentid, - body_rootid, - dof_bodyid, - subtree_com_in, - cdof_in, - cpos, - body2, - i, - worldid, - ) - for xyz in range(3): - if dimid < 3: - jac_dif = jac2p[xyz] - jac1p[xyz] - J += frame[dimid, xyz] * jac_dif - else: - jac_dif = jac2r[xyz] - jac1r[xyz] - J += frame[dimid - 3, xyz] * jac_dif - efc_J_out[worldid, efcid, i] = J - Jqvel += J * qvel_in[worldid, i] + # 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 + da = wp.max(da1, da2) + + for dofid in range(nv - 1, -1, -1): + if dofid == da: + # TODO(team): contact jacobian + jac1p, jac1r = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + con_pos, + body1, + dofid, + worldid, + ) + jac2p, jac2r = support.jac_dof( + body_parentid, + body_rootid, + dof_bodyid, + subtree_com_in, + cdof_in, + con_pos, + body2, + dofid, + worldid, + ) + + J = float(0.0) + for xyz in range(3): + if dimid < 3: + jac_dif = jac2p[xyz] - jac1p[xyz] + J += frame[dimid, xyz] * jac_dif + else: + jac_dif = jac2r[xyz] - jac1r[xyz] + J += frame[dimid - 3, xyz] * jac_dif + + efc_J_out[worldid, efcid, dofid] = J + Jqvel += J * qvel_in[worldid, dofid] + + # Advance tree pointers and recompute da for next iteration + if da1 == da: + da1 = dof_parentid[da1] + if da2 == da: + da2 = dof_parentid[da2] + da = wp.max(da1, da2) + else: + efc_J_out[worldid, efcid, dofid] = 0.0 body_invweight0_id = worldid % body_invweight0.shape[0] invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] @@ -1549,29 +1604,13 @@ def _efc_contact_elliptic( ) -@wp.kernel -def _num_equality( - # Data in: - ne_connect_in: wp.array(dtype=int), - ne_weld_in: wp.array(dtype=int), - ne_jnt_in: wp.array(dtype=int), - ne_ten_in: wp.array(dtype=int), - ne_flex_in: wp.array(dtype=int), - # Data out: - ne_out: wp.array(dtype=int), -): - worldid = wp.tid() - ne = ne_connect_in[worldid] + ne_weld_in[worldid] + ne_jnt_in[worldid] + ne_ten_in[worldid] + ne_flex_in[worldid] - ne_out[worldid] = ne - - @event_scope def make_constraint(m: types.Model, d: types.Data): """Creates constraint jacobians and other supporting data.""" wp.launch( _zero_constraint_counts, dim=d.nworld, - inputs=[d.ne, d.nf, d.nl, d.nefc, d.ne_connect, d.ne_weld, d.ne_jnt, d.ne_ten, d.ne_flex], + inputs=[d.ne, d.nf, d.nl, d.nefc], ) if not (m.opt.disableflags & types.DisableBit.CONSTRAINT): @@ -1608,6 +1647,7 @@ def make_constraint(m: types.Model, d: types.Data): refsafe, ], outputs=[ + d.ne, d.nefc, d.efc.type, d.efc.id, @@ -1618,7 +1658,6 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, - d.ne_connect, ], ) wp.launch( @@ -1653,6 +1692,7 @@ def make_constraint(m: types.Model, d: types.Data): refsafe, ], outputs=[ + d.ne, d.nefc, d.efc.type, d.efc.id, @@ -1663,7 +1703,6 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, - d.ne_weld, ], ) wp.launch( @@ -1689,6 +1728,7 @@ def make_constraint(m: types.Model, d: types.Data): refsafe, ], outputs=[ + d.ne, d.nefc, d.efc.type, d.efc.id, @@ -1699,7 +1739,6 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, - d.ne_jnt, ], ) wp.launch( @@ -1724,6 +1763,7 @@ def make_constraint(m: types.Model, d: types.Data): refsafe, ], outputs=[ + d.ne, d.nefc, d.efc.type, d.efc.id, @@ -1734,7 +1774,6 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, - d.ne_ten, ], ) @@ -1746,6 +1785,9 @@ def make_constraint(m: types.Model, d: types.Data): m.opt.timestep, 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, @@ -1756,6 +1798,7 @@ def make_constraint(m: types.Model, d: types.Data): refsafe, ], outputs=[ + d.ne, d.nefc, d.efc.type, d.efc.id, @@ -1766,17 +1809,9 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, - d.ne_flex, ], ) - wp.launch( - _num_equality, - dim=d.nworld, - inputs=[d.ne_connect, d.ne_weld, d.ne_jnt, d.ne_ten, d.ne_flex], - outputs=[d.ne], - ) - if not (m.opt.disableflags & types.DisableBit.FRICTIONLOSS): wp.launch( _efc_friction_dof, @@ -1957,8 +1992,12 @@ def make_constraint(m: types.Model, d: types.Data): 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, d.qvel, d.subtree_com, @@ -2002,8 +2041,12 @@ def make_constraint(m: types.Model, d: types.Data): 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, d.qvel, d.subtree_com, 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 84a2ce92..e4218855 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py @@ -25,7 +25,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import TileSet from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -from mujoco.mjx.third_party.mujoco_warp._src.warp_util import nested_kernel wp.set_module_options({"enable_backward": False}) @@ -80,12 +79,12 @@ def _qderiv_actuator_passive_vel( @cache_kernel def _qderiv_actuator_passive_actuation_dense(tile: TileSet, nu: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( # Data in: - vel_in: wp.array3d(dtype=float), actuator_moment_in: wp.array3d(dtype=float), # In: + vel_in: wp.array3d(dtype=float), adr: wp.array(dtype=int), # Out: qDeriv_out: wp.array3d(dtype=float), @@ -140,8 +139,8 @@ def _qderiv_actuator_passive( # Model: opt_timestep: wp.array(dtype=float), opt_disableflags: int, - opt_is_sparse: bool, dof_damping: wp.array2d(dtype=float), + is_sparse: bool, # Data in: qM_in: wp.array3d(dtype=float), # In: @@ -156,7 +155,7 @@ def _qderiv_actuator_passive( dofiid = qMi[elemid] dofjid = qMj[elemid] - if opt_is_sparse: + if is_sparse: qderiv = qDeriv_in[worldid, 0, elemid] else: qderiv = qDeriv_in[worldid, dofiid, dofjid] @@ -166,7 +165,7 @@ def _qderiv_actuator_passive( qderiv *= opt_timestep[worldid % opt_timestep.shape[0]] - if opt_is_sparse: + if is_sparse: qDeriv_out[worldid, 0, elemid] = qM_in[worldid, 0, elemid] - qderiv else: qM = qM_in[worldid, dofiid, dofjid] - qderiv @@ -181,8 +180,8 @@ def _qderiv_tendon_damping( # Model: ntendon: int, opt_timestep: wp.array(dtype=float), - opt_is_sparse: bool, tendon_damping: wp.array2d(dtype=float), + is_sparse: bool, # Data in: ten_J_in: wp.array3d(dtype=float), # In: @@ -202,7 +201,7 @@ def _qderiv_tendon_damping( qderiv *= opt_timestep[worldid % opt_timestep.shape[0]] - if opt_is_sparse: + if is_sparse: qDeriv_out[worldid, 0, elemid] -= qderiv else: qDeriv_out[worldid, dofiid, dofjid] -= qderiv @@ -245,7 +244,7 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d(dtype=float)): ], outputs=[vel], ) - if m.opt.is_sparse: + if m.is_sparse: wp.launch( _qderiv_actuator_passive_actuation_sparse, dim=(d.nworld, qMi.size), @@ -258,9 +257,9 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d(dtype=float)): wp.launch_tiled( _qderiv_actuator_passive_actuation_dense(tile, m.nu), dim=(d.nworld, tile.adr.size), - inputs=[vel_3d, d.actuator_moment, tile.adr], + inputs=[d.actuator_moment, vel_3d, tile.adr], outputs=[out], - block_dim=m.block_dim.mul_m_dense, + block_dim=m.block_dim.qderiv_actuator_dense, ) wp.launch( _qderiv_actuator_passive, @@ -268,8 +267,8 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d(dtype=float)): inputs=[ m.opt.timestep, m.opt.disableflags, - m.opt.is_sparse, m.dof_damping, + m.is_sparse, d.qM, qMi, qMj, @@ -285,7 +284,7 @@ def deriv_smooth_vel(m: Model, d: Data, out: wp.array2d(dtype=float)): wp.launch( _qderiv_tendon_damping, dim=(d.nworld, qMi.size), - inputs=[m.ntendon, m.opt.timestep, m.opt.is_sparse, m.tendon_damping, d.ten_J, qMi, qMj], + inputs=[m.ntendon, m.opt.timestep, 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 82f587da..e345c652 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py @@ -42,7 +42,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType from mujoco.mjx.third_party.mujoco_warp._src.types import vec10f from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -from mujoco.mjx.third_party.mujoco_warp._src.warp_util import nested_kernel wp.set_module_options({"enable_backward": False}) @@ -291,17 +290,17 @@ def _euler_damp_qfrc_sparse( @cache_kernel def _tile_euler_dense(tile: TileSet): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def euler_dense( # Model: - dof_damping: wp.array2d(dtype=float), opt_timestep: wp.array(dtype=float), + dof_damping: wp.array2d(dtype=float), # Data in: qM_in: wp.array3d(dtype=float), efc_Ma_in: wp.array2d(dtype=float), # In: adr_in: wp.array(dtype=int), - # Out: + # Data out: qacc_out: wp.array2d(dtype=float), ): worldid, nodeid = wp.tid() @@ -328,7 +327,7 @@ def euler(m: Model, d: Data): # integrate damping implicitly if not m.opt.disableflags & (DisableBit.EULERDAMP | DisableBit.DAMPER): qacc = wp.empty((d.nworld, m.nv), dtype=float) - if m.opt.is_sparse: + if m.is_sparse: qM = wp.clone(d.qM) qLD = wp.empty((d.nworld, 1, m.nC), dtype=float) qLDiagInv = wp.empty((d.nworld, m.nv), dtype=float) @@ -344,7 +343,7 @@ def euler(m: Model, d: Data): wp.launch_tiled( _tile_euler_dense(tile), dim=(d.nworld, tile.adr.size), - inputs=[m.dof_damping, m.opt.timestep, d.qM, d.efc.Ma, tile.adr], + inputs=[m.opt.timestep, m.dof_damping, d.qM, d.efc.Ma, tile.adr], outputs=[qacc], block_dim=m.block_dim.euler_dense, ) @@ -482,7 +481,7 @@ def rungekutta4(m: Model, d: Data): def implicit(m: Model, d: Data): """Integrates fully implicit in velocity.""" if ~(m.opt.disableflags | ~(DisableBit.ACTUATION | DisableBit.SPRING | DisableBit.DAMPER)): - if m.opt.is_sparse: + if m.is_sparse: qDeriv = wp.empty((d.nworld, 1, m.nM), dtype=float) qLD = wp.empty((d.nworld, 1, m.nC), dtype=float) else: @@ -524,7 +523,7 @@ def fwd_position(m: Model, d: Data, factorize: bool = True): # TODO(team): sparse actuator_moment version @cache_kernel def _actuator_velocity(nv: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def actuator_velocity( # Data in: qvel_in: wp.array2d(dtype=float), @@ -544,7 +543,7 @@ def _actuator_velocity(nv: int): @cache_kernel def _tendon_velocity(nv: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def tendon_velocity( # Data in: qvel_in: wp.array2d(dtype=float), diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py index 00449b86..9fb9242b 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py @@ -99,7 +99,7 @@ def discrete_acc(m: Model, d: Data, qacc: wp.array2d(dtype=float)): outputs=[qfrc], ) elif m.opt.integrator == IntegratorType.IMPLICITFAST: - if m.opt.is_sparse: + if m.is_sparse: qDeriv = wp.empty((d.nworld, 1, m.nM), dtype=float) else: qDeriv = wp.empty((d.nworld, m.nv, m.nv), dtype=float) @@ -120,10 +120,8 @@ def inv_constraint(m: Model, d: Data): d.qfrc_constraint.zero_() return - # update - h = wp.empty((d.nworld, 0, 0), dtype=float) # not used - hfactor = wp.empty((d.nworld, 0, 0), dtype=float) # not used - solver.create_context(m, d, h, hfactor, grad=False) + ctx = solver.create_inverse_context(m, d) + solver.init_context(m, d, ctx, grad=False) def inverse(m: Model, d: Data): 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 60ac1647..2febee1f 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -17,15 +17,17 @@ import dataclasses import importlib.metadata import re import warnings -from typing import Any, Optional, Sequence, Union +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 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.warp_util import nested_kernel def _is_mujoco_dev() -> bool: @@ -46,7 +48,7 @@ def _is_mujoco_dev() -> bool: BLEEDING_EDGE_MUJOCO = _is_mujoco_dev() -def _create_array(data: Any, spec: wp.array, sizes: dict[str, int]) -> Union[wp.array, None]: +def _create_array(data: Any, spec: wp.array, sizes: dict[str, int]) -> wp.array | None: """Creates a warp array and populates it with data. The array shape is determined by a field spec referencing MjModel / MjData array sizes. @@ -167,15 +169,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: return True return False - for objtype, objid, reftype, refid in zip( - mjm.sensor_objtype[is_collision_sensor], - mjm.sensor_objid[is_collision_sensor], - mjm.sensor_reftype[is_collision_sensor], - mjm.sensor_refid[is_collision_sensor], - ): - if not_implemented(objtype, objid, types.GeomType.BOX) and not_implemented(reftype, refid, types.GeomType.BOX): - raise NotImplementedError(f"Collision sensors with box-box collisions are not implemented.") - def _check_friction(name: str, id_: int, condim: int, friction, checks): for min_condim, indices in checks: if condim >= min_condim: @@ -203,11 +196,9 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: opt.tolerance = max(opt.tolerance, 1e-6) # warp only fields - opt.is_sparse = is_sparse(mjm) ls_parallel_id = mujoco.mj_name2id(mjm, mujoco.mjtObj.mjOBJ_NUMERIC, "ls_parallel") opt.ls_parallel = (ls_parallel_id > -1) and (mjm.numeric_data[mjm.numeric_adr[ls_parallel_id]] == 1) opt.ls_parallel_min_step = 1.0e-6 # TODO(team): determine good default setting - opt.has_fluid = mjm.opt.wind.any() or mjm.opt.density > 0 or mjm.opt.viscosity > 0 opt.broadphase = types.BroadphaseType.NXN opt.broadphase_filter = types.BroadphaseFilter.PLANE | types.BroadphaseFilter.SPHERE | types.BroadphaseFilter.OBB opt.graph_conditional = True @@ -226,7 +217,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: setattr(opt, f.name, f.type(getattr(opt, f.name))) # create stat - stat = types.Statistic(meaninertia=mjm.stat.meaninertia) + stat = types.Statistic(meaninertia=_create_array([mjm.stat.meaninertia], types.array("*", float), {"*": 1})) # create model m = types.Model(**{f.name: getattr(mjm, f.name, None) for f in dataclasses.fields(types.Model)}) @@ -245,14 +236,35 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.nmaxpyramid = np.maximum(1, 2 * (m.nmaxcondim - 1)) m.has_sdf_geom = (mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF).any() m.block_dim = types.BlockDim() + m.is_sparse = is_sparse(mjm) + m.has_fluid = mjm.opt.wind.any() or mjm.opt.density > 0 or mjm.opt.viscosity > 0 - # body ids grouped by tree level + # 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): body_depth[i] = body_depth[mjm.body_parentid[i]] + 1 bodies.setdefault(body_depth[i], []).append(i) m.body_tree = tuple(wp.array(bodies[i], dtype=int) for i in sorted(bodies)) + # branch-based traversal data + children_count = np.bincount(mjm.body_parentid[1:], minlength=mjm.nbody) + ancestor_chain = lambda b: ancestor_chain(mjm.body_parentid[b]) + [b] if b else [] + branches = [ancestor_chain(l) for l in np.where(children_count[1:] == 0)[0] + 1] + m.nbranch = len(branches) + + body_branches = [] + body_branch_start = [] + offset = 0 + + for branch in branches: + body_branches.extend(branch) + body_branch_start.append(offset) + offset += len(branch) + body_branch_start.append(offset) + + m.body_branches = np.array(body_branches, dtype=int) + m.body_branch_start = np.array(body_branch_start, dtype=int) + m.mocap_bodyid = np.arange(mjm.nbody)[mjm.body_mocapid >= 0] m.mocap_bodyid = m.mocap_bodyid[mjm.body_mocapid[mjm.body_mocapid >= 0].argsort()] m.body_fluid_ellipsoid = np.zeros(mjm.nbody, dtype=bool) @@ -365,18 +377,8 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: ) ) - # compute nmaxpolygon and nmaxmeshdeg given the geom pairs for the model - nboxbox = m.geom_pair_type_count[geom_trid_index(types.GeomType.BOX, types.GeomType.BOX)] - nboxmesh = m.geom_pair_type_count[geom_trid_index(types.GeomType.BOX, types.GeomType.MESH)] - nmeshmesh = m.geom_pair_type_count[geom_trid_index(types.GeomType.MESH, types.GeomType.MESH)] - # need at least 4 (square sides) if there's a box collision needing multiccd - m.nmaxpolygon = 4 * (nboxbox + nboxmesh > 0) - m.nmaxmeshdeg = 3 * (nboxbox + nboxmesh > 0) - # possibly need to allocate more memory if there's meshes - if nmeshmesh + nboxmesh > 0: - # TODO(kbayes): remove nboxbox or enable ccd for box-box collisions - m.nmaxpolygon = np.append(mjm.mesh_polyvertnum, m.nmaxpolygon).max() - m.nmaxmeshdeg = np.append(mjm.mesh_polymapnum, m.nmaxmeshdeg).max() + m.nmaxpolygon = np.append(mjm.mesh_polyvertnum, 0).max() + m.nmaxmeshdeg = np.append(mjm.mesh_polymapnum, 0).max() # filter plugins for only geom plugins, drop the rest m.plugin, m.plugin_attr = [], [] @@ -562,18 +564,44 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: m.qM_fullm_j.append(j) j = mjm.dof_parentid[j] - # indices for sparse qM mul_m (used in support) - m.qM_mulm_i, m.qM_mulm_j, m.qM_madr_ij = [], [], [] + # Gather-based sparse mul_m: for each row, all (col, madr) including diagonal + row_elements = [[] for _ in range(mjm.nv)] + + # Add diagonal + for i in range(mjm.nv): + row_elements[i].append((i, mjm.dof_Madr[i])) + + # Add off-diagonals: ancestors (lower) and descendants (upper) for i in range(mjm.nv): madr_ij, j = mjm.dof_Madr[i], i - while True: madr_ij, j = madr_ij + 1, mjm.dof_parentid[j] if j == -1: break - m.qM_mulm_i.append(i) - m.qM_mulm_j.append(j) - m.qM_madr_ij.append(madr_ij) + row_elements[i].append((j, madr_ij)) # row i gathers M[i,j] * vec[j] + row_elements[j].append((i, madr_ij)) # row j gathers M[j,i] * vec[i] + + # Flatten into CSR-like arrays + m.qM_mulm_rowadr = [0] + m.qM_mulm_col = [] + m.qM_mulm_madr = [] + for i in range(mjm.nv): + for col, madr in row_elements[i]: + m.qM_mulm_col.append(col) + m.qM_mulm_madr.append(madr) + m.qM_mulm_rowadr.append(len(m.qM_mulm_col)) + + # TODO(team): remove after mjwarp depends on mujoco > 3.4.0 in pyproject.toml + if BLEEDING_EDGE_MUJOCO: + m.flexedge_J_rownnz = mjm.flexedge_J_rownnz + m.flexedge_J_rowadr = mjm.flexedge_J_rowadr + m.flexedge_J_colind = mjm.flexedge_J_colind.reshape(-1) + else: + mjd = mujoco.MjData(mjm) + mujoco.mj_forward(mjm, mjd) + m.flexedge_J_rownnz = mjd.flexedge_J_rownnz + m.flexedge_J_rowadr = mjd.flexedge_J_rowadr + m.flexedge_J_colind = mjd.flexedge_J_colind.reshape(-1) # place m on device sizes = dict({"*": 1}, **{f.name: getattr(m, f.name) for f in dataclasses.fields(types.Model) if f.type is int}) @@ -628,8 +656,10 @@ def make_data( mjm: mujoco.MjModel, nworld: int = 1, nconmax: Optional[int] = None, + nccdmax: Optional[int] = None, njmax: Optional[int] = None, naconmax: Optional[int] = None, + naccdmax: Optional[int] = None, ) -> types.Data: """Creates a data object on device. @@ -638,9 +668,11 @@ def make_data( nworld: Number of worlds. nconmax: Number of contacts to allocate per world. Contacts exist in large heterogeneous arrays: one world may have more than nconmax contacts. + 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. naconmax: Number of contacts to allocate for all worlds. Overrides nconmax. + naccdmax: Maximum number of CCD contacts. Defaults to naconmax. Returns: The data object containing the current state and output arrays (device). @@ -649,21 +681,36 @@ def make_data( if nconmax is None: nconmax = _default_nconmax(mjm) + if nconmax < 0: + raise ValueError("nconmax must be >= 0") + + if nccdmax is None: + nccdmax = nconmax + elif nccdmax < 0: + raise ValueError("nccdmax must be >= 0") + elif nccdmax > nconmax: + raise ValueError(f"nccdmax ({nccdmax}) must be <= nconmax ({nconmax})") + if njmax is None: njmax = _default_njmax(mjm) + if njmax < 0: + raise ValueError("njmax must be >= 0") + if nworld < 1: raise ValueError(f"nworld must be >= 1") if naconmax is None: - if nconmax < 0: - raise ValueError("nconmax must be >= 0") naconmax = nworld * nconmax elif naconmax < 0: raise ValueError("naconmax must be >= 0") - if njmax < 0: - raise ValueError("njmax must be >= 0") + if naccdmax is None: + naccdmax = nworld * nccdmax + elif naccdmax < 0: + raise ValueError("naccdmax must be >= 0") + elif naccdmax > naconmax: + raise ValueError(f"naccdmax ({naccdmax}) must be <= naconmax ({naconmax})") sizes = dict({"*": 1}, **{f.name: getattr(mjm, f.name, None) for f in dataclasses.fields(types.Model) if f.type is int}) sizes["nmaxcondim"] = np.concatenate(([0], mjm.geom_condim, mjm.pair_dim)).max() @@ -693,6 +740,7 @@ def make_data( "efc": efc, "nworld": nworld, "naconmax": naconmax, + "naccdmax": naccdmax, "njmax": njmax, "qM": None, "qLD": None, @@ -710,6 +758,8 @@ def make_data( ), # equality constraints "eq_active": wp.array(np.tile(mjm.eq_active0.astype(bool), (nworld, 1)), shape=(nworld, mjm.neq), dtype=bool), + # flexedge + "flexedge_J": None, } for f in dataclasses.fields(types.Data): if f.name in d_kwargs: @@ -725,6 +775,8 @@ def make_data( d.qM = wp.zeros((nworld, sizes["nv_pad"], sizes["nv_pad"]), dtype=float) d.qLD = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float) + d.flexedge_J = wp.zeros((nworld, 1, mjd.flexedge_J.size), dtype=float) + return d @@ -733,8 +785,10 @@ def put_data( mjd: mujoco.MjData, nworld: int = 1, nconmax: Optional[int] = None, + nccdmax: Optional[int] = None, njmax: Optional[int] = None, naconmax: Optional[int] = None, + naccdmax: Optional[int] = None, ) -> types.Data: """Moves data from host to a device. @@ -744,9 +798,11 @@ def put_data( nworld: The number of worlds. nconmax: Number of contacts to allocate per world. Contacts exist in large heterogenous arrays: one world may have more than nconmax contacts. + 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. naconmax: Number of contacts to allocate for all worlds. Overrides nconmax. + naccdmax: Maximum number of CCD contacts. Defaults to naconmax. Returns: The data object containing the current state and output arrays (device). @@ -758,23 +814,38 @@ def put_data( if nconmax is None: nconmax = _default_nconmax(mjm, mjd) + if nconmax < 0: + raise ValueError("nconmax must be >= 0") + + if nccdmax is None: + nccdmax = nconmax + elif nccdmax < 0: + raise ValueError("nccdmax must be >= 0") + elif nccdmax > nconmax: + raise ValueError(f"nccdmax ({nccdmax}) must be <= nconmax ({nconmax})") + if njmax is None: njmax = _default_njmax(mjm, mjd) + if njmax < 0: + raise ValueError("njmax must be >= 0") + if nworld < 1: raise ValueError(f"nworld must be >= 1") if naconmax is None: - if nconmax < 0: - raise ValueError("nconmax must be >= 0") if mjd.ncon > nconmax: raise ValueError(f"nconmax overflow (nconmax must be >= {mjd.ncon})") naconmax = nworld * nconmax elif naconmax < mjd.ncon * nworld: raise ValueError(f"naconmax overflow (naconmax must be >= {mjd.ncon * nworld})") - if njmax < 0: - raise ValueError("njmax must be >= 0") + if naccdmax is None: + naccdmax = nworld * nccdmax + elif naccdmax < 0: + raise ValueError("naccdmax must be >= 0") + elif naccdmax > naconmax: + raise ValueError(f"naccdmax ({naccdmax}) must be <= naconmax ({naconmax})") if mjd.nefc > njmax: raise ValueError(f"njmax overflow (njmax must be >= {mjd.nefc})") @@ -850,6 +921,7 @@ def put_data( "efc": efc, "nworld": nworld, "naconmax": naconmax, + "naccdmax": naccdmax, "njmax": njmax, # fields set after initialization: "solver_niter": None, @@ -859,12 +931,6 @@ def put_data( "actuator_moment": None, "flexedge_J": None, "nacon": None, - "ne_connect": None, - "ne_weld": None, - "ne_jnt": None, - "ne_ten": None, - "ne_flex": None, - "nsolving": None, } for f in dataclasses.fields(types.Data): if f.name in d_kwargs: @@ -890,6 +956,8 @@ def put_data( d.qM = wp.array(np.full((nworld, sizes["nv_pad"], sizes["nv_pad"]), qM_padded), dtype=float) d.qLD = wp.array(np.full((nworld, mjm.nv, mjm.nv), qLD), dtype=float) + d.flexedge_J = wp.array(np.tile(mjd.flexedge_J.reshape(-1), (nworld, 1)).reshape((nworld, 1, -1)), dtype=float) + if mujoco.mj_isSparse(mjm): ten_J = np.zeros((mjm.ntendon, mjm.nv)) mujoco.mju_sparse2dense(ten_J, mjd.ten_J.reshape(-1), mjd.ten_J_rownnz, mjd.ten_J_rowadr, mjd.ten_J_colind.reshape(-1)) @@ -898,19 +966,6 @@ def put_data( 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) - flexedge_J = np.zeros((mjm.nflexedge, mjm.nv)) - if mjd.flexedge_J.size: - # TODO(team): remove after mjwarp depends on mujoco > 3.4.0 in pyproject.toml - if BLEEDING_EDGE_MUJOCO: - mujoco.mju_sparse2dense( - flexedge_J, mjd.flexedge_J.reshape(-1), mjm.flexedge_J_rownnz, mjm.flexedge_J_rowadr, mjm.flexedge_J_colind.reshape(-1) - ) - else: - mujoco.mju_sparse2dense( - flexedge_J, mjd.flexedge_J.reshape(-1), mjd.flexedge_J_rownnz, mjd.flexedge_J_rowadr, mjd.flexedge_J_colind.reshape(-1) - ) - d.flexedge_J = wp.array(np.full((nworld, mjm.nflexedge, mjm.nv), flexedge_J), dtype=float) - # TODO(taylorhowell): sparse actuator_moment actuator_moment = np.zeros((mjm.nu, mjm.nv)) mujoco.mju_sparse2dense(actuator_moment, mjd.actuator_moment, mjd.moment_rownnz, mjd.moment_rowadr, mjd.moment_colind) @@ -918,13 +973,6 @@ def put_data( d.nacon = wp.array([mjd.ncon * nworld], dtype=int) - d.ne_connect = wp.full(nworld, 3 * int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_CONNECT) & mjd.eq_active)), dtype=int) - d.ne_weld = wp.full(nworld, 6 * int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_WELD) & mjd.eq_active)), dtype=int) - d.ne_jnt = wp.full(nworld, np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_JOINT) & mjd.eq_active), dtype=int) - d.ne_ten = wp.full(nworld, np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_TENDON) & mjd.eq_active), dtype=int) - d.ne_flex = wp.full(nworld, np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_FLEX) & mjd.eq_active), dtype=int) - d.nsolving = wp.array([nworld], dtype=int) - return d @@ -1024,25 +1072,22 @@ def get_data_into( result.cdof[:] = d.cdof.numpy()[world_id] result.cinert[:] = d.cinert.numpy()[world_id] result.flexvert_xpos[:] = d.flexvert_xpos.numpy()[world_id] - flexedge_J = d.flexedge_J.numpy()[world_id] - if result.flexedge_J.size: + if mjm.nflexedge > 0: # TODO(team): remove after mjwarp depends on mujoco > 3.4.0 in pyproject.toml - if BLEEDING_EDGE_MUJOCO: - mujoco.mju_dense2sparse( - result.flexedge_J.reshape(-1), - flexedge_J, - mjm.flexedge_J_rownnz, - mjm.flexedge_J_rowadr, - mjm.flexedge_J_colind.reshape(-1), + if not BLEEDING_EDGE_MUJOCO: + m = put_model(mjm) + result.flexedge_J_rownnz[:] = m.flexedge_J_rownnz.numpy() + result.flexedge_J_rowadr[:] = m.flexedge_J_rowadr.numpy() + result.flexedge_J_colind[:, :] = m.flexedge_J_colind.numpy().reshape((mjm.nflexedge, mjm.nv)) + mujoco.mju_sparse2dense( + result.flexedge_J, + d.flexedge_J.numpy()[world_id].reshape(-1), + m.flexedge_J_rownnz.numpy(), + m.flexedge_J_rowadr.numpy(), + m.flexedge_J_colind.numpy(), ) else: - mujoco.mju_dense2sparse( - result.flexedge_J.reshape(-1), - flexedge_J, - result.flexedge_J_rownnz, - result.flexedge_J_rowadr, - result.flexedge_J_colind.reshape(-1), - ) + result.flexedge_J[:] = d.flexedge_J.numpy()[world_id].reshape(-1) result.flexedge_length[:] = d.flexedge_length.numpy()[world_id] result.flexedge_velocity[:] = d.flexedge_velocity.numpy()[world_id] result.actuator_length[:] = d.actuator_length.numpy()[world_id] @@ -1143,7 +1188,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): reset: Per-world bitmask. Reset if True. """ - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def reset_xfrc_applied(reset_in: wp.array(dtype=bool), xfrc_applied_out: wp.array2d(dtype=wp.spatial_vector)): worldid, bodyid, elemid = wp.tid() @@ -1153,7 +1198,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): xfrc_applied_out[worldid, bodyid][elemid] = 0.0 - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def reset_qM(reset_in: wp.array(dtype=bool), qM_out: wp.array3d(dtype=float)): worldid, elemid1, elemid2 = wp.tid() @@ -1163,7 +1208,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): qM_out[worldid, elemid1, elemid2] = 0.0 - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def reset_nworld( # Model: nq: int, @@ -1197,12 +1242,6 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): act_dot_out: wp.array2d(dtype=float), sensordata_out: wp.array2d(dtype=float), nacon_out: wp.array(dtype=int), - ne_connect_out: wp.array(dtype=int), - ne_weld_out: wp.array(dtype=int), - ne_jnt_out: wp.array(dtype=int), - ne_ten_out: wp.array(dtype=int), - ne_flex_out: wp.array(dtype=int), - nsolving_out: wp.array(dtype=int), ): worldid = wp.tid() @@ -1214,16 +1253,9 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): if worldid == 0: nacon_out[0] = 0 ne_out[worldid] = 0 - ne_connect_out[worldid] = 0 - ne_weld_out[worldid] = 0 - ne_jnt_out[worldid] = 0 - ne_ten_out[worldid] = 0 - ne_flex_out[worldid] = 0 nf_out[worldid] = 0 nl_out[worldid] = 0 nefc_out[worldid] = 0 - if worldid == 0: - nsolving_out[0] = nworld_in time_out[worldid] = 0.0 energy_out[worldid] = wp.vec2(0.0, 0.0) qpos0_id = worldid % qpos0.shape[0] @@ -1244,7 +1276,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): for i in range(nsensordata): sensordata_out[worldid, i] = 0.0 - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def reset_mocap( # Model: body_mocapid: wp.array(dtype=int), @@ -1268,7 +1300,7 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): mocap_pos_out[worldid, mocapid] = body_pos[worldid, bodyid] mocap_quat_out[worldid, mocapid] = body_quat[worldid, bodyid] - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def reset_contact( # Data in: nacon_in: wp.array(dtype=int), @@ -1382,17 +1414,604 @@ def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): d.act_dot, d.sensordata, d.nacon, - d.ne_connect, - d.ne_weld, - d.ne_jnt, - d.ne_ten, - d.ne_flex, - d.nsolving, ], ) -def override_model(model: Union[types.Model, mujoco.MjModel], overrides: Union[dict[str, Any], Sequence[str]]): +# kernel_analyzer: off +@wp.kernel +def _init_subtreemass( + body_mass_in: wp.array2d(dtype=float), + body_subtreemass_out: wp.array2d(dtype=float), +): + worldid, bodyid = wp.tid() + body_mass_id = worldid % body_mass_in.shape[0] + body_subtreemass_id = worldid % body_subtreemass_out.shape[0] + body_subtreemass_out[body_subtreemass_id, bodyid] = body_mass_in[body_mass_id, bodyid] + + +@wp.kernel +def _accumulate_subtreemass( + body_parentid: wp.array(dtype=int), + body_subtreemass_io: wp.array2d(dtype=float), + body_tree_: wp.array(dtype=int), +): + worldid, nodeid = wp.tid() + body_subtreemass_id = worldid % body_subtreemass_io.shape[0] + bodyid = body_tree_[nodeid] + parentid = body_parentid[bodyid] + if bodyid != 0: + wp.atomic_add(body_subtreemass_io, body_subtreemass_id, parentid, body_subtreemass_io[body_subtreemass_id, bodyid]) + + +@wp.kernel +def _copy_qpos0_to_qpos( + qpos0: wp.array2d(dtype=float), + qpos_out: wp.array2d(dtype=float), +): + worldid, i = wp.tid() + qpos0_id = worldid % qpos0.shape[0] + qpos_out[worldid, i] = qpos0[qpos0_id, i] + + +@wp.kernel +def _copy_tendon_length0( + ten_length_in: wp.array2d(dtype=float), + tendon_length0_out: wp.array2d(dtype=float), +): + worldid, tenid = wp.tid() + tendon_length0_id = worldid % tendon_length0_out.shape[0] + tendon_length0_out[tendon_length0_id, tenid] = ten_length_in[worldid, tenid] + + +@wp.kernel +def _compute_meaninertia( + nv: int, + is_sparse: bool, + dof_Madr_in: wp.array(dtype=int), + qM_in: wp.array3d(dtype=float), + meaninertia_out: wp.array(dtype=float), +): + """Compute mean diagonal inertia from qM at qpos0.""" + worldid = wp.tid() + + if nv == 0: + meaninertia_out[worldid % meaninertia_out.shape[0]] = 1.0 # Default from MuJoCo + return + + total = float(0.0) + for i in range(nv): + if is_sparse: + # Sparse: qM is flattened lower triangular, diagonal at dof_Madr[i] + madr = dof_Madr_in[i] + total += qM_in[worldid, 0, madr] + else: + # Dense: qM is 2D matrix, diagonal at [i,i] + total += qM_in[worldid, i, i] + + meaninertia_out[worldid % meaninertia_out.shape[0]] = total / float(nv) + + +@wp.kernel +def _set_unit_vector( + dofid_target: int, + unit_vec_out: wp.array2d(dtype=float), +): + worldid = wp.tid() + nv = unit_vec_out.shape[1] + for i in range(nv): + if i == dofid_target: + unit_vec_out[worldid, i] = 1.0 + else: + unit_vec_out[worldid, i] = 0.0 + + +@wp.kernel +def _extract_dof_A_diag( + dofid: int, + result_vec_in: wp.array2d(dtype=float), + dof_A_diag_out: wp.array2d(dtype=float), +): + worldid = wp.tid() + dof_A_diag_id = worldid % dof_A_diag_out.shape[0] + dof_A_diag_out[dof_A_diag_id, dofid] = result_vec_in[worldid, dofid] + + +@wp.kernel +def _finalize_dof_invweight0( + dof_jntid: wp.array(dtype=int), + jnt_type: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_A_diag_in: wp.array2d(dtype=float), + dof_invweight0_out: wp.array2d(dtype=float), +): + worldid, dofid = wp.tid() + dof_invweight0_id = worldid % dof_invweight0_out.shape[0] + dof_A_diag_id = worldid % dof_A_diag_in.shape[0] + + jntid = dof_jntid[dofid] + jtype = jnt_type[jntid] + dofadr = jnt_dofadr[jntid] + + if jtype == int(types.JointType.FREE.value): + # FREE joint: 6 DOFs, average first 3 (trans) and last 3 (rot) separately + if dofid < dofadr + 3: + avg = wp.static(1.0 / 3.0) * ( + dof_A_diag_in[dof_A_diag_id, dofadr + 0] + + dof_A_diag_in[dof_A_diag_id, dofadr + 1] + + dof_A_diag_in[dof_A_diag_id, dofadr + 2] + ) + else: + avg = wp.static(1.0 / 3.0) * ( + dof_A_diag_in[dof_A_diag_id, dofadr + 3] + + dof_A_diag_in[dof_A_diag_id, dofadr + 4] + + dof_A_diag_in[dof_A_diag_id, dofadr + 5] + ) + dof_invweight0_out[dof_invweight0_id, dofid] = avg + elif jtype == int(types.JointType.BALL.value): + # BALL joint: 3 DOFs, average all + avg = wp.static(1.0 / 3.0) * ( + dof_A_diag_in[dof_A_diag_id, dofadr + 0] + + dof_A_diag_in[dof_A_diag_id, dofadr + 1] + + dof_A_diag_in[dof_A_diag_id, dofadr + 2] + ) + dof_invweight0_out[dof_invweight0_id, dofid] = avg + else: + # HINGE/SLIDE: 1 DOF, no averaging + dof_invweight0_out[dof_invweight0_id, dofid] = dof_A_diag_in[dof_A_diag_id, dofid] + + +@wp.kernel +def _compute_body_jac_row( + nv: int, + bodyid_target: int, + row_idx: int, + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + body_dofadr: wp.array(dtype=int), + body_dofnum: wp.array(dtype=int), + dof_parentid: wp.array(dtype=int), + subtree_com_in: wp.array2d(dtype=wp.vec3), + xipos_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + body_jac_row_out: wp.array2d(dtype=float), +): + worldid = wp.tid() + + for i in range(nv): + body_jac_row_out[worldid, i] = 0.0 + + bodyid = bodyid_target + while bodyid > 0 and body_dofnum[bodyid] == 0: + bodyid = body_parentid[bodyid] + + if bodyid == 0: + return + + # Compute offset from point (xipos) to subtree_com of root body + point = xipos_in[worldid, bodyid_target] + offset = point - subtree_com_in[worldid, body_rootid[bodyid_target]] + + # Get last dof that affects this body + dofid = body_dofadr[bodyid] + body_dofnum[bodyid] - 1 + + # Backward pass over dof ancestor chain + while dofid >= 0: + cdof = cdof_in[worldid, dofid] + cdof_ang = wp.spatial_top(cdof) + cdof_lin = wp.spatial_bottom(cdof) + + if row_idx < 3: + tmp = wp.cross(cdof_ang, offset) + if row_idx == 0: + body_jac_row_out[worldid, dofid] = cdof_lin[0] + tmp[0] + elif row_idx == 1: + body_jac_row_out[worldid, dofid] = cdof_lin[1] + tmp[1] + else: + body_jac_row_out[worldid, dofid] = cdof_lin[2] + tmp[2] + else: + if row_idx == 3: + body_jac_row_out[worldid, dofid] = cdof_ang[0] + elif row_idx == 4: + body_jac_row_out[worldid, dofid] = cdof_ang[1] + else: + body_jac_row_out[worldid, dofid] = cdof_ang[2] + + dofid = dof_parentid[dofid] + + +@wp.kernel +def _compute_body_A_diag_entry( + nv: int, + bodyid_target: int, + row_idx: int, + body_jac_row_in: wp.array2d(dtype=float), + result_vec_in: wp.array2d(dtype=float), + body_A_diag_out: wp.array3d(dtype=float), +): + worldid = wp.tid() + body_A_diag_id = worldid % body_A_diag_out.shape[0] + # A[row,row] = J[row] · inv(M) · J[row]' = J[row] · result_vec + dot_prod = float(0.0) + for i in range(nv): + dot_prod += body_jac_row_in[worldid, i] * result_vec_in[worldid, i] + body_A_diag_out[body_A_diag_id, bodyid_target, row_idx] = dot_prod + + +@wp.kernel +def _finalize_body_invweight0( + body_weldid: wp.array(dtype=int), + body_A_diag_in: wp.array3d(dtype=float), + body_invweight0_out: wp.array2d(dtype=wp.vec2), +): + worldid, bodyid = wp.tid() + body_invweight0_id = worldid % body_invweight0_out.shape[0] + body_A_diag_id = worldid % body_A_diag_in.shape[0] + + # World body and static bodies have zero invweight + if bodyid == 0 or body_weldid[bodyid] == 0: + body_invweight0_out[body_invweight0_id, bodyid] = wp.vec2(0.0, 0.0) + return + + # Average diagonal: trans = (A[0,0]+A[1,1]+A[2,2])/3, rot = (A[3,3]+A[4,4]+A[5,5])/3 + inv_trans = wp.static(1.0 / 3.0) * ( + body_A_diag_in[body_A_diag_id, bodyid, 0] + + body_A_diag_in[body_A_diag_id, bodyid, 1] + + body_A_diag_in[body_A_diag_id, bodyid, 2] + ) + inv_rot = wp.static(1.0 / 3.0) * ( + body_A_diag_in[body_A_diag_id, bodyid, 3] + + body_A_diag_in[body_A_diag_id, bodyid, 4] + + body_A_diag_in[body_A_diag_id, bodyid, 5] + ) + + # Prevent degenerate constraints: if one component is near zero, use the other as fallback + if inv_trans < mujoco.mjMINVAL and inv_rot > mujoco.mjMINVAL: + inv_trans = inv_rot # use rotation as fallback for translation + elif inv_rot < mujoco.mjMINVAL and inv_trans > mujoco.mjMINVAL: + inv_rot = inv_trans # use translation as fallback for rotation + + body_invweight0_out[body_invweight0_id, bodyid] = wp.vec2(inv_trans, inv_rot) + + +@wp.kernel +def _copy_tendon_jacobian( + tenid_target: int, + ten_J_in: wp.array3d(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] + + +@wp.kernel +def _compute_tendon_dot_product( + tenid_target: int, + nv: int, + ten_J_in: wp.array3d(dtype=float), + result_vec_in: wp.array2d(dtype=float), + 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] + tendon_invweight0_out[tendon_invweight0_id, tenid_target] = dot_prod + + +@wp.kernel +def _compute_cam_pos0( + cam_bodyid: wp.array(dtype=int), + cam_targetbodyid: wp.array(dtype=int), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + cam_pos0_out: wp.array2d(dtype=wp.vec3), + cam_poscom0_out: wp.array2d(dtype=wp.vec3), + cam_mat0_out: wp.array2d(dtype=wp.mat33), +): + worldid, camid = wp.tid() + cam_pos0_id = worldid % cam_pos0_out.shape[0] + bodyid = cam_bodyid[camid] + targetid = cam_targetbodyid[camid] + cam_xpos = cam_xpos_in[worldid, camid] + + cam_pos0_out[cam_pos0_id, camid] = cam_xpos - xpos_in[worldid, bodyid] + if targetid >= 0: + cam_poscom0_out[cam_pos0_id, camid] = cam_xpos - subtree_com_in[worldid, targetid] + else: + cam_poscom0_out[cam_pos0_id, camid] = cam_xpos - subtree_com_in[worldid, bodyid] + cam_mat0_out[cam_pos0_id, camid] = cam_xmat_in[worldid, camid] + + +@wp.kernel +def _compute_light_pos0( + light_bodyid: wp.array(dtype=int), + light_targetbodyid: wp.array(dtype=int), + light_xpos_in: wp.array2d(dtype=wp.vec3), + light_xdir_in: wp.array2d(dtype=wp.vec3), + xpos_in: wp.array2d(dtype=wp.vec3), + subtree_com_in: wp.array2d(dtype=wp.vec3), + light_pos0_out: wp.array2d(dtype=wp.vec3), + light_poscom0_out: wp.array2d(dtype=wp.vec3), + light_dir0_out: wp.array2d(dtype=wp.vec3), +): + worldid, lightid = wp.tid() + light_pos0_id = worldid % light_pos0_out.shape[0] + bodyid = light_bodyid[lightid] + targetid = light_targetbodyid[lightid] + light_xpos = light_xpos_in[worldid, lightid] + + light_pos0_out[light_pos0_id, lightid] = light_xpos - xpos_in[worldid, bodyid] + if targetid >= 0: + light_poscom0_out[light_pos0_id, lightid] = light_xpos - subtree_com_in[worldid, targetid] + else: + light_poscom0_out[light_pos0_id, lightid] = light_xpos - subtree_com_in[worldid, bodyid] + light_dir0_out[light_pos0_id, lightid] = light_xdir_in[worldid, lightid] + + +@wp.kernel +def _copy_actuator_moment( + actid_target: int, + actuator_moment_in: wp.array3d(dtype=float), + act_moment_vec_out: wp.array2d(dtype=float), +): + worldid = wp.tid() + nv = actuator_moment_in.shape[2] + for i in range(nv): + act_moment_vec_out[worldid, i] = actuator_moment_in[worldid, actid_target, i] + + +@wp.kernel +def _compute_actuator_acc0( + actid_target: int, + nv: int, + result_vec_in: wp.array2d(dtype=float), + actuator_acc0_out: wp.array(dtype=float), +): + worldid = wp.tid() + norm_sq = float(0.0) + for i in range(nv): + norm_sq += result_vec_in[worldid, i] * result_vec_in[worldid, i] + actuator_acc0_out[actid_target] = wp.sqrt(norm_sq) + + +# kernel_analyzer: on + + +def set_const_fixed(m: types.Model, d: types.Data): + """Compute fixed quantities (independent of qpos0). + + Computes: + - body_subtreemass: mass of body and all descendants (depends on body_mass) + - ngravcomp: count of bodies with gravity compensation (depends on body_gravcomp) + + Args: + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + """ + wp.launch(_init_subtreemass, dim=(d.nworld, m.nbody), inputs=[m.body_mass], outputs=[m.body_subtreemass]) + for i in reversed(range(len(m.body_tree))): + body_tree = m.body_tree[i] + wp.launch( + _accumulate_subtreemass, + dim=(d.nworld, body_tree.size), + inputs=[m.body_parentid, m.body_subtreemass, body_tree], + ) + + # TODO(team): refactor for graph capture compatibility + body_gravcomp_np = m.body_gravcomp.numpy() + m.ngravcomp = int((body_gravcomp_np > 0.0).any(axis=0).sum()) + + +def set_const_0(m: types.Model, d: types.Data): + """Compute quantities that depend on qpos0. + + Computes: + - tendon_length0: tendon resting lengths + - dof_invweight0: inverse inertia for DOFs + - body_invweight0: inverse spatial inertia for bodies + - tendon_invweight0: inverse weight for tendons + - cam_pos0, cam_poscom0, cam_mat0: camera references + - light_pos0, light_poscom0, light_dir0: light references + - actuator_acc0: acceleration from unit actuator force + + Args: + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + """ + qpos_saved = wp.clone(d.qpos) + + wp.launch(_copy_qpos0_to_qpos, dim=(d.nworld, m.nq), inputs=[m.qpos0], outputs=[d.qpos]) + + smooth.kinematics(m, d) + smooth.com_pos(m, d) + smooth.camlight(m, d) + smooth.flex(m, d) + smooth.tendon(m, d) + smooth.crb(m, d) + smooth.tendon_armature(m, d) + smooth.factor_m(m, d) + smooth.transmission(m, d) + + # Compute meaninertia from qM diagonal at qpos0 + wp.launch( + _compute_meaninertia, + dim=d.nworld, + inputs=[m.nv, m.is_sparse, m.dof_Madr, d.qM], + outputs=[m.stat.meaninertia], + ) + + wp.launch(_copy_tendon_length0, dim=(d.nworld, m.ntendon), inputs=[d.ten_length], outputs=[m.tendon_length0]) + + # dof_invweight0: computed per joint with averaging for multi-DOF joints + # FREE: 6 DOFs, trans gets mean(A[0:3]), rot gets mean(A[3:6]) + # BALL: 3 DOFs, all get mean(A[0:3]) + # HINGE/SLIDE: 1 DOF, gets A[0,0] + if m.nv > 0: + unit_vec = wp.zeros((d.nworld, m.nv), dtype=float) + result_vec = wp.zeros((d.nworld, m.nv), dtype=float) + dof_A_diag = wp.zeros((d.nworld, m.nv), dtype=float) + + # TODO(team): more efficient approach instead of looping over nv? + for dofid in range(m.nv): + wp.launch(_set_unit_vector, dim=d.nworld, inputs=[dofid], outputs=[unit_vec]) + smooth.solve_m(m, d, result_vec, unit_vec) + wp.launch(_extract_dof_A_diag, dim=d.nworld, inputs=[dofid, result_vec], outputs=[dof_A_diag]) + + wp.launch( + _finalize_dof_invweight0, + dim=(d.nworld, m.nv), + inputs=[m.dof_jntid, m.jnt_type, m.jnt_dofadr, dof_A_diag], + outputs=[m.dof_invweight0], + ) + + # body_invweight0: computed as mean diagonal of J * inv(M) * J' + # where J is the 6xnv body Jacobian (3 rows translation, 3 rows rotation) + if m.nv > 0: + body_jac_row = wp.zeros((d.nworld, m.nv), dtype=float) + body_result_vec = wp.zeros((d.nworld, m.nv), dtype=float) + body_A_diag = wp.zeros((d.nworld, m.nbody, 6), dtype=float) + + # TODO(team): more efficient approach instead of nested iterations? + for bodyid in range(1, m.nbody): + for row_idx in range(6): + wp.launch( + _compute_body_jac_row, + dim=d.nworld, + inputs=[ + m.nv, + bodyid, + row_idx, + m.body_parentid, + m.body_rootid, + m.body_dofadr, + m.body_dofnum, + m.dof_parentid, + d.subtree_com, + d.xipos, + d.cdof, + ], + outputs=[body_jac_row], + ) + smooth.solve_m(m, d, body_result_vec, body_jac_row) + wp.launch( + _compute_body_A_diag_entry, + dim=d.nworld, + inputs=[m.nv, bodyid, row_idx, body_jac_row, body_result_vec], + outputs=[body_A_diag], + ) + + wp.launch( + _finalize_body_invweight0, + dim=(d.nworld, m.nbody), + inputs=[m.body_weldid, body_A_diag], + outputs=[m.body_invweight0], + ) + else: + m.body_invweight0.zero_() + + # 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) + + for tenid in range(m.ntendon): + wp.launch(_copy_tendon_jacobian, dim=d.nworld, inputs=[tenid, 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], + outputs=[m.tendon_invweight0], + ) + + wp.launch( + _compute_cam_pos0, + dim=(d.nworld, m.ncam), + inputs=[m.cam_bodyid, m.cam_targetbodyid, d.cam_xpos, d.cam_xmat, d.xpos, d.subtree_com], + outputs=[m.cam_pos0, m.cam_poscom0, m.cam_mat0], + ) + + wp.launch( + _compute_light_pos0, + dim=(d.nworld, m.nlight), + inputs=[m.light_bodyid, m.light_targetbodyid, d.light_xpos, d.light_xdir, d.xpos, d.subtree_com], + outputs=[m.light_pos0, m.light_poscom0, m.light_dir0], + ) + + # actuator_acc0[i] = ||inv(M) * actuator_moment[i]|| - acceleration from unit actuator force + if m.nu > 0 and m.nv > 0: + act_moment_vec = wp.zeros((d.nworld, m.nv), dtype=float) + act_result_vec = wp.zeros((d.nworld, m.nv), dtype=float) + + for actid in range(m.nu): + wp.launch(_copy_actuator_moment, dim=d.nworld, inputs=[actid, 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]) + + wp.copy(d.qpos, qpos_saved) + + +def set_const(m: types.Model, d: types.Data): + """Recomputes qpos0-dependent constant model fields. + + This function propagates changes from some model fields to derived fields, + allowing modifications that would otherwise be unsafe. It should be called + after modifying model parameters at runtime. + + Model fields that can be modified safely with set_const: + + Field | Notes + ---------------------------------|---------------------------------------------- + qpos0, qpos_spring | + 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. + dof_armature | + 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. + + For selective updates, use the sub-functions directly based on what changed: + + Modified Field | Call + ----------------|------------------ + body_mass | set_const + body_gravcomp | set_const_fixed + body_inertia | set_const_0 + qpos0 | set_const_0 + + Computes: + - Fixed quantities (via set_const_fixed): + - body_subtreemass: mass of body and all descendants + - ngravcomp: count of bodies with gravity compensation + - qpos0-dependent quantities (via set_const_0): + - tendon_length0: tendon resting lengths + - dof_invweight0: inverse inertia for DOFs + - body_invweight0: inverse spatial inertia for bodies + - tendon_invweight0: inverse weight for tendons + - cam_pos0, cam_poscom0, cam_mat0: camera references + - light_pos0, light_poscom0, light_dir0: light references + - actuator_acc0: acceleration from unit actuator force + + Skips: dof_M0, actuator_length0 (not in mjwarp). + + Args: + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + """ + set_const_fixed(m, d) + set_const_0(m, d) + + +def override_model(model: types.Model | mujoco.MjModel, overrides: dict[str, Any] | Sequence[str]): """Overrides model parameters. Overrides are of the format: @@ -1410,9 +2029,16 @@ def override_model(model: Union[types.Model, mujoco.MjModel], overrides: Union[d "opt.integrator": types.IntegratorType, "opt.solver": types.SolverType, } + # MuJoCo pybind11 enums don't support iteration, so we provide explicit mappings + mj_enum_fields = { + "opt.jacobian": { + "DENSE": mujoco.mjtJacobian.mjJAC_DENSE, + "SPARSE": mujoco.mjtJacobian.mjJAC_SPARSE, + "AUTO": mujoco.mjtJacobian.mjJAC_AUTO, + }, + } mjw_only_fields = {"opt.broadphase", "opt.broadphase_filter", "opt.ls_parallel", "opt.graph_conditional"} mj_only_fields = {"opt.jacobian"} - readonly_fields = {"opt.is_sparse"} if not isinstance(overrides, dict): overrides_dict = {} @@ -1430,9 +2056,6 @@ def override_model(model: Union[types.Model, mujoco.MjModel], overrides: Union[d if key in mj_only_fields and isinstance(model, types.Model): continue - if key in readonly_fields and isinstance(model, types.Model): - raise ValueError(f"Cannot override {key} on mjw.Model: field affects model initialization and has side effects") - obj, attrs = model, key.split(".") for i, attr in enumerate(attrs): if not hasattr(obj, attr): @@ -1443,7 +2066,12 @@ def override_model(model: Union[types.Model, mujoco.MjModel], overrides: Union[d typ = type(getattr(obj, attr)) - if key in enum_fields and isinstance(val, str): + if key in mj_enum_fields and isinstance(val, str): + enum_member = val.strip().upper() + if enum_member not in mj_enum_fields[key]: + raise ValueError(f"Unrecognized enum value for {key}: {enum_member}") + val = mj_enum_fields[key][enum_member] + elif key in enum_fields and isinstance(val, str): # special case: enum value enum_members = val.split("|") val = 0 @@ -1499,3 +2127,306 @@ def make_trajectory(model: mujoco.MjModel, keys: list[int]) -> np.ndarray: prev_time = time return np.array(ctrls) + + +@wp.kernel +def _build_rays( + # In: + offset: int, + img_w: int, + img_h: int, + projection: int, + fovy: float, + sensorsize: wp.vec2, + intrinsic: wp.vec4, + znear: float, + # Out: + ray_out: wp.array(dtype=wp.vec3), +): + xid, yid = wp.tid() + ray_out[offset + xid + yid * img_w] = render_util.compute_ray( + projection, fovy, sensorsize, intrinsic, img_w, img_h, xid, yid, znear + ) + + +def create_render_context( + mjm: mujoco.MjModel, + m: types.Model, + d: types.Data, + 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, + use_textures: bool = True, + use_shadows: bool = False, + enabled_geom_groups: list[int] = [0, 1, 2], + cam_active: list[bool] | None = None, + flex_render_smooth: bool = True, +) -> types.RenderContext: + """Creates a render context on device. + + Args: + mjm: The model containing kinematic and dynamic information on host. + m: The model on device. + d: The data on device. + cam_res: The width and height to render each camera image. If None, uses the + 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. + use_textures: Whether to use textures. + use_shadows: Whether to use shadows. + enabled_geom_groups: The geom groups to render. + cam_active: List of booleans indicating which cameras to include in rendering. + If None, all cameras are included. + flex_render_smooth: Whether to render flex meshes smoothly. + + Returns: + The render context containing rendering fields and output arrays on device. + """ + # 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 + + # Mesh BVHs + nmesh = mjm.nmesh + geom_enabled_mask = np.isin(mjm.geom_group, list(enabled_geom_groups)) + mesh_geom_mask = geom_enabled_mask & (mjm.geom_type == types.GeomType.MESH) & (mjm.geom_dataid >= 0) + used_mesh_id = set(mjm.geom_dataid[mesh_geom_mask].astype(int)) + geom_enabled_idx = np.nonzero(geom_enabled_mask)[0] + + mesh_registry = {} + mesh_bvh_id = [wp.uint64(0) for _ in range(nmesh)] + 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_registry[mesh.id] = mesh + mesh_bvh_id[mid] = mesh.id + mesh_bounds_size[mid] = half + + mesh_bvh_id_arr = wp.array(mesh_bvh_id, dtype=wp.uint64) + mesh_bounds_size_arr = wp.array(mesh_bounds_size, dtype=wp.vec3) + + # HField BVHs + nhfield = mjm.nhfield + hfield_geom_mask = geom_enabled_mask & (mjm.geom_type == types.GeomType.HFIELD) & (mjm.geom_dataid >= 0) + used_hfield_id = set(mjm.geom_dataid[hfield_geom_mask].astype(int)) + hfield_registry = {} + hfield_bvh_id = [wp.uint64(0) for _ in range(nhfield)] + 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) + hfield_registry[hmesh.id] = hmesh + hfield_bvh_id[hid] = hmesh.id + hfield_bounds_size[hid] = hhalf + + hfield_bvh_id_arr = wp.array(hfield_bvh_id, dtype=wp.uint64) + 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(d.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 + + if mjm.nflex > 0: + ( + fmesh, + face_point, + flex_group_roots, + flex_shell_data, + flex_faceadr_data, + flex_nface, + ) = bvh.build_flex_bvh(mjm, m, d) + + 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) + + 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) + + # Filter active cameras + if cam_active is not None: + assert len(cam_active) == mjm.ncam, f"cam_active must have length {mjm.ncam} (got {len(cam_active)})" + active_cam_indices = np.nonzero(cam_active)[0] + else: + active_cam_indices = list(range(mjm.ncam)) + + ncam = len(active_cam_indices) + + if cam_res is not None: + if isinstance(cam_res, tuple): + cam_res = [cam_res] * ncam + assert len(cam_res) == ncam, ( + f"Camera resolutions must be provided for all active cameras (got {len(cam_res)}, expected {ncam})" + ) + active_cam_res = cam_res + else: + active_cam_res = mjm.cam_resolution[active_cam_indices] + + cam_res_arr = wp.array(active_cam_res, dtype=wp.vec2i) + + if render_rgb and isinstance(render_rgb, bool): + render_rgb = [render_rgb] * ncam + elif render_rgb is None: + # TODO: remove after mjwarp depends on mujoco >= 3.4.1 in pyproject.toml + if BLEEDING_EDGE_MUJOCO: + render_rgb = [mjm.cam_output[i] & mujoco.mjtCamOutBit.mjCAMOUT_RGB for i in active_cam_indices] + else: + render_rgb = [True] * ncam + + if render_depth and isinstance(render_depth, bool): + render_depth = [render_depth] * ncam + elif render_depth is None: + # TODO: remove after mjwarp depends on mujoco >= 3.4.1 in pyproject.toml + if BLEEDING_EDGE_MUJOCO: + render_depth = [mjm.cam_output[i] & mujoco.mjtCamOutBit.mjCAMOUT_DEPTH for i in active_cam_indices] + else: + render_depth = [True] * ncam + + assert len(render_rgb) == ncam and len(render_depth) == ncam, ( + f"Render RGB and depth must be provided for all active cameras (got {len(render_rgb)}, {len(render_depth)}, expected {ncam})" + ) + + rgb_adr = -1 * np.ones(ncam, dtype=int) + depth_adr = -1 * np.ones(ncam, dtype=int) + cam_res_np = cam_res_arr.numpy() + ri = 0 + di = 0 + total = 0 + + for idx in range(ncam): + if render_rgb[idx]: + rgb_adr[idx] = ri + ri += cam_res_np[idx][0] * cam_res_np[idx][1] + if render_depth[idx]: + depth_adr[idx] = di + di += cam_res_np[idx][0] * cam_res_np[idx][1] + + total += cam_res_np[idx][0] * cam_res_np[idx][1] + + znear = mjm.vis.map.znear * mjm.stat.extent + + if m.cam_fovy.shape[0] > 1 or m.cam_intrinsic.shape[0] > 1: + ray = None + else: + ray = wp.zeros(int(total), dtype=wp.vec3) + + offset = 0 + for idx, cam_id in enumerate(active_cam_indices): + img_w = cam_res_np[idx][0] + img_h = cam_res_np[idx][1] + wp.launch( + kernel=_build_rays, + dim=(img_w, img_h), + inputs=[ + offset, + img_w, + img_h, + m.cam_projection.numpy()[cam_id].item(), + m.cam_fovy.numpy()[0, cam_id].item(), + wp.vec2(m.cam_sensorsize.numpy()[cam_id]), + wp.vec4(m.cam_intrinsic.numpy()[0, cam_id]), + znear, + ], + outputs=[ray], + ) + offset += img_w * img_h + + bvh_ngeom = len(geom_enabled_idx) + + rc = types.RenderContext( + nrender=ncam, + cam_res=cam_res_arr, + cam_id_map=wp.array(active_cam_indices, dtype=int), + use_textures=use_textures, + use_shadows=use_shadows, + background_color=render_util.pack_rgba_to_uint32(0.1 * 255.0, 0.1 * 255.0, 0.2 * 255.0, 1.0 * 255.0), + bvh_ngeom=bvh_ngeom, + enabled_geom_ids=wp.array(geom_enabled_idx, dtype=int), + mesh_registry=mesh_registry, + mesh_bvh_id=mesh_bvh_id_arr, + mesh_bounds_size=mesh_bounds_size_arr, + mesh_texcoord=wp.array(mjm.mesh_texcoord, dtype=wp.vec2), + mesh_texcoord_offsets=wp.array(mjm.mesh_texcoordadr, dtype=int), + mesh_facetexcoord=wp.array(mjm.mesh_facetexcoord, dtype=wp.vec3i), + textures=textures, + textures_registry=textures_registry, + hfield_registry=hfield_registry, + hfield_bvh_id=hfield_bvh_id_arr, + hfield_bounds_size=hfield_bounds_size_arr, + flex_mesh=flex_mesh, + 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_render_smooth=flex_render_smooth, + bvh=None, + bvh_id=None, + lower=wp.zeros(d.nworld * bvh_ngeom, dtype=wp.vec3), + upper=wp.zeros(d.nworld * bvh_ngeom, dtype=wp.vec3), + group=wp.zeros(d.nworld * bvh_ngeom, dtype=int), + group_root=wp.zeros(d.nworld, dtype=int), + ray=ray, + rgb_data=wp.zeros((d.nworld, ri), dtype=wp.uint32), + rgb_adr=wp.array(rgb_adr, dtype=int), + depth_data=wp.zeros((d.nworld, di), dtype=wp.float32), + depth_adr=wp.array(depth_adr, dtype=int), + render_rgb=wp.array(render_rgb, dtype=bool), + render_depth=wp.array(render_depth, dtype=bool), + znear=znear, + total_rays=int(total), + ) + + bvh.build_scene_bvh(m, d, rc) + + return rc diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py new file mode 100644 index 00000000..1b93db0c --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/island.py @@ -0,0 +1,178 @@ +# 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. +# ============================================================================== + +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 + + +@wp.kernel +def _tree_edges( + # Model: + nv: int, + body_treeid: wp.array(dtype=int), + jnt_dofadr: wp.array(dtype=int), + dof_treeid: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + site_bodyid: wp.array(dtype=int), + eq_type: wp.array(dtype=int), + eq_obj1id: wp.array(dtype=int), + eq_obj2id: wp.array(dtype=int), + eq_objtype: wp.array(dtype=int), + # Data in: + nefc_in: wp.array(dtype=int), + contact_geom_in: wp.array(dtype=wp.vec2i), + efc_type_in: wp.array2d(dtype=int), + efc_id_in: wp.array2d(dtype=int), + efc_J_in: wp.array3d(dtype=float), + njmax_in: int, + # Out: + tree_tree: wp.array3d(dtype=int), # kernel_analyzer: off +): + """Find tree edges. Launch: (nworld, njmax).""" + worldid, efcid = wp.tid() + + # skip if beyond active constraints + if efcid >= wp.min(njmax_in, nefc_in[worldid]): + return + + efc_type = efc_type_in[worldid, efcid] + efc_id = efc_id_in[worldid, efcid] + + tree0 = int(-1) + tree1 = int(-1) + use_generic = int(0) + + # equality (connect/weld) + if efc_type == ConstraintType.EQUALITY: + eq_t = eq_type[efc_id] + + if eq_t == EqType.CONNECT or eq_t == EqType.WELD: + b1 = eq_obj1id[efc_id] + b2 = eq_obj2id[efc_id] + + # site semantics + if eq_objtype[efc_id] == ObjType.SITE: + b1 = site_bodyid[b1] + b2 = site_bodyid[b2] + + tree0 = body_treeid[b1] + tree1 = body_treeid[b2] + else: + # JOINT, TENDON, FLEX + use_generic = 1 + + # joint friction + elif efc_type == ConstraintType.FRICTION_DOF: + tree0 = dof_treeid[efc_id] + + # joint limit + elif efc_type == ConstraintType.LIMIT_JOINT: + tree0 = dof_treeid[jnt_dofadr[efc_id]] + + # contact + elif ( + efc_type == ConstraintType.CONTACT_FRICTIONLESS + or efc_type == ConstraintType.CONTACT_PYRAMIDAL + or efc_type == ConstraintType.CONTACT_ELLIPTIC + ): + geom_pair = contact_geom_in[efc_id] + g1 = geom_pair[0] + g2 = geom_pair[1] + + # flex contacts have negative geom ids + if g1 >= 0 and g2 >= 0: + tree0 = body_treeid[geom_bodyid[g1]] + tree1 = body_treeid[geom_bodyid[g2]] + else: + use_generic = 1 + + # generic + else: + use_generic = 1 + + # handle static bodies + if use_generic == 0: + # swap so tree0 is non-negative if possible + if tree0 < 0 and tree1 >= 0: + tree0 = tree1 + tree1 = -1 + + # mark the edge + if tree0 >= 0: + if tree1 < 0 or tree0 == tree1: + # self-edge + wp.atomic_max(tree_tree, worldid, tree0, tree0, 1) + else: + # cross-tree edge + t1 = wp.min(tree0, tree1) + t2 = wp.max(tree0, tree1) + wp.atomic_max(tree_tree, worldid, t1, t2, 1) + wp.atomic_max(tree_tree, worldid, t2, t1, 1) + return + + # generic: scan Jacobian row + first_tree = int(-1) + has_cross_edge = int(0) + + for dof in range(nv): + # TODO(team): sparse efc_J + # TODO(team): tree dof skip + J_val = efc_J_in[worldid, efcid, dof] + if J_val != 0.0: + tree = dof_treeid[dof] + if tree < 0: + continue + if first_tree == -1: + first_tree = tree + elif tree != first_tree: + t1 = wp.min(first_tree, tree) + t2 = wp.max(first_tree, tree) + wp.atomic_max(tree_tree, worldid, t1, t2, 1) + has_cross_edge = 1 + + if first_tree >= 0 and has_cross_edge == 0: + wp.atomic_max(tree_tree, worldid, first_tree, first_tree, 1) + + +def tree_edges(m: types.Model, d: types.Data, tree_tree: wp.array3d(dtype=int)): + """Compute tree-tree adjacency matrix.""" + tree_tree.zero_() + wp.launch( + kernel=_tree_edges, + dim=(d.nworld, d.njmax), + inputs=[ + m.nv, + m.body_treeid, + m.jnt_dofadr, + m.dof_treeid, + m.geom_bodyid, + m.site_bodyid, + m.eq_type, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + d.nefc, + d.contact.geom, + d.efc.type, + d.efc.id, + d.efc.J, + d.njmax, + ], + outputs=[tree_tree], + ) 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 ed59f5ad..77ed86c0 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -258,7 +258,7 @@ def _gravity_force( if gravcomp: force = -gravity * body_mass[worldid % body_mass.shape[0], bodyid] * gravcomp pos = xipos_in[worldid, bodyid] - jac, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos, bodyid, dofid, worldid) + jac, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos, bodyid, dofid, worldid) wp.atomic_add(qfrc_gravcomp_out[worldid], dofid, wp.dot(jac, force)) @@ -526,9 +526,9 @@ def _fluid(m: Model, d: Data): @wp.kernel def _qfrc_passive( # Model: - opt_has_fluid: bool, jnt_actgravcomp: wp.array(dtype=int), dof_jntid: wp.array(dtype=int), + has_fluid: bool, # Data in: qfrc_spring_in: wp.array2d(dtype=float), qfrc_damper_in: wp.array2d(dtype=float), @@ -548,7 +548,7 @@ def _qfrc_passive( qfrc_passive += qfrc_gravcomp_in[worldid, dofid] # add fluid force - if opt_has_fluid: + if has_fluid: qfrc_passive += qfrc_fluid_in[worldid, dofid] qfrc_passive_out[worldid, dofid] = qfrc_passive @@ -703,7 +703,7 @@ def _flex_bending( force[i, x] -= flex_bending[edgeid, 16] * frc[i, x] for i in range(nvert): - bodyid = flex_vertbodyid[flex_vertadr[f] + v[i]] + bodyid = flex_vertbodyid[v[i]] for x in range(3): wp.atomic_add(qfrc_spring_out, worldid, body_dofadr[bodyid] + x, force[i, x]) @@ -826,16 +826,16 @@ def passive(m: Model, d: Data): outputs=[d.qfrc_gravcomp], ) - if m.opt.has_fluid: + if m.has_fluid: _fluid(m, d) wp.launch( _qfrc_passive, dim=(d.nworld, m.nv), inputs=[ - m.opt.has_fluid, m.jnt_actgravcomp, m.dof_jntid, + m.has_fluid, d.qfrc_spring, d.qfrc_damper, d.qfrc_gravcomp, 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 41267c73..befc9820 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py @@ -13,15 +13,17 @@ # limitations under the License. # ============================================================================== -from typing import Optional, Tuple +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_MAXVAL from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL 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 RenderContext from mujoco.mjx.third_party.mujoco_warp._src.types import vec6 wp.set_module_options({"enable_backward": False}) @@ -183,7 +185,7 @@ def _ray_triangle( @wp.func -def _ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: +def ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: """Returns the distance and normal at which a ray intersects with a plane.""" # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) @@ -207,7 +209,7 @@ def _ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp @wp.func -def _ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: +def ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: """Returns the distance and normal at which a ray intersects with a sphere.""" dif = pnt - pos @@ -224,11 +226,11 @@ def _ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> Tu @wp.func -def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: +def ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: """Returns the distance and normal at which a ray intersects with a capsule.""" # bounding sphere test ssz = size[0] + size[1] - dist_sphere, normal_sphere = _ray_sphere(pos, ssz * ssz, pnt, vec) + dist_sphere, normal_sphere = ray_sphere(pos, ssz * ssz, pnt, vec) if dist_sphere < 0: return -1.0, wp.vec3() @@ -248,8 +250,9 @@ def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: sol, xx = _ray_quad(a, b, c) part = 0 # -1: bottom, 0: cylinder, 1: top - # make sure round solution is between flat sides - if sol >= 0.0 and wp.abs(lpnt[2] + sol * vec[2]) <= size[1]: + # make sure round solution is between flat sides (must use local z component) + # TODO: We should add a test to catch this case. + if sol >= 0.0 and wp.abs(lpnt[2] + sol * lvec[2]) <= size[1]: if x < 0.0 or sol < x: x = sol @@ -260,7 +263,7 @@ def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: c = wp.dot(ldif, ldif) - sq_size0 _, xx = _ray_quad(a, b, c) - # accept only top half of sphere + # accept only top half of sphere (use local z component) for i in range(2): if xx[i] >= 0.0 and lpnt[2] + xx[i] * lvec[2] >= size[1]: if x < 0.0 or xx[i] < x: @@ -273,7 +276,7 @@ def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: c = wp.dot(ldif, ldif) - sq_size0 _, xx = _ray_quad(a, b, c) - # accept only bottom half of sphere + # accept only bottom half of sphere (use local z component) for i in range(2): if xx[i] >= 0.0 and lpnt[2] + xx[i] * lvec[2] <= -size[1]: if x < 0.0 or xx[i] < x: @@ -297,7 +300,7 @@ def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: @wp.func -def _ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: +def ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: """Returns the distance and normal at which a ray intersects with an ellipsoid.""" # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) @@ -328,11 +331,11 @@ def _ray_ellipsoid(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec @wp.func -def _ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: +def ray_cylinder(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, wp.vec3]: """Returns the distance and normal at which a ray intersects with a cylinder.""" # bounding sphere test ssz = size[0] * size[0] + size[1] * size[1] - dist_sphere, normal_sphere = _ray_sphere(pos, ssz, pnt, vec) + dist_sphere, normal_sphere = ray_sphere(pos, ssz, pnt, vec) if dist_sphere < 0: return -1.0, wp.vec3() @@ -392,13 +395,13 @@ _IFACE = wp.types.matrix((3, 2), dtype=int)(1, 2, 0, 2, 0, 1) @wp.func -def _ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, vec6, wp.vec3]: +def ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> Tuple[float, vec6, wp.vec3]: """Returns distance, per side information, and normal at which a ray intersects with a box.""" all = vec6(-1.0, -1.0, -1.0, -1.0, -1.0, -1.0) # bounding sphere test ssz = wp.dot(size, size) - dist_sphere, _ = _ray_sphere(pos, ssz, pnt, vec) + dist_sphere, _ = ray_sphere(pos, ssz, pnt, vec) if dist_sphere < 0: return -1.0, all, wp.vec3() @@ -446,7 +449,7 @@ def _ray_box(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.v @wp.func -def _ray_hfield( +def ray_hfield( # Model: geom_type: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), @@ -487,10 +490,10 @@ def _ray_hfield( top_pos = pos + mat_col * top_scale # init: intersection with base box - x, _, normal_base = _ray_box(base_pos, mat, base_size, pnt, vec) + x, _, normal_base = ray_box(base_pos, mat, base_size, pnt, vec) # check top box: done if no intersection - top_intersect, all, normal_top = _ray_box(top_pos, mat, top_size, pnt, vec) + top_intersect, all, normal_top = ray_box(top_pos, mat, top_size, pnt, vec) if top_intersect < 0.0: return x, normal_base @@ -627,10 +630,16 @@ def ray_mesh( data_id: int, pos: wp.vec3, mat: wp.mat33, + size: wp.vec3, pnt: wp.vec3, vec: wp.vec3, ) -> Tuple[float, wp.vec3]: """Returns the distance and normal for ray mesh intersections.""" + # bounding box test + dist_box, _all, _normal = ray_box(pos, mat, size, pnt, vec) + if dist_box < 0.0: + return -1.0, wp.vec3() + pnt, vec = _ray_map(pos, mat, pnt, vec) # compute orthogonal basis vectors @@ -687,6 +696,87 @@ def ray_mesh( return x, normal +@wp.func +def ray_mesh_with_bvh( + # In: + mesh_bvh_id: wp.array(dtype=wp.uint64), + mesh_geom_id: int, + pos: wp.vec3, + mat: wp.mat33, + pnt: wp.vec3, + vec: wp.vec3, + max_t: float, +) -> Tuple[float, wp.vec3, float, float, int, int]: + """Returns intersection information for ray mesh intersections. + + Requires wp.Mesh be constructed and their ids to be passed. + """ + t = float(-1.0) + u = float(0.0) + v = float(0.0) + sign = float(0.0) + n = wp.vec3(0.0, 0.0, 0.0) + f = int(-1) + + lpnt, lvec = _ray_map(pos, mat, pnt, vec) + hit = wp.mesh_query_ray(mesh_bvh_id[mesh_geom_id], lpnt, lvec, max_t, t, u, v, sign, n, f) + + if hit and wp.dot(lvec, n) < 0.0: # Backface culling in local space + normal = mat @ n + normal = wp.normalize(normal) + return t, normal, u, v, f, mesh_geom_id + + return -1.0, wp.vec3(0.0, 0.0, 0.0), 0.0, 0.0, -1, -1 + + +@wp.func +def ray_mesh_with_bvh_anyhit( + # In: + mesh_bvh_id: wp.array(dtype=wp.uint64), + mesh_geom_id: int, + pos: wp.vec3, + mat: wp.mat33, + pnt: wp.vec3, + vec: wp.vec3, + max_t: float, +) -> bool: + """Returns True if there is any hit for ray mesh intersections. + + Requires wp.Mesh be constructed and their ids to be passed. This variant is useful + for shadow ray casts where the only goal is if there is any ray hit. + """ + lpnt, lvec = _ray_map(pos, mat, pnt, vec) + return wp.mesh_query_ray_anyhit(mesh_bvh_id[mesh_geom_id], lpnt, lvec, max_t) + + +@wp.func +def ray_flex_with_bvh( + # In: + bvh_id: wp.uint64, + group_root: int, + pnt: wp.vec3, + vec: wp.vec3, + max_t: float, +) -> Tuple[float, wp.vec3, float, float, int]: + """Returns intersection information for flex intersections. + + Requires wp.Mesh be constructed and their ids to be passed. Flex are already in world space. + """ + t = float(-1.0) + u = float(0.0) + v = float(0.0) + sign = float(0.0) + 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) + + if hit: + return t, n, u, v, f + + return -1.0, wp.vec3(0.0, 0.0, 0.0), 0.0, 0.0, -1 + + @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. @@ -695,17 +785,17 @@ def ray_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.v """ # TODO(team): static loop unrolling to remove unnecessary branching if geomtype == GeomType.PLANE: - return _ray_plane(pos, mat, size, pnt, vec) + return ray_plane(pos, mat, size, pnt, vec) elif geomtype == GeomType.SPHERE: - return _ray_sphere(pos, size[0] * size[0], pnt, vec) + return ray_sphere(pos, size[0] * size[0], pnt, vec) elif geomtype == GeomType.CAPSULE: - return _ray_capsule(pos, mat, size, pnt, vec) + return ray_capsule(pos, mat, size, pnt, vec) elif geomtype == GeomType.ELLIPSOID: - return _ray_ellipsoid(pos, mat, size, pnt, vec) + return ray_ellipsoid(pos, mat, size, pnt, vec) elif geomtype == GeomType.CYLINDER: - return _ray_cylinder(pos, mat, size, pnt, vec) + return ray_cylinder(pos, mat, size, pnt, vec) elif geomtype == GeomType.BOX: - dist, _, normal = _ray_box(pos, mat, size, pnt, vec) + dist, _, normal = ray_box(pos, mat, size, pnt, vec) return dist, normal else: return -1.0, wp.vec3() @@ -771,11 +861,12 @@ def _ray_geom_mesh( geom_dataid[geomid], pos, mat, + geom_size[worldid % geom_size.shape[0], geomid], pnt, vec, ) elif type == GeomType.HFIELD: - return _ray_hfield( + return ray_hfield( geom_type, geom_dataid, hfield_size, @@ -836,7 +927,7 @@ def _ray( num_threads = wp.block_dim() - min_dist = float(wp.inf) + min_dist = float(MJ_MAXVAL) min_geomid = int(-1) min_normal = wp.vec3() @@ -874,9 +965,9 @@ def _ray( geomid, ) if dist < 0: - dist = wp.inf + dist = MJ_MAXVAL else: - dist = wp.inf + dist = MJ_MAXVAL normal = wp.vec3() tile_dist = wp.tile(dist) @@ -891,7 +982,162 @@ def _ray( min_geomid = tile_geomid[local_min_geomid[0]] min_normal = tile_normal[local_min_geomid[0]] - if wp.isinf(min_dist): + if min_dist >= MJ_MAXVAL: + dist_out[worldid, rayid] = -1.0 + else: + dist_out[worldid, rayid] = min_dist + geomid_out[worldid, rayid] = min_geomid + normal_out[worldid, rayid] = min_normal + + +@wp.func +def _ray_geom_mesh_bvh( + # Model: + body_weldid: wp.array(dtype=int), + geom_type: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_matid: wp.array2d(dtype=int), + geom_group: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + geom_rgba: wp.array2d(dtype=wp.vec4), + mat_rgba: wp.array2d(dtype=wp.vec4), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + worldid: int, + pnt: wp.vec3, + vec: wp.vec3, + geomgroup: vec6, + flg_static: bool, + bodyexclude: int, + geomid: int, + mesh_bvh_id: wp.array(dtype=wp.uint64), + hfield_bvh_id: wp.array(dtype=wp.uint64), + min_dist: float, +) -> Tuple[float, wp.vec3]: + if not _ray_eliminate( + body_weldid, + geom_bodyid, + geom_matid[worldid % geom_matid.shape[0]], + geom_group, + geom_rgba[worldid % geom_rgba.shape[0]], + mat_rgba[worldid % mat_rgba.shape[0]], + geomid, + geomgroup, + flg_static, + bodyexclude, + ): + pos = geom_xpos_in[worldid, geomid] + mat = geom_xmat_in[worldid, geomid] + gtype = geom_type[geomid] + + if gtype == GeomType.MESH or gtype == GeomType.HFIELD: + bvh_ids = mesh_bvh_id if gtype == GeomType.MESH else hfield_bvh_id + t, n, u, v, f, geom_mesh_id = ray_mesh_with_bvh( + bvh_ids, + geom_dataid[geomid], + pos, + mat, + pnt, + vec, + min_dist, + ) + if t >= 0.0 and t < min_dist: + return t, n + else: + return ray_geom( + pos, + mat, + geom_size[worldid % geom_size.shape[0], geomid], + pnt, + vec, + gtype, + ) + + return -1.0, wp.vec3() + + +@wp.kernel +def _ray_bvh( + # Model: + ngeom: int, + body_weldid: wp.array(dtype=int), + geom_type: wp.array(dtype=int), + geom_bodyid: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_matid: wp.array2d(dtype=int), + geom_group: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + geom_rgba: wp.array2d(dtype=wp.vec4), + mat_rgba: wp.array2d(dtype=wp.vec4), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + pnt: wp.array2d(dtype=wp.vec3), + vec: wp.array2d(dtype=wp.vec3), + geomgroup: vec6, + flg_static: bool, + bodyexclude: wp.array(dtype=int), + bvh_id: wp.uint64, + group_root: wp.array(dtype=int), + enabled_geom_ids: wp.array(dtype=int), + mesh_bvh_id: wp.array(dtype=wp.uint64), + hfield_bvh_id: wp.array(dtype=wp.uint64), + # Out: + dist_out: wp.array2d(dtype=float), + geomid_out: wp.array2d(dtype=int), + normal_out: wp.array2d(dtype=wp.vec3), +): + worldid, rayid = wp.tid() + + ray_origin = pnt[worldid, rayid] + ray_dir = vec[worldid, rayid] + body_exclude = bodyexclude[rayid] + + min_dist = float(MJ_MAXVAL) + min_geomid = int(-1) + min_normal = wp.vec3() + + query = wp.bvh_query_ray(bvh_id, ray_origin, ray_dir, group_root[worldid]) + bounds_nr = int(0) + + while wp.bvh_query_next(query, bounds_nr, min_dist): + bvh_local = bounds_nr - (worldid * ngeom) + geomid = enabled_geom_ids[bvh_local] + + dist, normal = _ray_geom_mesh_bvh( + body_weldid, + geom_type, + geom_bodyid, + geom_dataid, + geom_matid, + geom_group, + geom_size, + geom_rgba, + mat_rgba, + geom_xpos_in, + geom_xmat_in, + worldid, + pnt[worldid, rayid], + vec[worldid, rayid], + geomgroup, + flg_static, + body_exclude, + geomid, + mesh_bvh_id, + hfield_bvh_id, + min_dist, + ) + + if dist >= 0.0 and dist < min_dist: + min_dist = dist + min_geomid = geomid + min_normal = normal + + if min_dist >= MJ_MAXVAL: dist_out[worldid, rayid] = -1.0 else: dist_out[worldid, rayid] = min_dist @@ -904,9 +1150,10 @@ def ray( d: Data, pnt: wp.array2d(dtype=wp.vec3), vec: wp.array2d(dtype=wp.vec3), - geomgroup: Optional[vec6] = None, + geomgroup: vec6 | None = None, flg_static: bool = True, bodyexclude: int = -1, + rc: RenderContext | None = None, ) -> Tuple[wp.array, wp.array, wp.array]: """Returns the distance at which rays intersect with primitive geoms. @@ -915,9 +1162,11 @@ def ray( d: The data object containing the current state and output arrays (device). pnt: Ray origin points. vec: Ray directions. - geomgroup: Group inclusion/exclusion mask. If all are wp.inf, ignore. + geomgroup: Group inclusion/exclusion mask. flg_static: If True, allows rays to intersect with static geoms. bodyexclude: Ignore geoms on specified body id (-1 to disable). + rc: Optional Render context containing BVH information for BVH accelerated ray + intersections. Returns: Distances from ray origins to geom surfaces, IDs of intersected geoms (-1 if none), @@ -935,7 +1184,7 @@ def ray( ray_geomid = wp.empty((d.nworld, 1), dtype=int) ray_normal = wp.empty((d.nworld, 1), dtype=wp.vec3) - rays(m, d, pnt, vec, geomgroup, flg_static, ray_bodyexclude, ray_dist, ray_geomid, ray_normal) + rays(m, d, pnt, vec, geomgroup, flg_static, ray_bodyexclude, ray_dist, ray_geomid, ray_normal, rc) return ray_dist, ray_geomid, ray_normal @@ -951,6 +1200,7 @@ def rays( dist: wp.array2d(dtype=float), geomid: wp.array2d(dtype=int), normal: wp.array2d(dtype=wp.vec3), + rc: RenderContext | None = None, ): """Ray intersection for multiple worlds and multiple rays. @@ -968,41 +1218,75 @@ def rays( geomid: Output array for IDs of intersected geoms, shape (nworld, nray). -1 indicates no intersection. normal: Output array for normals at intersection points, shape (nworld, nray). + rc: Optional Render context containing BVH information for BVH accelerated ray + intersections. """ - wp.launch_tiled( - _ray, - dim=(d.nworld, pnt.shape[1]), - inputs=[ - m.ngeom, - m.nmeshface, - m.body_weldid, - m.geom_type, - m.geom_bodyid, - m.geom_dataid, - m.geom_matid, - m.geom_group, - m.geom_size, - m.geom_rgba, - m.mesh_vertadr, - m.mesh_faceadr, - m.mesh_vert, - m.mesh_face, - m.hfield_size, - m.hfield_nrow, - m.hfield_ncol, - m.hfield_adr, - m.hfield_data, - m.mat_rgba, - d.geom_xpos, - d.geom_xmat, - pnt, - vec, - geomgroup, - flg_static, - bodyexclude, - dist, - geomid, - normal, - ], - block_dim=m.block_dim.ray, - ) + # TODO: Investigate building rc if none and removing the non-accelerated path + if rc is None: + wp.launch_tiled( + _ray, + dim=(d.nworld, pnt.shape[1]), + inputs=[ + m.ngeom, + m.nmeshface, + m.body_weldid, + m.geom_type, + m.geom_bodyid, + m.geom_dataid, + m.geom_matid, + m.geom_group, + m.geom_size, + m.geom_rgba, + m.mesh_vertadr, + m.mesh_faceadr, + m.mesh_vert, + m.mesh_face, + m.hfield_size, + m.hfield_nrow, + m.hfield_ncol, + m.hfield_adr, + m.hfield_data, + m.mat_rgba, + d.geom_xpos, + d.geom_xmat, + pnt, + vec, + geomgroup, + flg_static, + bodyexclude, + dist, + geomid, + normal, + ], + block_dim=m.block_dim.ray, + ) + else: + wp.launch( + _ray_bvh, + dim=(d.nworld, pnt.shape[1]), + inputs=[ + rc.bvh_ngeom, + m.body_weldid, + m.geom_type, + m.geom_bodyid, + m.geom_dataid, + m.geom_matid, + m.geom_group, + m.geom_size, + m.geom_rgba, + m.mat_rgba, + d.geom_xpos, + d.geom_xmat, + pnt, + vec, + geomgroup, + flg_static, + bodyexclude, + rc.bvh_id, + rc.group_root, + rc.enabled_geom_ids, + rc.mesh_bvh_id, + rc.hfield_bvh_id, + ], + outputs=[dist, geomid, normal], + ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py new file mode 100644 index 00000000..43145ff2 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render.py @@ -0,0 +1,696 @@ +# 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. +# ============================================================================== + +from typing import Tuple + +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_box +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_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 +from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_sphere +from mujoco.mjx.third_party.mujoco_warp._src.render_util import compute_ray +from mujoco.mjx.third_party.mujoco_warp._src.render_util import pack_rgba_to_uint32 +from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import Data +from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType +from mujoco.mjx.third_party.mujoco_warp._src.types import Model +from mujoco.mjx.third_party.mujoco_warp._src.types import RenderContext +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: + geom_type: wp.array(dtype=int), + mesh_faceadr: wp.array(dtype=int), + # In: + geom_id: int, + tex_repeat: wp.vec2, + tex: TEXTURE_DTYPE, + pos: wp.vec3, + rot: wp.mat33, + mesh_facetexcoord: wp.array(dtype=wp.vec3i), + mesh_texcoord: wp.array(dtype=wp.vec2), + mesh_texcoord_offsets: wp.array(dtype=int), + hit_point: wp.vec3, + bary_u: float, + bary_v: float, + f: int, + mesh_id: int, +) -> wp.vec3: + uv = wp.vec2(0.0, 0.0) + + if geom_type[geom_id] == GeomType.PLANE: + local = wp.transpose(rot) @ (hit_point - pos) + uv = wp.vec2(local[0], local[1]) + + if geom_type[geom_id] == GeomType.MESH: + if f < 0 or mesh_id < 0: + return wp.vec3(0.0, 0.0, 0.0) + + face_adr = mesh_faceadr[mesh_id] + f + uv0 = mesh_texcoord[mesh_texcoord_offsets[mesh_id] + mesh_facetexcoord[face_adr][0]] + uv1 = mesh_texcoord[mesh_texcoord_offsets[mesh_id] + mesh_facetexcoord[face_adr][1]] + uv2 = mesh_texcoord[mesh_texcoord_offsets[mesh_id] + mesh_facetexcoord[face_adr][2]] + uv = uv0 * bary_u + uv1 * bary_v + uv2 * (1.0 - bary_u - bary_v) + + u = uv[0] * tex_repeat[0] + v = uv[1] * tex_repeat[1] + u = u - wp.floor(u) + v = v - wp.floor(v) + tex_color = wp.texture_sample(tex, wp.vec2(u, v), dtype=wp.vec4) + return wp.vec3(tex_color[0], tex_color[1], tex_color[2]) + + +# TODO: Investigate combining cast_ray and cast_ray_first_hit +@wp.func +def cast_ray( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + bvh_id: wp.uint64, + group_root: int, + world_id: int, + 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), + ray_origin_world: wp.vec3, + ray_dir_world: wp.vec3, +) -> Tuple[int, float, wp.vec3, float, float, int, int]: + dist = float(MJ_MAXVAL) + normal = wp.vec3(0.0, 0.0, 0.0) + geom_id = int(-1) + bary_u = float(0.0) + bary_v = float(0.0) + face_idx = int(-1) + geom_mesh_id = int(-1) + + query = wp.bvh_query_ray(bvh_id, ray_origin_world, ray_dir_world, group_root) + bounds_nr = int(0) + + 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] + + 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) + + # TODO: Investigate branch elimination with static loop unrolling + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + dist, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + dist, + ) + + if d >= 0.0 and d < dist: + dist = d + normal = n + geom_id = gi + bary_u = u + bary_v = v + face_idx = f + geom_mesh_id = hit_mesh_id + + return geom_id, dist, normal, bary_u, bary_v, face_idx, geom_mesh_id + + +@wp.func +def cast_ray_first_hit( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + bvh_id: wp.uint64, + group_root: int, + world_id: int, + 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), + ray_origin_world: wp.vec3, + ray_dir_world: wp.vec3, + max_dist: float, +) -> bool: + """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) + + 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] + + # TODO: Investigate branch elimination with static loop unrolling + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + max_dist, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + ray_origin_world, + ray_dir_world, + ) + if geom_type[gi] == 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], + 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 + + return False + + +@wp.func +def compute_lighting( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + # In: + use_shadows: bool, + bvh_id: wp.uint64, + group_root: int, + bvh_ngeom: int, + enabled_geom_ids: wp.array(dtype=int), + world_id: int, + mesh_bvh_id: wp.array(dtype=wp.uint64), + hfield_bvh_id: wp.array(dtype=wp.uint64), + lightactive: bool, + lighttype: int, + lightcastshadow: bool, + lightpos: wp.vec3, + lightdir: wp.vec3, + normal: wp.vec3, + hitpoint: wp.vec3, +) -> float: + light_contribution = float(0.0) + + # TODO: We should probably only be looping over active lights + # in the first place with a static loop of enabled light idx? + if not lightactive: + return light_contribution + + L = wp.vec3(0.0, 0.0, 0.0) + dist_to_light = float(MJ_MAXVAL) + attenuation = float(1.0) + + if lighttype == 1: # directional light + L = wp.normalize(-lightdir) + else: + L, dist_to_light = math.normalize_with_norm(lightpos - hitpoint) + attenuation = 1.0 / (1.0 + 0.02 * dist_to_light * dist_to_light) + if lighttype == 0: # spot light + spot_dir = wp.normalize(lightdir) + cos_theta = wp.dot(-L, spot_dir) + spot_factor = wp.min(1.0, wp.max(0.0, (cos_theta - 0.85) / (0.95 - 0.85))) + attenuation = attenuation * spot_factor + + ndotl = wp.max(0.0, wp.dot(normal, L)) + if ndotl == 0.0: + return light_contribution + + visible = float(1.0) + + if use_shadows and lightcastshadow: + # Nudge the origin slightly along the surface normal to avoid + # self-intersection when casting shadow rays + eps = 1.0e-4 + shadow_origin = hitpoint + normal * eps + # Distance-limited shadows: cap by dist_to_light (for non-directional) + max_t = float(dist_to_light - 1.0e-3) + if lighttype == 1: # directional light + max_t = float(1.0e8) + + shadow_hit = cast_ray_first_hit( + geom_type, + geom_dataid, + geom_size, + geom_xpos_in, + geom_xmat_in, + bvh_id, + group_root, + world_id, + bvh_ngeom, + enabled_geom_ids, + mesh_bvh_id, + hfield_bvh_id, + shadow_origin, + L, + max_t, + ) + + if shadow_hit: + visible = 0.3 + + return ndotl * attenuation * visible + + +@event_scope +def render(m: Model, d: Data, rc: RenderContext): + """Render the current frame. + + Outputs are stored in buffers within the render context. + + Args: + m: The model on device. + d: The data on device. + rc: The render context on device. + """ + rc.rgb_data.fill_(rc.background_color) + rc.depth_data.fill_(0.0) + + @wp.kernel(module="unique", enable_backward=False) + def _render_megakernel( + # Model: + geom_type: wp.array(dtype=int), + geom_dataid: wp.array(dtype=int), + geom_matid: wp.array2d(dtype=int), + geom_size: wp.array2d(dtype=wp.vec3), + geom_rgba: wp.array2d(dtype=wp.vec4), + cam_projection: wp.array(dtype=int), + cam_fovy: wp.array2d(dtype=float), + cam_sensorsize: wp.array(dtype=wp.vec2), + cam_intrinsic: wp.array2d(dtype=wp.vec4), + light_type: wp.array2d(dtype=int), + light_castshadow: wp.array2d(dtype=bool), + light_active: wp.array2d(dtype=bool), + mesh_faceadr: wp.array(dtype=int), + mat_texid: wp.array3d(dtype=int), + mat_texrepeat: wp.array2d(dtype=wp.vec2), + mat_rgba: wp.array2d(dtype=wp.vec4), + # Data in: + geom_xpos_in: wp.array2d(dtype=wp.vec3), + geom_xmat_in: wp.array2d(dtype=wp.mat33), + cam_xpos_in: wp.array2d(dtype=wp.vec3), + cam_xmat_in: wp.array2d(dtype=wp.mat33), + light_xpos_in: wp.array2d(dtype=wp.vec3), + light_xdir_in: wp.array2d(dtype=wp.vec3), + # In: + nrender: int, + use_shadows: bool, + bvh_ngeom: 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), + render_rgb: wp.array(dtype=bool), + render_depth: 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), + enabled_geom_ids: wp.array(dtype=int), + mesh_bvh_id: wp.array(dtype=wp.uint64), + mesh_facetexcoord: wp.array(dtype=wp.vec3i), + mesh_texcoord: wp.array(dtype=wp.vec2), + 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), + # Out: + rgb_out: wp.array2d(dtype=wp.uint32), + depth_out: wp.array2d(dtype=float), + ): + world_idx, ray_idx = wp.tid() + + # Map global ray_idx -> (cam_idx, ray_idx_local) using cumulative sizes + cam_idx = int(-1) + ray_idx_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: + cam_idx = i + ray_idx_local = ray_idx - accum + break + accum += num_i + if cam_idx == -1 or ray_idx_local < 0: + return + + if not render_rgb[cam_idx] and not render_depth[cam_idx]: + return + + # Map active camera index to MuJoCo camera ID + mujoco_cam_id = cam_id_map[cam_idx] + + if wp.static(rc.ray is None): + 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 + ray_dir_local_cam = compute_ray( + cam_projection[mujoco_cam_id], + cam_fovy[world_idx % cam_fovy.shape[0], mujoco_cam_id], + cam_sensorsize[mujoco_cam_id], + cam_intrinsic[world_idx % cam_intrinsic.shape[0], mujoco_cam_id], + img_w, + img_h, + px, + py, + wp.static(rc.znear), + ) + else: + ray_dir_local_cam = ray[ray_idx] + + 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] + + geom_id, dist, normal, u, v, f, mesh_id = cast_ray( + geom_type, + geom_dataid, + geom_size, + geom_xpos_in, + geom_xmat_in, + bvh_id, + group_root[world_idx], + world_idx, + bvh_ngeom, + enabled_geom_ids, + mesh_bvh_id, + hfield_bvh_id, + 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 + + # Early Out + if geom_id == -1: + return + + if render_depth[cam_idx]: + depth_out[world_idx, depth_adr[cam_idx] + ray_idx_local] = dist + + if not render_rgb[cam_idx]: + return + + # Shade the pixel + 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] + else: + color = mat_rgba[world_idx % mat_rgba.shape[0], geom_matid[world_idx % 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] + if mat_id >= 0: + tex_id = mat_texid[world_idx % 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], + textures[tex_id], + geom_xpos_in[world_idx, geom_id], + geom_xmat_in[world_idx, geom_id], + mesh_facetexcoord, + mesh_texcoord, + mesh_texcoord_offsets, + hit_point, + u, + v, + f, + mesh_id, + ) + base_color = wp.cw_mul(base_color, tex_color) + + len_n = wp.length(normal) + n = normal if len_n > 0.0 else wp.vec3(0.0, 0.0, 1.0) + n = wp.normalize(n) + hemispheric = 0.5 * (n[2] + 1.0) + ambient_color = wp.vec3(0.4, 0.4, 0.45) * hemispheric + wp.vec3(0.1, 0.1, 0.12) * (1.0 - hemispheric) + result = 0.5 * wp.cw_mul(base_color, ambient_color) + + # Apply lighting and shadows + for l in range(wp.static(m.nlight)): + light_contribution = compute_lighting( + geom_type, + geom_dataid, + geom_size, + geom_xpos_in, + geom_xmat_in, + use_shadows, + bvh_id, + group_root[world_idx], + bvh_ngeom, + enabled_geom_ids, + world_idx, + 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], + normal, + hit_point, + ) + result = result + base_color * light_contribution + + 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( + hit_color[0] * 255.0, + hit_color[1] * 255.0, + hit_color[2] * 255.0, + 255.0, + ) + + wp.launch( + kernel=_render_megakernel, + dim=(d.nworld, rc.total_rays), + inputs=[ + m.geom_type, + m.geom_dataid, + m.geom_matid, + m.geom_size, + m.geom_rgba, + m.cam_projection, + m.cam_fovy, + m.cam_sensorsize, + m.cam_intrinsic, + m.light_type, + m.light_castshadow, + m.light_active, + m.mesh_faceadr, + m.mat_texid, + m.mat_texrepeat, + m.mat_rgba, + d.geom_xpos, + d.geom_xmat, + d.cam_xpos, + d.cam_xmat, + d.light_xpos, + d.light_xdir, + rc.nrender, + rc.use_shadows, + rc.bvh_ngeom, + rc.cam_res, + rc.cam_id_map, + rc.ray, + rc.rgb_adr, + rc.depth_adr, + rc.render_rgb, + rc.render_depth, + rc.bvh_id, + rc.group_root, + rc.flex_bvh_id, + rc.flex_group_root, + rc.enabled_geom_ids, + rc.mesh_bvh_id, + rc.mesh_facetexcoord, + rc.mesh_texcoord, + rc.mesh_texcoord_offsets, + rc.hfield_bvh_id, + rc.flex_rgba, + rc.textures, + ], + outputs=[ + rc.rgb_data, + rc.depth_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 new file mode 100644 index 00000000..260cbc77 --- /dev/null +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/render_util.py @@ -0,0 +1,130 @@ +# 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. +# ============================================================================== + +import mujoco +import warp as wp + +from mujoco.mjx.third_party.mujoco_warp._src.types import ProjectionType + +wp.set_module_options({"enable_backward": False}) + + +@wp.kernel +def _convert_texture_data( + # In: + width: int, + adr: int, + nc: int, + tex_data_in: wp.array(dtype=wp.uint8), + # Out: + tex_data_out: wp.array3d(dtype=float), +): + """Convert uint8 texture data to vec4 format for efficient sampling.""" + x, y = wp.tid() + offset = adr + (y * width + x) * nc + r = tex_data_in[offset + 0] if nc > 0 else wp.uint8(0) + g = tex_data_in[offset + 1] if nc > 1 else wp.uint8(0) + b = tex_data_in[offset + 2] if nc > 2 else wp.uint8(0) + a = wp.uint8(255) + + tex_data_out[y, x, 0] = float(r) * wp.static(1.0 / 255.0) + tex_data_out[y, x, 1] = float(g) * wp.static(1.0 / 255.0) + tex_data_out[y, x, 2] = float(b) * wp.static(1.0 / 255.0) + tex_data_out[y, x, 3] = float(a) * wp.static(1.0 / 255.0) + + +def create_warp_texture(mjm: mujoco.MjModel, tex_id: int) -> wp.array: + """Create a Warp texture from a MuJoCo model texture data.""" + tex_adr = mjm.tex_adr[tex_id] + tex_width = mjm.tex_width[tex_id] + tex_height = mjm.tex_height[tex_id] + nchannel = mjm.tex_nchannel[tex_id] + tex_data = wp.zeros((tex_height, tex_width, 4), dtype=float) + + wp.launch( + _convert_texture_data, + dim=(tex_width, tex_height), + inputs=[tex_width, tex_adr, nchannel, wp.array(mjm.tex_data, dtype=wp.uint8)], + outputs=[tex_data], + ) + return wp.Texture2D(tex_data, filter_mode=wp.TextureFilterMode.LINEAR) + + +@wp.func +def compute_ray( + # In: + projection: int, + fovy: float, + sensorsize: wp.vec2, + intrinsic: wp.vec4, + img_w: int, + img_h: int, + px: int, + py: int, + znear: float, +) -> wp.vec3: + """Compute ray direction for a pixel with per-world camera parameters. + + This combines _camera_frustum_bounds and build_primary_rays logic for use + inside a kernel when camera parameters are batched/randomized across worlds. + """ + if projection == ProjectionType.ORTHOGRAPHIC: + return wp.vec3(0.0, 0.0, -1.0) + + aspect = float(img_w) / float(img_h) + sensor_h = sensorsize[1] + + # Check if we have intrinsics (sensorsize[1] != 0) + if sensor_h != 0.0: + fx = intrinsic[0] + fy = intrinsic[1] + cx = intrinsic[2] + cy = intrinsic[3] + sensor_w = sensorsize[0] + + target_aspect = float(img_w) / float(img_h) + sensor_aspect = sensor_w / sensor_h + if target_aspect > sensor_aspect: + sensor_h = sensor_w / target_aspect + elif target_aspect < sensor_aspect: + sensor_w = sensor_h * target_aspect + + inv_fx_znear = znear / fx + inv_fy_znear = znear / fy + left = -inv_fx_znear * (sensor_w * 0.5 - cx) + right = inv_fx_znear * (sensor_w * 0.5 + cx) + top = inv_fy_znear * (sensor_h * 0.5 - cy) + bottom = -inv_fy_znear * (sensor_h * 0.5 + cy) + else: + fovy_rad = fovy * wp.static(wp.pi / 180.0) + half_height = znear * wp.tan(0.5 * fovy_rad) + half_width = half_height * aspect + left = -half_width + right = half_width + top = half_height + bottom = -half_height + + u = (float(px) + 0.5) / float(img_w) + v = (float(py) + 0.5) / float(img_h) + x = left + (right - left) * u + y = top + (bottom - top) * v + + return wp.normalize(wp.vec3(x, y, -znear)) + + +@wp.func +def pack_rgba_to_uint32(r: float, g: float, b: float, a: float) -> wp.uint32: + """Pack RGBA values into a single uint32 for efficient memory access.""" + return wp.uint32((int(a) << int(24)) | (int(r) << int(16)) | (int(g) << int(8)) | int(b)) 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 76b29385..a152f468 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -23,6 +23,7 @@ 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_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 @@ -42,7 +43,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import vec8i from mujoco.mjx.third_party.mujoco_warp._src.util_misc import inside_geom from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope -from mujoco.mjx.third_party.mujoco_warp._src.warp_util import nested_kernel wp.set_module_options({"enable_backward": False}) @@ -124,10 +124,10 @@ def _magnetometer( @wp.func def _cam_projection( # Model: - cam_fovy: wp.array(dtype=float), + cam_fovy: wp.array2d(dtype=float), cam_resolution: wp.array(dtype=wp.vec2i), cam_sensorsize: wp.array(dtype=wp.vec2), - cam_intrinsic: wp.array(dtype=wp.vec4), + cam_intrinsic: wp.array2d(dtype=wp.vec4), # Data in: site_xpos_in: wp.array2d(dtype=wp.vec3), cam_xpos_in: wp.array2d(dtype=wp.vec3), @@ -138,8 +138,8 @@ def _cam_projection( refid: int, ) -> wp.vec2: sensorsize = cam_sensorsize[refid] - intrinsic = cam_intrinsic[refid] - fovy = cam_fovy[refid] + intrinsic = cam_intrinsic[worldid % cam_intrinsic.shape[0], refid] + fovy = cam_fovy[worldid % cam_fovy.shape[0], refid] res = cam_resolution[refid] target_xpos = site_xpos_in[worldid, objid] @@ -470,10 +470,10 @@ def _sensor_pos( site_quat: wp.array2d(dtype=wp.quat), cam_bodyid: wp.array(dtype=int), cam_quat: wp.array2d(dtype=wp.quat), - cam_fovy: wp.array(dtype=float), + cam_fovy: wp.array2d(dtype=float), cam_resolution: wp.array(dtype=wp.vec2i), cam_sensorsize: wp.array(dtype=wp.vec2), - cam_intrinsic: wp.array(dtype=wp.vec4), + cam_intrinsic: wp.array2d(dtype=wp.vec4), sensor_type: wp.array(dtype=int), sensor_datatype: wp.array(dtype=int), sensor_objtype: wp.array(dtype=int), @@ -782,7 +782,7 @@ def sensor_pos(m: Model, d: Data): d, rangefinder_pnt, rangefinder_vec, - vec6(wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf), + vec6(MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL, MJ_MAXVAL), True, m.sensor_rangefinder_bodyid, rangefinder_dist, @@ -797,8 +797,7 @@ def sensor_pos(m: Model, d: Data): energy_vel(m, d) # collision sensors (distance, normal, fromto) - sensor_collision = wp.empty((d.nworld, m.nsensorcollision, 8, 7), dtype=float) - sensor_collision.fill_(1.0e32) + sensor_collision = wp.full((d.nworld, m.nsensorcollision, 8, 7), 1.0e32, dtype=float) if m.nsensorcollision: wp.launch( _sensor_collision, @@ -1787,38 +1786,12 @@ def _sensor_acc( nmatch = sensor_contact_nmatch_in[worldid, contactsensorid] if reduce == 3: # netforce - # compute point: force-weighted centroid of contact position + # Single-pass computation: first compute centroid, then wrench about centroid + # Pass 1: compute force-weighted centroid of contact positions net_pos = wp.vec3(0.0) - total_force_magnitude = float(0.0) - - for i in range(nmatch): - cid = sensor_contact_matchid_in[worldid, contactsensorid, i] - - contact_forcetorque = support.contact_force_fn( - opt_cone, - contact_frame_in, - contact_friction_in, - contact_dim_in, - contact_efc_address_in, - efc_force_in, - njmax_in, - nacon_in, - worldid, - cid, - False, - ) - - weight = wp.norm_l2(wp.spatial_top(contact_forcetorque)) - net_pos += weight * contact_pos_in[cid] - total_force_magnitude += weight - - net_pos /= wp.max(total_force_magnitude, MJ_MINVAL) - - # TODO(team): iterate over matches once - - # compute total wrench about point, in the global frame net_force = wp.vec3(0.0) net_torque = wp.vec3(0.0) + total_force_magnitude = float(0.0) for i in range(nmatch): cid = sensor_contact_matchid_in[worldid, contactsensorid, i] @@ -1837,8 +1810,15 @@ def _sensor_acc( cid, False, ) - contact_forcetorque *= dir + # Accumulate for centroid computation (unsigned force magnitude) + weight = wp.norm_l2(wp.spatial_top(contact_forcetorque)) + contact_pos = contact_pos_in[cid] + net_pos += weight * contact_pos + total_force_magnitude += weight + + # Apply direction and transform to global frame + contact_forcetorque *= dir force_local = wp.spatial_top(contact_forcetorque) torque_local = wp.spatial_bottom(contact_forcetorque) @@ -1848,12 +1828,18 @@ def _sensor_acc( force_global = frameT @ force_local torque_global = frameT @ torque_local - # add to total force, torque + # Accumulate force and torque (about origin for now) net_force += force_global net_torque += torque_global + # Accumulate moment contribution: will adjust after centroid is computed + net_torque += wp.cross(contact_pos, force_global) - # add induced moment: torque += (pos - point) x force - net_torque += wp.cross(contact_pos_in[cid] - net_pos, force_global) + # Finalize centroid + net_pos /= wp.max(total_force_magnitude, MJ_MINVAL) + + # Adjust torque: subtract moment from centroid (since we accumulated about origin) + # torque_about_centroid = torque_about_origin - centroid x total_force + net_torque -= wp.cross(net_pos, net_force) adr_slot = adr @@ -1888,7 +1874,8 @@ def _sensor_acc( out[adr_slot + 1] = 1.0 out[adr_slot + 2] = 0.0 else: - for i in range(wp.min(nmatch, num)): + nslots = wp.min(nmatch, num) + for i in range(nslots): # sorted contact id cid = sensor_contact_matchid_in[worldid, contactsensorid, i] @@ -2165,7 +2152,15 @@ 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, plugin, plugin_attr, contact_type, geom_size[worldid, geom], plugin_id, mesh_id + oct_child, + oct_aabb, + oct_coeff, + plugin, + plugin_attr, + contact_type, + geom_size[worldid % geom_size.shape[0], geom], + plugin_id, + mesh_id, ) depth = wp.min(sdf(contact_type, lpos, plugin_attributes, plugin_index, volume_data, mesh_data), 0.0) @@ -2357,16 +2352,16 @@ def _contact_match( @cache_kernel def _contact_sort(maxmatch: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def contact_sort( # Model: sensor_intprm: wp.array2d(dtype=int), sensor_contact_adr: wp.array(dtype=int), - # Data in: + # In: sensor_contact_nmatch_in: wp.array2d(dtype=int), sensor_contact_matchid_in: wp.array3d(dtype=int), sensor_contact_criteria_in: wp.array3d(dtype=float), - # Data out: + # Out: sensor_contact_matchid_out: wp.array3d(dtype=int), ): worldid, contactsensorid = wp.tid() @@ -2819,13 +2814,13 @@ def energy_pos(m: Model, d: Data): @cache_kernel def _energy_vel_kinetic(nv: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def energy_vel_kinetic( # Data in: qvel_in: wp.array2d(dtype=float), # In: Mqvel: wp.array2d(dtype=float), - # Out: + # Data out: energy_out: wp.array(dtype=wp.vec2), ): worldid = wp.tid() @@ -2849,12 +2844,13 @@ def energy_vel(m: Model, d: Data): # kinetic energy: 0.5 * qvel.T @ M @ qvel # M @ qvel - support.mul_m(m, d, d.efc.mv, d.qvel) + mv = wp.zeros((d.nworld, m.nv), dtype=float) + support.mul_m(m, d, mv, d.qvel) wp.launch_tiled( _energy_vel_kinetic(m.nv), dim=d.nworld, - inputs=[d.qvel, d.efc.mv], + inputs=[d.qvel, mv], outputs=[d.energy], block_dim=m.block_dim.energy_vel_kinetic, ) 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 2bc68aff..ff5ddbe7 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py @@ -19,11 +19,13 @@ 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 Model from mujoco.mjx.third_party.mujoco_warp._src.types import ObjType @@ -35,13 +37,12 @@ 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.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 nested_kernel wp.set_module_options({"enable_backward": False}) @wp.kernel -def _kinematics_level( +def _kinematics_branch( # Model: qpos0: wp.array2d(dtype=float), body_parentid: wp.array(dtype=int), @@ -50,56 +51,51 @@ def _kinematics_level( body_jntadr: wp.array(dtype=int), body_pos: wp.array2d(dtype=wp.vec3), body_quat: wp.array2d(dtype=wp.quat), - body_ipos: wp.array2d(dtype=wp.vec3), - body_iquat: wp.array2d(dtype=wp.quat), jnt_type: wp.array(dtype=int), jnt_qposadr: wp.array(dtype=int), jnt_pos: wp.array2d(dtype=wp.vec3), jnt_axis: wp.array2d(dtype=wp.vec3), + body_branches: wp.array(dtype=int), + body_branch_start: wp.array(dtype=int), # Data in: qpos_in: wp.array2d(dtype=float), mocap_pos_in: wp.array2d(dtype=wp.vec3), mocap_quat_in: wp.array2d(dtype=wp.quat), - xpos_in: wp.array2d(dtype=wp.vec3), - xquat_in: wp.array2d(dtype=wp.quat), - xmat_in: wp.array2d(dtype=wp.mat33), - # In: - body_tree_: wp.array(dtype=int), # Data out: xpos_out: wp.array2d(dtype=wp.vec3), xquat_out: wp.array2d(dtype=wp.quat), - xmat_out: wp.array2d(dtype=wp.mat33), - xipos_out: wp.array2d(dtype=wp.vec3), - ximat_out: wp.array2d(dtype=wp.mat33), xanchor_out: wp.array2d(dtype=wp.vec3), xaxis_out: wp.array2d(dtype=wp.vec3), ): - worldid, nodeid = wp.tid() - bodyid = body_tree_[nodeid] - jntadr = body_jntadr[bodyid] - jntnum = body_jntnum[bodyid] + worldid, branchid = wp.tid() + + start = body_branch_start[branchid] + end = body_branch_start[branchid + 1] + qpos = qpos_in[worldid] - body_pos_id = worldid % body_pos.shape[0] - body_quat_id = worldid % body_quat.shape[0] - jnt_axis_id = worldid % jnt_axis.shape[0] - free_joint = False - if jntnum == 1: - jnt_type_ = jnt_type[jntadr] - free_joint = jnt_type_ == JointType.FREE + for i in range(start, end): + bodyid = body_branches[i] + pid = body_parentid[bodyid] + jntadr = body_jntadr[bodyid] + jntnum = body_jntnum[bodyid] + + if jntnum == 1: + jnt_type_ = jnt_type[jntadr] + if jnt_type_ == JointType.FREE: + qadr = jnt_qposadr[jntadr] + xpos = wp.vec3(qpos[qadr], qpos[qadr + 1], qpos[qadr + 2]) + xquat = wp.quat(qpos[qadr + 3], qpos[qadr + 4], qpos[qadr + 5], qpos[qadr + 6]) + xquat = wp.normalize(xquat) + + xpos_out[worldid, bodyid] = xpos + xquat_out[worldid, bodyid] = xquat + xanchor_out[worldid, jntadr] = xpos + xaxis_out[worldid, jntadr] = jnt_axis[worldid % jnt_axis.shape[0], jntadr] + continue - if free_joint: - # free joint - qadr = jnt_qposadr[jntadr] - xpos = wp.vec3(qpos[qadr], qpos[qadr + 1], qpos[qadr + 2]) - xquat = wp.quat(qpos[qadr + 3], qpos[qadr + 4], qpos[qadr + 5], qpos[qadr + 6]) - xquat = wp.normalize(xquat) - xanchor_out[worldid, jntadr] = xpos - xaxis_out[worldid, jntadr] = jnt_axis[jnt_axis_id, jntadr] - else: # regular or no joints # apply fixed translation and rotation relative to parent - qpos0_id = worldid % qpos0.shape[0] jnt_pos_id = worldid % jnt_pos.shape[0] pid = body_parentid[bodyid] @@ -109,17 +105,17 @@ def _kinematics_level( xpos = mocap_pos_in[worldid, mocapid] xquat = mocap_quat_in[worldid, mocapid] else: - xpos = body_pos[body_pos_id, bodyid] - xquat = body_quat[body_quat_id, bodyid] + xpos = body_pos[worldid % body_pos.shape[0], bodyid] + xquat = body_quat[worldid % body_quat.shape[0], bodyid] if pid >= 0: - xpos = xmat_in[worldid, pid] @ xpos + xpos_in[worldid, pid] - xquat = math.mul_quat(xquat_in[worldid, pid], xquat) + xpos = math.rot_vec_quat(xpos, xquat_out[worldid, pid]) + xpos_out[worldid, pid] + xquat = math.mul_quat(xquat_out[worldid, pid], xquat) for _ in range(jntnum): qadr = jnt_qposadr[jntadr] jnt_type_ = jnt_type[jntadr] - jnt_axis_ = jnt_axis[jnt_axis_id, jntadr] + jnt_axis_ = jnt_axis[worldid % jnt_axis.shape[0], jntadr] xanchor = math.rot_vec_quat(jnt_pos[jnt_pos_id, jntadr], xquat) + xpos xaxis = math.rot_vec_quat(jnt_axis_, xquat) @@ -130,9 +126,9 @@ def _kinematics_level( # correct for off-center rotation xpos = xanchor - math.rot_vec_quat(jnt_pos[jnt_pos_id, jntadr], xquat) elif jnt_type_ == JointType.SLIDE: - xpos += xaxis * (qpos[qadr] - qpos0[qpos0_id, qadr]) + xpos += xaxis * (qpos[qadr] - qpos0[worldid % qpos0.shape[0], qadr]) elif jnt_type_ == JointType.HINGE: - qpos0_ = qpos0[qpos0_id, qadr] + qpos0_ = qpos0[worldid % qpos0.shape[0], qadr] qloc_ = math.axis_angle_to_quat(jnt_axis_, qpos[qadr] - qpos0_) xquat = math.mul_quat(xquat, qloc_) # correct for off-center rotation @@ -142,14 +138,41 @@ def _kinematics_level( xaxis_out[worldid, jntadr] = xaxis jntadr += 1 - xpos_out[worldid, bodyid] = xpos - xquat = wp.normalize(xquat) - xquat_out[worldid, bodyid] = xquat - xmat_out[worldid, bodyid] = math.quat_to_mat(xquat) + xquat = wp.normalize(xquat) + xpos_out[worldid, bodyid] = xpos + xquat_out[worldid, bodyid] = xquat + + +@wp.kernel +def _compute_body_inertial_frames( + # Model: + body_ipos: wp.array2d(dtype=wp.vec3), + body_iquat: wp.array2d(dtype=wp.quat), + # Data in: + xpos_in: wp.array2d(dtype=wp.vec3), + xquat_in: wp.array2d(dtype=wp.quat), + # Data out: + xipos_out: wp.array2d(dtype=wp.vec3), + ximat_out: wp.array2d(dtype=wp.mat33), +): + worldid, bodyid = wp.tid() + xpos = xpos_in[worldid, bodyid] + xquat = xquat_in[worldid, bodyid] xipos_out[worldid, bodyid] = xpos + math.rot_vec_quat(body_ipos[worldid % body_ipos.shape[0], bodyid], xquat) ximat_out[worldid, bodyid] = math.quat_to_mat(math.mul_quat(xquat, body_iquat[worldid % body_iquat.shape[0], bodyid])) +@wp.kernel +def _compute_body_matrices( + # Data in: + xquat_in: wp.array2d(dtype=wp.quat), + # Data out: + xmat_out: wp.array2d(dtype=wp.mat33), +): + worldid, bodyid = wp.tid() + xmat_out[worldid, bodyid] = math.quat_to_mat(xquat_in[worldid, bodyid]) + + @wp.kernel def _geom_local_to_global( # Model: @@ -217,7 +240,6 @@ def _flex_vertices( @wp.kernel def _flex_edges( # Model: - nv: int, nflex: int, body_parentid: wp.array(dtype=int), body_rootid: wp.array(dtype=int), @@ -228,6 +250,8 @@ def _flex_edges( 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), @@ -244,28 +268,81 @@ def _flex_edges( if locid >= 0 and locid < flex_edgenum[i]: f = i break + vbase = flex_vertadr[f] v = flex_edge[edgeid] - pos1 = flexvert_xpos_in[worldid, vbase + v[0]] - pos2 = flexvert_xpos_in[worldid, vbase + v[1]] + vbase0 = vbase + v[0] + vbase1 = vbase + v[1] + + pos1 = flexvert_xpos_in[worldid, vbase0] + pos2 = flexvert_xpos_in[worldid, vbase1] vec = pos2 - pos1 - vecnorm = wp.length(vec) - flexedge_length_out[worldid, edgeid] = vecnorm + edge, edge_length = math.normalize_with_norm(vec) + flexedge_length_out[worldid, edgeid] = edge_length # TODO(quaglino): use Jacobian - b1 = flex_vertbodyid[vbase + v[0]] - b2 = flex_vertbodyid[vbase + v[1]] - i = body_dofadr[b1] - j = body_dofadr[b2] - vel1 = wp.vec3(qvel_in[worldid, i], qvel_in[worldid, i + 1], qvel_in[worldid, i + 2]) - vel2 = wp.vec3(qvel_in[worldid, j], qvel_in[worldid, j + 1], qvel_in[worldid, j + 2]) - edge = wp.normalize(vec) + b1 = flex_vertbodyid[vbase0] + b2 = flex_vertbodyid[vbase1] + + dofi = body_dofadr[b1] + dofj = body_dofadr[b2] + dofi0 = dofi + 0 + dofi1 = dofi + 1 + dofi2 = dofi + 2 + dofj0 = dofj + 0 + dofj1 = dofj + 1 + dofj2 = dofj + 2 + + vel1 = wp.vec3(qvel_in[worldid, dofi0], qvel_in[worldid, dofi1], qvel_in[worldid, dofi2]) + vel2 = wp.vec3(qvel_in[worldid, dofj0], qvel_in[worldid, dofj1], qvel_in[worldid, dofj2]) flexedge_velocity_out[worldid, edgeid] = wp.dot(vel2 - vel1, edge) - # Edge jacobian - for k in range(nv): - jacp1, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos1, b1, k, worldid) - jacp2, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos2, b2, k, worldid) - jacdif = jacp2 - jacp1 - flexedge_J_out[worldid, edgeid, k] = wp.dot(jacdif, edge) + + rowadr = flexedge_J_rowadr[edgeid] + + sparseid0 = rowadr + 0 + sparseid1 = rowadr + 1 + sparseid2 = rowadr + 2 + sparseid3 = rowadr + 3 + sparseid4 = rowadr + 4 + sparseid5 = rowadr + 5 + + # TODO(team): jacdif + + jacp1, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos1, b1, dofi0, worldid) + jacp2, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos2, b2, dofi0, worldid) + jacdif = jacp2 - jacp1 + Ji0 = wp.dot(jacdif, edge) + + jacp1, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos1, b1, dofi1, worldid) + jacp2, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos2, b2, dofi1, worldid) + jacdif = jacp2 - jacp1 + Ji1 = wp.dot(jacdif, edge) + + jacp1, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos1, b1, dofi2, worldid) + jacp2, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos2, b2, dofi2, worldid) + jacdif = jacp2 - jacp1 + Ji2 = wp.dot(jacdif, edge) + + jacp1, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos1, b1, dofj0, worldid) + jacp2, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos2, b2, dofj0, worldid) + jacdif = jacp2 - jacp1 + Jj0 = wp.dot(jacdif, edge) + + jacp1, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos1, b1, dofj1, worldid) + jacp2, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos2, b2, dofj1, worldid) + jacdif = jacp2 - jacp1 + Jj1 = wp.dot(jacdif, edge) + + jacp1, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos1, b1, dofj2, worldid) + jacp2, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pos2, b2, dofj2, worldid) + jacdif = jacp2 - jacp1 + Jj2 = wp.dot(jacdif, edge) + + flexedge_J_out[worldid, 0, sparseid0] = Ji0 + flexedge_J_out[worldid, 0, sparseid1] = Ji1 + flexedge_J_out[worldid, 0, sparseid2] = Ji2 + flexedge_J_out[worldid, 0, sparseid3] = Jj0 + flexedge_J_out[worldid, 0, sparseid4] = Jj1 + flexedge_J_out[worldid, 0, sparseid5] = Jj2 @event_scope @@ -276,35 +353,43 @@ def kinematics(m: Model, d: Data): derived positions and orientations of geoms, sites, and flexible elements, based on the current joint positions and any attached mocap bodies. """ - for i in range(1, len(m.body_tree)): - body_tree = m.body_tree[i] - wp.launch( - _kinematics_level, - dim=(d.nworld, body_tree.size), - inputs=[ - m.qpos0, - m.body_parentid, - m.body_mocapid, - m.body_jntnum, - m.body_jntadr, - m.body_pos, - m.body_quat, - m.body_ipos, - m.body_iquat, - m.jnt_type, - m.jnt_qposadr, - m.jnt_pos, - m.jnt_axis, - d.qpos, - d.mocap_pos, - d.mocap_quat, - d.xpos, - d.xquat, - d.xmat, - body_tree, - ], - outputs=[d.xpos, d.xquat, d.xmat, d.xipos, d.ximat, d.xanchor, d.xaxis], - ) + wp.launch( + _kinematics_branch, + dim=(d.nworld, m.nbranch), + inputs=[ + m.qpos0, + m.body_parentid, + m.body_mocapid, + m.body_jntnum, + m.body_jntadr, + m.body_pos, + m.body_quat, + m.jnt_type, + m.jnt_qposadr, + m.jnt_pos, + m.jnt_axis, + m.body_branches, + m.body_branch_start, + d.qpos, + d.mocap_pos, + d.mocap_quat, + ], + outputs=[d.xpos, d.xquat, d.xanchor, d.xaxis], + ) + + wp.launch( + _compute_body_matrices, + dim=(d.nworld, m.nbody), + inputs=[d.xquat], + outputs=[d.xmat], + ) + + wp.launch( + _compute_body_inertial_frames, + dim=(d.nworld, m.nbody), + inputs=[m.body_ipos, m.body_iquat, d.xpos, d.xquat], + outputs=[d.xipos, d.ximat], + ) wp.launch( _geom_local_to_global, @@ -328,7 +413,6 @@ def flex(m: Model, d: Data): _flex_edges, dim=(d.nworld, m.nflexedge), inputs=[ - m.nv, m.nflex, m.body_parentid, m.body_rootid, @@ -339,12 +423,18 @@ def flex(m: Model, d: Data): m.flex_edgenum, m.flex_vertbodyid, m.flex_edge, + m.flexedge_J_rowadr, + m.flexedge_J_colind, d.qvel, d.subtree_com, d.cdof, d.flexvert_xpos, ], - outputs=[d.flexedge_J, d.flexedge_length, d.flexedge_velocity], + outputs=[ + d.flexedge_J, + d.flexedge_length, + d.flexedge_velocity, + ], ) @@ -787,7 +877,7 @@ def crb(m: Model, d: Data): wp.launch(_crb_accumulate, dim=(d.nworld, body_tree.size), inputs=[m.body_parentid, d.crb, body_tree], outputs=[d.crb]) d.qM.zero_() - if m.opt.is_sparse: + if m.is_sparse: wp.launch( _qM_sparse, dim=(d.nworld, m.nv), @@ -803,10 +893,10 @@ def crb(m: Model, d: Data): @wp.kernel def _tendon_armature( # Model: - opt_is_sparse: bool, dof_parentid: wp.array(dtype=int), dof_Madr: wp.array(dtype=int), tendon_armature: wp.array2d(dtype=float), + is_sparse: bool, # Data in: ten_J_in: wp.array3d(dtype=float), # Data out: @@ -814,7 +904,7 @@ def _tendon_armature( ): worldid, tenid, dofid = wp.tid() - if opt_is_sparse: # opt_is_sparse is not batched + if is_sparse: # is_sparse is not batched madr_ij = dof_Madr[dofid] armature = tendon_armature[worldid, tenid] @@ -837,7 +927,7 @@ def _tendon_armature( qMij = armature * ten_Jj * ten_Ji - if opt_is_sparse: + if is_sparse: wp.atomic_add(qM_out[worldid, 0], madr_ij, qMij) madr_ij += 1 else: @@ -854,7 +944,7 @@ def tendon_armature(m: Model, d: Data): wp.launch( _tendon_armature, dim=(d.nworld, m.ntendon, m.nv), - inputs=[m.opt.is_sparse, m.dof_parentid, m.dof_Madr, m.tendon_armature, d.ten_J], + inputs=[m.dof_parentid, m.dof_Madr, m.tendon_armature, m.is_sparse, d.ten_J], outputs=[d.qM], ) @@ -927,9 +1017,9 @@ def _factor_i_sparse(m: Model, d: Data, M: wp.array3d(dtype=float), L: wp.array3 def _tile_cholesky_factorize(tile: TileSet): """Returns a kernel for dense Cholesky factorization of a tile.""" - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def cholesky_factorize( - # Data In: + # Data in: qM_in: wp.array3d(dtype=float), # In: adr: wp.array(dtype=int), @@ -962,7 +1052,7 @@ def _factor_i_dense(m: Model, d: Data, M: wp.array, L: wp.array): @event_scope def factor_m(m: Model, d: Data): """Factorization of inertia-like matrix M, assumed spd.""" - if m.opt.is_sparse: + if m.is_sparse: _factor_i_sparse(m, d, d.qM, d.qLD, d.qLDiagInv) else: _factor_i_dense(m, d, d.qM, d.qLD) @@ -987,44 +1077,60 @@ def _rne_cacc_world(m: Model, d: Data): @wp.kernel -def _cacc( +def _cacc_branch( # Model: body_parentid: wp.array(dtype=int), body_dofnum: wp.array(dtype=int), body_dofadr: wp.array(dtype=int), + body_branches: wp.array(dtype=int), + body_branch_start: wp.array(dtype=int), # Data in: qvel_in: wp.array2d(dtype=float), qacc_in: wp.array2d(dtype=float), cdof_in: wp.array2d(dtype=wp.spatial_vector), cdof_dot_in: wp.array2d(dtype=wp.spatial_vector), - cacc_in: wp.array2d(dtype=wp.spatial_vector), # In: - body_tree_: wp.array(dtype=int), flg_acc: bool, # Data out: cacc_out: wp.array2d(dtype=wp.spatial_vector), ): - worldid, nodeid = wp.tid() - bodyid = body_tree_[nodeid] - dofnum = body_dofnum[bodyid] + worldid, branchid = wp.tid() + + start = body_branch_start[branchid] + end = body_branch_start[branchid + 1] + + bodyid = body_branches[start] pid = body_parentid[bodyid] - dofadr = body_dofadr[bodyid] - local_cacc = cacc_in[worldid, pid] - for i in range(dofnum): - local_cacc += cdof_dot_in[worldid, dofadr + i] * qvel_in[worldid, dofadr + i] - if flg_acc: - local_cacc += cdof_in[worldid, dofadr + i] * qacc_in[worldid, dofadr + i] - cacc_out[worldid, bodyid] = local_cacc + local_cacc = cacc_out[worldid, pid] + for i in range(start, end): + bodyid = body_branches[i] + dofnum = body_dofnum[bodyid] + dofadr = body_dofadr[bodyid] + for j in range(dofnum): + local_cacc += cdof_dot_in[worldid, dofadr + j] * qvel_in[worldid, dofadr + j] + if flg_acc: + local_cacc += cdof_in[worldid, dofadr + j] * qacc_in[worldid, dofadr + j] + cacc_out[worldid, bodyid] = local_cacc def _rne_cacc_forward(m: Model, d: Data, flg_acc: bool = False): - for body_tree in m.body_tree: - wp.launch( - _cacc, - dim=(d.nworld, body_tree.size), - inputs=[m.body_parentid, m.body_dofnum, m.body_dofadr, d.qvel, d.qacc, d.cdof, d.cdof_dot, d.cacc, body_tree, flg_acc], - outputs=[d.cacc], - ) + wp.launch( + _cacc_branch, + dim=(d.nworld, m.nbranch), + inputs=[ + m.body_parentid, + m.body_dofnum, + m.body_dofadr, + m.body_branches, + m.body_branch_start, + d.qvel, + d.qacc, + d.cdof, + d.cdof_dot, + flg_acc, + ], + outputs=[d.cacc], + ) @wp.kernel @@ -1138,6 +1244,36 @@ def _cfrc_ext( cfrc_ext_out[worldid, bodyid] = support.transform_force(xfrc_applied, subtree_com - xipos) +@wp.kernel +def _count_equality_constraints( + # Model: + eq_type: wp.array(dtype=int), + # Data in: + ne_in: wp.array(dtype=int), + efc_type_in: wp.array2d(dtype=int), + efc_id_in: wp.array2d(dtype=int), + # Out: + ne_connect_out: wp.array(dtype=int), + ne_weld_out: wp.array(dtype=int), +): + """Counts connect and weld equality constraints from efc data.""" + worldid, efcid = wp.tid() + + # Only process rows within the equality constraint range + if efcid >= ne_in[worldid]: + return + + # Get the equality constraint ID and its type + eq_id = efc_id_in[worldid, efcid] + eq_constraint_type = eq_type[eq_id] + + # Count by type (each connect has 3 rows, each weld has 6 rows) + if eq_constraint_type == EqType.CONNECT: + wp.atomic_add(ne_connect_out, worldid, 1) + elif eq_constraint_type == EqType.WELD: + wp.atomic_add(ne_weld_out, worldid, 1) + + @wp.kernel def _cfrc_ext_equality( # Model: @@ -1154,6 +1290,7 @@ def _cfrc_ext_equality( subtree_com_in: wp.array2d(dtype=wp.vec3), efc_id_in: wp.array2d(dtype=int), efc_force_in: wp.array2d(dtype=float), + # In: ne_connect_in: wp.array(dtype=int), ne_weld_in: wp.array(dtype=int), # Data out: @@ -1324,27 +1461,40 @@ def rne_postconstraint(m: Model, d: Data): outputs=[d.cfrc_ext], ) - wp.launch( - _cfrc_ext_equality, - dim=(d.nworld, m.neq), - inputs=[ - m.body_rootid, - m.site_bodyid, - m.site_pos, - m.eq_obj1id, - m.eq_obj2id, - m.eq_objtype, - m.eq_data, - d.xpos, - d.xmat, - d.subtree_com, - d.efc.id, - d.efc.force, - d.ne_connect, - d.ne_weld, - ], - outputs=[d.cfrc_ext], - ) + # Equality constraint forces - only if model has equality constraints + if m.neq > 0: + # Allocate inline counters and count from efc data + ne_connect = wp.zeros((d.nworld,), dtype=int) + ne_weld = wp.zeros((d.nworld,), dtype=int) + + wp.launch( + _count_equality_constraints, + dim=(d.nworld, d.njmax), # TODO(team): launch over max equality constraints + inputs=[m.eq_type, d.ne, d.efc.type, d.efc.id], + outputs=[ne_connect, ne_weld], + ) + + wp.launch( + _cfrc_ext_equality, + dim=(d.nworld, m.neq), + inputs=[ + m.body_rootid, + m.site_bodyid, + m.site_pos, + m.eq_obj1id, + m.eq_obj2id, + m.eq_objtype, + m.eq_data, + d.xpos, + d.xmat, + d.subtree_com, + d.efc.id, + d.efc.force, + ne_connect, + ne_weld, + ], + outputs=[d.cfrc_ext], + ) # cfrc_ext += contacts wp.launch( @@ -1482,7 +1632,7 @@ def _tendon_dot( # get endpoint Jacobian time derivatives, subtract # TODO(team): parallelize? for i in range(nv): - jac1, _ = support.jac_dot( + jac1, _ = support.jac_dot_dof( body_parentid, body_rootid, jnt_type, @@ -1498,7 +1648,7 @@ def _tendon_dot( i, worldid, ) - jac2, _ = support.jac_dot( + jac2, _ = support.jac_dot_dof( body_parentid, body_rootid, jnt_type, @@ -1520,7 +1670,7 @@ def _tendon_dot( Jdot = wp.dot(jacdif, dpnt) # get endpoint Jacobians, subtract - jac1, _ = support.jac( + jac1, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, @@ -1531,7 +1681,7 @@ def _tendon_dot( i, worldid, ) - jac2, _ = support.jac( + jac2, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, @@ -1663,72 +1813,76 @@ def _comvel_root(cvel_out: wp.array2d(dtype=wp.spatial_vector)): @wp.kernel -def _comvel_level( +def _comvel_branch( # Model: body_parentid: wp.array(dtype=int), body_jntnum: wp.array(dtype=int), body_jntadr: wp.array(dtype=int), body_dofadr: wp.array(dtype=int), jnt_type: wp.array(dtype=int), + body_branches: wp.array(dtype=int), + body_branch_start: wp.array(dtype=int), # Data in: qvel_in: wp.array2d(dtype=float), cdof_in: wp.array2d(dtype=wp.spatial_vector), - cvel_in: wp.array2d(dtype=wp.spatial_vector), - # In: - body_tree_: wp.array(dtype=int), # Data out: cvel_out: wp.array2d(dtype=wp.spatial_vector), cdof_dot_out: wp.array2d(dtype=wp.spatial_vector), ): - worldid, nodeid = wp.tid() - bodyid = body_tree_[nodeid] - dofid = body_dofadr[bodyid] - jntid = body_jntadr[bodyid] - jntnum = body_jntnum[bodyid] - pid = body_parentid[bodyid] + worldid, branchid = wp.tid() - if jntnum == 0: - cvel_out[worldid, bodyid] = cvel_in[worldid, pid] - return + start = body_branch_start[branchid] + end = body_branch_start[branchid + 1] - cvel = cvel_in[worldid, pid] qvel = qvel_in[worldid] cdof = cdof_in[worldid] - for j in range(jntid, jntid + jntnum): - jnttype = jnt_type[j] + for i in range(start, end): + bodyid = body_branches[i] + pid = body_parentid[bodyid] + cvel = cvel_out[worldid, pid] + dofid = body_dofadr[bodyid] + jntid = body_jntadr[bodyid] + jntnum = body_jntnum[bodyid] - if jnttype == JointType.FREE: - cvel += cdof[dofid + 0] * qvel[dofid + 0] - cvel += cdof[dofid + 1] * qvel[dofid + 1] - cvel += cdof[dofid + 2] * qvel[dofid + 2] + if jntnum == 0: + cvel_out[worldid, bodyid] = cvel + continue - cdof_dot_out[worldid, dofid + 3] = math.motion_cross(cvel, cdof[dofid + 3]) - cdof_dot_out[worldid, dofid + 4] = math.motion_cross(cvel, cdof[dofid + 4]) - cdof_dot_out[worldid, dofid + 5] = math.motion_cross(cvel, cdof[dofid + 5]) + for j in range(jntid, jntid + jntnum): + jnttype = jnt_type[j] - cvel += cdof[dofid + 3] * qvel[dofid + 3] - cvel += cdof[dofid + 4] * qvel[dofid + 4] - cvel += cdof[dofid + 5] * qvel[dofid + 5] + if jnttype == JointType.FREE: + cvel += cdof[dofid + 0] * qvel[dofid + 0] + cvel += cdof[dofid + 1] * qvel[dofid + 1] + cvel += cdof[dofid + 2] * qvel[dofid + 2] - dofid += 6 - elif jnttype == JointType.BALL: - cdof_dot_out[worldid, dofid + 0] = math.motion_cross(cvel, cdof[dofid + 0]) - cdof_dot_out[worldid, dofid + 1] = math.motion_cross(cvel, cdof[dofid + 1]) - cdof_dot_out[worldid, dofid + 2] = math.motion_cross(cvel, cdof[dofid + 2]) + cdof_dot_out[worldid, dofid + 3] = math.motion_cross(cvel, cdof[dofid + 3]) + cdof_dot_out[worldid, dofid + 4] = math.motion_cross(cvel, cdof[dofid + 4]) + cdof_dot_out[worldid, dofid + 5] = math.motion_cross(cvel, cdof[dofid + 5]) - cvel += cdof[dofid + 0] * qvel[dofid + 0] - cvel += cdof[dofid + 1] * qvel[dofid + 1] - cvel += cdof[dofid + 2] * qvel[dofid + 2] + cvel += cdof[dofid + 3] * qvel[dofid + 3] + cvel += cdof[dofid + 4] * qvel[dofid + 4] + cvel += cdof[dofid + 5] * qvel[dofid + 5] - dofid += 3 - else: - cdof_dot_out[worldid, dofid] = math.motion_cross(cvel, cdof[dofid]) - cvel += cdof[dofid] * qvel[dofid] + dofid += 6 + elif jnttype == JointType.BALL: + cdof_dot_out[worldid, dofid + 0] = math.motion_cross(cvel, cdof[dofid + 0]) + cdof_dot_out[worldid, dofid + 1] = math.motion_cross(cvel, cdof[dofid + 1]) + cdof_dot_out[worldid, dofid + 2] = math.motion_cross(cvel, cdof[dofid + 2]) - dofid += 1 + cvel += cdof[dofid + 0] * qvel[dofid + 0] + cvel += cdof[dofid + 1] * qvel[dofid + 1] + cvel += cdof[dofid + 2] * qvel[dofid + 2] - cvel_out[worldid, bodyid] = cvel + dofid += 3 + else: + cdof_dot_out[worldid, dofid] = math.motion_cross(cvel, cdof[dofid]) + cvel += cdof[dofid] * qvel[dofid] + + dofid += 1 + + cvel_out[worldid, bodyid] = cvel @event_scope @@ -1740,13 +1894,22 @@ def com_vel(m: Model, d: Data): """ wp.launch(_comvel_root, dim=(d.nworld, 6), inputs=[], outputs=[d.cvel]) - for body_tree in m.body_tree: - wp.launch( - _comvel_level, - dim=(d.nworld, body_tree.size), - inputs=[m.body_parentid, m.body_jntnum, m.body_jntadr, m.body_dofadr, m.jnt_type, d.qvel, d.cdof, d.cvel, body_tree], - outputs=[d.cvel, d.cdof_dot], - ) + wp.launch( + _comvel_branch, + dim=(d.nworld, m.nbranch), + inputs=[ + m.body_parentid, + m.body_jntnum, + m.body_jntadr, + m.body_dofadr, + m.jnt_type, + m.body_branches, + m.body_branch_start, + d.qvel, + d.cdof, + ], + outputs=[d.cvel, d.cdof_dot], + ) @wp.kernel @@ -1869,14 +2032,14 @@ def _transmission( for i in range(nv): # get Jacobians of axis(jacA) and vec(jac) # mj_jacPointAxis - jacp, jacr = support.jac( + jacp, jacr = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_idslider, site_bodyid[idslider], i, worldid ) jacS = jacp jacA = wp.cross(jacr, axis) # mj_jacSite - jac, _ = support.jac( + jac, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_id, site_bodyid[id], i, worldid ) jac -= jacS @@ -1929,7 +2092,7 @@ def _transmission( # moment: global Jacobian projected on wrench # TODO(team): parallelize for i in range(nv): - jacp, jacr = support.jac( + jacp, jacr = support.jac_dof( body_parentid, body_rootid, dof_bodyid, @@ -2001,12 +2164,12 @@ def _transmission( # TODO(team): parallelize for i in range(nv): - jacp, jacr = support.jac( + jacp, jacr = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos, site_bodyid[siteid], i, worldid ) # jacref: global Jacobian of reference site - jacpref, jacrref = support.jac( + jacpref, jacrref = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, ref_xpos, site_bodyid[refid], i, worldid ) @@ -2119,8 +2282,8 @@ def _transmission_body_moment( normal = wp.vec3(contact_frame[0, 0], contact_frame[0, 1], contact_frame[0, 2]) # get Jacobian difference - jacp1, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b1, dofid, worldid) - jacp2, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b2, dofid, worldid) + jacp1, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b1, dofid, worldid) + jacp2, _ = support.jac_dof(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b2, dofid, worldid) jacdif = jacp2 - jacp1 # project Jacobian along the normal of the contact frame @@ -2152,6 +2315,7 @@ def transmission(m: Model, d: Data): Updates the actuator length and moments for all actuators in the model, including joint and tendon transmissions. """ + d.actuator_moment.zero_() wp.launch( _transmission, dim=[d.nworld, m.nu], @@ -2290,7 +2454,7 @@ def _solve_LD_sparse( def _tile_cholesky_solve(tile: TileSet): """Returns a kernel for dense Cholesky backsubstitution of a tile.""" - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def cholesky_solve( # In: L: wp.array3d(dtype=float), @@ -2345,7 +2509,7 @@ def solve_LD( x: Output array for the solution. y: Input right-hand side array. """ - if m.opt.is_sparse: + if m.is_sparse: _solve_LD_sparse(m, d, L, D, x, y) else: _solve_LD_dense(m, d, L, x, y) @@ -2368,7 +2532,7 @@ def solve_m(m: Model, d: Data, x: wp.array2d(dtype=float), y: wp.array2d(dtype=f def _tile_cholesky_factorize_solve(tile: TileSet): """Returns a kernel for dense Cholesky factorization and backsubstitution of a tile.""" - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def cholesky_factorize_solve( # In: M: wp.array3d(dtype=float), @@ -2428,7 +2592,7 @@ def factor_solve_i(m, d, M, L, D, x, y): x: Output array for the solution. y: Input right-hand side array. """ - if m.opt.is_sparse: + if m.is_sparse: _factor_i_sparse(m, d, M, L, D) _solve_LD_sparse(m, d, L, D, x, y) else: @@ -2449,6 +2613,7 @@ def _subtree_vel_forward( # Data out: subtree_linvel_out: wp.array2d(dtype=wp.vec3), subtree_angmom_out: wp.array2d(dtype=wp.vec3), + # Out: subtree_bodyvel_out: wp.array2d(dtype=wp.spatial_vector), ): worldid, bodyid = wp.tid() @@ -2504,8 +2669,8 @@ def _angular_momentum( xipos_in: wp.array2d(dtype=wp.vec3), subtree_com_in: wp.array2d(dtype=wp.vec3), subtree_linvel_in: wp.array2d(dtype=wp.vec3), - subtree_bodyvel_in: wp.array2d(dtype=wp.spatial_vector), # In: + subtree_bodyvel_in: wp.array2d(dtype=wp.spatial_vector), body_tree_: wp.array(dtype=int), # Data out: subtree_angmom_out: wp.array2d(dtype=wp.vec3), @@ -2553,12 +2718,14 @@ def subtree_vel(m: Model, d: Data): Computes the linear momentum and angular momentum for each subtree, accumulating contributions up the kinematic tree. """ + subtree_bodyvel = wp.empty((d.nworld, m.nbody), dtype=wp.spatial_vector) + # bodywise quantities wp.launch( _subtree_vel_forward, dim=(d.nworld, m.nbody), inputs=[m.body_rootid, m.body_mass, m.body_inertia, d.xipos, d.ximat, d.subtree_com, d.cvel], - outputs=[d.subtree_linvel, d.subtree_angmom, d.subtree_bodyvel], + outputs=[d.subtree_linvel, d.subtree_angmom, subtree_bodyvel], ) # sum body linear momentum recursively up the kinematic tree @@ -2581,7 +2748,7 @@ def subtree_vel(m: Model, d: Data): d.xipos, d.subtree_com, d.subtree_linvel, - d.subtree_bodyvel, + subtree_bodyvel, body_tree, ], outputs=[d.subtree_angmom], @@ -2666,8 +2833,8 @@ def _spatial_site_tendon( if body0 != body1: # TODO(team): parallelize for i in range(nv): - jacp1, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pnt0, body0, i, worldid) - jacp2, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, pnt1, body1, i, worldid) + 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: @@ -2734,7 +2901,7 @@ def _spatial_geom_tendon( if sideid >= 0: side = site_xpos_in[worldid, sideid] else: - side = wp.vec3(wp.inf) + side = wp.vec3(MJ_MAXVAL) # compute geom wrap length and connect points (if wrap occurs) length_geomgeom, geom_pnt0, geom_pnt1 = util_misc.wrap(site_pnt0, site_pnt1, geom_xpos, geom_xmat, geomsize, geom_type, side) @@ -2769,11 +2936,11 @@ def _spatial_geom_tendon( J = float(0.0) # site-geom if dif_body_sitegeom: - jacp_site0, _ = support.jac( + jacp_site0, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt0, bodyid_site0, i, worldid ) - jacp_geom0, _ = support.jac( + jacp_geom0, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, geom_pnt0, bodyid_geom, i, worldid ) @@ -2781,11 +2948,11 @@ def _spatial_geom_tendon( # geom-site if dif_body_geomsite: - jacp_geom1, _ = support.jac( + 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( + jacp_site1, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt1, bodyid_site1, i, worldid ) @@ -2808,10 +2975,10 @@ def _spatial_geom_tendon( if bodyid_site0 != bodyid_site1: # TODO(team): parallelize for i in range(nv): - jacp1, _ = support.jac( + jacp1, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt0, bodyid_site0, i, worldid ) - jacp2, _ = support.jac( + jacp2, _ = support.jac_dof( body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_pnt1, bodyid_site1, i, worldid ) @@ -2892,7 +3059,7 @@ def _spatial_tendon_wrap( wrapid = id1 id1 = wrap_objid[adr + j + 2] - if wp.norm_l2(wpnt_geom0) < wp.inf: + if wp.norm_l2(wpnt_geom0) < MJ_MAXVAL: wpnt_geom1 = wp.spatial_bottom(wrap_geom_xpos) wpnt_site1 = site_xpos_in[worldid, id1] 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 1398976a..8eabb52f 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py @@ -13,9 +13,9 @@ # limitations under the License. # ============================================================================== +import dataclasses from math import ceil from math import sqrt -from typing import Tuple import warp as wp @@ -27,16 +27,119 @@ from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocke from mujoco.mjx.third_party.mujoco_warp._src.block_cholesky import create_blocked_cholesky_solve_func 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 nested_kernel wp.set_module_options({"enable_backward": False}) _BLOCK_CHOLESKY_DIM = 32 +@dataclasses.dataclass +class InverseContext: + """Workspace arrays for inverse dynamics.""" + + Jaref: wp.array2d(dtype=float) + search_dot: wp.array(dtype=float) + gauss: wp.array(dtype=float) + cost: wp.array(dtype=float) + prev_cost: wp.array(dtype=float) + done: wp.array(dtype=bool) + + +@dataclasses.dataclass +class SolverContext: + """Workspace arrays for constraint solver.""" + + Jaref: wp.array2d(dtype=float) + search_dot: wp.array(dtype=float) + gauss: wp.array(dtype=float) + cost: wp.array(dtype=float) + prev_cost: wp.array(dtype=float) + done: wp.array(dtype=bool) + grad: wp.array2d(dtype=float) + grad_dot: wp.array(dtype=float) + Mgrad: wp.array2d(dtype=float) + search: wp.array2d(dtype=float) + mv: wp.array2d(dtype=float) + jv: wp.array2d(dtype=float) + quad: wp.array2d(dtype=wp.vec3) + quad_gauss: wp.array(dtype=wp.vec3) + alpha: wp.array(dtype=float) + prev_grad: wp.array2d(dtype=float) + prev_Mgrad: wp.array2d(dtype=float) + beta: wp.array(dtype=float) + h: wp.array3d(dtype=float) + hfactor: wp.array3d(dtype=float) + + +def create_inverse_context(m: types.Model, d: types.Data) -> InverseContext: + """Create an InverseContext with allocated workspace arrays. + + Args: + m: Model containing nv, nv_pad, and solver type. + d: Data containing nworld and njmax. + + Returns: + InverseContext with allocated arrays. + """ + nworld = d.nworld + 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), + ) + + +def create_solver_context(m: types.Model, d: types.Data) -> SolverContext: + """Create a SolverContext with allocated workspace arrays. + + Args: + m: Model containing nv, nv_pad, and solver type. + d: Data containing nworld and njmax. + + Returns: + SolverContext with allocated arrays. + """ + nworld = d.nworld + nv = m.nv + nv_pad = m.nv_pad + njmax = d.njmax + + # Newton solver needs h; hfactor only needed if nv > _BLOCK_CHOLESKY_DIM + alloc_h = m.opt.solver == types.SolverType.NEWTON + 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), + ) + + @wp.func -def _rescale(nv: int, stat_meaninertia: float, value: float) -> float: - return value / (stat_meaninertia * float(nv)) +def _rescale(nv: int, meaninertia: float, value: float) -> float: + return value / (meaninertia * float(nv)) @wp.func @@ -44,6 +147,39 @@ def _in_bracket(x: wp.vec3, y: wp.vec3) -> bool: return (x[1] < y[1] and y[1] < 0.0) or (x[1] > y[1] and y[1] > 0.0) +@wp.func +def _eval_pt_direct(jaref: float, jv: float, d: float, alpha: float) -> wp.vec3: + """Eval quadratic constraint, return (cost, grad, hessian).""" + x = jaref + alpha * jv + jvD = jv * d + return wp.vec3(0.5 * d * x * x, jvD * x, jv * jvD) + + +@wp.func +def _eval_pt_direct_alpha_zero(jaref: float, jv: float, d: float) -> wp.vec3: + """Eval quadratic constraint at alpha=0.""" + jvD = jv * d + return wp.vec3(0.5 * d * jaref * jaref, jvD * jaref, jv * jvD) + + +@wp.func +def _eval_pt_direct_3alphas( + jaref: float, jv: float, d: float, lo_alpha: float, hi_alpha: float, mid_alpha: float +) -> tuple[wp.vec3, wp.vec3, wp.vec3]: + """Eval quadratic constraint for 3 alphas.""" + x_lo = jaref + lo_alpha * jv + x_hi = jaref + hi_alpha * jv + x_mid = jaref + mid_alpha * jv + jvD = jv * d + hessian = jv * jvD + half_d = 0.5 * d + return ( + wp.vec3(half_d * x_lo * x_lo, jvD * x_lo, hessian), + wp.vec3(half_d * x_hi * x_hi, jvD * x_hi, hessian), + wp.vec3(half_d * x_mid * x_mid, jvD * x_mid, hessian), + ) + + @wp.func def _eval_cost(quad: wp.vec3, alpha: float) -> float: return alpha * alpha * quad[2] + alpha * quad[1] + quad[0] @@ -51,32 +187,67 @@ def _eval_cost(quad: wp.vec3, alpha: float) -> float: @wp.func def _eval_pt(quad: wp.vec3, alpha: float) -> wp.vec3: + """Eval quad polynomial at alpha, return (cost, grad, hessian).""" + aq2 = alpha * quad[2] return wp.vec3( - _eval_cost(quad, alpha), - 2.0 * alpha * quad[2] + quad[1], + alpha * aq2 + alpha * quad[1] + quad[0], + 2.0 * aq2 + quad[1], 2.0 * quad[2], ) @wp.func -def _eval_frictionloss( - # In: - x: float, - f: float, - rf: float, - Jaref: float, - jv: float, - quad: wp.vec3, -) -> wp.vec3: - # -bound < x < bound : quadratic +def _eval_pt_3alphas(quad: wp.vec3, lo_alpha: float, hi_alpha: float, mid_alpha: float) -> tuple[wp.vec3, wp.vec3, wp.vec3]: + """Eval quad polynomial for 3 alphas.""" + q0, q1, q2 = quad[0], quad[1], quad[2] + hessian = 2.0 * q2 + lo_aq2 = lo_alpha * q2 + hi_aq2 = hi_alpha * q2 + mid_aq2 = mid_alpha * q2 + return ( + wp.vec3(lo_alpha * lo_aq2 + lo_alpha * q1 + q0, 2.0 * lo_aq2 + q1, hessian), + wp.vec3(hi_alpha * hi_aq2 + hi_alpha * q1 + q0, 2.0 * hi_aq2 + q1, hessian), + wp.vec3(mid_alpha * mid_aq2 + mid_alpha * q1 + q0, 2.0 * mid_aq2 + q1, hessian), + ) + + +@wp.func +def _eval_frictionloss_pt(x: float, f: float, rf: float, jv: float, d: float) -> wp.vec3: + """Eval frictionloss and return (cost, grad, hessian). x = Jaref + alpha * jv.""" if (-rf < x) and (x < rf): - return quad - # x < -bound: linear negative + jvD = jv * d + return wp.vec3(0.5 * d * x * x, jvD * x, jv * jvD) elif x <= -rf: - return wp.vec3(f * (-0.5 * rf - Jaref), -f * jv, 0.0) - # bound < x : linear positive + return wp.vec3(f * (-0.5 * rf - x), -f * jv, 0.0) else: - return wp.vec3(f * (-0.5 * rf + Jaref), f * jv, 0.0) + return wp.vec3(f * (-0.5 * rf + x), f * jv, 0.0) + + +@wp.func +def _eval_frictionloss_pt_one(x: float, f: float, rf: float, half_d: float, jvD: float, hessian: float, f_jv: float) -> wp.vec3: + """Eval frictionloss with precomputed shared values.""" + if (-rf < x) and (x < rf): + return wp.vec3(half_d * x * x, jvD * x, hessian) + elif x <= -rf: + return wp.vec3(f * (-0.5 * rf - x), -f_jv, 0.0) + else: + return wp.vec3(f * (-0.5 * rf + x), f_jv, 0.0) + + +@wp.func +def _eval_frictionloss_pt_3alphas( + x_lo: float, x_hi: float, x_mid: float, f: float, rf: float, jv: float, d: float +) -> tuple[wp.vec3, wp.vec3, wp.vec3]: + """Eval frictionloss for 3 x values with shared precomputation.""" + jvD = jv * d + half_d = 0.5 * d + hessian = jv * jvD + f_jv = f * jv + return ( + _eval_frictionloss_pt_one(x_lo, f, rf, half_d, jvD, hessian, f_jv), + _eval_frictionloss_pt_one(x_hi, f, rf, half_d, jvD, hessian, f_jv), + _eval_frictionloss_pt_one(x_mid, f, rf, half_d, jvD, hessian, f_jv), + ) @wp.func @@ -141,354 +312,6 @@ def _eval_elliptic( return wp.vec3(0.0, 0.0, 0.0) -@wp.func -def _eval_init( - # Data in: - contact_friction_in: wp.array(dtype=types.vec5), - contact_efc_address_in: wp.array2d(dtype=int), - # In: - ne_clip: int, - nef_clip: int, - nefc_clip: int, - impratio_invsqrt: float, - type_in: wp.array(dtype=int), - id_in: wp.array(dtype=int), - D_in: wp.array(dtype=float), - frictionloss_in: wp.array(dtype=float), - Jaref_in: wp.array(dtype=float), - jv_in: wp.array(dtype=float), - quad_in: wp.array(dtype=wp.vec3), - alpha: float, -) -> wp.vec3: - lo = wp.vec3(0.0, 0.0, 0.0) - for efcid in range(ne_clip): - quad = quad_in[efcid] - lo += _eval_pt(quad, alpha) - - for efcid in range(ne_clip, nef_clip): - D = D_in[efcid] - f = frictionloss_in[efcid] - Jaref = Jaref_in[efcid] - jv = jv_in[efcid] - - # search point, friction loss, bound (rf) - x = Jaref + alpha * jv - rf = math.safe_div(f, D) - - quad_f = _eval_frictionloss(x, f, rf, Jaref, jv, quad_in[efcid]) - lo += _eval_pt(quad_f, alpha) - - for efcid in range(nef_clip, nefc_clip): - if type_in[efcid] == types.ConstraintType.CONTACT_ELLIPTIC: - conid = id_in[efcid] - - efcid0 = contact_efc_address_in[conid, 0] - if efcid != efcid0: - continue - - efcid1 = contact_efc_address_in[conid, 1] - efcid2 = contact_efc_address_in[conid, 2] - efc_quad0 = quad_in[efcid0] - efc_quad1 = quad_in[efcid1] - efc_quad2 = quad_in[efcid2] - friction = contact_friction_in[conid] - - lo += _eval_elliptic(impratio_invsqrt, friction, efc_quad0, efc_quad1, efc_quad2, alpha) - else: - Jaref = Jaref_in[efcid] - jv = jv_in[efcid] - quad = quad_in[efcid] - - x = Jaref + alpha * jv - res = _eval_pt(quad, alpha) - lo += res * float(x < 0.0) - - return lo - - -@wp.func -def _eval( - # Data in: - contact_friction_in: wp.array(dtype=types.vec5), - contact_efc_address_in: wp.array2d(dtype=int), - # In: - ne_clip: int, - nef_clip: int, - nefc_clip: int, - impratio_invsqrt: float, - type_in: wp.array(dtype=int), - id_in: wp.array(dtype=int), - D_in: wp.array(dtype=float), - frictionloss_in: wp.array(dtype=float), - Jaref_in: wp.array(dtype=float), - jv_in: wp.array(dtype=float), - quad_in: wp.array(dtype=wp.vec3), - lo_alpha: float, - hi_alpha: float, - mid_alpha: float, -) -> Tuple[wp.vec3, wp.vec3, wp.vec3]: - lo = wp.vec3(0.0, 0.0, 0.0) - hi = wp.vec3(0.0, 0.0, 0.0) - mid = wp.vec3(0.0, 0.0, 0.0) - for efcid in range(ne_clip): - quad = quad_in[efcid] - lo += _eval_pt(quad, lo_alpha) - hi += _eval_pt(quad, hi_alpha) - mid += _eval_pt(quad, mid_alpha) - - for efcid in range(ne_clip, nef_clip): - quad = quad_in[efcid] - D = D_in[efcid] - f = frictionloss_in[efcid] - Jaref = Jaref_in[efcid] - jv = jv_in[efcid] - - # search point, friction loss, bound (rf) - rf = math.safe_div(f, D) - x_lo = Jaref + lo_alpha * jv - x_hi = Jaref + hi_alpha * jv - x_mid = Jaref + mid_alpha * jv - - quad_f = _eval_frictionloss(x_lo, f, rf, Jaref, jv, quad) - lo += _eval_pt(quad_f, lo_alpha) - quad_f = _eval_frictionloss(x_hi, f, rf, Jaref, jv, quad) - hi += _eval_pt(quad_f, hi_alpha) - quad_f = _eval_frictionloss(x_mid, f, rf, Jaref, jv, quad) - mid += _eval_pt(quad_f, mid_alpha) - - for efcid in range(nef_clip, nefc_clip): - if type_in[efcid] == types.ConstraintType.CONTACT_ELLIPTIC: - conid = id_in[efcid] - - efcid0 = contact_efc_address_in[conid, 0] - if efcid != efcid0: - continue - - efcid1 = contact_efc_address_in[conid, 1] - efcid2 = contact_efc_address_in[conid, 2] - efc_quad0 = quad_in[efcid0] - efc_quad1 = quad_in[efcid1] - efc_quad2 = quad_in[efcid2] - friction = contact_friction_in[conid] - - lo += _eval_elliptic(impratio_invsqrt, friction, efc_quad0, efc_quad1, efc_quad2, lo_alpha) - hi += _eval_elliptic(impratio_invsqrt, friction, efc_quad0, efc_quad1, efc_quad2, hi_alpha) - mid += _eval_elliptic(impratio_invsqrt, friction, efc_quad0, efc_quad1, efc_quad2, mid_alpha) - else: - Jaref = Jaref_in[efcid] - jv = jv_in[efcid] - quad = quad_in[efcid] - - x_lo = Jaref + lo_alpha * jv - x_hi = Jaref + hi_alpha * jv - x_mid = Jaref + mid_alpha * jv - lo += _eval_pt(quad, lo_alpha) * float(x_lo < 0.0) - hi += _eval_pt(quad, hi_alpha) * float(x_hi < 0.0) - mid += _eval_pt(quad, mid_alpha) * float(x_mid < 0.0) - - return lo, hi, mid - - -@wp.kernel -def linesearch_iterative( - # Model: - nv: int, - opt_tolerance: wp.array(dtype=float), - opt_ls_tolerance: wp.array(dtype=float), - opt_ls_iterations: int, - opt_impratio_invsqrt: wp.array(dtype=float), - stat_meaninertia: 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_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), - efc_Jaref_in: wp.array2d(dtype=float), - efc_search_dot_in: wp.array(dtype=float), - efc_jv_in: wp.array2d(dtype=float), - efc_quad_in: wp.array2d(dtype=wp.vec3), - efc_quad_gauss_in: wp.array(dtype=wp.vec3), - efc_done_in: wp.array(dtype=bool), - njmax_in: int, - # Data out: - efc_alpha_out: wp.array(dtype=float), -): - worldid = wp.tid() - - if efc_done_in[worldid]: - return - - impratio_invsqrt = opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] - efc_type = efc_type_in[worldid] - efc_id = efc_id_in[worldid] - efc_D = efc_D_in[worldid] - efc_frictionloss = efc_frictionloss_in[worldid] - efc_Jaref = efc_Jaref_in[worldid] - efc_jv = efc_jv_in[worldid] - efc_quad = efc_quad_in[worldid] - efc_quad_gauss = efc_quad_gauss_in[worldid] - tolerance = opt_tolerance[worldid % opt_tolerance.shape[0]] - ls_tolerance = opt_ls_tolerance[worldid % opt_ls_tolerance.shape[0]] - ne_clip = min(njmax_in, ne_in[worldid]) - nef_clip = min(njmax_in, ne_clip + nf_in[worldid]) - nefc_clip = min(njmax_in, nefc_in[worldid]) - - # Calculate p0 - snorm = wp.sqrt(efc_search_dot_in[worldid]) - scale = stat_meaninertia * wp.float(nv) - gtol = tolerance * ls_tolerance * snorm * scale - p0 = wp.vec3(efc_quad_gauss[0], efc_quad_gauss[1], 2.0 * efc_quad_gauss[2]) - p0 += _eval_init( - contact_friction_in, - contact_efc_address_in, - ne_clip, - nef_clip, - nefc_clip, - impratio_invsqrt, - efc_type, - efc_id, - efc_D, - efc_frictionloss, - efc_Jaref, - efc_jv, - efc_quad, - 0.0, - ) - - # Calculate lo bound - lo_alpha_in = -math.safe_div(p0[1], p0[2]) - lo_in = _eval_pt(efc_quad_gauss, lo_alpha_in) - lo_in += _eval_init( - contact_friction_in, - contact_efc_address_in, - ne_clip, - nef_clip, - nefc_clip, - impratio_invsqrt, - efc_type, - efc_id, - efc_D, - efc_frictionloss, - efc_Jaref, - efc_jv, - efc_quad, - lo_alpha_in, - ) - - # Initialize bounds - lo_less = lo_in[1] < p0[1] - lo = wp.where(lo_less, lo_in, p0) - lo_alpha = wp.where(lo_less, lo_alpha_in, 0.0) - hi = wp.where(lo_less, p0, lo_in) - hi_alpha = wp.where(lo_less, 0.0, lo_alpha_in) - - # Launch main linesearch iterative loop - alpha = float(0.0) - for _ in range(opt_ls_iterations): - lo_next_alpha = lo_alpha - math.safe_div(lo[1], lo[2]) - hi_next_alpha = hi_alpha - math.safe_div(hi[1], hi[2]) - mid_alpha = 0.5 * (lo_alpha + hi_alpha) - - lo_next, hi_next, mid = _eval( - contact_friction_in, - contact_efc_address_in, - ne_clip, - nef_clip, - nefc_clip, - impratio_invsqrt, - efc_type, - efc_id, - efc_D, - efc_frictionloss, - efc_Jaref, - efc_jv, - efc_quad, - lo_next_alpha, - hi_next_alpha, - mid_alpha, - ) - lo_next += _eval_pt(efc_quad_gauss, lo_next_alpha) - hi_next += _eval_pt(efc_quad_gauss, hi_next_alpha) - mid += _eval_pt(efc_quad_gauss, mid_alpha) - - # swap lo: - swap_lo_lo_next = _in_bracket(lo, lo_next) - lo = wp.where(swap_lo_lo_next, lo_next, lo) - lo_alpha = wp.where(swap_lo_lo_next, lo_next_alpha, lo_alpha) - swap_lo_mid = _in_bracket(lo, mid) - lo = wp.where(swap_lo_mid, mid, lo) - lo_alpha = wp.where(swap_lo_mid, mid_alpha, lo_alpha) - swap_lo_hi_next = _in_bracket(lo, hi_next) - lo = wp.where(swap_lo_hi_next, hi_next, lo) - lo_alpha = wp.where(swap_lo_hi_next, hi_next_alpha, lo_alpha) - swap_lo = swap_lo_lo_next or swap_lo_mid or swap_lo_hi_next - - # swap hi: - swap_hi_hi_next = _in_bracket(hi, hi_next) - hi = wp.where(swap_hi_hi_next, hi_next, hi) - hi_alpha = wp.where(swap_hi_hi_next, hi_next_alpha, hi_alpha) - swap_hi_mid = _in_bracket(hi, mid) - hi = wp.where(swap_hi_mid, mid, hi) - hi_alpha = wp.where(swap_hi_mid, mid_alpha, hi_alpha) - swap_hi_lo_next = _in_bracket(hi, lo_next) - hi = wp.where(swap_hi_lo_next, lo_next, hi) - hi_alpha = wp.where(swap_hi_lo_next, lo_next_alpha, hi_alpha) - swap_hi = swap_hi_hi_next or swap_hi_mid or swap_hi_lo_next - - # if we did not adjust the interval, we are done - # also done if either low or hi slope is nearly flat - ls_done = (not swap_lo and not swap_hi) or (lo[1] < 0 and lo[1] > -gtol) or (hi[1] > 0 and hi[1] < gtol) - - # update alpha if we have an improvement - improved = lo[0] < p0[0] or hi[0] < p0[0] - lo_better = lo[0] < hi[0] - alpha = wp.where(improved and lo_better, lo_alpha, alpha) - alpha = wp.where(improved and not lo_better, hi_alpha, alpha) - if ls_done: - break - - efc_alpha_out[worldid] = alpha - - -def _linesearch_iterative(m: types.Model, d: types.Data): - """Iterative linesearch.""" - wp.launch( - linesearch_iterative, - dim=d.nworld, - inputs=[ - m.nv, - m.opt.tolerance, - m.opt.ls_tolerance, - m.opt.ls_iterations, - m.opt.impratio_invsqrt, - m.stat.meaninertia, - d.ne, - d.nf, - d.nefc, - d.contact.friction, - d.contact.efc_address, - d.efc.type, - d.efc.id, - d.efc.D, - d.efc.frictionloss, - d.efc.Jaref, - d.efc.search_dot, - d.efc.jv, - d.efc.quad, - d.efc.quad_gauss, - d.efc.done, - d.njmax, - ], - outputs=[d.efc.alpha], - block_dim=m.block_dim.linesearch_iterative, - ) - - @wp.func def _log_scale(min_value: float, max_value: float, num_values: int, i: int) -> float: step = (wp.log(max_value) - wp.log(min_value)) / wp.max(1.0, float(num_values - 1)) @@ -511,24 +334,25 @@ def linesearch_parallel_fused( efc_id_in: wp.array2d(dtype=int), efc_D_in: wp.array2d(dtype=float), efc_frictionloss_in: wp.array2d(dtype=float), - efc_Jaref_in: wp.array2d(dtype=float), - efc_jv_in: wp.array2d(dtype=float), - efc_quad_in: wp.array2d(dtype=wp.vec3), - efc_quad_gauss_in: wp.array(dtype=wp.vec3), - efc_done_in: wp.array(dtype=bool), njmax_in: int, nacon_in: wp.array(dtype=int), + # In: + ctx_Jaref_in: wp.array2d(dtype=float), + ctx_jv_in: wp.array2d(dtype=float), + ctx_quad_in: wp.array2d(dtype=wp.vec3), + ctx_quad_gauss_in: wp.array(dtype=wp.vec3), + ctx_done_in: wp.array(dtype=bool), # Out: cost_out: wp.array2d(dtype=float), ): worldid, alphaid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return alpha = _log_scale(opt_ls_parallel_min_step, 1.0, opt_ls_iterations, alphaid) - out = _eval_cost(efc_quad_gauss_in[worldid], alpha) + out = _eval_cost(ctx_quad_gauss_in[worldid], alpha) ne = ne_in[worldid] nf = nf_in[worldid] @@ -537,19 +361,19 @@ def linesearch_parallel_fused( for efcid in range(min(njmax_in, nefc_in[worldid])): # equality if efcid < ne: - out += _eval_cost(efc_quad_in[worldid, efcid], alpha) + out += _eval_cost(ctx_quad_in[worldid, efcid], alpha) # friction elif efcid < ne + nf: # search point, friction loss, bound (rf) - start = efc_Jaref_in[worldid, efcid] - dir = efc_jv_in[worldid, efcid] + start = ctx_Jaref_in[worldid, efcid] + dir = ctx_jv_in[worldid, efcid] x = start + alpha * dir f = efc_frictionloss_in[worldid, efcid] rf = math.safe_div(f, efc_D_in[worldid, efcid]) # -bound < x < bound : quadratic if (-rf < x) and (x < rf): - quad = efc_quad_in[worldid, efcid] + quad = ctx_quad_in[worldid, efcid] # x < -bound: linear negative elif x <= -rf: quad = wp.vec3(f * (-0.5 * rf - start), -f * dir, 0.0) @@ -576,12 +400,12 @@ def linesearch_parallel_fused( # unpack quad efcid1 = contact_efc_address_in[conid, 1] efcid2 = contact_efc_address_in[conid, 2] - u0 = efc_quad_in[worldid, efcid1][0] - v0 = efc_quad_in[worldid, efcid1][1] - uu = efc_quad_in[worldid, efcid1][2] - uv = efc_quad_in[worldid, efcid2][0] - vv = efc_quad_in[worldid, efcid2][1] - dm = efc_quad_in[worldid, efcid2][2] + u0 = ctx_quad_in[worldid, efcid1][0] + v0 = ctx_quad_in[worldid, efcid1][1] + uu = ctx_quad_in[worldid, efcid1][2] + uv = ctx_quad_in[worldid, efcid2][0] + vv = ctx_quad_in[worldid, efcid2][1] + dm = ctx_quad_in[worldid, efcid2][2] # compute N, Tsqr N = u0 + alpha * v0 @@ -591,7 +415,7 @@ def linesearch_parallel_fused( if Tsqr <= 0.0: # bottom zone: quadratic cost if N < 0.0: - out += _eval_cost(efc_quad_in[worldid, efcid], alpha) + out += _eval_cost(ctx_quad_in[worldid, efcid], alpha) # otherwise regular processing else: # tangential force @@ -603,17 +427,17 @@ def linesearch_parallel_fused( pass # mu * N + T <= 0 : bottom zone elif mu * N + T <= 0.0: - out += _eval_cost(efc_quad_in[worldid, efcid], alpha) + out += _eval_cost(ctx_quad_in[worldid, efcid], alpha) # otherwise middle zone else: out += 0.5 * dm * (N - mu * T) * (N - mu * T) else: # search point - x = efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid] + x = ctx_Jaref_in[worldid, efcid] + alpha * ctx_jv_in[worldid, efcid] # active if x < 0.0: - out += _eval_cost(efc_quad_in[worldid, efcid], alpha) + out += _eval_cost(ctx_quad_in[worldid, efcid], alpha) cost_out[worldid, alphaid] = out @@ -623,30 +447,66 @@ def linesearch_parallel_best_alpha( # Model: opt_ls_iterations: int, opt_ls_parallel_min_step: float, - # Data in: - efc_done_in: wp.array(dtype=bool), # In: + ctx_done_in: wp.array(dtype=bool), cost_in: wp.array2d(dtype=float), - # Data out: - efc_alpha_out: wp.array(dtype=float), + # Out: + ctx_alpha_out: wp.array(dtype=float), ): worldid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return bestid = int(0) - best_cost = float(wp.inf) + best_cost = float(types.MJ_MAXVAL) for i in range(opt_ls_iterations): cost = cost_in[worldid, i] if cost < best_cost: best_cost = cost bestid = i - efc_alpha_out[worldid] = _log_scale(opt_ls_parallel_min_step, 1.0, opt_ls_iterations, bestid) + ctx_alpha_out[worldid] = _log_scale(opt_ls_parallel_min_step, 1.0, opt_ls_iterations, bestid) -def _linesearch_parallel(m: types.Model, d: types.Data, cost: wp.array2d(dtype=float)): +def _linesearch_parallel(m: types.Model, d: types.Data, ctx: SolverContext, cost: wp.array2d(dtype=float)): + """Parallel linesearch with setup and teardown kernels.""" + dofs_per_thread = 20 if m.nv > 50 else 50 + threads_per_efc = ceil(m.nv / dofs_per_thread) + + # quad_gauss = [gauss, search.T @ Ma - search.T @ qfrc_smooth, 0.5 * search.T @ mv] + if threads_per_efc > 1: + ctx.quad_gauss.zero_() + + wp.launch( + linesearch_prepare_gauss(m.nv, dofs_per_thread), + dim=(d.nworld, threads_per_efc), + inputs=[d.qfrc_smooth, d.efc.Ma, ctx.search, ctx.gauss, ctx.mv, ctx.done], + outputs=[ctx.quad_gauss], + ) + + # quad = [0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D] + + wp.launch( + linesearch_prepare_quad, + dim=(d.nworld, d.njmax), + inputs=[ + m.opt.impratio_invsqrt, + d.nefc, + d.contact.friction, + d.contact.dim, + d.contact.efc_address, + d.efc.type, + d.efc.id, + d.efc.D, + d.nacon, + ctx.Jaref, + ctx.jv, + ctx.done, + ], + outputs=[ctx.quad], + ) + wp.launch( linesearch_parallel_fused, dim=(d.nworld, m.opt.ls_iterations), @@ -663,13 +523,13 @@ def _linesearch_parallel(m: types.Model, d: types.Data, cost: wp.array2d(dtype=f d.efc.id, d.efc.D, d.efc.frictionloss, - d.efc.Jaref, - d.efc.jv, - d.efc.quad, - d.efc.quad_gauss, - d.efc.done, d.njmax, d.nacon, + ctx.Jaref, + ctx.jv, + ctx.quad, + ctx.quad_gauss, + ctx.done, ], outputs=[cost], ) @@ -677,8 +537,832 @@ def _linesearch_parallel(m: types.Model, d: types.Data, cost: wp.array2d(dtype=f wp.launch( linesearch_parallel_best_alpha, dim=(d.nworld), - inputs=[m.opt.ls_iterations, m.opt.ls_parallel_min_step, d.efc.done, cost], - outputs=[d.efc.alpha], + inputs=[m.opt.ls_iterations, m.opt.ls_parallel_min_step, ctx.done, cost], + outputs=[ctx.alpha], + ) + + # Teardown: update qacc, Ma, Jaref + wp.launch( + linesearch_qacc_ma, + dim=(d.nworld, m.nv), + inputs=[ctx.search, ctx.mv, ctx.alpha, ctx.done], + outputs=[d.qacc, d.efc.Ma], + ) + + wp.launch( + linesearch_jaref, + dim=(d.nworld, d.njmax), + inputs=[d.nefc, ctx.jv, ctx.alpha, ctx.done], + outputs=[ctx.Jaref], + ) + + +# kernel_analyzer: off +@wp.func +def _compute_efc_eval_pt_pyramidal( + efcid: int, + alpha: float, + ne: int, + nf: int, + # Per-row data: + efc_D: float, + efc_frictionloss: wp.array(dtype=float), + ctx_Jaref: float, + ctx_jv: float, +) -> wp.vec3: + """Compute for pyramidal cones (no elliptic contact data needed).""" + # Limit/other constraint + if efcid >= ne + nf: + x = ctx_Jaref + alpha * ctx_jv + if x < 0.0: + return _eval_pt_direct(ctx_Jaref, ctx_jv, efc_D, alpha) + return wp.vec3(0.0) + + # Friction constraint - needs quad for frictionloss computation + if efcid >= ne: + f = efc_frictionloss[efcid] + x = ctx_Jaref + alpha * ctx_jv + rf = math.safe_div(f, efc_D) + return _eval_frictionloss_pt(x, f, rf, ctx_jv, efc_D) + + # Equality constraint + return _eval_pt_direct(ctx_Jaref, ctx_jv, efc_D, alpha) + + +@wp.func +def _compute_efc_eval_pt_elliptic( + efcid: int, + alpha: float, + ne: int, + nf: int, + impratio_invsqrt: float, + # Per-row data (arrays for deferred load): + efc_type: int, + efc_D_in: wp.array(dtype=float), + efc_frictionloss: wp.array(dtype=float), + ctx_Jaref: float, + ctx_jv: float, + ctx_quad: wp.vec3, + # Contact data (for elliptic): + contact_friction: types.vec5, + efc_address0: int, + quad1: wp.vec3, + quad2: wp.vec3, +) -> wp.vec3: + """Compute for elliptic cones (includes elliptic contact data).""" + # Contact/limit/other constraints + if efcid >= ne + nf: + # Contact elliptic + if efc_type == types.ConstraintType.CONTACT_ELLIPTIC: + if efcid != efc_address0: # Not primary row + return wp.vec3(0.0) + return _eval_elliptic(impratio_invsqrt, contact_friction, ctx_quad, quad1, quad2, alpha) + + # Limit/other constraint — direct eval (no quad read) + x = ctx_Jaref + alpha * ctx_jv + if x < 0.0: + return _eval_pt_direct(ctx_Jaref, ctx_jv, efc_D_in[efcid], alpha) + return wp.vec3(0.0) + + # Friction constraint - load D and frictionloss only here + if efcid >= ne: + efc_D = efc_D_in[efcid] + f = efc_frictionloss[efcid] + x = ctx_Jaref + alpha * ctx_jv + rf = math.safe_div(f, efc_D) + return _eval_frictionloss_pt(x, f, rf, ctx_jv, efc_D) + + # Equality constraint — direct eval (no quad read) + return _eval_pt_direct(ctx_Jaref, ctx_jv, efc_D_in[efcid], alpha) + + +@wp.func +def _compute_efc_eval_pt_alpha_zero_pyramidal( + efcid: int, + ne: int, + nf: int, + # Per-row data: + efc_D: float, + efc_frictionloss: wp.array(dtype=float), + ctx_Jaref: float, + ctx_jv: float, +) -> wp.vec3: + """Optimized version for alpha=0.0, pyramidal cones.""" + # Limit/other constraint + if efcid >= ne + nf: + if ctx_Jaref < 0.0: + return _eval_pt_direct_alpha_zero(ctx_Jaref, ctx_jv, efc_D) + return wp.vec3(0.0) + + # Friction constraint - needs quad for frictionloss computation + if efcid >= ne: + f = efc_frictionloss[efcid] + rf = math.safe_div(f, efc_D) + return _eval_frictionloss_pt(ctx_Jaref, f, rf, ctx_jv, efc_D) + + # Equality constraint + return _eval_pt_direct_alpha_zero(ctx_Jaref, ctx_jv, efc_D) + + +@wp.func +def _compute_efc_eval_pt_alpha_zero_elliptic( + efcid: int, + ne: int, + nf: int, + impratio_invsqrt: float, + # Per-row data (arrays for deferred load): + efc_type: int, + efc_D_in: wp.array(dtype=float), + efc_frictionloss: wp.array(dtype=float), + ctx_Jaref: float, + ctx_jv: float, + ctx_quad: wp.vec3, + # Contact data (for elliptic): + contact_friction: types.vec5, + efc_address0: int, + quad1: wp.vec3, + quad2: wp.vec3, +) -> wp.vec3: + """Optimized version for alpha=0.0, elliptic cones.""" + # Contact/limit/other constraints + if efcid >= ne + nf: + # Contact elliptic + if efc_type == types.ConstraintType.CONTACT_ELLIPTIC: + if efcid != efc_address0: # Not primary row + return wp.vec3(0.0) + return _eval_elliptic(impratio_invsqrt, contact_friction, ctx_quad, quad1, quad2, 0.0) + + # Limit/other constraint — direct eval (no quad read) + if ctx_Jaref < 0.0: + return _eval_pt_direct_alpha_zero(ctx_Jaref, ctx_jv, efc_D_in[efcid]) + return wp.vec3(0.0) + + # Friction constraint - load D and frictionloss only here + if efcid >= ne: + efc_D = efc_D_in[efcid] + f = efc_frictionloss[efcid] + rf = math.safe_div(f, efc_D) + return _eval_frictionloss_pt(ctx_Jaref, f, rf, ctx_jv, efc_D) + + # Equality constraint — direct eval (no quad read) + return _eval_pt_direct_alpha_zero(ctx_Jaref, ctx_jv, efc_D_in[efcid]) + + +@wp.func +def _compute_efc_eval_pt_3alphas_pyramidal( + efcid: int, + lo_alpha: float, + hi_alpha: float, + mid_alpha: float, + ne: int, + nf: int, + # Per-row data: + efc_D: float, + efc_frictionloss: wp.array(dtype=float), + ctx_Jaref: float, + ctx_jv: float, +) -> tuple[wp.vec3, wp.vec3, wp.vec3]: + """Compute (cost, gradient, hessian) for 3 alphas, pyramidal cones. + + Returns a tuple of 3 vec3s for (lo_alpha, hi_alpha, mid_alpha). + Constraint types checked in order: limit/other -> friction -> equality. + """ + # Limit/other constraints: active only when x < 0 + if efcid >= ne + nf: + x_lo = ctx_Jaref + lo_alpha * ctx_jv + x_hi = ctx_Jaref + hi_alpha * ctx_jv + x_mid = ctx_Jaref + mid_alpha * ctx_jv + pt_lo, pt_hi, pt_mid = _eval_pt_direct_3alphas(ctx_Jaref, ctx_jv, efc_D, lo_alpha, hi_alpha, mid_alpha) + r_lo = wp.where(x_lo < 0.0, pt_lo, wp.vec3(0.0)) + r_hi = wp.where(x_hi < 0.0, pt_hi, wp.vec3(0.0)) + r_mid = wp.where(x_mid < 0.0, pt_mid, wp.vec3(0.0)) + return (r_lo, r_hi, r_mid) + + # Friction constraint - needs quad for frictionloss computation + if efcid >= ne: + x_lo = ctx_Jaref + lo_alpha * ctx_jv + x_hi = ctx_Jaref + hi_alpha * ctx_jv + x_mid = ctx_Jaref + mid_alpha * ctx_jv + f = efc_frictionloss[efcid] + rf = math.safe_div(f, efc_D) + return _eval_frictionloss_pt_3alphas(x_lo, x_hi, x_mid, f, rf, ctx_jv, efc_D) + + # Equality constraint: always active + return _eval_pt_direct_3alphas(ctx_Jaref, ctx_jv, efc_D, lo_alpha, hi_alpha, mid_alpha) + + +@wp.func +def _compute_efc_eval_pt_3alphas_elliptic( + efcid: int, + lo_alpha: float, + hi_alpha: float, + mid_alpha: float, + ne: int, + nf: int, + impratio_invsqrt: float, + # Per-row data (arrays for deferred load): + efc_type: int, + efc_D_in: wp.array(dtype=float), + efc_frictionloss: wp.array(dtype=float), + ctx_Jaref: float, + ctx_jv: float, + ctx_quad: wp.vec3, + # Contact data (for elliptic): + contact_friction: types.vec5, + efc_address0: int, + quad1: wp.vec3, + quad2: wp.vec3, +) -> tuple[wp.vec3, wp.vec3, wp.vec3]: + """Compute (cost, gradient, hessian) for 3 alphas, elliptic cones. + + Returns a tuple of 3 vec3s for (lo_alpha, hi_alpha, mid_alpha). + Constraint types checked in order: contact elliptic/limit/other -> friction -> equality. + """ + # x = search point, needed for friction and limit constraints + x_lo = ctx_Jaref + lo_alpha * ctx_jv + x_hi = ctx_Jaref + hi_alpha * ctx_jv + x_mid = ctx_Jaref + mid_alpha * ctx_jv + + # Contact/limit/other constraints + if efcid >= ne + nf: + # Contact elliptic: uses special elliptic cone evaluation + if efc_type == types.ConstraintType.CONTACT_ELLIPTIC: + if efcid != efc_address0: # secondary rows contribute nothing + return (wp.vec3(0.0), wp.vec3(0.0), wp.vec3(0.0)) + return ( + _eval_elliptic(impratio_invsqrt, contact_friction, ctx_quad, quad1, quad2, lo_alpha), + _eval_elliptic(impratio_invsqrt, contact_friction, ctx_quad, quad1, quad2, hi_alpha), + _eval_elliptic(impratio_invsqrt, contact_friction, ctx_quad, quad1, quad2, mid_alpha), + ) + + # Limit/other constraints — direct eval (no quad read) + efc_D = efc_D_in[efcid] + pt_lo, pt_hi, pt_mid = _eval_pt_direct_3alphas(ctx_Jaref, ctx_jv, efc_D, lo_alpha, hi_alpha, mid_alpha) + r_lo = wp.where(x_lo < 0.0, pt_lo, wp.vec3(0.0)) + r_hi = wp.where(x_hi < 0.0, pt_hi, wp.vec3(0.0)) + r_mid = wp.where(x_mid < 0.0, pt_mid, wp.vec3(0.0)) + return (r_lo, r_hi, r_mid) + + # Friction constraint - load D and frictionloss only here + if efcid >= ne: + efc_D = efc_D_in[efcid] + f = efc_frictionloss[efcid] + rf = math.safe_div(f, efc_D) + return _eval_frictionloss_pt_3alphas(x_lo, x_hi, x_mid, f, rf, ctx_jv, efc_D) + + # Equality constraint — direct eval (no quad read) + return _eval_pt_direct_3alphas(ctx_Jaref, ctx_jv, efc_D_in[efcid], lo_alpha, hi_alpha, mid_alpha) + + +# kernel_analyzer: on + +# ============================================================================= +# Iterative Linesearch +# ============================================================================= +# +# Iterative linesearch implementation using Warp's tiled execution model with +# parallel reductions over constraint (EFC) rows. +# +# Key optimizations: +# +# 1. KERNEL FUSION - Reduces kernel launch overhead by combining: +# - linesearch_jv_fused: jv = J @ search (for small nv <= 50) +# - linesearch_prepare_quad: quad coefficients (pyramidal: computed directly, +# elliptic: computed in a prepare phase with __syncthreads barrier) +# - linesearch_prepare_gauss: quad_gauss via tile reduction over DOFs +# - linesearch_qacc_ma: qacc and Ma updates at kernel end +# - linesearch_jaref: Jaref update at kernel end +# +# 2. PARALLEL REDUCTIONS - Uses wp.tile_reduce for summing cost/gradient/hessian +# contributions across EFC rows within each world. The main iteration loop +# packs 3 vec3 reductions into a single mat33 reduction for efficiency. +# +# 3. COMPILE-TIME SPECIALIZATION via factory parameters: +# - cone_type: Eliminates elliptic cone branches for pyramidal-only models +# - ls_iterations: Enables loop unrolling for the main bracket search +# - fuse_jv: Conditionally includes jv computation based on nv size +# +# 4. DIRECT EVALUATION (pyramidal only) - For equality and limit constraints, +# computes cost/gradient/hessian directly from (Jaref, jv, efc_D, alpha) +# without intermediate quad coefficients, using _eval_pt_direct functions. +# +# 5. BATCHED 3-ALPHA EVALUATION - The main iteration loop evaluates 3 alpha +# values per iteration (lo_next, hi_next, mid). Instead of calling +# _compute_efc_eval_pt 3 times per constraint row (which would repeat +# constraint type checks and data loads), we use _compute_efc_eval_pt_3alphas +# which: +# - Performs constraint type branching once per row +# - Loads efc_D, efc_frictionloss, contact data once +# - Computes x = Jaref + alpha * jv for all 3 alphas +# - For pyramidal direct evaluation: shares jvD = jv * efc_D, hessian = jv * jvD +# - For quad-based evaluation: uses _eval_pt_3alphas which computes the +# constant hessian (2.0 * quad[2]) once and reuses for all 3 alphas +# +# 6. DEFERRED DATA LOADING - efc_D and efc_frictionloss are only loaded +# inside the constraint branches where they're needed, reducing register +# pressure for other constraint types. +# +# Trade-offs: +# - Requires block synchronization (__syncthreads) for elliptic quad preparation +# - Separate kernel compilation for each (block_dim, ls_iterations, cone_type, +# fuse_jv) combination (cached by Warp) +# +# Optimizations attempted but not beneficial: +# - Caching EFC data (Jaref, jv, quad, etc.) in shared memory tiles for reuse +# across the p0, lo_in, and main iteration loops. +# +# ============================================================================= + + +@cache_kernel +def linesearch_iterative(ls_iterations: int, cone_type: types.ConeType, fuse_jv: 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). + """ + LS_ITERATIONS = ls_iterations + IS_ELLIPTIC = cone_type == types.ConeType.ELLIPTIC + FUSE_JV = fuse_jv + + # Native snippet for CUDA __syncthreads() + @wp.func_native(snippet="WP_TILE_SYNC();") + def _syncthreads(): + pass + + # Select specialized helper functions based on cone type + if IS_ELLIPTIC: + _compute_efc_eval_pt = _compute_efc_eval_pt_elliptic + _compute_efc_eval_pt_alpha_zero = _compute_efc_eval_pt_alpha_zero_elliptic + _compute_efc_eval_pt_3alphas = _compute_efc_eval_pt_3alphas_elliptic + else: + _compute_efc_eval_pt = _compute_efc_eval_pt_pyramidal + _compute_efc_eval_pt_alpha_zero = _compute_efc_eval_pt_alpha_zero_pyramidal + _compute_efc_eval_pt_3alphas = _compute_efc_eval_pt_3alphas_pyramidal + + @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_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() + + if ctx_done_in[worldid]: + return + + ne = ne_in[worldid] + nf = nf_in[worldid] + nefc = wp.min(njmax_in, nefc_in[worldid]) + + # jv = J @ search (fused for small nv) + if wp.static(FUSE_JV): + for efcid in range(tid, nefc, wp.block_dim()): + jv = float(0.0) + for i in range(nv): + jv += efc_J_in[worldid, efcid, i] * ctx_search_in[worldid, i] + ctx_jv_out[worldid, efcid] = jv + + _syncthreads() # ensure all jv values are written before reading + + # quad coefficients (elliptic contacts only, requires barrier sync) + # Non-elliptic constraints (equality, friction, limit) now use direct + # evaluation from (Jaref, jv, efc_D), avoiding quad reads entirely. + if wp.static(IS_ELLIPTIC): + # elliptic-only config values + impratio_invsqrt = opt_impratio_invsqrt[worldid % opt_impratio_invsqrt.shape[0]] + nacon = nacon_in[0] + + for efcid in range(tid, nefc, wp.block_dim()): + # Only compute and store quad for CONTACT_ELLIPTIC (needs inter-row data) + if efc_type_in[worldid, efcid] == types.ConstraintType.CONTACT_ELLIPTIC: + conid = efc_id_in[worldid, efcid] + if conid < nacon: + efcid0 = contact_efc_address_in[conid, 0] + if efcid == efcid0: + Jaref = ctx_Jaref_in[worldid, efcid] + jv = ctx_jv_in[worldid, efcid] + efc_D = efc_D_in[worldid, efcid] + + jvD = jv * efc_D + quad = wp.vec3(0.5 * Jaref * Jaref * efc_D, jvD * Jaref, 0.5 * jv * jvD) + + # primary row: accumulate secondary rows and write quad, quad1, quad2 + dim = contact_dim_in[conid] + friction = contact_friction_in[conid] + mu = friction[0] * impratio_invsqrt + + u0 = Jaref * mu + v0 = jv * mu + + uu = float(0.0) + uv = float(0.0) + vv = float(0.0) + for j in range(1, dim): + efcidj = contact_efc_address_in[conid, j] + if efcidj >= 0: + jvj = ctx_jv_in[worldid, efcidj] + jarefj = ctx_Jaref_in[worldid, efcidj] + dj = efc_D_in[worldid, efcidj] + DJj = dj * jarefj + + quad += wp.vec3(0.5 * jarefj * DJj, jvj * DJj, 0.5 * jvj * dj * jvj) + + # rescale to make primal cone circular + frictionj = friction[j - 1] + uj = jarefj * frictionj + vj = jvj * frictionj + + uu += uj * uj + uv += uj * vj + vv += vj * vj + + ctx_quad_out[worldid, efcid] = quad + + efcid1 = contact_efc_address_in[conid, 1] + ctx_quad_out[worldid, efcid1] = wp.vec3(u0, v0, uu) + + mu2 = mu * mu + efcid2 = contact_efc_address_in[conid, 2] + ctx_quad_out[worldid, efcid2] = wp.vec3(uv, vv, efc_D / (mu2 * (1.0 + mu2))) + + _syncthreads() # ensure all quads are written before reading + + # gtol (tolerance values loaded here, deferred from kernel start) + tolerance = opt_tolerance[worldid % opt_tolerance.shape[0]] + ls_tolerance = opt_ls_tolerance[worldid % opt_ls_tolerance.shape[0]] + snorm = wp.sqrt(ctx_search_dot_in[worldid]) + meaninertia = stat_meaninertia[worldid % stat_meaninertia.shape[0]] + scale = meaninertia * wp.float(nv) + gtol = tolerance * ls_tolerance * snorm * scale + + # p0 via parallel reduction + local_p0 = wp.vec3(0.0) + for efcid in range(tid, nefc, wp.block_dim()): + if wp.static(IS_ELLIPTIC): + efc_type = efc_type_in[worldid, efcid] + efc_id = 0 + contact_friction = types.vec5(0.0) + efc_addr0 = int(0) + ctx_quad = wp.vec3(0.0) + quad1 = wp.vec3(0.0) + quad2 = wp.vec3(0.0) + + if efc_type == types.ConstraintType.CONTACT_ELLIPTIC: + efc_id = efc_id_in[worldid, efcid] + contact_friction = contact_friction_in[efc_id] + efc_addr0 = contact_efc_address_in[efc_id, 0] + efc_addr1 = contact_efc_address_in[efc_id, 1] + efc_addr2 = contact_efc_address_in[efc_id, 2] + ctx_quad = ctx_quad_in[worldid, efcid] + quad1 = ctx_quad_in[worldid, efc_addr1] + quad2 = ctx_quad_in[worldid, efc_addr2] + + local_p0 += _compute_efc_eval_pt_alpha_zero( + efcid, + ne, + nf, + impratio_invsqrt, + efc_type, + efc_D_in[worldid], + efc_frictionloss_in[worldid], + ctx_Jaref_in[worldid, efcid], + ctx_jv_in[worldid, efcid], + ctx_quad, + contact_friction, + efc_addr0, + quad1, + quad2, + ) + else: + # direct evaluation for pyramidal cones (no intermediate quad) + local_p0 += _compute_efc_eval_pt_alpha_zero( + efcid, + ne, + nf, + efc_D_in[worldid, efcid], + efc_frictionloss_in[worldid], + ctx_Jaref_in[worldid, efcid], + ctx_jv_in[worldid, efcid], + ) + + # at this point, every thread has computed some contributions to p0 in local_p0 + # we now create a tile of all local_p0 contributions and reduce them to a single value + # this is done in parallel using a tile reduction + p0_tile = wp.tile(local_p0, preserve_type=True) + p0_sum = wp.tile_reduce(wp.add, p0_tile) + + # quad_gauss = [gauss, search.T @ Ma - search.T @ qfrc_smooth, 0.5 * search.T @ mv] + local_gauss = wp.vec2(0.0) # vec2 since component 0 is constant (ctx_gauss_in) + for dofid in range(tid, nv, wp.block_dim()): + search = ctx_search_in[worldid, dofid] + local_gauss += wp.vec2( + search * (efc_Ma_out[worldid, dofid] - qfrc_smooth_in[worldid, dofid]), + 0.5 * search * ctx_mv_in[worldid, dofid], + ) + + gauss_tile = wp.tile(local_gauss, preserve_type=True) + gauss_sum = wp.tile_reduce(wp.add, gauss_tile) + gauss_reduced = gauss_sum[0] + ctx_quad_gauss = wp.vec3(ctx_gauss_in[worldid], gauss_reduced[0], gauss_reduced[1]) + + # add quad_gauss contribution to p0 + p0 = wp.vec3(ctx_quad_gauss[0], ctx_quad_gauss[1], 2.0 * ctx_quad_gauss[2]) + p0_sum[0] + + # lo_in at lo_alpha_in = -p0[1] / p0[2] + lo_alpha_in = -math.safe_div(p0[1], p0[2]) + + local_lo_in = wp.vec3(0.0) + for efcid in range(tid, nefc, wp.block_dim()): + if wp.static(IS_ELLIPTIC): + efc_type = efc_type_in[worldid, efcid] + efc_id = 0 + contact_friction = types.vec5(0.0) + efc_addr0 = int(0) + ctx_quad = wp.vec3(0.0) + quad1 = wp.vec3(0.0) + quad2 = wp.vec3(0.0) + + if efc_type == types.ConstraintType.CONTACT_ELLIPTIC: + efc_id = efc_id_in[worldid, efcid] + contact_friction = contact_friction_in[efc_id] + efc_addr0 = contact_efc_address_in[efc_id, 0] + efc_addr1 = contact_efc_address_in[efc_id, 1] + efc_addr2 = contact_efc_address_in[efc_id, 2] + ctx_quad = ctx_quad_in[worldid, efcid] + quad1 = ctx_quad_in[worldid, efc_addr1] + quad2 = ctx_quad_in[worldid, efc_addr2] + + local_lo_in += _compute_efc_eval_pt( + efcid, + lo_alpha_in, + ne, + nf, + impratio_invsqrt, + efc_type, + efc_D_in[worldid], + efc_frictionloss_in[worldid], + ctx_Jaref_in[worldid, efcid], + ctx_jv_in[worldid, efcid], + ctx_quad, + contact_friction, + efc_addr0, + quad1, + quad2, + ) + else: + # direct evaluation for pyramidal cones (no intermediate quad) + local_lo_in += _compute_efc_eval_pt( + efcid, + lo_alpha_in, + ne, + nf, + efc_D_in[worldid, efcid], + efc_frictionloss_in[worldid], + ctx_Jaref_in[worldid, efcid], + ctx_jv_in[worldid, efcid], + ) + + lo_in_tile = wp.tile(local_lo_in, preserve_type=True) + lo_in_sum = wp.tile_reduce(wp.add, lo_in_tile) + lo_in = _eval_pt(ctx_quad_gauss, lo_alpha_in) + lo_in_sum[0] + + # check for initial convergence: if |derivative| < gtol, accept Newton step immediately + initial_converged = wp.abs(lo_in[1]) < gtol + + # main iterative loop - skip if already converged + if not initial_converged: + alpha = float(0.0) + + # initialize bounds + lo_less = lo_in[1] < p0[1] + lo = wp.where(lo_less, lo_in, p0) + lo_alpha = wp.where(lo_less, lo_alpha_in, 0.0) + hi = wp.where(lo_less, p0, lo_in) + hi_alpha = wp.where(lo_less, 0.0, lo_alpha_in) + + for _ in range(LS_ITERATIONS): + lo_next_alpha = lo_alpha - math.safe_div(lo[1], lo[2]) + hi_next_alpha = hi_alpha - math.safe_div(hi[1], hi[2]) + mid_alpha = 0.5 * (lo_alpha + hi_alpha) + + local_lo = wp.vec3(0.0) + local_hi = wp.vec3(0.0) + local_mid = wp.vec3(0.0) + + for efcid in range(tid, nefc, wp.block_dim()): + if wp.static(IS_ELLIPTIC): + efc_type = efc_type_in[worldid, efcid] + efc_id = 0 + contact_friction = types.vec5(0.0) + efc_addr0 = int(0) + ctx_quad = wp.vec3(0.0) + quad1 = wp.vec3(0.0) + quad2 = wp.vec3(0.0) + + if efc_type == types.ConstraintType.CONTACT_ELLIPTIC: + efc_id = efc_id_in[worldid, efcid] + contact_friction = contact_friction_in[efc_id] + efc_addr0 = contact_efc_address_in[efc_id, 0] + efc_addr1 = contact_efc_address_in[efc_id, 1] + efc_addr2 = contact_efc_address_in[efc_id, 2] + ctx_quad = ctx_quad_in[worldid, efcid] + quad1 = ctx_quad_in[worldid, efc_addr1] + quad2 = ctx_quad_in[worldid, efc_addr2] + + r_lo, r_hi, r_mid = _compute_efc_eval_pt_3alphas( + efcid, + lo_next_alpha, + hi_next_alpha, + mid_alpha, + ne, + nf, + impratio_invsqrt, + efc_type, + efc_D_in[worldid], + efc_frictionloss_in[worldid], + ctx_Jaref_in[worldid, efcid], + ctx_jv_in[worldid, efcid], + ctx_quad, + contact_friction, + efc_addr0, + quad1, + quad2, + ) + else: + # direct evaluation for pyramidal cones (no intermediate quad) + r_lo, r_hi, r_mid = _compute_efc_eval_pt_3alphas( + efcid, + lo_next_alpha, + hi_next_alpha, + mid_alpha, + ne, + nf, + efc_D_in[worldid, efcid], + efc_frictionloss_in[worldid], + ctx_Jaref_in[worldid, efcid], + ctx_jv_in[worldid, efcid], + ) + local_lo += r_lo + local_hi += r_hi + local_mid += r_mid + + # reduce with packed mat33 (3 vec3s into columns: col0=lo, col1=hi, col2=mid) + local_combined = wp.mat33( + local_lo[0], + local_hi[0], + local_mid[0], + local_lo[1], + local_hi[1], + local_mid[1], + local_lo[2], + local_hi[2], + local_mid[2], + ) + + # reduce with packed mat33 (3 vec3s into columns: col0=lo, col1=hi, col2=mid) + # this is faster than 3 vec3 reductions because it avoids synchronization barriers + combined_tile = wp.tile(local_combined, preserve_type=True) + combined_sum = wp.tile_reduce(wp.add, combined_tile) + result = combined_sum[0] + + # extract columns back to vec3s and add quad_gauss contributions + gauss_lo, gauss_hi, gauss_mid = _eval_pt_3alphas(ctx_quad_gauss, lo_next_alpha, hi_next_alpha, mid_alpha) + lo_next = gauss_lo + wp.vec3(result[0, 0], result[1, 0], result[2, 0]) + hi_next = gauss_hi + wp.vec3(result[0, 1], result[1, 1], result[2, 1]) + mid = gauss_mid + wp.vec3(result[0, 2], result[1, 2], result[2, 2]) + + # bracket swapping + # swap lo: + swap_lo_lo_next = _in_bracket(lo, lo_next) + lo = wp.where(swap_lo_lo_next, lo_next, lo) + lo_alpha = wp.where(swap_lo_lo_next, lo_next_alpha, lo_alpha) + swap_lo_mid = _in_bracket(lo, mid) + lo = wp.where(swap_lo_mid, mid, lo) + lo_alpha = wp.where(swap_lo_mid, mid_alpha, lo_alpha) + swap_lo_hi_next = _in_bracket(lo, hi_next) + lo = wp.where(swap_lo_hi_next, hi_next, lo) + lo_alpha = wp.where(swap_lo_hi_next, hi_next_alpha, lo_alpha) + swap_lo = swap_lo_lo_next or swap_lo_mid or swap_lo_hi_next + + # swap hi: + swap_hi_hi_next = _in_bracket(hi, hi_next) + hi = wp.where(swap_hi_hi_next, hi_next, hi) + hi_alpha = wp.where(swap_hi_hi_next, hi_next_alpha, hi_alpha) + swap_hi_mid = _in_bracket(hi, mid) + hi = wp.where(swap_hi_mid, mid, hi) + hi_alpha = wp.where(swap_hi_mid, mid_alpha, hi_alpha) + swap_hi_lo_next = _in_bracket(hi, lo_next) + hi = wp.where(swap_hi_lo_next, lo_next, hi) + hi_alpha = wp.where(swap_hi_lo_next, lo_next_alpha, hi_alpha) + swap_hi = swap_hi_hi_next or swap_hi_mid or swap_hi_lo_next + + # check for convergence + ls_done = (not swap_lo and not swap_hi) or (lo[1] < 0.0 and lo[1] > -gtol) or (hi[1] > 0.0 and hi[1] < gtol) + + # update alpha if improved + improved = lo[0] < p0[0] or hi[0] < p0[0] + lo_better = lo[0] < hi[0] + alpha = wp.where(improved and lo_better, lo_alpha, alpha) + alpha = wp.where(improved and not lo_better, hi_alpha, alpha) + + if ls_done: + break + else: + alpha = lo_alpha_in + + # qacc and Ma update + for dofid in range(tid, nv, wp.block_dim()): + qacc_out[worldid, dofid] += alpha * ctx_search_in[worldid, dofid] + efc_Ma_out[worldid, dofid] += alpha * ctx_mv_in[worldid, dofid] + + # Jaref update + for efcid in range(tid, nefc, wp.block_dim()): + ctx_Jaref_out[worldid, efcid] += alpha * ctx_jv_in[worldid, efcid] + + return kernel + + +def _linesearch_iterative(m: types.Model, d: types.Data, ctx: SolverContext, fuse_jv: bool): + """Iterative linesearch with parallel reductions over efc rows and dofs. + + Args: + m: Model. + d: Data. + ctx: SolverContext. + 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, + 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, ) @@ -686,75 +1370,78 @@ def _linesearch_parallel(m: types.Model, d: types.Data, cost: wp.array2d(dtype=f def linesearch_zero_jv( # Data in: nefc_in: wp.array(dtype=int), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_jv_out: wp.array2d(dtype=float), + # In: + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_jv_out: wp.array2d(dtype=float), ): worldid, efcid = wp.tid() if efcid >= nefc_in[worldid]: return - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - efc_jv_out[worldid, efcid] = 0.0 + ctx_jv_out[worldid, efcid] = 0.0 @cache_kernel def linesearch_jv_fused(nv: int, dofs_per_thread: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( # Data in: nefc_in: wp.array(dtype=int), efc_J_in: wp.array3d(dtype=float), - efc_search_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_jv_out: wp.array2d(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() if efcid >= nefc_in[worldid]: return - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return jv_out = float(0.0) if wp.static(dofs_per_thread >= nv): for i in range(wp.static(min(dofs_per_thread, nv))): - jv_out += efc_J_in[worldid, efcid, i] * efc_search_in[worldid, i] - efc_jv_out[worldid, efcid] = jv_out + jv_out += efc_J_in[worldid, efcid, i] * ctx_search_in[worldid, i] + ctx_jv_out[worldid, efcid] = jv_out else: for i in range(wp.static(dofs_per_thread)): ii = dofstart * wp.static(dofs_per_thread) + i if ii < nv: - jv_out += efc_J_in[worldid, efcid, ii] * efc_search_in[worldid, ii] - wp.atomic_add(efc_jv_out, worldid, efcid, jv_out) + jv_out += efc_J_in[worldid, efcid, ii] * ctx_search_in[worldid, ii] + wp.atomic_add(ctx_jv_out, worldid, efcid, jv_out) return kernel @cache_kernel def linesearch_prepare_gauss(nv: int, dofs_per_thread: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( # Data in: qfrc_smooth_in: wp.array2d(dtype=float), efc_Ma_in: wp.array2d(dtype=float), - efc_search_in: wp.array2d(dtype=float), - efc_gauss_in: wp.array(dtype=float), - efc_mv_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_quad_gauss_out: wp.array(dtype=wp.vec3), + # In: + ctx_search_in: wp.array2d(dtype=float), + ctx_gauss_in: wp.array(dtype=float), + ctx_mv_in: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_quad_gauss_out: wp.array(dtype=wp.vec3), ): worldid, dofstart = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return quad_gauss_1 = float(0.0) @@ -762,26 +1449,26 @@ def linesearch_prepare_gauss(nv: int, dofs_per_thread: int): if wp.static(dofs_per_thread >= nv): for i in range(wp.static(nv)): - search = efc_search_in[worldid, i] + search = ctx_search_in[worldid, i] quad_gauss_1 += search * (efc_Ma_in[worldid, i] - qfrc_smooth_in[worldid, i]) - quad_gauss_2 += 0.5 * search * efc_mv_in[worldid, i] + quad_gauss_2 += 0.5 * search * ctx_mv_in[worldid, i] - quad_gauss_0 = efc_gauss_in[worldid] - efc_quad_gauss_out[worldid] = wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2) + quad_gauss_0 = ctx_gauss_in[worldid] + ctx_quad_gauss_out[worldid] = wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2) else: for i in range(wp.static(dofs_per_thread)): ii = dofstart * wp.static(dofs_per_thread) + i if ii < nv: - search = efc_search_in[worldid, ii] + search = ctx_search_in[worldid, ii] quad_gauss_1 += search * (efc_Ma_in[worldid, ii] - qfrc_smooth_in[worldid, ii]) - quad_gauss_2 += 0.5 * search * efc_mv_in[worldid, ii] + quad_gauss_2 += 0.5 * search * ctx_mv_in[worldid, ii] if dofstart == 0: - quad_gauss_0 = efc_gauss_in[worldid] - wp.atomic_add(efc_quad_gauss_out, worldid, wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2)) + quad_gauss_0 = ctx_gauss_in[worldid] + wp.atomic_add(ctx_quad_gauss_out, worldid, wp.vec3(quad_gauss_0, quad_gauss_1, quad_gauss_2)) else: - wp.atomic_add(efc_quad_gauss_out, worldid, wp.vec3(0.0, quad_gauss_1, quad_gauss_2)) + wp.atomic_add(ctx_quad_gauss_out, worldid, wp.vec3(0.0, quad_gauss_1, quad_gauss_2)) return kernel @@ -798,23 +1485,24 @@ def linesearch_prepare_quad( efc_type_in: wp.array2d(dtype=int), efc_id_in: wp.array2d(dtype=int), efc_D_in: wp.array2d(dtype=float), - efc_Jaref_in: wp.array2d(dtype=float), - efc_jv_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), nacon_in: wp.array(dtype=int), - # Data out: - efc_quad_out: wp.array2d(dtype=wp.vec3), + # In: + ctx_Jaref_in: wp.array2d(dtype=float), + ctx_jv_in: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_quad_out: wp.array2d(dtype=wp.vec3), ): worldid, efcid = wp.tid() if efcid >= nefc_in[worldid]: return - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - Jaref = efc_Jaref_in[worldid, efcid] - jv = efc_jv_in[worldid, efcid] + Jaref = ctx_Jaref_in[worldid, efcid] + jv = ctx_jv_in[worldid, efcid] efc_D = efc_D_in[worldid, efcid] # init with scalar quadratic @@ -848,8 +1536,8 @@ def linesearch_prepare_quad( efcidj = contact_efc_address_in[conid, j] if efcidj < 0: return - jvj = efc_jv_in[worldid, efcidj] - jarefj = efc_Jaref_in[worldid, efcidj] + jvj = ctx_jv_in[worldid, efcidj] + jarefj = ctx_Jaref_in[worldid, efcidj] dj = efc_D_in[worldid, efcidj] DJj = dj * jarefj @@ -871,170 +1559,129 @@ def linesearch_prepare_quad( quad1 = wp.vec3(u0, v0, uu) efcid1 = contact_efc_address_in[conid, 1] - efc_quad_out[worldid, efcid1] = quad1 + ctx_quad_out[worldid, efcid1] = quad1 mu2 = mu * mu quad2 = wp.vec3(uv, vv, efc_D / (mu2 * (1.0 + mu2))) efcid2 = contact_efc_address_in[conid, 2] - efc_quad_out[worldid, efcid2] = quad2 + ctx_quad_out[worldid, efcid2] = quad2 - efc_quad_out[worldid, efcid] = quad + ctx_quad_out[worldid, efcid] = quad @wp.kernel def linesearch_qacc_ma( - # Data in: - efc_search_in: wp.array2d(dtype=float), - efc_mv_in: wp.array2d(dtype=float), - efc_alpha_in: wp.array(dtype=float), - efc_done_in: wp.array(dtype=bool), + # In: + ctx_search_in: wp.array2d(dtype=float), + ctx_mv_in: wp.array2d(dtype=float), + ctx_alpha_in: wp.array(dtype=float), + ctx_done_in: wp.array(dtype=bool), # Data out: qacc_out: wp.array2d(dtype=float), efc_Ma_out: wp.array2d(dtype=float), ): worldid, dofid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - alpha = efc_alpha_in[worldid] - qacc_out[worldid, dofid] += alpha * efc_search_in[worldid, dofid] - efc_Ma_out[worldid, dofid] += alpha * efc_mv_in[worldid, dofid] + alpha = ctx_alpha_in[worldid] + qacc_out[worldid, dofid] += alpha * ctx_search_in[worldid, dofid] + efc_Ma_out[worldid, dofid] += alpha * ctx_mv_in[worldid, dofid] @wp.kernel def linesearch_jaref( # Data in: nefc_in: wp.array(dtype=int), - efc_jv_in: wp.array2d(dtype=float), - efc_alpha_in: wp.array(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_Jaref_out: wp.array2d(dtype=float), + # In: + ctx_jv_in: wp.array2d(dtype=float), + ctx_alpha_in: wp.array(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_Jaref_out: wp.array2d(dtype=float), ): worldid, efcid = wp.tid() if efcid >= nefc_in[worldid]: return - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - efc_Jaref_out[worldid, efcid] += efc_alpha_in[worldid] * efc_jv_in[worldid, efcid] + ctx_Jaref_out[worldid, efcid] += ctx_alpha_in[worldid] * ctx_jv_in[worldid, efcid] @event_scope -def _linesearch(m: types.Model, d: types.Data, cost: wp.array2d(dtype=float)): - # mv = qM @ search - support.mul_m(m, d, d.efc.mv, d.efc.search, skip=d.efc.done) +def _linesearch(m: types.Model, d: types.Data, ctx: SolverContext, cost: wp.array2d(dtype=float)): + """Linesearch for constraint solver. - # jv = efc_J @ search - # TODO(team): is there a better way of doing batched matmuls with dynamic array sizes? + Args: + m: Model + d: Data + ctx: SolverContext + cost: Scratch array for storing costs per (world, alpha) - used for parallel mode + """ + # mv = qM @ search (common to both parallel and iterative) + support.mul_m(m, d, ctx.mv, ctx.search, skip=ctx.done) - # if we are only using 1 thread, it makes sense to do more dofs as we can also skip the - # init kernel. For more than 1 thread, dofs_per_thread is lower for better load balancing. + # Fuse jv computation in-kernel for small nv (iterative only) + # Parallel linesearch always requires jv pre-computed + fuse_jv = m.nv <= 50 and not m.opt.ls_parallel - if m.nv > 50: - dofs_per_thread = 20 - else: - dofs_per_thread = 50 + # jv = J @ search (when not fused into iterative kernel) + if not fuse_jv: + dofs_per_thread = 20 if m.nv > 50 else 50 + threads_per_efc = ceil(m.nv / dofs_per_thread) + + if threads_per_efc > 1: + wp.launch( + linesearch_zero_jv, + dim=(d.nworld, d.njmax), + inputs=[d.nefc, ctx.done], + outputs=[ctx.jv], + ) - threads_per_efc = ceil(m.nv / dofs_per_thread) - # we need to clear the jv array if we're doing atomic adds. - if threads_per_efc > 1: wp.launch( - linesearch_zero_jv, - dim=(d.nworld, d.njmax), - inputs=[d.nefc, d.efc.done], - outputs=[d.efc.jv], + linesearch_jv_fused(m.nv, dofs_per_thread), + dim=(d.nworld, d.njmax, threads_per_efc), + inputs=[d.nefc, d.efc.J, ctx.search, ctx.done], + outputs=[ctx.jv], ) - wp.launch( - linesearch_jv_fused(m.nv, dofs_per_thread), - dim=(d.nworld, d.njmax, threads_per_efc), - inputs=[d.nefc, d.efc.J, d.efc.search, d.efc.done], - outputs=[d.efc.jv], - ) - - # prepare quadratics - # quad_gauss = [gauss, search.T @ Ma - search.T @ qfrc_smooth, 0.5 * search.T @ mv] - if threads_per_efc > 1: - d.efc.quad_gauss.zero_() - - wp.launch( - linesearch_prepare_gauss(m.nv, dofs_per_thread), - dim=(d.nworld, threads_per_efc), - inputs=[d.qfrc_smooth, d.efc.Ma, d.efc.search, d.efc.gauss, d.efc.mv, d.efc.done], - outputs=[d.efc.quad_gauss], - ) - - # quad = [0.5 * Jaref * Jaref * efc_D, jv * Jaref * efc_D, 0.5 * jv * jv * efc_D] - wp.launch( - linesearch_prepare_quad, - dim=(d.nworld, d.njmax), - inputs=[ - m.opt.impratio_invsqrt, - d.nefc, - d.contact.friction, - d.contact.dim, - d.contact.efc_address, - d.efc.type, - d.efc.id, - d.efc.D, - d.efc.Jaref, - d.efc.jv, - d.efc.done, - d.nacon, - ], - outputs=[d.efc.quad], - ) - if m.opt.ls_parallel: - _linesearch_parallel(m, d, cost) + _linesearch_parallel(m, d, ctx, cost) else: - _linesearch_iterative(m, d) - - wp.launch( - linesearch_qacc_ma, - dim=(d.nworld, m.nv), - inputs=[d.efc.search, d.efc.mv, d.efc.alpha, d.efc.done], - outputs=[d.qacc, d.efc.Ma], - ) - - wp.launch( - linesearch_jaref, - dim=(d.nworld, d.njmax), - inputs=[d.nefc, d.efc.jv, d.efc.alpha, d.efc.done], - outputs=[d.efc.Jaref], - ) + _linesearch_iterative(m, d, ctx, fuse_jv) @wp.kernel def solve_init_efc( # Data out: solver_niter_out: wp.array(dtype=int), - efc_search_dot_out: wp.array(dtype=float), - efc_cost_out: wp.array(dtype=float), - efc_done_out: wp.array(dtype=bool), + # Out: + ctx_search_dot_out: wp.array(dtype=float), + ctx_cost_out: wp.array(dtype=float), + ctx_done_out: wp.array(dtype=bool), ): worldid = wp.tid() - efc_cost_out[worldid] = wp.inf + ctx_cost_out[worldid] = types.MJ_MAXVAL solver_niter_out[worldid] = 0 - efc_done_out[worldid] = False - efc_search_dot_out[worldid] = 0.0 + ctx_done_out[worldid] = False + ctx_search_dot_out[worldid] = 0.0 @cache_kernel def solve_init_jaref(nv: int, dofs_per_thread: int): - @nested_kernel(module="unique", enable_backward=False) + @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_in: wp.array3d(dtype=float), efc_aref_in: wp.array2d(dtype=float), - # Data out: - efc_Jaref_out: wp.array2d(dtype=float), + # Out: + ctx_Jaref_out: wp.array2d(dtype=float), ): worldid, efcid, dofstart = wp.tid() @@ -1046,7 +1693,7 @@ def solve_init_jaref(nv: int, dofs_per_thread: int): if wp.static(dofs_per_thread >= nv): for i in range(wp.static(min(dofs_per_thread, nv))): jaref += efc_J_in[worldid, efcid, i] * qacc_in[worldid, i] - efc_Jaref_out[worldid, efcid] = jaref - efc_aref_in[worldid, efcid] + ctx_Jaref_out[worldid, efcid] = jaref - efc_aref_in[worldid, efcid] else: for i in range(wp.static(dofs_per_thread)): @@ -1055,45 +1702,45 @@ def solve_init_jaref(nv: int, dofs_per_thread: int): jaref += efc_J_in[worldid, efcid, ii] * qacc_in[worldid, ii] if dofstart == 0: - wp.atomic_add(efc_Jaref_out, worldid, efcid, jaref - efc_aref_in[worldid, efcid]) + wp.atomic_add(ctx_Jaref_out, worldid, efcid, jaref - efc_aref_in[worldid, efcid]) else: - wp.atomic_add(efc_Jaref_out, worldid, efcid, jaref) + wp.atomic_add(ctx_Jaref_out, worldid, efcid, jaref) return kernel @wp.kernel def solve_init_search( - # Data in: - efc_Mgrad_in: wp.array2d(dtype=float), - # Data out: - efc_search_out: wp.array2d(dtype=float), - efc_search_dot_out: wp.array(dtype=float), + # In: + ctx_Mgrad_in: wp.array2d(dtype=float), + # Out: + ctx_search_out: wp.array2d(dtype=float), + ctx_search_dot_out: wp.array(dtype=float), ): worldid, dofid = wp.tid() - search = -1.0 * efc_Mgrad_in[worldid, dofid] - efc_search_out[worldid, dofid] = search - wp.atomic_add(efc_search_dot_out, worldid, search * search) + search = -1.0 * ctx_Mgrad_in[worldid, dofid] + ctx_search_out[worldid, dofid] = search + wp.atomic_add(ctx_search_dot_out, worldid, search * search) @wp.kernel def update_constraint_init_cost( - # Data in: - efc_cost_in: wp.array(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_gauss_out: wp.array(dtype=float), - efc_cost_out: wp.array(dtype=float), - efc_prev_cost_out: wp.array(dtype=float), + # In: + ctx_cost_in: wp.array(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_gauss_out: wp.array(dtype=float), + ctx_cost_out: wp.array(dtype=float), + ctx_prev_cost_out: wp.array(dtype=float), ): worldid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - efc_gauss_out[worldid] = 0.0 - efc_prev_cost_out[worldid] = efc_cost_in[worldid] - efc_cost_out[worldid] = 0.0 + ctx_gauss_out[worldid] = 0.0 + ctx_prev_cost_out[worldid] = ctx_cost_in[worldid] + ctx_cost_out[worldid] = 0.0 @wp.kernel @@ -1111,24 +1758,26 @@ def update_constraint_efc( efc_id_in: wp.array2d(dtype=int), efc_D_in: wp.array2d(dtype=float), efc_frictionloss_in: wp.array2d(dtype=float), - efc_Jaref_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), 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_cost_out: wp.array(dtype=float), efc_state_out: wp.array2d(dtype=int), + # Out: + ctx_cost_out: wp.array(dtype=float), ): worldid, efcid = wp.tid() if efcid >= nefc_in[worldid]: return - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return efc_D = efc_D_in[worldid, efcid] - Jaref = efc_Jaref_in[worldid, efcid] + Jaref = ctx_Jaref_in[worldid, efcid] ne = ne_in[worldid] nf = nf_in[worldid] @@ -1137,7 +1786,7 @@ def update_constraint_efc( # equality efc_force_out[worldid, efcid] = -efc_D * Jaref efc_state_out[worldid, efcid] = types.ConstraintState.QUADRATIC - wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) + wp.atomic_add(ctx_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) elif efcid < ne + nf: # friction f = efc_frictionloss_in[worldid, efcid] @@ -1145,15 +1794,15 @@ def update_constraint_efc( if Jaref <= -rf: efc_force_out[worldid, efcid] = f efc_state_out[worldid, efcid] = types.ConstraintState.LINEARNEG - wp.atomic_add(efc_cost_out, worldid, -f * (0.5 * rf + Jaref)) + wp.atomic_add(ctx_cost_out, worldid, -f * (0.5 * rf + Jaref)) elif Jaref >= rf: efc_force_out[worldid, efcid] = -f efc_state_out[worldid, efcid] = types.ConstraintState.LINEARPOS - wp.atomic_add(efc_cost_out, worldid, -f * (0.5 * rf - Jaref)) + wp.atomic_add(ctx_cost_out, worldid, -f * (0.5 * rf - Jaref)) else: efc_force_out[worldid, efcid] = -efc_D * Jaref efc_state_out[worldid, efcid] = types.ConstraintState.QUADRATIC - wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) + wp.atomic_add(ctx_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) elif efc_type_in[worldid, efcid] != types.ConstraintType.CONTACT_ELLIPTIC: # limit, frictionless contact, pyramidal friction cone contact if Jaref >= 0.0: @@ -1162,7 +1811,7 @@ def update_constraint_efc( else: efc_force_out[worldid, efcid] = -efc_D * Jaref efc_state_out[worldid, efcid] = types.ConstraintState.QUADRATIC - wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) + wp.atomic_add(ctx_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) else: # elliptic friction cone contact conid = efc_id_in[worldid, efcid] @@ -1177,7 +1826,7 @@ def update_constraint_efc( if efcid0 < 0: return - N = efc_Jaref_in[worldid, efcid0] * mu + N = ctx_Jaref_in[worldid, efcid0] * mu ufrictionj = float(0.0) TT = float(0.0) @@ -1186,7 +1835,7 @@ def update_constraint_efc( if efcidj < 0: return frictionj = friction[j - 1] - uj = efc_Jaref_in[worldid, efcidj] * frictionj + uj = ctx_Jaref_in[worldid, efcidj] * frictionj TT += uj * uj if efcid == efcidj: ufrictionj = uj * frictionj @@ -1204,7 +1853,7 @@ def update_constraint_efc( elif (mu * N + T <= 0.0) or ((T <= 0.0) and (N < 0.0)): efc_force_out[worldid, efcid] = -efc_D * Jaref efc_state_out[worldid, efcid] = types.ConstraintState.QUADRATIC - wp.atomic_add(efc_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) + wp.atomic_add(ctx_cost_out, worldid, 0.5 * efc_D * Jaref * Jaref) # middle zone else: dm = math.safe_div(efc_D_in[worldid, efcid0], mu * mu * (1.0 + mu * mu)) @@ -1214,7 +1863,7 @@ def update_constraint_efc( if efcid == efcid0: efc_force_out[worldid, efcid] = force - wp.atomic_add(efc_cost_out, worldid, 0.5 * dm * nmt * nmt) + wp.atomic_add(ctx_cost_out, worldid, 0.5 * dm * nmt * nmt) else: efc_force_out[worldid, efcid] = -math.safe_div(force, T) * ufrictionj @@ -1227,14 +1876,15 @@ def update_constraint_init_qfrc_constraint( nefc_in: wp.array(dtype=int), efc_J_in: wp.array3d(dtype=float), efc_force_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), njmax_in: int, + # In: + ctx_done_in: wp.array(dtype=bool), # Data out: qfrc_constraint_out: wp.array2d(dtype=float), ): worldid, dofid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return sum_qfrc = float(0.0) @@ -1248,21 +1898,22 @@ def update_constraint_init_qfrc_constraint( @cache_kernel def update_constraint_gauss_cost(nv: int, dofs_per_thread: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( # Data in: qacc_in: wp.array2d(dtype=float), qfrc_smooth_in: wp.array2d(dtype=float), qacc_smooth_in: wp.array2d(dtype=float), efc_Ma_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_gauss_out: wp.array(dtype=float), - efc_cost_out: wp.array(dtype=float), + # In: + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_gauss_out: wp.array(dtype=float), + ctx_cost_out: wp.array(dtype=float), ): worldid, dofstart = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return gauss_cost = float(0.0) @@ -1270,8 +1921,8 @@ def update_constraint_gauss_cost(nv: int, dofs_per_thread: int): if wp.static(dofs_per_thread >= nv): for i in range(wp.static(min(dofs_per_thread, nv))): gauss_cost += (efc_Ma_in[worldid, i] - qfrc_smooth_in[worldid, i]) * (qacc_in[worldid, i] - qacc_smooth_in[worldid, i]) - efc_gauss_out[worldid] += 0.5 * gauss_cost - efc_cost_out[worldid] += 0.5 * gauss_cost + ctx_gauss_out[worldid] += 0.5 * gauss_cost + ctx_cost_out[worldid] += 0.5 * gauss_cost else: for i in range(wp.static(dofs_per_thread)): @@ -1280,19 +1931,19 @@ 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(efc_gauss_out, worldid, gauss_cost) - wp.atomic_add(efc_cost_out, worldid, gauss_cost) + wp.atomic_add(ctx_gauss_out, worldid, gauss_cost) + wp.atomic_add(ctx_cost_out, worldid, gauss_cost) return kernel -def _update_constraint(m: types.Model, d: types.Data): +def _update_constraint(m: types.Model, d: types.Data, ctx: SolverContext | InverseContext): """Update constraint arrays after each solve iteration.""" wp.launch( update_constraint_init_cost, dim=(d.nworld), - inputs=[d.efc.cost, d.efc.done], - outputs=[d.efc.gauss, d.efc.cost, d.efc.prev_cost], + inputs=[ctx.cost, ctx.done], + outputs=[ctx.gauss, ctx.cost, ctx.prev_cost], ) wp.launch( @@ -1310,18 +1961,18 @@ def _update_constraint(m: types.Model, d: types.Data): d.efc.id, d.efc.D, d.efc.frictionloss, - d.efc.Jaref, - d.efc.done, d.nacon, + ctx.Jaref, + ctx.done, ], - outputs=[d.efc.force, d.efc.cost, d.efc.state], + outputs=[d.efc.force, d.efc.state, ctx.cost], ) # qfrc_constraint = efc_J.T @ efc_force wp.launch( update_constraint_init_qfrc_constraint, dim=(d.nworld, m.nv), - inputs=[d.nefc, d.efc.J, d.efc.force, d.efc.done, d.njmax], + inputs=[d.nefc, d.efc.J, d.efc.force, d.njmax, ctx.done], outputs=[d.qfrc_constraint], ) @@ -1338,24 +1989,24 @@ def _update_constraint(m: types.Model, d: types.Data): wp.launch( update_constraint_gauss_cost(m.nv, dofs_per_thread), dim=(d.nworld, threads_per_efc), - inputs=[d.qacc, d.qfrc_smooth, d.qacc_smooth, d.efc.Ma, d.efc.done], - outputs=[d.efc.gauss, d.efc.cost], + inputs=[d.qacc, d.qfrc_smooth, d.qacc_smooth, d.efc.Ma, ctx.done], + outputs=[ctx.gauss, ctx.cost], ) @wp.kernel def update_gradient_zero_grad_dot( - # Data in: - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_grad_dot_out: wp.array(dtype=float), + # In: + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_grad_dot_out: wp.array(dtype=float), ): worldid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - efc_grad_dot_out[worldid] = 0.0 + ctx_grad_dot_out[worldid] = 0.0 @wp.kernel @@ -1364,19 +2015,20 @@ def update_gradient_grad( qfrc_smooth_in: wp.array2d(dtype=float), qfrc_constraint_in: wp.array2d(dtype=float), efc_Ma_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_grad_out: wp.array2d(dtype=float), - efc_grad_dot_out: wp.array(dtype=float), + # In: + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_grad_out: wp.array2d(dtype=float), + ctx_grad_dot_out: wp.array(dtype=float), ): worldid, dofid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return grad = efc_Ma_in[worldid, dofid] - qfrc_smooth_in[worldid, dofid] - qfrc_constraint_in[worldid, dofid] - efc_grad_out[worldid, dofid] = grad - wp.atomic_add(efc_grad_dot_out, worldid, grad * grad) + ctx_grad_out[worldid, dofid] = grad + wp.atomic_add(ctx_grad_dot_out, worldid, grad * grad) @wp.kernel @@ -1386,18 +2038,19 @@ def update_gradient_set_h_qM_lower_sparse( qM_fullm_j: wp.array(dtype=int), # Data in: qM_in: wp.array3d(dtype=float), - efc_done_in: wp.array(dtype=bool), + # In: + ctx_done_in: wp.array(dtype=bool), # Out: - h_out: wp.array3d(dtype=float), + ctx_h_out: wp.array3d(dtype=float), ): worldid, elementid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return i = qM_fullm_i[elementid] j = qM_fullm_j[elementid] - h_out[worldid, i, j] += qM_in[worldid, 0, elementid] + ctx_h_out[worldid, i, j] += qM_in[worldid, 0, elementid] @wp.func @@ -1420,20 +2073,21 @@ def active_check(tid: int, threshold: int) -> float: def update_gradient_JTDAJ_sparse_tiled(tile_size: int, njmax: int): TILE_SIZE = tile_size - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( # Data in: nefc_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), - efc_done_in: wp.array(dtype=bool), + # In: + ctx_done_in: wp.array(dtype=bool), # Out: - h_out: wp.array3d(dtype=float), + ctx_h_out: wp.array3d(dtype=float), ): worldid, elementid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return nefc = nefc_in[worldid] @@ -1480,7 +2134,7 @@ def update_gradient_JTDAJ_sparse_tiled(tile_size: int, njmax: int): # AD: setting bounds_check to True explicitly here because for some reason it was # slower to disable it. - wp.tile_store(h_out[worldid], sum_val, offset=(offset_i, offset_j), bounds_check=True) + wp.tile_store(ctx_h_out[worldid], sum_val, offset=(offset_i, offset_j), bounds_check=True) return kernel @@ -1492,7 +2146,7 @@ def update_gradient_JTDAJ_dense_tiled(nv_padded: int, tile_size: int, njmax: int TILE_SIZE_K = tile_size - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( # Data in: nefc_in: wp.array(dtype=int), @@ -1500,13 +2154,14 @@ def update_gradient_JTDAJ_dense_tiled(nv_padded: int, tile_size: int, njmax: int efc_J_in: wp.array3d(dtype=float), efc_D_in: wp.array2d(dtype=float), efc_state_in: wp.array2d(dtype=int), - efc_done_in: wp.array(dtype=bool), + # In: + ctx_done_in: wp.array(dtype=bool), # Out: - h_out: wp.array3d(dtype=float), + ctx_h_out: wp.array3d(dtype=float), ): worldid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return nefc = nefc_in[worldid] @@ -1521,8 +2176,7 @@ def update_gradient_JTDAJ_dense_tiled(nv_padded: 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_ki = wp.tile_load(efc_J_in[worldid], shape=(TILE_SIZE_K, nv_padded), offset=(k, 0), bounds_check=False) - J_kj = J_ki + J_kj = wp.tile_load(efc_J_in[worldid], shape=(TILE_SIZE_K, nv_padded), 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) @@ -1537,11 +2191,11 @@ def update_gradient_JTDAJ_dense_tiled(nv_padded: 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_ki), wp.tile_broadcast(D_k, shape=(nv_padded, TILE_SIZE_K))) + J_ki = wp.tile_map(wp.mul, wp.tile_transpose(J_kj), wp.tile_broadcast(D_k, shape=(nv_padded, TILE_SIZE_K))) sum_val += wp.tile_matmul(J_ki, J_kj) - wp.tile_store(h_out[worldid], sum_val, bounds_check=False) + wp.tile_store(ctx_h_out[worldid], sum_val, bounds_check=False) return kernel @@ -1562,16 +2216,16 @@ def update_gradient_JTCJ( contact_worldid_in: wp.array(dtype=int), efc_J_in: wp.array3d(dtype=float), efc_D_in: wp.array2d(dtype=float), - efc_Jaref_in: wp.array2d(dtype=float), efc_state_in: wp.array2d(dtype=int), - efc_done_in: wp.array(dtype=bool), 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), + ctx_h_out: wp.array3d(dtype=float), ): conid_start, elementid = wp.tid() @@ -1585,7 +2239,7 @@ def update_gradient_JTCJ( return worldid = contact_worldid_in[conid] - if efc_done_in[worldid]: + if ctx_done_in[worldid]: continue condim = contact_dim_in[conid] @@ -1610,13 +2264,13 @@ def update_gradient_JTCJ( if dm == 0.0: continue - n = efc_Jaref_in[worldid, efcid0] * mu + n = ctx_Jaref_in[worldid, efcid0] * mu u = types.vec6(n, 0.0, 0.0, 0.0, 0.0, 0.0) tt = float(0.0) for j in range(1, condim): efcidj = contact_efc_address_in[conid, j] - uj = efc_Jaref_in[worldid, efcidj] * fri[j - 1] + uj = ctx_Jaref_in[worldid, efcidj] * fri[j - 1] tt += uj * uj u[j] = uj @@ -1684,51 +2338,51 @@ def update_gradient_JTCJ( if dim1id != dim2id: h += hcone * efc_J12 * efc_J21 - h_out[worldid, dof1id, dof2id] += h + ctx_h_out[worldid, dof1id, dof2id] += h @cache_kernel def update_gradient_cholesky(tile_size: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( - # Data in: - efc_grad_in: wp.array2d(dtype=float), + # In: + ctx_grad_in: wp.array2d(dtype=float), h_in: wp.array3d(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_Mgrad_out: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_Mgrad_out: wp.array2d(dtype=float), ): worldid = wp.tid() TILE_SIZE = wp.static(tile_size) - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return mat_tile = wp.tile_load(h_in[worldid], shape=(TILE_SIZE, TILE_SIZE)) fact_tile = wp.tile_cholesky(mat_tile) - input_tile = wp.tile_load(efc_grad_in[worldid], shape=TILE_SIZE) + input_tile = wp.tile_load(ctx_grad_in[worldid], shape=TILE_SIZE) output_tile = wp.tile_cholesky_solve(fact_tile, input_tile) - wp.tile_store(efc_Mgrad_out[worldid], output_tile) + wp.tile_store(ctx_Mgrad_out[worldid], output_tile) return kernel @cache_kernel def update_gradient_cholesky_blocked(tile_size: int, matrix_size: int): - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def kernel( - # Data in: - efc_grad_in: wp.array3d(dtype=float), - h_in: wp.array3d(dtype=float), - efc_done_in: wp.array(dtype=bool), - hfactor: wp.array3d(dtype=float), - # Data out: - efc_Mgrad_out: wp.array3d(dtype=float), + # In: + ctx_done_in: wp.array(dtype=bool), + ctx_grad_in: wp.array3d(dtype=float), + ctx_h_in: wp.array3d(dtype=float), + ctx_hfactor: wp.array3d(dtype=float), + # Out: + ctx_Mgrad_out: wp.array3d(dtype=float), ): worldid = wp.tid() TILE_SIZE = wp.static(tile_size) - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return # We need matrix size both as a runtime input as well as a static input: @@ -1736,41 +2390,41 @@ def update_gradient_cholesky_blocked(tile_size: int, matrix_size: int): # runtime input is needed for the loop bounds, otherwise warp will unroll # unconditionally leading to shared memory capacity issues. - wp.static(create_blocked_cholesky_func(TILE_SIZE))(h_in[worldid], matrix_size, hfactor[worldid]) + wp.static(create_blocked_cholesky_func(TILE_SIZE))(ctx_h_in[worldid], matrix_size, ctx_hfactor[worldid]) wp.static(create_blocked_cholesky_solve_func(TILE_SIZE, matrix_size))( - hfactor[worldid], efc_grad_in[worldid], matrix_size, efc_Mgrad_out[worldid] + ctx_hfactor[worldid], ctx_grad_in[worldid], matrix_size, ctx_Mgrad_out[worldid] ) return kernel @wp.kernel -def padding_h(nv: int, efc_done_in: wp.array(dtype=bool), h_out: wp.array3d(dtype=float)): +def padding_h(nv: int, ctx_done_in: wp.array(dtype=bool), ctx_h_out: wp.array3d(dtype=float)): worldid, elementid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return dofid = nv + elementid - h_out[worldid, dofid, dofid] = 1.0 + ctx_h_out[worldid, dofid, dofid] = 1.0 -def _update_gradient(m: types.Model, d: types.Data, h: wp.array3d(dtype=float), hfactor: wp.array3d(dtype=float)): +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=[d.efc.done], outputs=[d.efc.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, d.efc.done], - outputs=[d.efc.grad, d.efc.grad_dot], + inputs=[d.qfrc_smooth, d.qfrc_constraint, d.efc.Ma, ctx.done], + outputs=[ctx.grad, ctx.grad_dot], ) if m.opt.solver == types.SolverType.CG: - smooth.solve_m(m, d, d.efc.Mgrad, d.efc.grad) + 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 m.opt.is_sparse: + if 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) wp.launch_tiled( @@ -1781,17 +2435,17 @@ def _update_gradient(m: types.Model, d: types.Data, h: wp.array3d(dtype=float), d.efc.J, d.efc.D, d.efc.state, - d.efc.done, + ctx.done, ], - outputs=[h], + 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), - inputs=[m.qM_fullm_i, m.qM_fullm_j, d.qM, d.efc.done], - outputs=[h], + inputs=[m.qM_fullm_i, m.qM_fullm_j, d.qM, ctx.done], + outputs=[ctx.h], ) else: nv_padded = d.efc.J.shape[2] @@ -1804,9 +2458,9 @@ def _update_gradient(m: types.Model, d: types.Data, h: wp.array3d(dtype=float), d.efc.J, d.efc.D, d.efc.state, - d.efc.done, + ctx.done, ], - outputs=[h], + outputs=[ctx.h], block_dim=m.block_dim.update_gradient_JTDAJ_dense, ) @@ -1849,39 +2503,38 @@ def _update_gradient(m: types.Model, d: types.Data, h: wp.array3d(dtype=float), d.contact.worldid, d.efc.J, d.efc.D, - d.efc.Jaref, d.efc.state, - d.efc.done, d.naconmax, d.nacon, + ctx.Jaref, + ctx.done, nblocks_perblock, dim_block, ], - outputs=[h], + outputs=[ctx.h], ) - # TODO(team): Define good threshold for blocked vs non-blocked cholesky if m.nv <= _BLOCK_CHOLESKY_DIM: wp.launch_tiled( update_gradient_cholesky(m.nv), dim=d.nworld, - inputs=[d.efc.grad, h, d.efc.done], - outputs=[d.efc.Mgrad], + 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, d.efc.done], - outputs=[h], + 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=[d.efc.grad.reshape(shape=(d.nworld, d.efc.grad.shape[1], 1)), h, d.efc.done, hfactor], - outputs=[d.efc.Mgrad.reshape(shape=(d.nworld, d.efc.Mgrad.shape[1], 1))], + 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, ) else: @@ -1890,91 +2543,91 @@ def _update_gradient(m: types.Model, d: types.Data, h: wp.array3d(dtype=float), @wp.kernel def solve_prev_grad_Mgrad( - # Data in: - efc_grad_in: wp.array2d(dtype=float), - efc_Mgrad_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_prev_grad_out: wp.array2d(dtype=float), - efc_prev_Mgrad_out: wp.array2d(dtype=float), + # In: + ctx_grad_in: wp.array2d(dtype=float), + ctx_Mgrad_in: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_prev_grad_out: wp.array2d(dtype=float), + ctx_prev_Mgrad_out: wp.array2d(dtype=float), ): worldid, dofid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - efc_prev_grad_out[worldid, dofid] = efc_grad_in[worldid, dofid] - efc_prev_Mgrad_out[worldid, dofid] = efc_Mgrad_in[worldid, dofid] + ctx_prev_grad_out[worldid, dofid] = ctx_grad_in[worldid, dofid] + ctx_prev_Mgrad_out[worldid, dofid] = ctx_Mgrad_in[worldid, dofid] @wp.kernel def solve_beta( # Model: nv: int, - # Data in: - efc_grad_in: wp.array2d(dtype=float), - efc_Mgrad_in: wp.array2d(dtype=float), - efc_prev_grad_in: wp.array2d(dtype=float), - efc_prev_Mgrad_in: wp.array2d(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_beta_out: wp.array(dtype=float), + # In: + ctx_grad_in: wp.array2d(dtype=float), + ctx_Mgrad_in: wp.array2d(dtype=float), + ctx_prev_grad_in: wp.array2d(dtype=float), + ctx_prev_Mgrad_in: wp.array2d(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_beta_out: wp.array(dtype=float), ): worldid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return beta_num = float(0.0) beta_den = float(0.0) for dofid in range(nv): - prev_Mgrad = efc_prev_Mgrad_in[worldid][dofid] - beta_num += efc_grad_in[worldid, dofid] * (efc_Mgrad_in[worldid, dofid] - prev_Mgrad) - beta_den += efc_prev_grad_in[worldid, dofid] * prev_Mgrad + prev_Mgrad = ctx_prev_Mgrad_in[worldid][dofid] + beta_num += ctx_grad_in[worldid, dofid] * (ctx_Mgrad_in[worldid, dofid] - prev_Mgrad) + beta_den += ctx_prev_grad_in[worldid, dofid] * prev_Mgrad - efc_beta_out[worldid] = wp.max(0.0, beta_num / wp.max(types.MJ_MINVAL, beta_den)) + ctx_beta_out[worldid] = wp.max(0.0, beta_num / wp.max(types.MJ_MINVAL, beta_den)) @wp.kernel def solve_zero_search_dot( - # Data in: - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_search_dot_out: wp.array(dtype=float), + # In: + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_search_dot_out: wp.array(dtype=float), ): worldid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - efc_search_dot_out[worldid] = 0.0 + ctx_search_dot_out[worldid] = 0.0 @wp.kernel def solve_search_update( # Model: opt_solver: int, - # Data in: - efc_Mgrad_in: wp.array2d(dtype=float), - efc_search_in: wp.array2d(dtype=float), - efc_beta_in: wp.array(dtype=float), - efc_done_in: wp.array(dtype=bool), - # Data out: - efc_search_out: wp.array2d(dtype=float), - efc_search_dot_out: wp.array(dtype=float), + # In: + ctx_Mgrad_in: wp.array2d(dtype=float), + ctx_search_in: wp.array2d(dtype=float), + ctx_beta_in: wp.array(dtype=float), + ctx_done_in: wp.array(dtype=bool), + # Out: + ctx_search_out: wp.array2d(dtype=float), + ctx_search_dot_out: wp.array(dtype=float), ): worldid, dofid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return - search = -1.0 * efc_Mgrad_in[worldid, dofid] + search = -1.0 * ctx_Mgrad_in[worldid, dofid] if opt_solver == types.SolverType.CG: - search += efc_beta_in[worldid] * efc_search_in[worldid, dofid] + search += ctx_beta_in[worldid] * ctx_search_in[worldid, dofid] - efc_search_out[worldid, dofid] = search - wp.atomic_add(efc_search_dot_out, worldid, search * search) + ctx_search_out[worldid, dofid] = search + wp.atomic_add(ctx_search_dot_out, worldid, search * search) @wp.kernel @@ -1983,32 +2636,34 @@ def solve_done( nv: int, opt_tolerance: wp.array(dtype=float), opt_iterations: int, - stat_meaninertia: float, - # Data in: - efc_grad_dot_in: wp.array(dtype=float), - efc_cost_in: wp.array(dtype=float), - efc_prev_cost_in: wp.array(dtype=float), - efc_done_in: wp.array(dtype=bool), + stat_meaninertia: wp.array(dtype=float), + # In: + ctx_grad_dot_in: wp.array(dtype=float), + ctx_cost_in: wp.array(dtype=float), + ctx_prev_cost_in: wp.array(dtype=float), + ctx_done_in: wp.array(dtype=bool), # Data out: solver_niter_out: wp.array(dtype=int), - efc_done_out: wp.array(dtype=bool), + # Out: nsolving_out: wp.array(dtype=int), + ctx_done_out: wp.array(dtype=bool), ): worldid = wp.tid() - if efc_done_in[worldid]: + if ctx_done_in[worldid]: return solver_niter_out[worldid] += 1 tolerance = opt_tolerance[worldid % opt_tolerance.shape[0]] + meaninertia = stat_meaninertia[worldid % stat_meaninertia.shape[0]] - improvement = _rescale(nv, stat_meaninertia, efc_prev_cost_in[worldid] - efc_cost_in[worldid]) - gradient = _rescale(nv, stat_meaninertia, wp.sqrt(efc_grad_dot_in[worldid])) + improvement = _rescale(nv, meaninertia, ctx_prev_cost_in[worldid] - ctx_cost_in[worldid]) + gradient = _rescale(nv, meaninertia, wp.sqrt(ctx_grad_dot_in[worldid])) done = (improvement < tolerance) or (gradient < tolerance) if done or solver_niter_out[worldid] == opt_iterations: # if the solver has converged or the maximum number of iterations has been reached then # mark this world as done and remove it from the number of unconverged worlds - efc_done_out[worldid] = True + ctx_done_out[worldid] = True wp.atomic_add(nsolving_out, 0, -1) @@ -2016,39 +2671,39 @@ def solve_done( def _solver_iteration( m: types.Model, d: types.Data, - h: wp.array3d(dtype=float), - hfactor: wp.array3d(dtype=float), + ctx: SolverContext, step_size_cost: wp.array2d(dtype=float), + nsolving: wp.array(dtype=int), ): - _linesearch(m, d, step_size_cost) + _linesearch(m, d, ctx, step_size_cost) if m.opt.solver == types.SolverType.CG: wp.launch( solve_prev_grad_Mgrad, dim=(d.nworld, m.nv), - inputs=[d.efc.grad, d.efc.Mgrad, d.efc.done], - outputs=[d.efc.prev_grad, d.efc.prev_Mgrad], + inputs=[ctx.grad, ctx.Mgrad, ctx.done], + outputs=[ctx.prev_grad, ctx.prev_Mgrad], ) - _update_constraint(m, d) - _update_gradient(m, d, h, hfactor) + _update_constraint(m, d, ctx) + _update_gradient(m, d, ctx) # polak-ribiere if m.opt.solver == types.SolverType.CG: wp.launch( solve_beta, dim=d.nworld, - inputs=[m.nv, d.efc.grad, d.efc.Mgrad, d.efc.prev_grad, d.efc.prev_Mgrad, d.efc.done], - outputs=[d.efc.beta], + inputs=[m.nv, ctx.grad, ctx.Mgrad, ctx.prev_grad, ctx.prev_Mgrad, ctx.done], + outputs=[ctx.beta], ) - wp.launch(solve_zero_search_dot, dim=(d.nworld), inputs=[d.efc.done], outputs=[d.efc.search_dot]) + wp.launch(solve_zero_search_dot, dim=(d.nworld), inputs=[ctx.done], outputs=[ctx.search_dot]) wp.launch( solve_search_update, dim=(d.nworld, m.nv), - inputs=[m.opt.solver, d.efc.Mgrad, d.efc.search, d.efc.beta, d.efc.done], - outputs=[d.efc.search, d.efc.search_dot], + inputs=[m.opt.solver, ctx.Mgrad, ctx.search, ctx.beta, ctx.done], + outputs=[ctx.search, ctx.search_dot], ) wp.launch( @@ -2059,23 +2714,21 @@ def _solver_iteration( m.opt.tolerance, m.opt.iterations, m.stat.meaninertia, - d.efc.grad_dot, - d.efc.cost, - d.efc.prev_cost, - d.efc.done, + ctx.grad_dot, + ctx.cost, + ctx.prev_cost, + ctx.done, ], - outputs=[d.solver_niter, d.efc.done, d.nsolving], + outputs=[d.solver_niter, nsolving, ctx.done], ) -def create_context( - m: types.Model, d: types.Data, h: wp.array3d(dtype=float), hfactor: wp.array3d(dtype=float), grad: bool = True -): +def init_context(m: types.Model, d: types.Data, ctx: SolverContext | InverseContext, grad: bool = True): # initialize some efc arrays wp.launch( solve_init_efc, dim=(d.nworld), - outputs=[d.solver_niter, d.efc.search_dot, d.efc.cost, d.efc.done], + outputs=[d.solver_niter, ctx.search_dot, ctx.cost, ctx.done], ) # jaref = d.efc_J @ d.qacc - d.efc_aref @@ -2091,22 +2744,22 @@ def create_context( threads_per_efc = ceil(m.nv / dofs_per_thread) # we need to clear the jaref array if we're doing atomic adds. if threads_per_efc > 1: - d.efc.Jaref.zero_() + ctx.Jaref.zero_() wp.launch( solve_init_jaref(m.nv, dofs_per_thread), dim=(d.nworld, d.njmax, threads_per_efc), inputs=[d.nefc, d.qacc, d.efc.J, d.efc.aref], - outputs=[d.efc.Jaref], + outputs=[ctx.Jaref], ) # Ma = qM @ qacc - support.mul_m(m, d, d.efc.Ma, d.qacc, skip=d.efc.done) + support.mul_m(m, d, d.efc.Ma, d.qacc, skip=ctx.done) - _update_constraint(m, d) + _update_constraint(m, d, ctx) if grad: - _update_gradient(m, d, h, hfactor) + _update_gradient(m, d, ctx) @event_scope @@ -2115,40 +2768,31 @@ def solve(m: types.Model, d: types.Data): wp.copy(d.qacc, d.qacc_smooth) d.solver_niter.fill_(0) else: - _solve(m, d) + ctx = create_solver_context(m, d) + _solve(m, d, ctx) -def _solve(m: types.Model, d: types.Data): +def _solve(m: types.Model, d: types.Data, ctx: SolverContext): """Finds forces that satisfy constraints.""" if not (m.opt.disableflags & types.DisableBit.WARMSTART): wp.copy(d.qacc, d.qacc_warmstart) else: wp.copy(d.qacc, d.qacc_smooth) - # Newton solver Hessian - if m.opt.solver == types.SolverType.NEWTON: - h = wp.zeros((d.nworld, m.nv_pad, m.nv_pad), dtype=float) - if m.nv > _BLOCK_CHOLESKY_DIM: - hfactor = wp.zeros((d.nworld, m.nv_pad, m.nv_pad), dtype=float) - else: - hfactor = wp.empty((d.nworld, 0, 0), dtype=float) - else: - h = wp.empty((d.nworld, 0, 0), dtype=float) - hfactor = wp.empty((d.nworld, 0, 0), dtype=float) - - # create context - create_context(m, d, h, hfactor, grad=True) + # context + init_context(m, d, ctx, grad=True) # search = -Mgrad wp.launch( solve_init_search, dim=(d.nworld, m.nv), - inputs=[d.efc.Mgrad], - outputs=[d.efc.search, d.efc.search_dot], + inputs=[ctx.Mgrad], + outputs=[ctx.search, ctx.search_dot], ) step_size_cost = wp.empty((d.nworld, m.opt.ls_iterations if m.opt.ls_parallel else 0), dtype=float) + nsolving = wp.full(shape=(1,), value=d.nworld, dtype=int) if m.opt.iterations != 0 and m.opt.graph_conditional: # Note: the iteration kernel (indicated by while_body) is repeatedly launched # as long as condition_iteration is not zero. @@ -2157,11 +2801,12 @@ def _solve(m: types.Model, d: types.Data): # When the number of iterations reaches m.opt.iterations, solver_niter # becomes zero and all worlds are marked as converged to avoid an infinite loop. # note: we only launch the iteration kernel if everything is not done - d.nsolving.fill_(d.nworld) - wp.capture_while(d.nsolving, while_body=_solver_iteration, m=m, d=d, h=h, hfactor=hfactor, step_size_cost=step_size_cost) + wp.capture_while( + nsolving, while_body=_solver_iteration, m=m, d=d, ctx=ctx, step_size_cost=step_size_cost, nsolving=nsolving + ) else: # This branch is mostly for when JAX is used as it is currently not compatible # with CUDA graph conditional. # It should be removed when JAX becomes compatible. for _ in range(m.opt.iterations): - _solver_iteration(m, d, h, hfactor, step_size_cost) + _solver_iteration(m, d, ctx, step_size_cost, nsolving) 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 429c7a73..b45472c4 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py @@ -23,21 +23,21 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import Data 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 TileSet 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 -from mujoco.mjx.third_party.mujoco_warp._src.warp_util import nested_kernel wp.set_module_options({"enable_backward": False}) @cache_kernel -def mul_m_sparse_diag(check_skip: bool): - @nested_kernel(module="unique", enable_backward=False) - def _mul_m_sparse_diag( +def mul_m_sparse(check_skip: bool): + @wp.kernel(module="unique") + def _mul_m_sparse( # Model: - dof_Madr: wp.array(dtype=int), + qM_mulm_rowadr: wp.array(dtype=int), + qM_mulm_col: wp.array(dtype=int), + qM_mulm_madr: wp.array(dtype=int), # Data in: qM_in: wp.array3d(dtype=float), # In: @@ -46,26 +46,33 @@ def mul_m_sparse_diag(check_skip: bool): # Out: res: wp.array2d(dtype=float), ): - """Diagonal update for sparse matmul.""" + """Sparse matmul: one thread per DOF, gather-based (no atomics).""" worldid, dofid = wp.tid() if wp.static(check_skip): if skip[worldid]: return - res[worldid, dofid] = qM_in[worldid, 0, dof_Madr[dofid]] * vec[worldid, dofid] + # Gather all contributions (diagonal + off-diagonal) + acc = float(0.0) + start = qM_mulm_rowadr[dofid] + end = qM_mulm_rowadr[dofid + 1] + for k in range(start, end): + col = qM_mulm_col[k] + madr = qM_mulm_madr[k] + acc += qM_in[worldid, 0, madr] * vec[worldid, col] - return _mul_m_sparse_diag + res[worldid, dofid] = acc + + return _mul_m_sparse @cache_kernel -def mul_m_sparse_ij(check_skip: bool): - @nested_kernel(module="unique", enable_backward=False) - def _mul_m_sparse_ij( - # Model: - qM_mulm_i: wp.array(dtype=int), - qM_mulm_j: wp.array(dtype=int), - qM_madr_ij: wp.array(dtype=int), +def mul_m_dense(nv: int, check_skip: bool): + """Simple SIMT dense matmul: one thread per output element.""" + + @wp.kernel(module="unique") + def _mul_m_dense( # Data in: qM_in: wp.array3d(dtype=float), # In: @@ -74,52 +81,16 @@ def mul_m_sparse_ij(check_skip: bool): # Out: res: wp.array2d(dtype=float), ): - """Off-diagonal update for sparse matmul.""" - worldid, elementid = wp.tid() + worldid, i = wp.tid() if wp.static(check_skip): if skip[worldid]: return - i = qM_mulm_i[elementid] - j = qM_mulm_j[elementid] - madr_ij = qM_madr_ij[elementid] - - qM_ij = qM_in[worldid, 0, madr_ij] - - wp.atomic_add(res[worldid], i, qM_ij * vec[worldid, j]) - wp.atomic_add(res[worldid], j, qM_ij * vec[worldid, i]) - - return _mul_m_sparse_ij - - -@cache_kernel -def mul_m_dense(tile: TileSet, check_skip: bool): - """Returns a matmul kernel for some tile size.""" - - @nested_kernel(module="unique", enable_backward=False) - def _mul_m_dense( - # Data In: - qM_in: wp.array3d(dtype=float), - # In: - adr: wp.array(dtype=int), - vec: wp.array3d(dtype=float), - skip: wp.array(dtype=bool), - # Out: - res: wp.array3d(dtype=float), - ): - worldid, nodeid = wp.tid() - TILE_SIZE = wp.static(tile.size) - - if wp.static(check_skip): - if skip[worldid]: - return - - dofid = adr[nodeid] - qM_tile = wp.tile_load(qM_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid), bounds_check=False) - vec_tile = wp.tile_load(vec[worldid], shape=(TILE_SIZE, 1), offset=(dofid, 0), bounds_check=False) - res_tile = wp.tile_matmul(qM_tile, vec_tile) - wp.tile_store(res[worldid], res_tile, offset=(dofid, 0), bounds_check=False) + acc = float(0.0) + for j in range(wp.static(nv)): + acc += qM_in[worldid, i, j] * vec[worldid, j] + res[worldid, i] = acc return _mul_m_dense @@ -149,36 +120,21 @@ def mul_m( if M is None: M = d.qM - if m.opt.is_sparse: + if m.is_sparse: wp.launch( - mul_m_sparse_diag(check_skip), + mul_m_sparse(check_skip), dim=(d.nworld, m.nv), - inputs=[m.dof_Madr, M, vec, skip], - outputs=[res], - ) - - wp.launch( - mul_m_sparse_ij(check_skip), - dim=(d.nworld, m.qM_madr_ij.size), - inputs=[m.qM_mulm_i, m.qM_mulm_j, m.qM_madr_ij, M, vec, skip], + inputs=[m.qM_mulm_rowadr, m.qM_mulm_col, m.qM_mulm_madr, M, vec, skip], outputs=[res], ) else: - for tile in m.qM_tiles: - wp.launch_tiled( - mul_m_dense(tile, check_skip), - dim=(d.nworld, tile.adr.size), - inputs=[ - M, - tile.adr, - # note reshape: tile_matmul expects 2d input - vec.reshape(vec.shape + (1,)), - skip, - ], - outputs=[res.reshape(res.shape + (1,))], - block_dim=m.block_dim.mul_m_dense, - ) + wp.launch( + mul_m_dense(m.nv, check_skip), + dim=(d.nworld, m.nv), + inputs=[M, vec, skip], + outputs=[res], + ) @wp.kernel @@ -408,7 +364,7 @@ def transform_force(frc: wp.spatial_vector, offset: wp.vec3) -> wp.spatial_vecto @wp.func -def jac( +def jac_dof( # Model: body_parentid: wp.array(dtype=int), body_rootid: wp.array(dtype=int), @@ -446,8 +402,77 @@ def jac( return jacp, jacr +@cache_kernel +def _make_jac_kernel(has_jacp: bool, has_jacr: bool): + @wp.kernel(module="unique", enable_backward=False) + def _jac( + # Model: + body_parentid: wp.array(dtype=int), + body_rootid: wp.array(dtype=int), + dof_bodyid: wp.array(dtype=int), + # Data in: + subtree_com_in: wp.array2d(dtype=wp.vec3), + cdof_in: wp.array2d(dtype=wp.spatial_vector), + # In: + point_in: wp.array(dtype=wp.vec3), + bodyid_in: wp.array(dtype=int), + # Out: + jacp_out: wp.array3d(dtype=float), + jacr_out: wp.array3d(dtype=float), + ): + worldid, dofid = wp.tid() + + jacp_val, jacr_val = jac_dof( + body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, point_in[worldid], bodyid_in[worldid], dofid, worldid + ) + + if wp.static(has_jacp): + jacp_out[worldid, 0, dofid] = jacp_val[0] + jacp_out[worldid, 1, dofid] = jacp_val[1] + jacp_out[worldid, 2, dofid] = jacp_val[2] + + if wp.static(has_jacr): + jacr_out[worldid, 0, dofid] = jacr_val[0] + jacr_out[worldid, 1, dofid] = jacr_val[1] + jacr_out[worldid, 2, dofid] = jacr_val[2] + + return _jac + + +@event_scope +def jac( + m: Model, + d: Data, + jacp: wp.array | None, # wp.array3d(dtype=float) + jacr: wp.array | None, # wp.array3d(dtype=float) + point: wp.array(dtype=wp.vec3), + body: wp.array(dtype=int), +): + """Compute translational and rotational Jacobian for point on body. + + Args: + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state (device). + jacp: Output translational Jacobian (optional). + jacr: Output rotational Jacobian (optional). + point: 3D point in global coordinates. + body: Body ID for each world. + """ + kernel = _make_jac_kernel(jacp is not None, jacr is not None) + + jacp_arr = jacp or wp.empty((0, 0, 0), dtype=float) + jacr_arr = jacr or wp.empty((0, 0, 0), dtype=float) + + wp.launch( + kernel, + dim=(d.nworld, m.nv), + inputs=[m.body_parentid, m.body_rootid, m.dof_bodyid, d.subtree_com, d.cdof, point, body], + outputs=[jacp_arr, jacr_arr], + ) + + @wp.func -def jac_dot( +def jac_dot_dof( # Model: body_parentid: wp.array(dtype=int), body_rootid: wp.array(dtype=int), @@ -529,7 +554,7 @@ def get_state(m: Model, d: Data, state: wp.array2d(dtype=float), sig: int, activ if sig >= (1 << State.NSTATE): raise ValueError(f"invalid state signature {sig} >= 2^mjNSTATE") - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def _get_state( # Model: nq: int, @@ -668,7 +693,7 @@ def set_state(m: Model, d: Data, state: wp.array2d(dtype=float), sig: int, activ if sig >= (1 << State.NSTATE): raise ValueError(f"invalid state signature {sig} >= 2^mjNSTATE") - @nested_kernel(module="unique", enable_backward=False) + @wp.kernel(module="unique", enable_backward=False) def _set_state( # Model: nq: int, 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 99cdbb97..8d50fc8f 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py @@ -32,6 +32,9 @@ MJ_MAX_EPAFACES = 5 TILE_SIZE_JTDAJ_SPARSE = 16 TILE_SIZE_JTDAJ_DENSE = 16 +# TODO(team): remove after mjwarp depends on warp-lang >= 1.12 in pyproject.toml +TEXTURE_DTYPE = wp.Texture2D if hasattr(wp, "Texture2D") else int + # TODO(team): add check that all wp.launch_tiled 'block_dim' settings are configurable @dataclasses.dataclass @@ -61,9 +64,9 @@ class BlockDim: update_gradient_cholesky_blocked: int = 32 update_gradient_JTDAJ_sparse: int = 64 update_gradient_JTDAJ_dense: int = 96 - linesearch_iterative: int = 64 - # support - mul_m_dense: int = 32 + linesearch_iterative: int = 32 + # derivative + qderiv_actuator_dense: int = 32 class BroadphaseType(enum.IntEnum): @@ -114,6 +117,23 @@ class CamLightType(enum.IntEnum): TARGETBODYCOM = mujoco.mjtCamLight.mjCAMLIGHT_TARGETBODYCOM +class ProjectionType(enum.IntEnum): + """Type of camera projection. + + Attributes: + PERSPECTIVE: perspective projection + ORTHOGRAPHIC: orthographic projection + """ + + # TODO(team): remove after mjwarp depends on mujoco > 3.4.0 in pyproject.toml + if hasattr(mujoco, "mjtProjection"): + PERSPECTIVE = mujoco.mjtProjection.mjPROJ_PERSPECTIVE + ORTHOGRAPHIC = mujoco.mjtProjection.mjPROJ_ORTHOGRAPHIC + else: + PERSPECTIVE = 0 + ORTHOGRAPHIC = 1 + + class DataType(enum.IntFlag): """Sensor data types. @@ -164,7 +184,8 @@ class DisableBit(enum.IntFlag): REFSAFE = mujoco.mjtDisableBit.mjDSBL_REFSAFE SENSOR = mujoco.mjtDisableBit.mjDSBL_SENSOR EULERDAMP = mujoco.mjtDisableBit.mjDSBL_EULERDAMP - # unsupported: MIDPHASE, AUTORESET, NATIVECCD, ISLAND + NATIVECCD = mujoco.mjtDisableBit.mjDSBL_NATIVECCD + # unsupported: MIDPHASE, AUTORESET, ISLAND class EnableBit(enum.IntFlag): @@ -321,6 +342,20 @@ class GeomType(enum.IntEnum): # unsupported: NGEOMTYPES, ARROW*, LINE, SKIN, LABEL, NONE +class CollisionType(enum.IntEnum): + """Type of narrowphase collision. + + Attributes: + PRIMITIVE: primitive collision + CONVEX: convex collision (CCD) + SDF: sdf collision + """ + + PRIMITIVE = 0 + CONVEX = 1 + SDF = 2 + + class SolverType(enum.IntEnum): """Constraint solver algorithm. @@ -665,10 +700,8 @@ class Option: warp only fields: impratio_invsqrt: ratio of friction-to-normal contact impedance (stored as inverse square root) - is_sparse: whether to use sparse representations ls_parallel: evaluate engine solver step sizes in parallel ls_parallel_min_step: minimum step size for solver linesearch - has_fluid: True if wind, density, or viscosity are non-zero at put_model time broadphase: broadphase type (BroadphaseType) broadphase_filter: broadphase filter bitflag (BroadphaseFilter) graph_conditional: flag to use cuda graph conditional @@ -700,10 +733,8 @@ class Option: sdf_iterations: int # warp only fields: impratio_invsqrt: array("*", float) - is_sparse: bool ls_parallel: bool ls_parallel_min_step: float - has_fluid: bool broadphase: BroadphaseType broadphase_filter: BroadphaseFilter graph_conditional: bool @@ -716,10 +747,10 @@ class Statistic: """Model statistics (in qpos0). Attributes: - meaninertia: mean diagonal inertia + meaninertia: mean diagonal inertia (per-world) """ - meaninertia: float + meaninertia: array("*", float) @dataclasses.dataclass @@ -749,6 +780,7 @@ class Model: nbody: number of bodies noct: number of total octree cells in all meshes njnt: number of joints + ntree: number of kinematic trees nM: number of non-zeros in sparse inertia matrix nC: number of non-zeros in sparse body-dof matrix ngeom: number of geoms @@ -794,6 +826,7 @@ class Model: body_jntadr: start addr of joints; -1: no joints (nbody,) body_dofnum: number of motion degrees of freedom (nbody,) body_dofadr: start addr of dofs; -1: no dofs (nbody,) + body_treeid: id of body's tree; -1: static (nbody,) body_geomnum: number of geoms (nbody,) body_geomadr: start addr of geoms; -1: no geoms (nbody,) body_pos: position offset rel. to parent body (*, nbody, 3) @@ -828,6 +861,7 @@ class Model: dof_bodyid: id of dof's body (nv,) dof_jntid: id of dof's joint (nv,) dof_parentid: id of dof's parent; -1: none (nv,) + dof_treeid: id of dof's tree (nv,) dof_Madr: dof address in M-diagonal (nv,) dof_solref: constraint solver reference: frictionloss (*, nv, NREF) dof_solimp: constraint solver impedance: frictionloss (*, nv, NIMP) @@ -835,6 +869,9 @@ class Model: dof_armature: dof armature inertia/mass (*, nv) dof_damping: damping coefficient (*, nv) dof_invweight0: diag. inverse inertia in qpos0 (*, nv) + tree_bodynum: number of bodies in tree (incl. root) (ntree,) + tree_dofadr: start address of tree's dofs (ntree,) + tree_dofnum: number of dofs in tree (ntree,) geom_type: geometric type (GeomType) (ngeom,) geom_contype: geom contact type (ngeom,) geom_conaffinity: geom contact affinity (ngeom,) @@ -870,10 +907,11 @@ class Model: cam_poscom0: global position rel. to sub-com in qpos0 (*, ncam, 3) cam_pos0: global position rel. to body in qpos0 (*, ncam, 3) cam_mat0: global orientation in qpos0 (*, ncam, 3, 3) - cam_fovy: y field-of-view (ortho ? len : deg) (ncam,) + cam_projection: projection type (ProjectionType) (ncam,) + cam_fovy: y field-of-view (ortho ? len : deg) (*, ncam) cam_resolution: resolution: pixels [width, height] (ncam, 2) cam_sensorsize: sensor size: length [width, height] (ncam, 2) - cam_intrinsic: [focal length; principal point] (ncam, 4) + cam_intrinsic: [focal length; principal point] (*, ncam, 4) light_mode: light tracking mode (CamLightType) (nlight,) light_bodyid: id of light's body (nlight,) light_targetbodyid: id of targeted body; -1: none (nlight,) @@ -903,6 +941,9 @@ class Model: flex_stiffness: finite element stiffness matrix (nflexelem, 21) flex_bending: bending stiffness (nflexedge, 17) flex_damping: Rayleigh's damping coefficient (nflex,) + flexedge_J_rownnz: number of nonzeros in Jacobian row (nflexedge,) + flexedge_J_rowadr: row start address in colind array (nflexedge,) + flexedge_J_colind: column indices in sparse Jacobian (nJfe,) mesh_vertadr: first vertex address (nmesh,) mesh_vertnum: number of vertices (nmesh,) mesh_faceadr: first face address (nmesh,) @@ -927,7 +968,7 @@ class Model: hfield_ncol: number of columns in grid (nhfield,) hfield_adr: start address in hfield_data (nhfield,) hfield_data: elevation data (nhfielddata,) - mat_texid: texture id for rendering (nmat, mjNTEXROLE) + mat_texid: texture id for rendering (*, nmat, mjNTEXROLE) mat_texrepeat: texture repeat for rendering (*, nmat, 2) mat_rgba: rgba (*, nmat, 4) pair_dim: contact dimensionality (npair,) @@ -1008,6 +1049,7 @@ class Model: mapM2M: index mapping from M (legacy) to M (CSR) (nC) warp only fields: + nbranch: number of branches (leaf-to-root paths) nv_pad: number of degrees of freedom + padding nacttrnbody: number of actuators with body transmission nsensorcollision: number of unique collisions for @@ -1019,9 +1061,13 @@ class Model: nmaxpyramid: maximum number of pyramid directions nmaxpolygon: maximum number of verts per polygon nmaxmeshdeg: maximum number of polygons per vert + is_sparse: whether to use sparse representations + has_fluid: True if wind, density, or viscosity are non-zero at put_model time has_sdf_geom: whether the model contains SDF geoms block_dim: block dim options body_tree: list of body ids by tree level + body_branches: flattened body ids for all branches + body_branch_start: start index in body_branches for each branch (nbranch + 1,) mocap_bodyid: id of body for mocap (nmocap,) body_fluid_ellipsoid: does body use ellipsoid fluid (nbody,) jnt_limited_slide_hinge_adr: limited/slide/hinge jntadr @@ -1085,9 +1131,9 @@ class Model: qLD_updates: tuple of index triples for sparse factorization qM_fullm_i: sparse mass matrix addressing qM_fullm_j: sparse mass matrix addressing - qM_mulm_i: sparse matmul addressing - qM_mulm_j: sparse matmul addressing - qM_madr_ij: sparse matmul addressing + qM_mulm_rowadr: sparse matmul row pointers + qM_mulm_col: sparse matmul column indices + qM_mulm_madr: sparse matmul matrix addresses """ nq: int @@ -1097,6 +1143,7 @@ class Model: nbody: int noct: int njnt: int + ntree: int nM: int nC: int ngeom: int @@ -1142,6 +1189,7 @@ class Model: body_jntadr: array("nbody", int) body_dofnum: array("nbody", int) body_dofadr: array("nbody", int) + body_treeid: array("nbody", int) body_geomnum: array("nbody", int) body_geomadr: array("nbody", int) body_pos: array("*", "nbody", wp.vec3) @@ -1176,6 +1224,7 @@ class Model: dof_bodyid: array("nv", int) dof_jntid: array("nv", int) dof_parentid: array("nv", int) + dof_treeid: array("nv", int) dof_Madr: array("nv", int) dof_solref: array("*", "nv", wp.vec2) dof_solimp: array("*", "nv", vec5) @@ -1183,6 +1232,9 @@ class Model: dof_armature: array("*", "nv", float) dof_damping: array("*", "nv", float) dof_invweight0: array("*", "nv", float) + tree_bodynum: array("ntree", int) + tree_dofadr: array("ntree", int) + tree_dofnum: array("ntree", int) geom_type: array("ngeom", int) geom_contype: array("ngeom", int) geom_conaffinity: array("ngeom", int) @@ -1218,10 +1270,11 @@ class Model: cam_poscom0: array("*", "ncam", wp.vec3) cam_pos0: array("*", "ncam", wp.vec3) cam_mat0: array("*", "ncam", wp.mat33) - cam_fovy: array("ncam", float) + cam_projection: array("ncam", int) + cam_fovy: array("*", "ncam", float) cam_resolution: array("ncam", wp.vec2i) cam_sensorsize: array("ncam", wp.vec2) - cam_intrinsic: array("ncam", wp.vec4) + cam_intrinsic: array("*", "ncam", wp.vec4) light_mode: array("nlight", int) light_bodyid: array("nlight", int) light_targetbodyid: array("nlight", int) @@ -1251,6 +1304,9 @@ class Model: flex_stiffness: array("nflexelem", 21, float) flex_bending: array("nflexedge", 17, float) flex_damping: array("nflex", float) + flexedge_J_rownnz: array("nflexedge", int) + flexedge_J_rowadr: array("nflexedge", int) + flexedge_J_colind: wp.array(dtype=int) mesh_vertadr: array("nmesh", int) mesh_vertnum: array("nmesh", int) mesh_faceadr: array("nmesh", int) @@ -1275,7 +1331,7 @@ class Model: hfield_ncol: array("nhfield", int) hfield_adr: array("nhfield", int) hfield_data: array("nhfielddata", float) - mat_texid: array("nmat", 10, int) + mat_texid: array("*", "nmat", 10, int) mat_texrepeat: array("*", "nmat", wp.vec2) mat_rgba: array("*", "nmat", wp.vec4) pair_dim: array("npair", int) @@ -1355,6 +1411,7 @@ class Model: M_colind: array("nC", int) mapM2M: array("nC", int) # warp only fields: + nbranch: int nv_pad: int nacttrnbody: int nsensorcollision: int @@ -1365,9 +1422,13 @@ class Model: nmaxpyramid: int nmaxpolygon: int nmaxmeshdeg: int + is_sparse: bool + has_fluid: bool has_sdf_geom: bool block_dim: BlockDim body_tree: tuple[wp.array(dtype=int), ...] + body_branches: wp.array(dtype=int) + body_branch_start: wp.array(dtype=int) mocap_bodyid: array("nmocap", int) body_fluid_ellipsoid: array("nbody", bool) jnt_limited_slide_hinge_adr: wp.array(dtype=int) @@ -1422,9 +1483,10 @@ class Model: qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...] qM_fullm_i: wp.array(dtype=int) qM_fullm_j: wp.array(dtype=int) - qM_mulm_i: wp.array(dtype=int) - qM_mulm_j: wp.array(dtype=int) - qM_madr_ij: wp.array(dtype=int) + # Gather-based sparse mul_m indices (thread per DOF, no atomics) + qM_mulm_rowadr: wp.array(dtype=int) # start address for each row [nv+1] + qM_mulm_col: wp.array(dtype=int) # column index to gather from + qM_mulm_madr: wp.array(dtype=int) # matrix address to read class ContactType(enum.IntFlag): @@ -1492,26 +1554,9 @@ class Constraint: aref: reference pseudo-acceleration (nworld, njmax) frictionloss: frictionloss (friction) (nworld, njmax) force: constraint force in constraint space (nworld, njmax) - Jaref: Jac*qacc - aref (nworld, njmax) - Ma: M*qacc (nworld, nv) - grad: gradient of master cost (nworld, nv_pad) - grad_dot: dot(grad, grad) (nworld,) - Mgrad: M / grad (nworld, nv_pad) - search: linesearch vector (nworld, nv) - search_dot: dot(search, search) (nworld,) - gauss: Gauss Cost (nworld,) - cost: constraint + Gauss cost (nworld,) - prev_cost: cost from previous iter (nworld,) state: constraint state (nworld, njmax_pad) - mv: qM @ search (nworld, nv) - jv: efc_J @ search (nworld, njmax) - quad: quadratic cost coefficients (nworld, njmax, 3) - quad_gauss: quadratic cost Gauss coefficients (nworld, 3) - alpha: line search step size (nworld,) - prev_grad: previous grad (nworld, nv) - prev_Mgrad: previous Mgrad (nworld, nv) - beta: Polak-Ribiere beta (nworld,) - done: solver done (nworld,) + warp only fields: + Ma: M*qacc (nworld, nv) """ type: array("nworld", "njmax", int) @@ -1524,26 +1569,8 @@ class Constraint: aref: array("nworld", "njmax", float) frictionloss: array("nworld", "njmax", float) force: array("nworld", "njmax", float) - Jaref: array("nworld", "njmax", float) - Ma: array("nworld", "nv", float) - grad: array("nworld", "nv_pad", float) - grad_dot: array("nworld", float) - Mgrad: array("nworld", "nv_pad", float) - search: array("nworld", "nv", float) - search_dot: array("nworld", float) - gauss: array("nworld", float) - cost: array("nworld", float) - prev_cost: array("nworld", float) state: array("nworld", "njmax_pad", int) - mv: array("nworld", "nv", float) - jv: array("nworld", "njmax", float) - quad: array("nworld", "njmax", wp.vec3) - quad_gauss: array("nworld", wp.vec3) - alpha: array("nworld", float) - prev_grad: array("nworld", "nv", float) - prev_Mgrad: array("nworld", "nv", float) - beta: array("nworld", float) - done: array("nworld", bool) + Ma: array("nworld", "nv", float) @dataclasses.dataclass @@ -1590,7 +1617,7 @@ class Data: cdof: com-based motion axis of each dof (rot:lin) (nworld, nv, 6) cinert: com-based body inertia and mass (nworld, nbody, 10) flexvert_xpos: cartesian flex vertex positions (nworld, nflexvert, 3) - flexedge_J: edge length Jacobian (nworld, nflexedge, nv) + flexedge_J: edge length Jacobian (nworld, 1, nflexedge*6) flexedge_length: flex edge lengths (nworld, nflexedge, 1) ten_wrapadr: start address of tendon's path (nworld, ntendon) ten_wrapnum: number of wrap points in path (nworld, ntendon) @@ -1606,7 +1633,7 @@ class Data: qLD: L'*D*L factorization of M (nworld, nv, nv) if dense (nworld, 1, nC) if sparse qLDiagInv: 1/diag(D) (nworld, nv) - flexedge_velocity: flex edge velocities (nworld, nflexedge,) + flexedge_velocity: flex edge velocities (nworld, nflexedge) ten_velocity: tendon velocities (nworld, ntendon) actuator_velocity: actuator velocities (nworld, nu) cvel: com-based velocity (rot:lin) (nworld, nbody, 6) @@ -1636,18 +1663,9 @@ class Data: warp only fields: nworld: number of worlds naconmax: maximum number of contacts (shared across all worlds) + naccdmax: maximum number of contacts for CCD (all worlds) njmax: maximum number of constraints per world nacon: number of detected contacts (across all worlds) (1,) - ne_connect: number of equality connect constraints (nworld,) - ne_weld: number of equality weld constraints (nworld,) - ne_jnt: number of equality joint constraints (nworld,) - ne_ten: number of equality tendon constraints (nworld,) - ne_flex: number of flex edge equality constraints (nworld,) - nsolving: number of unconverged worlds (1,) - subtree_bodyvel: subtree body velocity (ang, vel) (nworld, nbody, 6) - collision_pair: collision pairs from broadphase (naconmax, 2) - collision_pairid: ids from broadphase (naconmax, 2) - collision_worldid: collision world ids from broadphase (naconmax,) ncollision: collision count from broadphase (1,) """ @@ -1690,7 +1708,7 @@ class Data: cdof: array("nworld", "nv", wp.spatial_vector) cinert: array("nworld", "nbody", vec10) flexvert_xpos: array("nworld", "nflexvert", wp.vec3) - flexedge_J: array("nworld", "nflexedge", "nv", float) + flexedge_J: wp.array3d(dtype=float) flexedge_length: array("nworld", "nflexedge", float) ten_wrapadr: array("nworld", "ntendon", int) ten_wrapnum: array("nworld", "ntendon", int) @@ -1732,18 +1750,131 @@ class Data: # warp only fields: nworld: int naconmax: int + naccdmax: int njmax: int nacon: array(1, int) - ne_connect: array("nworld", int) - ne_weld: array("nworld", int) - ne_jnt: array("nworld", int) - ne_ten: array("nworld", int) - ne_flex: array("nworld", int) - nsolving: array(1, int) - subtree_bodyvel: array("nworld", "nbody", wp.spatial_vector) - - # warp only: collision driver - collision_pair: array("naconmax", wp.vec2i) - collision_pairid: array("naconmax", wp.vec2i) - collision_worldid: array("naconmax", int) ncollision: array(1, int) + + +@dataclasses.dataclass +class CollisionContext: + """Collision driver intermediate arrays. + + Attributes: + collision_pair: collision pairs from broadphase (naconmax, 2) + collision_pairid: ids from broadphase (naconmax, 2) + collision_worldid: collision world ids from broadphase (naconmax,) + """ + + collision_pair: wp.array + collision_pairid: wp.array + collision_worldid: wp.array + + +@dataclasses.dataclass +class RenderContext: + """Context for rendering. + + Attributes: + nrender: number of actively rendering cameras + cam_res: camera resolution for actively rendering cameras + cam_id_map: camera id map + use_textures: whether to use textures + use_shadows: whether to use shadows + bvh_ngeom: number of geometries in the BVH + enabled_geom_ids: enabled geometry ids + mesh_registry: mesh BVH id to warp mesh mapping + mesh_bvh_id: mesh BVH ids + mesh_bounds_size: mesh bounds size + mesh_texcoord: mesh texture coordinates + mesh_texcoord_offsets: mesh texture coordinate offsets + mesh_facetexcoord: mesh face texture coordinates + textures: textures + textures_registry: texture registry + hfield_registry: hfield BVH id to warp mesh mapping + hfield_bvh_id: hfield BVH ids + hfield_bounds_size: hfield bounds size + flex_mesh: flex mesh + flex_rgba: flex rgba + flex_bvh_id: flex BVH id + flex_face_point: flex face points + flex_faceadr: flex face addresses + flex_nface: number of flex faces + flex_nwork: total flex work items for refit + flex_group_root: flex group roots + flex_elemdataadr: flex element data addresses + flex_shell: flex shell data + flex_shelldataadr: flex shell data addresses + flex_radius: flex radius + flex_workadr: flex work item addresses for refit + flex_worknum: flex work item counts for refit + flex_render_smooth: whether to render flex meshes smoothly + bvh: scene BVH + bvh_id: scene BVH id + lower: lower bounds + upper: upper bounds + group: groups + group_root: group roots + ray: rays + rgb_data: RGB data + rgb_adr: RGB addresses + rgb_size: per-camera RGB buffer sizes + depth_data: depth data + depth_adr: depth addresses + depth_size: per-camera depth buffer sizes + render_rgb: per-camera RGB render flags + render_depth: per-camera depth render flags + znear: near plane distance + total_rays: total number of rays + """ + + nrender: int + cam_res: array("ncam", wp.vec2i) + cam_id_map: array("ncam", int) + use_textures: bool + use_shadows: bool + background_color: wp.uint32 + bvh_ngeom: int + enabled_geom_ids: array("*", int) + mesh_registry: dict + mesh_bvh_id: array("nmesh", wp.uint64) + mesh_bounds_size: array("nmesh", wp.vec3) + mesh_texcoord: array("*", wp.vec2) + mesh_texcoord_offsets: array("nmesh", int) + mesh_facetexcoord: array("nmeshface", wp.vec3i) + # TODO(team): remove after mjwarp depends on warp-lang >= 1.12 in pyproject.toml + textures: array("*", TEXTURE_DTYPE) + textures_registry: list[TEXTURE_DTYPE] + hfield_registry: dict + hfield_bvh_id: array("nhfield", wp.uint64) + hfield_bounds_size: array("nhfield", wp.vec3) + flex_mesh: wp.Mesh + 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_render_smooth: bool + bvh: wp.Bvh + bvh_id: wp.uint64 + lower: array("*", wp.vec3) + upper: array("*", wp.vec3) + group: array("*", int) + group_root: array("*", int) + ray: array("*", wp.vec3) + rgb_data: array("*", wp.uint32) + rgb_adr: array("ncam", int) + depth_data: array("*", wp.float32) + depth_adr: array("ncam", int) + render_rgb: array("ncam", bool) + render_depth: array("ncam", bool) + znear: float + total_rays: int diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py index 521e0054..258da2ee 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/util_misc.py @@ -20,6 +20,7 @@ from typing import Tuple import warp as wp from mujoco.mjx.third_party.mujoco_warp._src import math +from mujoco.mjx.third_party.mujoco_warp._src.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 GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import WrapType @@ -105,13 +106,13 @@ def wrap_circle(end: wp.vec4, side: wp.vec2, radius: float) -> Tuple[float, wp.v Args: end: Two 2D points. - side: Optional 2D side point, no side point: wp.vec2(wp.inf). + side: Optional 2D side point, no side point: wp.vec2(MJ_MAXVAL). radius: Circle radius. Returns: Length of circular wrap or -1.0 if no wrap, pair of 2D wrap points. """ - valid_side = wp.norm_l2(side) < wp.inf + valid_side = wp.norm_l2(side) < MJ_MAXVAL end0 = wp.vec2(end[0], end[1]) end1 = wp.vec2(end[2], end[3]) @@ -122,13 +123,13 @@ def wrap_circle(end: wp.vec4, side: wp.vec2, radius: float) -> Tuple[float, wp.v # either point inside circle or circle too small: no wrap if (sqlen0 < sqrad) or (sqlen1 < sqrad) or (radius < MJ_MINVAL): - return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + return -1.0, wp.vec2(MJ_MAXVAL), wp.vec2(MJ_MAXVAL) # points too close: no wrap dif = end1 - end0 dd = wp.dot(dif, dif) if dd < MJ_MINVAL: - return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + return -1.0, wp.vec2(MJ_MAXVAL), wp.vec2(MJ_MAXVAL) # find nearest point on line segment to origin: a * dif + d0 a = -wp.dot(dif, end0) / dd @@ -137,7 +138,7 @@ def wrap_circle(end: wp.vec4, side: wp.vec2, radius: float) -> Tuple[float, wp.v # check for intersection and side tmp = a * dif + end0 if (wp.dot(tmp, tmp) > sqrad) and (not valid_side or wp.dot(side, tmp) >= 0.0): - return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + return -1.0, wp.vec2(MJ_MAXVAL), wp.vec2(MJ_MAXVAL) sqrt0 = wp.sqrt(sqlen0 - sqrad) sqrt1 = wp.sqrt(sqlen1 - sqrad) @@ -191,7 +192,7 @@ def wrap_circle(end: wp.vec4, side: wp.vec2, radius: float) -> Tuple[float, wp.v # check for intersection if is_intersect(end0, pnt0, end1, pnt1): - return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + return -1.0, wp.vec2(MJ_MAXVAL), wp.vec2(MJ_MAXVAL) # return curve length return length_circle(pnt0, pnt1, ind, radius), pnt0, pnt1 @@ -230,7 +231,7 @@ def wrap_inside( # either point inside circle or circle too small: no wrap if (len0 <= radius) or (len1 <= radius) or (radius < MJ_MINVAL) or (len0 < MJ_MINVAL) or (len1 < MJ_MINVAL): - return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + return -1.0, wp.vec2(MJ_MAXVAL), wp.vec2(MJ_MAXVAL) # segment-circle intersection: no wrap if dd > MJ_MINVAL: @@ -241,7 +242,7 @@ def wrap_inside( if (a > 0.0) and (a < 1.0): tmp = end0 + a * dif if wp.norm_l2(tmp) <= radius: - return -1.0, wp.vec2(wp.inf), wp.vec2(wp.inf) + return -1.0, wp.vec2(MJ_MAXVAL), wp.vec2(MJ_MAXVAL) # prepare default in case of numerical failure: average pnt = 0.5 * (end0 + end1) @@ -335,14 +336,14 @@ def wrap( mat: Orientation of geom. radius: Geom radius. geomtype: Wrap type (mjtWrap). - side: 3D position for sidesite, no side point: wp.vec3(wp.inf). + side: 3D position for sidesite, no side point: wp.vec3(MJ_MAXVAL). Returns: Length of circular wrap else -1.0 if no wrap, pair of 3D wrap points. """ # check object type if geomtype != WrapType.SPHERE and geomtype != WrapType.CYLINDER: - return wp.inf, wp.vec3(wp.inf), wp.vec3(wp.inf) + return MJ_MAXVAL, wp.vec3(MJ_MAXVAL), wp.vec3(MJ_MAXVAL) # map sites to wrap object's local frame matT = wp.transpose(mat) @@ -351,7 +352,7 @@ def wrap( # too close to origin: return if (wp.norm_l2(p0) < MJ_MINVAL) or (wp.norm_l2(p1) < MJ_MINVAL): - return -1.0, wp.vec3(wp.inf), wp.vec3(wp.inf) + return -1.0, wp.vec3(MJ_MAXVAL), wp.vec3(MJ_MAXVAL) # construct 2D frame for circle wrap if geomtype == WrapType.SPHERE: @@ -399,7 +400,7 @@ def wrap( ) # handle sidesite - valid_side = wp.norm_l2(side) < wp.inf + valid_side = wp.norm_l2(side) < MJ_MAXVAL if valid_side: # side point: apply same projection as x0, x1 @@ -414,7 +415,7 @@ def wrap( sidepnt_proj, _ = math.normalize_with_norm(sidepnt_proj) sidepnt_proj *= radius else: - sidepnt_proj = wp.vec2(wp.inf) + sidepnt_proj = wp.vec2(MJ_MAXVAL) # apply inside wrap if valid_side and wp.norm_l2(sidepnt) < radius: @@ -424,7 +425,7 @@ def wrap( # no wrap: return if wlen < 0.0: - return -1.0, wp.vec3(wp.inf), wp.vec3(wp.inf) + return -1.0, wp.vec3(MJ_MAXVAL), wp.vec3(MJ_MAXVAL) # reconstruct 3D points in local frame: res res0 = axis0 * pnt0[0] + axis1 * pnt0[1] 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 0301a386..7657346a 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 @@ -15,7 +15,6 @@ import functools import inspect -from typing import Callable, Optional import warp as wp @@ -119,82 +118,6 @@ def event_scope(fn, name: str = ""): return wrapper -# @nested_kernel decorator to automatically set up modules based on nested -# function names -def nested_kernel( - f: Optional[Callable] = None, - *, - enable_backward: Optional[bool] = None, - module: Optional[wp.Module] = None, -): - """Decorator to register a Warp kernel from a Python function. - - The function must be defined with type annotations for all arguments. - The function must not return anything. - - Example:: - - @nested_kernel - def my_kernel(a: wp.array(dtype=float), b: wp.array(dtype=float)): - tid = wp.tid() - b[tid] = a[tid] + 1.0 - - - @nested_kernel(enable_backward=False) - def my_kernel_no_backward(a: wp.array(dtype=float, ndim=2), x: float): - # the backward pass will not be generated - i, j = wp.tid() - a[i, j] = x - - - @nested_kernel(module="unique") - def my_kernel_unique_module(a: wp.array(dtype=float), b: wp.array(dtype=float)): - # the kernel will be registered in new unique module created just for this - # kernel and its dependent functions and structs - tid = wp.tid() - b[tid] = a[tid] + 1.0 - - - @neste_kernel(enable_backward=False, module=None) - def my_kernel_with_args(a: wp.array(dtype=float), b: wp.array(dtype=float)): - # can now use arguments even when module=None - tid = wp.tid() - b[tid] = a[tid] + 1.0 - - Args: - f: The function to be registered as a kernel. - enable_backward: If False, the backward pass will not be generated. - module: The :class:`warp.Module` to which the kernel belongs. Alternatively, - if a string `"unique"` is provided, the kernel is assigned to a new module - named after the kernel name and hash. If None, the module is inferred from - the function's module. - - Returns: - The registered kernel. - """ - - def decorator(func): - if module is None: - # create a module name based on the name of the nested function - # get the qualified name, e.g. "main..nested_kernel" - qualname = func.__qualname__ - parts = [part for part in qualname.split(".") if part != ""] - outer_functions = parts[:-1] - module_name = wp.get_module(".".join([func.__module__] + outer_functions)) - else: - module_name = module - - return wp.kernel(func, enable_backward=enable_backward, module=module_name) - - # Handle both @kernel and @kernel(...) usage patterns - if f is None: - # Called with arguments: @kernel(enable_backward=False) - return decorator - else: - # Called without arguments: @kernel - return decorator(f) - - _KERNEL_CACHE = {} @@ -221,4 +144,4 @@ def check_toolkit_driver(): wp.init() if wp.get_device().is_cuda: if not wp.is_conditional_graph_supported(): - RuntimeError("Minimum supported CUDA version: 12.4.") + raise RuntimeError("Minimum supported CUDA version: 12.4.") diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml index 94d6f7ea..c90938a0 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml @@ -54,7 +54,7 @@ dev = [ "ruff", "pygls>=1.0.0,<2.0.0", "lsprotocol>=2023.0.1,<2024.0.0", - "mujoco>=3.3.7.dev0", + "mujoco>=3.4.1.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 3cdb2c7f..5f1de2d4 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py @@ -24,6 +24,7 @@ Example: import copy import enum import logging +import shutil import sys import time from typing import Sequence @@ -51,10 +52,11 @@ class EngineOptions(enum.IntEnum): C = 1 -_CLEAR_KERNEL_CACHE = flags.DEFINE_bool("clear_kernel_cache", False, "Clear kernel cache (to calculate full JIT time)") +_CLEAR_WARP_CACHE = flags.DEFINE_bool("clear_warp_cache", False, "Clear warp caches (kernel, LTO, CUDA compute)") _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.") +_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.") _DEVICE = flags.DEFINE_string("device", None, "override the default Warp device") @@ -134,27 +136,37 @@ def _main(argv: Sequence[str]) -> None: else: wp.config.quiet = flags.FLAGS["verbosity"].value < 1 wp.init() - if _CLEAR_KERNEL_CACHE.value: + wp.set_device(_DEVICE.value) + if _CLEAR_WARP_CACHE.value: wp.clear_kernel_cache() + wp.clear_lto_cache() + # Clear CUDA compute cache for truly cold start JIT + compute_cache = epath.Path("~/.nv/ComputeCache").expanduser() + if compute_cache.exists(): + shutil.rmtree(compute_cache) + compute_cache.mkdir() - with wp.ScopedDevice(_DEVICE.value): - m = mjw.put_model(mjm) - override_model(m, _OVERRIDE.value) - broadphase, filter = mjw.BroadphaseType(m.opt.broadphase).name, mjw.BroadphaseFilter(m.opt.broadphase_filter).name - solver, cone = mjw.SolverType(m.opt.solver).name, mjw.ConeType(m.opt.cone).name - integrator = mjw.IntegratorType(m.opt.integrator).name - iterations, ls_iterations = m.opt.iterations, m.opt.ls_iterations - ls_str = f"{'parallel' if m.opt.ls_parallel else 'iterative'} linesearch iterations: {ls_iterations}" - print( - f" nbody: {m.nbody} nv: {m.nv} ngeom: {m.ngeom} nu: {m.nu} is_sparse: {m.opt.is_sparse}\n" - f" broadphase: {broadphase} broadphase_filter: {filter}\n" - f" solver: {solver} cone: {cone} iterations: {iterations} {ls_str}\n" - f" integrator: {integrator} graph_conditional: {m.opt.graph_conditional}" - ) - d = mjw.put_data(mjm, mjd, nconmax=_NCONMAX.value, njmax=_NJMAX.value) - print(f"Data\n nworld: {d.nworld} nconmax: {d.naconmax / d.nworld} njmax: {d.njmax}\n") - graph = _compile_step(m, d) - print(f"MuJoCo Warp simulating with dt = {m.opt.timestep.numpy()[0]:.3f}...") + 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) + graph = _compile_step(m, d) if wp.get_device().is_cuda else None + if graph is None: + mjw.step(m, d) # warmup step + print("Running Warp unoptimized on CPU.") + broadphase, filter = mjw.BroadphaseType(m.opt.broadphase).name, mjw.BroadphaseFilter(m.opt.broadphase_filter).name + solver, cone = mjw.SolverType(m.opt.solver).name, mjw.ConeType(m.opt.cone).name + integrator = mjw.IntegratorType(m.opt.integrator).name + iterations, ls_iterations = m.opt.iterations, m.opt.ls_iterations + ls_str = f"{'parallel' if m.opt.ls_parallel else 'iterative'} linesearch iterations: {ls_iterations}" + print( + f" nbody: {m.nbody} nv: {m.nv} ngeom: {m.ngeom} nu: {m.nu} is_sparse: {m.is_sparse}\n" + f" broadphase: {broadphase} broadphase_filter: {filter}\n" + f" solver: {solver} cone: {cone} iterations: {iterations} {ls_str}\n" + f" integrator: {integrator} graph_conditional: {m.opt.graph_conditional}" + ) + print(f"Data\n nworld: {d.nworld} nconmax: {int(d.naconmax / d.nworld)} njmax: {d.njmax}\n") + print(f"MuJoCo Warp simulating with dt = {m.opt.timestep.numpy()[0]:.3f}...") with mujoco.viewer.launch_passive(mjm, mjd, key_callback=key_callback) as viewer: opt = copy.copy(mjm.opt) @@ -175,22 +187,19 @@ def _main(argv: Sequence[str]) -> None: wp.copy(d.qpos, wp.array([mjd.qpos.astype(np.float32)])) wp.copy(d.qvel, wp.array([mjd.qvel.astype(np.float32)])) wp.copy(d.time, wp.array([mjd.time], dtype=wp.float32)) - # if the user changed an option in the MuJoCo Simulate UI, go ahead and recompile the step # TODO: update memory tied to option max iterations if mjm.opt != opt: opt = copy.copy(mjm.opt) m = mjw.put_model(mjm) - graph = _compile_step(m, d) - - if _VIEWER_GLOBAL_STATE["running"]: - wp.capture_launch(graph) - wp.synchronize() - elif _VIEWER_GLOBAL_STATE["step_once"]: + graph = _compile_step(m, d) if wp.get_device().is_cuda else None + if _VIEWER_GLOBAL_STATE["running"] or _VIEWER_GLOBAL_STATE["step_once"]: _VIEWER_GLOBAL_STATE["step_once"] = False - wp.capture_launch(graph) - wp.synchronize() - + if graph is None: + mjw.step(m, d) + else: + wp.capture_launch(graph) + wp.synchronize() mjw.get_data_into(mjd, mjm, d) viewer.sync() diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index 66b6d7ac..f4be0dae 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -42,6 +42,7 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _collision_shim( # Model @@ -112,10 +113,8 @@ def _collision_shim( opt__sdf_initpoints: int, opt__sdf_iterations: int, # Data + naccdmax: int, naconmax: int, - collision_pair: wp.array(dtype=wp.vec2i), - collision_pairid: wp.array(dtype=wp.vec2i), - collision_worldid: wp.array(dtype=int), geom_xmat: wp.array2d(dtype=wp.mat33), geom_xpos: wp.array2d(dtype=wp.vec3), nacon: wp.array(dtype=int), @@ -203,9 +202,6 @@ def _collision_shim( _m.pair_solreffriction = pair_solreffriction _m.plugin = plugin _m.plugin_attr = plugin_attr - _d.collision_pair = collision_pair - _d.collision_pairid = collision_pairid - _d.collision_worldid = collision_worldid _d.contact.dim = contact__dim _d.contact.dist = contact__dist _d.contact.frame = contact__frame @@ -221,6 +217,7 @@ def _collision_shim( _d.contact.worldid = contact__worldid _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos + _d.naccdmax = naccdmax _d.nacon = nacon _d.naconmax = naconmax _d.ncollision = ncollision @@ -230,9 +227,6 @@ def _collision_shim( def _collision_jax_impl(m: types.Model, d: types.Data): output_dims = { - 'collision_pair': d._impl.collision_pair.shape, - 'collision_pairid': d._impl.collision_pairid.shape, - 'collision_worldid': d._impl.collision_worldid.shape, 'nacon': d._impl.nacon.shape, 'ncollision': d._impl.ncollision.shape, 'contact__dim': d._impl.contact__dim.shape, @@ -251,13 +245,10 @@ def _collision_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _collision_shim, - num_outputs=18, + num_outputs=15, output_dims=output_dims, vmap_method=None, in_out_argnames={ - 'collision_pair', - 'collision_pairid', - 'collision_worldid', 'nacon', 'ncollision', 'contact__dim', @@ -364,10 +355,8 @@ def _collision_jax_impl(m: types.Model, d: types.Data): m.opt.enableflags, m.opt._impl.sdf_initpoints, m.opt._impl.sdf_iterations, + d._impl.naccdmax, d._impl.naconmax, - d._impl.collision_pair, - d._impl.collision_pairid, - d._impl.collision_worldid, d.geom_xmat, d.geom_xpos, d._impl.nacon, @@ -387,24 +376,21 @@ def _collision_jax_impl(m: types.Model, d: types.Data): d._impl.contact__worldid, ) d = d.tree_replace({ - '_impl.collision_pair': out[0], - '_impl.collision_pairid': out[1], - '_impl.collision_worldid': out[2], - '_impl.nacon': out[3], - '_impl.ncollision': out[4], - '_impl.contact__dim': out[5], - '_impl.contact__dist': out[6], - '_impl.contact__frame': out[7], - '_impl.contact__friction': out[8], - '_impl.contact__geom': out[9], - '_impl.contact__geomcollisionid': out[10], - '_impl.contact__includemargin': out[11], - '_impl.contact__pos': out[12], - '_impl.contact__solimp': out[13], - '_impl.contact__solref': out[14], - '_impl.contact__solreffriction': out[15], - '_impl.contact__type': out[16], - '_impl.contact__worldid': out[17], + '_impl.nacon': out[0], + '_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], }) return d diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index 4ccad4f5..2604244d 100644 --- a/mjx/mujoco/mjx/warp/forward.py +++ b/mjx/mujoco/mjx/warp/forward.py @@ -42,6 +42,7 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _forward_shim( # Model @@ -71,6 +72,8 @@ def _forward_shim( actuator_trntype: wp.array(dtype=int), actuator_trntype_body_adr: wp.array(dtype=int), block_dim: mjwp_types.BlockDim, + body_branch_start: wp.array(dtype=int), + body_branches: wp.array(dtype=int), body_dofadr: wp.array(dtype=int), body_dofnum: wp.array(dtype=int), body_fluid_ellipsoid: wp.array(dtype=bool), @@ -93,8 +96,8 @@ def _forward_shim( body_tree: tuple[wp.array(dtype=int), ...], body_weldid: wp.array(dtype=int), cam_bodyid: wp.array(dtype=int), - cam_fovy: wp.array(dtype=float), - cam_intrinsic: wp.array(dtype=wp.vec4), + cam_fovy: wp.array2d(dtype=float), + cam_intrinsic: wp.array2d(dtype=wp.vec4), cam_mat0: wp.array2d(dtype=wp.mat33), cam_mode: wp.array(dtype=int), cam_pos: wp.array2d(dtype=wp.vec3), @@ -126,6 +129,7 @@ def _forward_shim( eq_solimp: wp.array2d(dtype=mjwp_types.vec5), eq_solref: wp.array2d(dtype=wp.vec2), eq_ten_adr: wp.array(dtype=int), + eq_type: wp.array(dtype=int), eq_wld_adr: wp.array(dtype=int), flex_bending: wp.array2d(dtype=float), flex_damping: wp.array(dtype=float), @@ -142,6 +146,9 @@ def _forward_shim( flex_stiffness: wp.array2d(dtype=float), flex_vertadr: wp.array(dtype=int), flex_vertbodyid: 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), flexedge_invweight0: wp.array(dtype=float), flexedge_length0: wp.array(dtype=float), geom_aabb: wp.array3d(dtype=wp.vec3), @@ -166,12 +173,14 @@ def _forward_shim( geom_solmix: wp.array2d(dtype=float), geom_solref: wp.array2d(dtype=wp.vec2), geom_type: wp.array(dtype=int), + has_fluid: bool, has_sdf_geom: bool, hfield_adr: wp.array(dtype=int), hfield_data: wp.array(dtype=float), hfield_ncol: wp.array(dtype=int), hfield_nrow: wp.array(dtype=int), hfield_size: wp.array(dtype=wp.vec4), + is_sparse: bool, jnt_actfrclimited: wp.array(dtype=bool), jnt_actfrcrange: wp.array2d(dtype=wp.vec2), jnt_actgravcomp: wp.array(dtype=int), @@ -221,6 +230,7 @@ def _forward_shim( na: int, nacttrnbody: int, nbody: int, + nbranch: int, ncam: int, neq: int, nflex: int, @@ -264,9 +274,9 @@ def _forward_shim( qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...], qM_fullm_i: wp.array(dtype=int), qM_fullm_j: wp.array(dtype=int), - qM_madr_ij: wp.array(dtype=int), - qM_mulm_i: wp.array(dtype=int), - qM_mulm_j: wp.array(dtype=int), + qM_mulm_col: wp.array(dtype=int), + qM_mulm_madr: wp.array(dtype=int), + qM_mulm_rowadr: wp.array(dtype=int), qM_tiles: tuple[mjwp_types.TileSet, ...], qpos0: wp.array2d(dtype=float), qpos_spring: wp.array2d(dtype=float), @@ -343,9 +353,7 @@ def _forward_shim( opt__enableflags: int, opt__graph_conditional: bool, opt__gravity: wp.array(dtype=wp.vec3), - opt__has_fluid: bool, opt__impratio_invsqrt: wp.array(dtype=float), - opt__is_sparse: bool, opt__iterations: int, opt__ls_iterations: int, opt__ls_parallel: bool, @@ -360,8 +368,9 @@ def _forward_shim( opt__tolerance: wp.array(dtype=float), opt__viscosity: wp.array(dtype=float), opt__wind: wp.array(dtype=wp.vec3), - stat__meaninertia: float, + stat__meaninertia: wp.array(dtype=float), # Data + naccdmax: int, naconmax: int, njmax: int, act: wp.array2d(dtype=float), @@ -378,9 +387,6 @@ def _forward_shim( cfrc_ext: wp.array2d(dtype=wp.spatial_vector), cfrc_int: wp.array2d(dtype=wp.spatial_vector), cinert: wp.array2d(dtype=mjwp_types.vec10), - collision_pair: wp.array(dtype=wp.vec2i), - collision_pairid: wp.array(dtype=wp.vec2i), - collision_worldid: wp.array(dtype=int), crb: wp.array2d(dtype=mjwp_types.vec10), ctrl: wp.array2d(dtype=float), cvel: wp.array2d(dtype=wp.spatial_vector), @@ -399,15 +405,9 @@ def _forward_shim( nacon: wp.array(dtype=int), ncollision: wp.array(dtype=int), ne: wp.array(dtype=int), - ne_connect: wp.array(dtype=int), - ne_flex: wp.array(dtype=int), - ne_jnt: wp.array(dtype=int), - ne_ten: wp.array(dtype=int), - ne_weld: wp.array(dtype=int), nefc: wp.array(dtype=int), nf: wp.array(dtype=int), nl: wp.array(dtype=int), - nsolving: wp.array(dtype=int), qLD: wp.array3d(dtype=float), qLDiagInv: wp.array2d(dtype=float), qM: wp.array3d(dtype=float), @@ -431,7 +431,6 @@ def _forward_shim( site_xpos: wp.array2d(dtype=wp.vec3), solver_niter: wp.array(dtype=int), subtree_angmom: wp.array2d(dtype=wp.vec3), - subtree_bodyvel: wp.array2d(dtype=wp.spatial_vector), subtree_com: wp.array2d(dtype=wp.vec3), subtree_linvel: wp.array2d(dtype=wp.vec3), ten_J: wp.array3d(dtype=float), @@ -466,31 +465,13 @@ def _forward_shim( contact__worldid: wp.array(dtype=int), efc__D: wp.array2d(dtype=float), efc__J: wp.array3d(dtype=float), - efc__Jaref: wp.array2d(dtype=float), efc__Ma: wp.array2d(dtype=float), - efc__Mgrad: wp.array2d(dtype=float), - efc__alpha: wp.array(dtype=float), efc__aref: wp.array2d(dtype=float), - efc__beta: wp.array(dtype=float), - efc__cost: wp.array(dtype=float), - efc__done: wp.array(dtype=bool), efc__force: wp.array2d(dtype=float), efc__frictionloss: wp.array2d(dtype=float), - efc__gauss: wp.array(dtype=float), - efc__grad: wp.array2d(dtype=float), - efc__grad_dot: wp.array(dtype=float), efc__id: wp.array2d(dtype=int), - efc__jv: wp.array2d(dtype=float), efc__margin: wp.array2d(dtype=float), - efc__mv: wp.array2d(dtype=float), efc__pos: wp.array2d(dtype=float), - efc__prev_Mgrad: wp.array2d(dtype=float), - efc__prev_cost: wp.array(dtype=float), - efc__prev_grad: wp.array2d(dtype=float), - efc__quad: wp.array2d(dtype=wp.vec3), - efc__quad_gauss: wp.array(dtype=wp.vec3), - efc__search: wp.array2d(dtype=float), - efc__search_dot: wp.array(dtype=float), efc__state: wp.array2d(dtype=int), efc__type: wp.array2d(dtype=int), efc__vel: wp.array2d(dtype=float), @@ -524,6 +505,8 @@ def _forward_shim( _m.actuator_trntype = actuator_trntype _m.actuator_trntype_body_adr = actuator_trntype_body_adr _m.block_dim = block_dim + _m.body_branch_start = body_branch_start + _m.body_branches = body_branches _m.body_dofadr = body_dofadr _m.body_dofnum = body_dofnum _m.body_fluid_ellipsoid = body_fluid_ellipsoid @@ -579,6 +562,7 @@ def _forward_shim( _m.eq_solimp = eq_solimp _m.eq_solref = eq_solref _m.eq_ten_adr = eq_ten_adr + _m.eq_type = eq_type _m.eq_wld_adr = eq_wld_adr _m.flex_bending = flex_bending _m.flex_damping = flex_damping @@ -595,6 +579,9 @@ def _forward_shim( _m.flex_stiffness = flex_stiffness _m.flex_vertadr = flex_vertadr _m.flex_vertbodyid = flex_vertbodyid + _m.flexedge_J_colind = flexedge_J_colind + _m.flexedge_J_rowadr = flexedge_J_rowadr + _m.flexedge_J_rownnz = flexedge_J_rownnz _m.flexedge_invweight0 = flexedge_invweight0 _m.flexedge_length0 = flexedge_length0 _m.geom_aabb = geom_aabb @@ -619,12 +606,14 @@ def _forward_shim( _m.geom_solmix = geom_solmix _m.geom_solref = geom_solref _m.geom_type = geom_type + _m.has_fluid = has_fluid _m.has_sdf_geom = has_sdf_geom _m.hfield_adr = hfield_adr _m.hfield_data = hfield_data _m.hfield_ncol = hfield_ncol _m.hfield_nrow = hfield_nrow _m.hfield_size = hfield_size + _m.is_sparse = is_sparse _m.jnt_actfrclimited = jnt_actfrclimited _m.jnt_actfrcrange = jnt_actfrcrange _m.jnt_actgravcomp = jnt_actgravcomp @@ -674,6 +663,7 @@ def _forward_shim( _m.na = na _m.nacttrnbody = nacttrnbody _m.nbody = nbody + _m.nbranch = nbranch _m.ncam = ncam _m.neq = neq _m.nflex = nflex @@ -716,9 +706,7 @@ def _forward_shim( _m.opt.enableflags = opt__enableflags _m.opt.graph_conditional = opt__graph_conditional _m.opt.gravity = opt__gravity - _m.opt.has_fluid = opt__has_fluid _m.opt.impratio_invsqrt = opt__impratio_invsqrt - _m.opt.is_sparse = opt__is_sparse _m.opt.iterations = opt__iterations _m.opt.ls_iterations = opt__ls_iterations _m.opt.ls_parallel = opt__ls_parallel @@ -745,9 +733,9 @@ def _forward_shim( _m.qLD_updates = qLD_updates _m.qM_fullm_i = qM_fullm_i _m.qM_fullm_j = qM_fullm_j - _m.qM_madr_ij = qM_madr_ij - _m.qM_mulm_i = qM_mulm_i - _m.qM_mulm_j = qM_mulm_j + _m.qM_mulm_col = qM_mulm_col + _m.qM_mulm_madr = qM_mulm_madr + _m.qM_mulm_rowadr = qM_mulm_rowadr _m.qM_tiles = qM_tiles _m.qpos0 = qpos0 _m.qpos_spring = qpos_spring @@ -828,9 +816,6 @@ def _forward_shim( _d.cfrc_ext = cfrc_ext _d.cfrc_int = cfrc_int _d.cinert = cinert - _d.collision_pair = collision_pair - _d.collision_pairid = collision_pairid - _d.collision_worldid = collision_worldid _d.contact.dim = contact__dim _d.contact.dist = contact__dist _d.contact.efc_address = contact__efc_address @@ -850,31 +835,13 @@ def _forward_shim( _d.cvel = cvel _d.efc.D = efc__D _d.efc.J = efc__J - _d.efc.Jaref = efc__Jaref _d.efc.Ma = efc__Ma - _d.efc.Mgrad = efc__Mgrad - _d.efc.alpha = efc__alpha _d.efc.aref = efc__aref - _d.efc.beta = efc__beta - _d.efc.cost = efc__cost - _d.efc.done = efc__done _d.efc.force = efc__force _d.efc.frictionloss = efc__frictionloss - _d.efc.gauss = efc__gauss - _d.efc.grad = efc__grad - _d.efc.grad_dot = efc__grad_dot _d.efc.id = efc__id - _d.efc.jv = efc__jv _d.efc.margin = efc__margin - _d.efc.mv = efc__mv _d.efc.pos = efc__pos - _d.efc.prev_Mgrad = efc__prev_Mgrad - _d.efc.prev_cost = efc__prev_cost - _d.efc.prev_grad = efc__prev_grad - _d.efc.quad = efc__quad - _d.efc.quad_gauss = efc__quad_gauss - _d.efc.search = efc__search - _d.efc.search_dot = efc__search_dot _d.efc.state = efc__state _d.efc.type = efc__type _d.efc.vel = efc__vel @@ -890,20 +857,15 @@ def _forward_shim( _d.light_xpos = light_xpos _d.mocap_pos = mocap_pos _d.mocap_quat = mocap_quat + _d.naccdmax = naccdmax _d.nacon = nacon _d.naconmax = naconmax _d.ncollision = ncollision _d.ne = ne - _d.ne_connect = ne_connect - _d.ne_flex = ne_flex - _d.ne_jnt = ne_jnt - _d.ne_ten = ne_ten - _d.ne_weld = ne_weld _d.nefc = nefc _d.nf = nf _d.njmax = njmax _d.nl = nl - _d.nsolving = nsolving _d.qLD = qLD _d.qLDiagInv = qLDiagInv _d.qM = qM @@ -927,7 +889,6 @@ def _forward_shim( _d.site_xpos = site_xpos _d.solver_niter = solver_niter _d.subtree_angmom = subtree_angmom - _d.subtree_bodyvel = subtree_bodyvel _d.subtree_com = subtree_com _d.subtree_linvel = subtree_linvel _d.ten_J = ten_J @@ -965,9 +926,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'cfrc_ext': d._impl.cfrc_ext.shape, 'cfrc_int': d._impl.cfrc_int.shape, 'cinert': d._impl.cinert.shape, - 'collision_pair': d._impl.collision_pair.shape, - 'collision_pairid': d._impl.collision_pairid.shape, - 'collision_worldid': d._impl.collision_worldid.shape, 'crb': d._impl.crb.shape, 'cvel': d.cvel.shape, 'energy': d._impl.energy.shape, @@ -982,15 +940,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'nacon': d._impl.nacon.shape, 'ncollision': d._impl.ncollision.shape, 'ne': d._impl.ne.shape, - 'ne_connect': d._impl.ne_connect.shape, - 'ne_flex': d._impl.ne_flex.shape, - 'ne_jnt': d._impl.ne_jnt.shape, - 'ne_ten': d._impl.ne_ten.shape, - 'ne_weld': d._impl.ne_weld.shape, 'nefc': d._impl.nefc.shape, 'nf': d._impl.nf.shape, 'nl': d._impl.nl.shape, - 'nsolving': d._impl.nsolving.shape, 'qLD': d._impl.qLD.shape, 'qLDiagInv': d._impl.qLDiagInv.shape, 'qM': d._impl.qM.shape, @@ -1011,7 +963,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'site_xpos': d.site_xpos.shape, 'solver_niter': d._impl.solver_niter.shape, 'subtree_angmom': d._impl.subtree_angmom.shape, - 'subtree_bodyvel': d._impl.subtree_bodyvel.shape, 'subtree_com': d.subtree_com.shape, 'subtree_linvel': d._impl.subtree_linvel.shape, 'ten_J': d._impl.ten_J.shape, @@ -1044,38 +995,20 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__worldid': d._impl.contact__worldid.shape, 'efc__D': d._impl.efc__D.shape, 'efc__J': d._impl.efc__J.shape, - 'efc__Jaref': d._impl.efc__Jaref.shape, 'efc__Ma': d._impl.efc__Ma.shape, - 'efc__Mgrad': d._impl.efc__Mgrad.shape, - 'efc__alpha': d._impl.efc__alpha.shape, 'efc__aref': d._impl.efc__aref.shape, - 'efc__beta': d._impl.efc__beta.shape, - 'efc__cost': d._impl.efc__cost.shape, - 'efc__done': d._impl.efc__done.shape, 'efc__force': d._impl.efc__force.shape, 'efc__frictionloss': d._impl.efc__frictionloss.shape, - 'efc__gauss': d._impl.efc__gauss.shape, - 'efc__grad': d._impl.efc__grad.shape, - 'efc__grad_dot': d._impl.efc__grad_dot.shape, 'efc__id': d._impl.efc__id.shape, - 'efc__jv': d._impl.efc__jv.shape, 'efc__margin': d._impl.efc__margin.shape, - 'efc__mv': d._impl.efc__mv.shape, 'efc__pos': d._impl.efc__pos.shape, - 'efc__prev_Mgrad': d._impl.efc__prev_Mgrad.shape, - 'efc__prev_cost': d._impl.efc__prev_cost.shape, - 'efc__prev_grad': d._impl.efc__prev_grad.shape, - 'efc__quad': d._impl.efc__quad.shape, - 'efc__quad_gauss': d._impl.efc__quad_gauss.shape, - 'efc__search': d._impl.efc__search.shape, - 'efc__search_dot': d._impl.efc__search_dot.shape, 'efc__state': d._impl.efc__state.shape, 'efc__type': d._impl.efc__type.shape, 'efc__vel': d._impl.efc__vel.shape, } jf = ffi.jax_callable_variadic_tuple( _forward_shim, - num_outputs=120, + num_outputs=92, output_dims=output_dims, vmap_method=None, in_out_argnames={ @@ -1092,9 +1025,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'cfrc_ext', 'cfrc_int', 'cinert', - 'collision_pair', - 'collision_pairid', - 'collision_worldid', 'crb', 'cvel', 'energy', @@ -1109,15 +1039,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'nacon', 'ncollision', 'ne', - 'ne_connect', - 'ne_flex', - 'ne_jnt', - 'ne_ten', - 'ne_weld', 'nefc', 'nf', 'nl', - 'nsolving', 'qLD', 'qLDiagInv', 'qM', @@ -1138,7 +1062,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'site_xpos', 'solver_niter', 'subtree_angmom', - 'subtree_bodyvel', 'subtree_com', 'subtree_linvel', 'ten_J', @@ -1171,31 +1094,13 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__worldid', 'efc__D', 'efc__J', - 'efc__Jaref', 'efc__Ma', - 'efc__Mgrad', - 'efc__alpha', 'efc__aref', - 'efc__beta', - 'efc__cost', - 'efc__done', 'efc__force', 'efc__frictionloss', - 'efc__gauss', - 'efc__grad', - 'efc__grad_dot', 'efc__id', - 'efc__jv', 'efc__margin', - 'efc__mv', 'efc__pos', - 'efc__prev_Mgrad', - 'efc__prev_cost', - 'efc__prev_grad', - 'efc__quad', - 'efc__quad_gauss', - 'efc__search', - 'efc__search_dot', 'efc__state', 'efc__type', 'efc__vel', @@ -1222,6 +1127,8 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'body_pos', 'body_quat', 'body_subtreemass', + 'cam_fovy', + 'cam_intrinsic', 'cam_mat0', 'cam_pos', 'cam_pos0', @@ -1398,6 +1305,8 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.actuator_trntype, m._impl.actuator_trntype_body_adr, m._impl.block_dim, + m._impl.body_branch_start, + m._impl.body_branches, m.body_dofadr, m.body_dofnum, m._impl.body_fluid_ellipsoid, @@ -1453,6 +1362,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.eq_solimp, m.eq_solref, m._impl.eq_ten_adr, + m.eq_type, m._impl.eq_wld_adr, m._impl.flex_bending, m._impl.flex_damping, @@ -1469,6 +1379,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.flex_stiffness, m._impl.flex_vertadr, m._impl.flex_vertbodyid, + m._impl.flexedge_J_colind, + m._impl.flexedge_J_rowadr, + m._impl.flexedge_J_rownnz, m._impl.flexedge_invweight0, m._impl.flexedge_length0, m.geom_aabb, @@ -1493,12 +1406,14 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.geom_solmix, m.geom_solref, m.geom_type, + m._impl.has_fluid, m._impl.has_sdf_geom, m.hfield_adr, m.hfield_data, m.hfield_ncol, m.hfield_nrow, m.hfield_size, + m._impl.is_sparse, m.jnt_actfrclimited, m.jnt_actfrcrange, m.jnt_actgravcomp, @@ -1548,6 +1463,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.na, m._impl.nacttrnbody, m.nbody, + m._impl.nbranch, m.ncam, m.neq, m._impl.nflex, @@ -1591,9 +1507,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.qLD_updates, m._impl.qM_fullm_i, m._impl.qM_fullm_j, - m._impl.qM_madr_ij, - m._impl.qM_mulm_i, - m._impl.qM_mulm_j, + m._impl.qM_mulm_col, + m._impl.qM_mulm_madr, + m._impl.qM_mulm_rowadr, m._impl.qM_tiles, m.qpos0, m.qpos_spring, @@ -1670,9 +1586,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.opt.enableflags, m.opt._impl.graph_conditional, m.opt.gravity, - m.opt._impl.has_fluid, m.opt._impl.impratio_invsqrt, - m.opt._impl.is_sparse, m.opt.iterations, m.opt.ls_iterations, m.opt._impl.ls_parallel, @@ -1688,6 +1602,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.opt.viscosity, m.opt.wind, m.stat.meaninertia, + d._impl.naccdmax, d._impl.naconmax, d._impl.njmax, d.act, @@ -1704,9 +1619,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.cfrc_ext, d._impl.cfrc_int, d._impl.cinert, - d._impl.collision_pair, - d._impl.collision_pairid, - d._impl.collision_worldid, d._impl.crb, d.ctrl, d.cvel, @@ -1725,15 +1637,9 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.nacon, d._impl.ncollision, d._impl.ne, - d._impl.ne_connect, - d._impl.ne_flex, - d._impl.ne_jnt, - d._impl.ne_ten, - d._impl.ne_weld, d._impl.nefc, d._impl.nf, d._impl.nl, - d._impl.nsolving, d._impl.qLD, d._impl.qLDiagInv, d._impl.qM, @@ -1757,7 +1663,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d.site_xpos, d._impl.solver_niter, d._impl.subtree_angmom, - d._impl.subtree_bodyvel, d.subtree_com, d._impl.subtree_linvel, d._impl.ten_J, @@ -1792,31 +1697,13 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.contact__worldid, d._impl.efc__D, d._impl.efc__J, - d._impl.efc__Jaref, d._impl.efc__Ma, - d._impl.efc__Mgrad, - d._impl.efc__alpha, d._impl.efc__aref, - d._impl.efc__beta, - d._impl.efc__cost, - d._impl.efc__done, d._impl.efc__force, d._impl.efc__frictionloss, - d._impl.efc__gauss, - d._impl.efc__grad, - d._impl.efc__grad_dot, d._impl.efc__id, - d._impl.efc__jv, d._impl.efc__margin, - d._impl.efc__mv, d._impl.efc__pos, - d._impl.efc__prev_Mgrad, - d._impl.efc__prev_cost, - d._impl.efc__prev_grad, - d._impl.efc__quad, - d._impl.efc__quad_gauss, - d._impl.efc__search, - d._impl.efc__search_dot, d._impl.efc__state, d._impl.efc__type, d._impl.efc__vel, @@ -1835,113 +1722,85 @@ def _forward_jax_impl(m: types.Model, d: types.Data): '_impl.cfrc_ext': out[10], '_impl.cfrc_int': out[11], '_impl.cinert': out[12], - '_impl.collision_pair': out[13], - '_impl.collision_pairid': out[14], - '_impl.collision_worldid': out[15], - '_impl.crb': out[16], - 'cvel': out[17], - '_impl.energy': out[18], - '_impl.flexedge_J': out[19], - '_impl.flexedge_length': out[20], - '_impl.flexedge_velocity': out[21], - '_impl.flexvert_xpos': out[22], - 'geom_xmat': out[23], - 'geom_xpos': out[24], - '_impl.light_xdir': out[25], - '_impl.light_xpos': out[26], - '_impl.nacon': out[27], - '_impl.ncollision': out[28], - '_impl.ne': out[29], - '_impl.ne_connect': out[30], - '_impl.ne_flex': out[31], - '_impl.ne_jnt': out[32], - '_impl.ne_ten': out[33], - '_impl.ne_weld': out[34], - '_impl.nefc': out[35], - '_impl.nf': out[36], - '_impl.nl': out[37], - '_impl.nsolving': out[38], - '_impl.qLD': out[39], - '_impl.qLDiagInv': out[40], - '_impl.qM': out[41], - 'qacc': out[42], - 'qacc_smooth': out[43], - 'qfrc_actuator': out[44], - 'qfrc_bias': out[45], - 'qfrc_constraint': out[46], - '_impl.qfrc_damper': out[47], - 'qfrc_fluid': out[48], - 'qfrc_gravcomp': out[49], - 'qfrc_passive': out[50], - 'qfrc_smooth': out[51], - '_impl.qfrc_spring': out[52], - 'qvel': out[53], - 'sensordata': out[54], - 'site_xmat': out[55], - 'site_xpos': out[56], - '_impl.solver_niter': out[57], - '_impl.subtree_angmom': out[58], - '_impl.subtree_bodyvel': out[59], - 'subtree_com': out[60], - '_impl.subtree_linvel': out[61], - '_impl.ten_J': out[62], - 'ten_length': out[63], - '_impl.ten_velocity': out[64], - '_impl.ten_wrapadr': out[65], - '_impl.ten_wrapnum': out[66], - '_impl.wrap_obj': out[67], - '_impl.wrap_xpos': out[68], - 'xanchor': out[69], - 'xaxis': out[70], - 'ximat': out[71], - 'xipos': out[72], - 'xmat': out[73], - 'xpos': out[74], - 'xquat': out[75], - '_impl.contact__dim': out[76], - '_impl.contact__dist': out[77], - '_impl.contact__efc_address': 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__worldid': out[89], - '_impl.efc__D': out[90], - '_impl.efc__J': out[91], - '_impl.efc__Jaref': out[92], - '_impl.efc__Ma': out[93], - '_impl.efc__Mgrad': out[94], - '_impl.efc__alpha': out[95], - '_impl.efc__aref': out[96], - '_impl.efc__beta': out[97], - '_impl.efc__cost': out[98], - '_impl.efc__done': out[99], - '_impl.efc__force': out[100], - '_impl.efc__frictionloss': out[101], - '_impl.efc__gauss': out[102], - '_impl.efc__grad': out[103], - '_impl.efc__grad_dot': out[104], - '_impl.efc__id': out[105], - '_impl.efc__jv': out[106], - '_impl.efc__margin': out[107], - '_impl.efc__mv': out[108], - '_impl.efc__pos': out[109], - '_impl.efc__prev_Mgrad': out[110], - '_impl.efc__prev_cost': out[111], - '_impl.efc__prev_grad': out[112], - '_impl.efc__quad': out[113], - '_impl.efc__quad_gauss': out[114], - '_impl.efc__search': out[115], - '_impl.efc__search_dot': out[116], - '_impl.efc__state': out[117], - '_impl.efc__type': out[118], - '_impl.efc__vel': out[119], + '_impl.crb': out[13], + 'cvel': out[14], + '_impl.energy': out[15], + '_impl.flexedge_J': out[16], + '_impl.flexedge_length': out[17], + '_impl.flexedge_velocity': out[18], + '_impl.flexvert_xpos': out[19], + 'geom_xmat': out[20], + 'geom_xpos': out[21], + '_impl.light_xdir': out[22], + '_impl.light_xpos': out[23], + '_impl.nacon': out[24], + '_impl.ncollision': out[25], + '_impl.ne': out[26], + '_impl.nefc': out[27], + '_impl.nf': out[28], + '_impl.nl': out[29], + '_impl.qLD': out[30], + '_impl.qLDiagInv': out[31], + '_impl.qM': out[32], + 'qacc': out[33], + 'qacc_smooth': out[34], + 'qfrc_actuator': out[35], + 'qfrc_bias': out[36], + 'qfrc_constraint': out[37], + '_impl.qfrc_damper': out[38], + 'qfrc_fluid': out[39], + 'qfrc_gravcomp': out[40], + 'qfrc_passive': out[41], + 'qfrc_smooth': out[42], + '_impl.qfrc_spring': out[43], + 'qvel': out[44], + 'sensordata': out[45], + 'site_xmat': out[46], + 'site_xpos': out[47], + '_impl.solver_niter': out[48], + '_impl.subtree_angmom': out[49], + 'subtree_com': out[50], + '_impl.subtree_linvel': out[51], + '_impl.ten_J': out[52], + 'ten_length': out[53], + '_impl.ten_velocity': out[54], + '_impl.ten_wrapadr': out[55], + '_impl.ten_wrapnum': out[56], + '_impl.wrap_obj': out[57], + '_impl.wrap_xpos': out[58], + 'xanchor': out[59], + 'xaxis': out[60], + 'ximat': out[61], + 'xipos': out[62], + 'xmat': out[63], + 'xpos': out[64], + 'xquat': out[65], + '_impl.contact__dim': out[66], + '_impl.contact__dist': out[67], + '_impl.contact__efc_address': out[68], + '_impl.contact__frame': out[69], + '_impl.contact__friction': out[70], + '_impl.contact__geom': out[71], + '_impl.contact__geomcollisionid': out[72], + '_impl.contact__includemargin': out[73], + '_impl.contact__pos': out[74], + '_impl.contact__solimp': out[75], + '_impl.contact__solref': out[76], + '_impl.contact__solreffriction': out[77], + '_impl.contact__type': out[78], + '_impl.contact__worldid': out[79], + '_impl.efc__D': out[80], + '_impl.efc__J': out[81], + '_impl.efc__Ma': out[82], + '_impl.efc__aref': out[83], + '_impl.efc__force': out[84], + '_impl.efc__frictionloss': out[85], + '_impl.efc__id': out[86], + '_impl.efc__margin': out[87], + '_impl.efc__pos': out[88], + '_impl.efc__state': out[89], + '_impl.efc__type': out[90], + '_impl.efc__vel': out[91], }) return d @@ -1986,6 +1845,8 @@ def _step_shim( actuator_trntype: wp.array(dtype=int), actuator_trntype_body_adr: wp.array(dtype=int), block_dim: mjwp_types.BlockDim, + body_branch_start: wp.array(dtype=int), + body_branches: wp.array(dtype=int), body_dofadr: wp.array(dtype=int), body_dofnum: wp.array(dtype=int), body_fluid_ellipsoid: wp.array(dtype=bool), @@ -2008,8 +1869,8 @@ def _step_shim( body_tree: tuple[wp.array(dtype=int), ...], body_weldid: wp.array(dtype=int), cam_bodyid: wp.array(dtype=int), - cam_fovy: wp.array(dtype=float), - cam_intrinsic: wp.array(dtype=wp.vec4), + cam_fovy: wp.array2d(dtype=float), + cam_intrinsic: wp.array2d(dtype=wp.vec4), cam_mat0: wp.array2d(dtype=wp.mat33), cam_mode: wp.array(dtype=int), cam_pos: wp.array2d(dtype=wp.vec3), @@ -2041,6 +1902,7 @@ def _step_shim( eq_solimp: wp.array2d(dtype=mjwp_types.vec5), eq_solref: wp.array2d(dtype=wp.vec2), eq_ten_adr: wp.array(dtype=int), + eq_type: wp.array(dtype=int), eq_wld_adr: wp.array(dtype=int), flex_bending: wp.array2d(dtype=float), flex_damping: wp.array(dtype=float), @@ -2057,6 +1919,9 @@ def _step_shim( flex_stiffness: wp.array2d(dtype=float), flex_vertadr: wp.array(dtype=int), flex_vertbodyid: 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), flexedge_invweight0: wp.array(dtype=float), flexedge_length0: wp.array(dtype=float), geom_aabb: wp.array3d(dtype=wp.vec3), @@ -2081,12 +1946,14 @@ def _step_shim( geom_solmix: wp.array2d(dtype=float), geom_solref: wp.array2d(dtype=wp.vec2), geom_type: wp.array(dtype=int), + has_fluid: bool, has_sdf_geom: bool, hfield_adr: wp.array(dtype=int), hfield_data: wp.array(dtype=float), hfield_ncol: wp.array(dtype=int), hfield_nrow: wp.array(dtype=int), hfield_size: wp.array(dtype=wp.vec4), + is_sparse: bool, jnt_actfrclimited: wp.array(dtype=bool), jnt_actfrcrange: wp.array2d(dtype=wp.vec2), jnt_actgravcomp: wp.array(dtype=int), @@ -2137,6 +2004,7 @@ def _step_shim( na: int, nacttrnbody: int, nbody: int, + nbranch: int, ncam: int, neq: int, nflex: int, @@ -2180,9 +2048,9 @@ def _step_shim( qLD_updates: tuple[wp.array(dtype=wp.vec3i), ...], qM_fullm_i: wp.array(dtype=int), qM_fullm_j: wp.array(dtype=int), - qM_madr_ij: wp.array(dtype=int), - qM_mulm_i: wp.array(dtype=int), - qM_mulm_j: wp.array(dtype=int), + qM_mulm_col: wp.array(dtype=int), + qM_mulm_madr: wp.array(dtype=int), + qM_mulm_rowadr: wp.array(dtype=int), qM_tiles: tuple[mjwp_types.TileSet, ...], qpos0: wp.array2d(dtype=float), qpos_spring: wp.array2d(dtype=float), @@ -2259,10 +2127,8 @@ def _step_shim( opt__enableflags: int, opt__graph_conditional: bool, opt__gravity: wp.array(dtype=wp.vec3), - opt__has_fluid: bool, opt__impratio_invsqrt: wp.array(dtype=float), opt__integrator: int, - opt__is_sparse: bool, opt__iterations: int, opt__ls_iterations: int, opt__ls_parallel: bool, @@ -2277,8 +2143,9 @@ def _step_shim( opt__tolerance: wp.array(dtype=float), opt__viscosity: wp.array(dtype=float), opt__wind: wp.array(dtype=wp.vec3), - stat__meaninertia: float, + stat__meaninertia: wp.array(dtype=float), # Data + naccdmax: int, naconmax: int, njmax: int, act: wp.array2d(dtype=float), @@ -2295,9 +2162,6 @@ def _step_shim( cfrc_ext: wp.array2d(dtype=wp.spatial_vector), cfrc_int: wp.array2d(dtype=wp.spatial_vector), cinert: wp.array2d(dtype=mjwp_types.vec10), - collision_pair: wp.array(dtype=wp.vec2i), - collision_pairid: wp.array(dtype=wp.vec2i), - collision_worldid: wp.array(dtype=int), crb: wp.array2d(dtype=mjwp_types.vec10), ctrl: wp.array2d(dtype=float), cvel: wp.array2d(dtype=wp.spatial_vector), @@ -2316,15 +2180,9 @@ def _step_shim( nacon: wp.array(dtype=int), ncollision: wp.array(dtype=int), ne: wp.array(dtype=int), - ne_connect: wp.array(dtype=int), - ne_flex: wp.array(dtype=int), - ne_jnt: wp.array(dtype=int), - ne_ten: wp.array(dtype=int), - ne_weld: wp.array(dtype=int), nefc: wp.array(dtype=int), nf: wp.array(dtype=int), nl: wp.array(dtype=int), - nsolving: wp.array(dtype=int), qLD: wp.array3d(dtype=float), qLDiagInv: wp.array2d(dtype=float), qM: wp.array3d(dtype=float), @@ -2348,7 +2206,6 @@ def _step_shim( site_xpos: wp.array2d(dtype=wp.vec3), solver_niter: wp.array(dtype=int), subtree_angmom: wp.array2d(dtype=wp.vec3), - subtree_bodyvel: wp.array2d(dtype=wp.spatial_vector), subtree_com: wp.array2d(dtype=wp.vec3), subtree_linvel: wp.array2d(dtype=wp.vec3), ten_J: wp.array3d(dtype=float), @@ -2383,31 +2240,13 @@ def _step_shim( contact__worldid: wp.array(dtype=int), efc__D: wp.array2d(dtype=float), efc__J: wp.array3d(dtype=float), - efc__Jaref: wp.array2d(dtype=float), efc__Ma: wp.array2d(dtype=float), - efc__Mgrad: wp.array2d(dtype=float), - efc__alpha: wp.array(dtype=float), efc__aref: wp.array2d(dtype=float), - efc__beta: wp.array(dtype=float), - efc__cost: wp.array(dtype=float), - efc__done: wp.array(dtype=bool), efc__force: wp.array2d(dtype=float), efc__frictionloss: wp.array2d(dtype=float), - efc__gauss: wp.array(dtype=float), - efc__grad: wp.array2d(dtype=float), - efc__grad_dot: wp.array(dtype=float), efc__id: wp.array2d(dtype=int), - efc__jv: wp.array2d(dtype=float), efc__margin: wp.array2d(dtype=float), - efc__mv: wp.array2d(dtype=float), efc__pos: wp.array2d(dtype=float), - efc__prev_Mgrad: wp.array2d(dtype=float), - efc__prev_cost: wp.array(dtype=float), - efc__prev_grad: wp.array2d(dtype=float), - efc__quad: wp.array2d(dtype=wp.vec3), - efc__quad_gauss: wp.array(dtype=wp.vec3), - efc__search: wp.array2d(dtype=float), - efc__search_dot: wp.array(dtype=float), efc__state: wp.array2d(dtype=int), efc__type: wp.array2d(dtype=int), efc__vel: wp.array2d(dtype=float), @@ -2441,6 +2280,8 @@ def _step_shim( _m.actuator_trntype = actuator_trntype _m.actuator_trntype_body_adr = actuator_trntype_body_adr _m.block_dim = block_dim + _m.body_branch_start = body_branch_start + _m.body_branches = body_branches _m.body_dofadr = body_dofadr _m.body_dofnum = body_dofnum _m.body_fluid_ellipsoid = body_fluid_ellipsoid @@ -2496,6 +2337,7 @@ def _step_shim( _m.eq_solimp = eq_solimp _m.eq_solref = eq_solref _m.eq_ten_adr = eq_ten_adr + _m.eq_type = eq_type _m.eq_wld_adr = eq_wld_adr _m.flex_bending = flex_bending _m.flex_damping = flex_damping @@ -2512,6 +2354,9 @@ def _step_shim( _m.flex_stiffness = flex_stiffness _m.flex_vertadr = flex_vertadr _m.flex_vertbodyid = flex_vertbodyid + _m.flexedge_J_colind = flexedge_J_colind + _m.flexedge_J_rowadr = flexedge_J_rowadr + _m.flexedge_J_rownnz = flexedge_J_rownnz _m.flexedge_invweight0 = flexedge_invweight0 _m.flexedge_length0 = flexedge_length0 _m.geom_aabb = geom_aabb @@ -2536,12 +2381,14 @@ def _step_shim( _m.geom_solmix = geom_solmix _m.geom_solref = geom_solref _m.geom_type = geom_type + _m.has_fluid = has_fluid _m.has_sdf_geom = has_sdf_geom _m.hfield_adr = hfield_adr _m.hfield_data = hfield_data _m.hfield_ncol = hfield_ncol _m.hfield_nrow = hfield_nrow _m.hfield_size = hfield_size + _m.is_sparse = is_sparse _m.jnt_actfrclimited = jnt_actfrclimited _m.jnt_actfrcrange = jnt_actfrcrange _m.jnt_actgravcomp = jnt_actgravcomp @@ -2592,6 +2439,7 @@ def _step_shim( _m.na = na _m.nacttrnbody = nacttrnbody _m.nbody = nbody + _m.nbranch = nbranch _m.ncam = ncam _m.neq = neq _m.nflex = nflex @@ -2634,10 +2482,8 @@ def _step_shim( _m.opt.enableflags = opt__enableflags _m.opt.graph_conditional = opt__graph_conditional _m.opt.gravity = opt__gravity - _m.opt.has_fluid = opt__has_fluid _m.opt.impratio_invsqrt = opt__impratio_invsqrt _m.opt.integrator = opt__integrator - _m.opt.is_sparse = opt__is_sparse _m.opt.iterations = opt__iterations _m.opt.ls_iterations = opt__ls_iterations _m.opt.ls_parallel = opt__ls_parallel @@ -2664,9 +2510,9 @@ def _step_shim( _m.qLD_updates = qLD_updates _m.qM_fullm_i = qM_fullm_i _m.qM_fullm_j = qM_fullm_j - _m.qM_madr_ij = qM_madr_ij - _m.qM_mulm_i = qM_mulm_i - _m.qM_mulm_j = qM_mulm_j + _m.qM_mulm_col = qM_mulm_col + _m.qM_mulm_madr = qM_mulm_madr + _m.qM_mulm_rowadr = qM_mulm_rowadr _m.qM_tiles = qM_tiles _m.qpos0 = qpos0 _m.qpos_spring = qpos_spring @@ -2747,9 +2593,6 @@ def _step_shim( _d.cfrc_ext = cfrc_ext _d.cfrc_int = cfrc_int _d.cinert = cinert - _d.collision_pair = collision_pair - _d.collision_pairid = collision_pairid - _d.collision_worldid = collision_worldid _d.contact.dim = contact__dim _d.contact.dist = contact__dist _d.contact.efc_address = contact__efc_address @@ -2769,31 +2612,13 @@ def _step_shim( _d.cvel = cvel _d.efc.D = efc__D _d.efc.J = efc__J - _d.efc.Jaref = efc__Jaref _d.efc.Ma = efc__Ma - _d.efc.Mgrad = efc__Mgrad - _d.efc.alpha = efc__alpha _d.efc.aref = efc__aref - _d.efc.beta = efc__beta - _d.efc.cost = efc__cost - _d.efc.done = efc__done _d.efc.force = efc__force _d.efc.frictionloss = efc__frictionloss - _d.efc.gauss = efc__gauss - _d.efc.grad = efc__grad - _d.efc.grad_dot = efc__grad_dot _d.efc.id = efc__id - _d.efc.jv = efc__jv _d.efc.margin = efc__margin - _d.efc.mv = efc__mv _d.efc.pos = efc__pos - _d.efc.prev_Mgrad = efc__prev_Mgrad - _d.efc.prev_cost = efc__prev_cost - _d.efc.prev_grad = efc__prev_grad - _d.efc.quad = efc__quad - _d.efc.quad_gauss = efc__quad_gauss - _d.efc.search = efc__search - _d.efc.search_dot = efc__search_dot _d.efc.state = efc__state _d.efc.type = efc__type _d.efc.vel = efc__vel @@ -2809,20 +2634,15 @@ def _step_shim( _d.light_xpos = light_xpos _d.mocap_pos = mocap_pos _d.mocap_quat = mocap_quat + _d.naccdmax = naccdmax _d.nacon = nacon _d.naconmax = naconmax _d.ncollision = ncollision _d.ne = ne - _d.ne_connect = ne_connect - _d.ne_flex = ne_flex - _d.ne_jnt = ne_jnt - _d.ne_ten = ne_ten - _d.ne_weld = ne_weld _d.nefc = nefc _d.nf = nf _d.njmax = njmax _d.nl = nl - _d.nsolving = nsolving _d.qLD = qLD _d.qLDiagInv = qLDiagInv _d.qM = qM @@ -2846,7 +2666,6 @@ def _step_shim( _d.site_xpos = site_xpos _d.solver_niter = solver_niter _d.subtree_angmom = subtree_angmom - _d.subtree_bodyvel = subtree_bodyvel _d.subtree_com = subtree_com _d.subtree_linvel = subtree_linvel _d.ten_J = ten_J @@ -2885,9 +2704,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'cfrc_ext': d._impl.cfrc_ext.shape, 'cfrc_int': d._impl.cfrc_int.shape, 'cinert': d._impl.cinert.shape, - 'collision_pair': d._impl.collision_pair.shape, - 'collision_pairid': d._impl.collision_pairid.shape, - 'collision_worldid': d._impl.collision_worldid.shape, 'crb': d._impl.crb.shape, 'cvel': d.cvel.shape, 'energy': d._impl.energy.shape, @@ -2902,15 +2718,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'nacon': d._impl.nacon.shape, 'ncollision': d._impl.ncollision.shape, 'ne': d._impl.ne.shape, - 'ne_connect': d._impl.ne_connect.shape, - 'ne_flex': d._impl.ne_flex.shape, - 'ne_jnt': d._impl.ne_jnt.shape, - 'ne_ten': d._impl.ne_ten.shape, - 'ne_weld': d._impl.ne_weld.shape, 'nefc': d._impl.nefc.shape, 'nf': d._impl.nf.shape, 'nl': d._impl.nl.shape, - 'nsolving': d._impl.nsolving.shape, 'qLD': d._impl.qLD.shape, 'qLDiagInv': d._impl.qLDiagInv.shape, 'qM': d._impl.qM.shape, @@ -2933,7 +2743,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'site_xpos': d.site_xpos.shape, 'solver_niter': d._impl.solver_niter.shape, 'subtree_angmom': d._impl.subtree_angmom.shape, - 'subtree_bodyvel': d._impl.subtree_bodyvel.shape, 'subtree_com': d.subtree_com.shape, 'subtree_linvel': d._impl.subtree_linvel.shape, 'ten_J': d._impl.ten_J.shape, @@ -2967,38 +2776,20 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__worldid': d._impl.contact__worldid.shape, 'efc__D': d._impl.efc__D.shape, 'efc__J': d._impl.efc__J.shape, - 'efc__Jaref': d._impl.efc__Jaref.shape, 'efc__Ma': d._impl.efc__Ma.shape, - 'efc__Mgrad': d._impl.efc__Mgrad.shape, - 'efc__alpha': d._impl.efc__alpha.shape, 'efc__aref': d._impl.efc__aref.shape, - 'efc__beta': d._impl.efc__beta.shape, - 'efc__cost': d._impl.efc__cost.shape, - 'efc__done': d._impl.efc__done.shape, 'efc__force': d._impl.efc__force.shape, 'efc__frictionloss': d._impl.efc__frictionloss.shape, - 'efc__gauss': d._impl.efc__gauss.shape, - 'efc__grad': d._impl.efc__grad.shape, - 'efc__grad_dot': d._impl.efc__grad_dot.shape, 'efc__id': d._impl.efc__id.shape, - 'efc__jv': d._impl.efc__jv.shape, 'efc__margin': d._impl.efc__margin.shape, - 'efc__mv': d._impl.efc__mv.shape, 'efc__pos': d._impl.efc__pos.shape, - 'efc__prev_Mgrad': d._impl.efc__prev_Mgrad.shape, - 'efc__prev_cost': d._impl.efc__prev_cost.shape, - 'efc__prev_grad': d._impl.efc__prev_grad.shape, - 'efc__quad': d._impl.efc__quad.shape, - 'efc__quad_gauss': d._impl.efc__quad_gauss.shape, - 'efc__search': d._impl.efc__search.shape, - 'efc__search_dot': d._impl.efc__search_dot.shape, 'efc__state': d._impl.efc__state.shape, 'efc__type': d._impl.efc__type.shape, 'efc__vel': d._impl.efc__vel.shape, } jf = ffi.jax_callable_variadic_tuple( _step_shim, - num_outputs=124, + num_outputs=96, output_dims=output_dims, vmap_method=None, in_out_argnames={ @@ -3016,9 +2807,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'cfrc_ext', 'cfrc_int', 'cinert', - 'collision_pair', - 'collision_pairid', - 'collision_worldid', 'crb', 'cvel', 'energy', @@ -3033,15 +2821,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'nacon', 'ncollision', 'ne', - 'ne_connect', - 'ne_flex', - 'ne_jnt', - 'ne_ten', - 'ne_weld', 'nefc', 'nf', 'nl', - 'nsolving', 'qLD', 'qLDiagInv', 'qM', @@ -3064,7 +2846,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'site_xpos', 'solver_niter', 'subtree_angmom', - 'subtree_bodyvel', 'subtree_com', 'subtree_linvel', 'ten_J', @@ -3098,31 +2879,13 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__worldid', 'efc__D', 'efc__J', - 'efc__Jaref', 'efc__Ma', - 'efc__Mgrad', - 'efc__alpha', 'efc__aref', - 'efc__beta', - 'efc__cost', - 'efc__done', 'efc__force', 'efc__frictionloss', - 'efc__gauss', - 'efc__grad', - 'efc__grad_dot', 'efc__id', - 'efc__jv', 'efc__margin', - 'efc__mv', 'efc__pos', - 'efc__prev_Mgrad', - 'efc__prev_cost', - 'efc__prev_grad', - 'efc__quad', - 'efc__quad_gauss', - 'efc__search', - 'efc__search_dot', 'efc__state', 'efc__type', 'efc__vel', @@ -3149,6 +2912,8 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'body_pos', 'body_quat', 'body_subtreemass', + 'cam_fovy', + 'cam_intrinsic', 'cam_mat0', 'cam_pos', 'cam_pos0', @@ -3329,6 +3094,8 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.actuator_trntype, m._impl.actuator_trntype_body_adr, m._impl.block_dim, + m._impl.body_branch_start, + m._impl.body_branches, m.body_dofadr, m.body_dofnum, m._impl.body_fluid_ellipsoid, @@ -3384,6 +3151,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.eq_solimp, m.eq_solref, m._impl.eq_ten_adr, + m.eq_type, m._impl.eq_wld_adr, m._impl.flex_bending, m._impl.flex_damping, @@ -3400,6 +3168,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.flex_stiffness, m._impl.flex_vertadr, m._impl.flex_vertbodyid, + m._impl.flexedge_J_colind, + m._impl.flexedge_J_rowadr, + m._impl.flexedge_J_rownnz, m._impl.flexedge_invweight0, m._impl.flexedge_length0, m.geom_aabb, @@ -3424,12 +3195,14 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.geom_solmix, m.geom_solref, m.geom_type, + m._impl.has_fluid, m._impl.has_sdf_geom, m.hfield_adr, m.hfield_data, m.hfield_ncol, m.hfield_nrow, m.hfield_size, + m._impl.is_sparse, m.jnt_actfrclimited, m.jnt_actfrcrange, m.jnt_actgravcomp, @@ -3480,6 +3253,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.na, m._impl.nacttrnbody, m.nbody, + m._impl.nbranch, m.ncam, m.neq, m._impl.nflex, @@ -3523,9 +3297,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.qLD_updates, m._impl.qM_fullm_i, m._impl.qM_fullm_j, - m._impl.qM_madr_ij, - m._impl.qM_mulm_i, - m._impl.qM_mulm_j, + m._impl.qM_mulm_col, + m._impl.qM_mulm_madr, + m._impl.qM_mulm_rowadr, m._impl.qM_tiles, m.qpos0, m.qpos_spring, @@ -3602,10 +3376,8 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.opt.enableflags, m.opt._impl.graph_conditional, m.opt.gravity, - m.opt._impl.has_fluid, m.opt._impl.impratio_invsqrt, m.opt.integrator, - m.opt._impl.is_sparse, m.opt.iterations, m.opt.ls_iterations, m.opt._impl.ls_parallel, @@ -3621,6 +3393,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.opt.viscosity, m.opt.wind, m.stat.meaninertia, + d._impl.naccdmax, d._impl.naconmax, d._impl.njmax, d.act, @@ -3637,9 +3410,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.cfrc_ext, d._impl.cfrc_int, d._impl.cinert, - d._impl.collision_pair, - d._impl.collision_pairid, - d._impl.collision_worldid, d._impl.crb, d.ctrl, d.cvel, @@ -3658,15 +3428,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.nacon, d._impl.ncollision, d._impl.ne, - d._impl.ne_connect, - d._impl.ne_flex, - d._impl.ne_jnt, - d._impl.ne_ten, - d._impl.ne_weld, d._impl.nefc, d._impl.nf, d._impl.nl, - d._impl.nsolving, d._impl.qLD, d._impl.qLDiagInv, d._impl.qM, @@ -3690,7 +3454,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): d.site_xpos, d._impl.solver_niter, d._impl.subtree_angmom, - d._impl.subtree_bodyvel, d.subtree_com, d._impl.subtree_linvel, d._impl.ten_J, @@ -3725,31 +3488,13 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.contact__worldid, d._impl.efc__D, d._impl.efc__J, - d._impl.efc__Jaref, d._impl.efc__Ma, - d._impl.efc__Mgrad, - d._impl.efc__alpha, d._impl.efc__aref, - d._impl.efc__beta, - d._impl.efc__cost, - d._impl.efc__done, d._impl.efc__force, d._impl.efc__frictionloss, - d._impl.efc__gauss, - d._impl.efc__grad, - d._impl.efc__grad_dot, d._impl.efc__id, - d._impl.efc__jv, d._impl.efc__margin, - d._impl.efc__mv, d._impl.efc__pos, - d._impl.efc__prev_Mgrad, - d._impl.efc__prev_cost, - d._impl.efc__prev_grad, - d._impl.efc__quad, - d._impl.efc__quad_gauss, - d._impl.efc__search, - d._impl.efc__search_dot, d._impl.efc__state, d._impl.efc__type, d._impl.efc__vel, @@ -3769,116 +3514,88 @@ def _step_jax_impl(m: types.Model, d: types.Data): '_impl.cfrc_ext': out[11], '_impl.cfrc_int': out[12], '_impl.cinert': out[13], - '_impl.collision_pair': out[14], - '_impl.collision_pairid': out[15], - '_impl.collision_worldid': out[16], - '_impl.crb': out[17], - 'cvel': out[18], - '_impl.energy': out[19], - '_impl.flexedge_J': out[20], - '_impl.flexedge_length': out[21], - '_impl.flexedge_velocity': out[22], - '_impl.flexvert_xpos': out[23], - 'geom_xmat': out[24], - 'geom_xpos': out[25], - '_impl.light_xdir': out[26], - '_impl.light_xpos': out[27], - '_impl.nacon': out[28], - '_impl.ncollision': out[29], - '_impl.ne': out[30], - '_impl.ne_connect': out[31], - '_impl.ne_flex': out[32], - '_impl.ne_jnt': out[33], - '_impl.ne_ten': out[34], - '_impl.ne_weld': out[35], - '_impl.nefc': out[36], - '_impl.nf': out[37], - '_impl.nl': out[38], - '_impl.nsolving': out[39], - '_impl.qLD': out[40], - '_impl.qLDiagInv': out[41], - '_impl.qM': out[42], - 'qacc': out[43], - 'qacc_smooth': out[44], - 'qacc_warmstart': out[45], - 'qfrc_actuator': out[46], - 'qfrc_bias': out[47], - 'qfrc_constraint': out[48], - '_impl.qfrc_damper': out[49], - 'qfrc_fluid': out[50], - 'qfrc_gravcomp': out[51], - 'qfrc_passive': out[52], - 'qfrc_smooth': out[53], - '_impl.qfrc_spring': out[54], - 'qpos': out[55], - 'qvel': out[56], - 'sensordata': out[57], - 'site_xmat': out[58], - 'site_xpos': out[59], - '_impl.solver_niter': out[60], - '_impl.subtree_angmom': out[61], - '_impl.subtree_bodyvel': out[62], - 'subtree_com': out[63], - '_impl.subtree_linvel': out[64], - '_impl.ten_J': out[65], - 'ten_length': out[66], - '_impl.ten_velocity': out[67], - '_impl.ten_wrapadr': out[68], - '_impl.ten_wrapnum': out[69], - 'time': out[70], - '_impl.wrap_obj': out[71], - '_impl.wrap_xpos': out[72], - 'xanchor': out[73], - 'xaxis': out[74], - 'ximat': out[75], - 'xipos': out[76], - 'xmat': out[77], - 'xpos': out[78], - 'xquat': out[79], - '_impl.contact__dim': out[80], - '_impl.contact__dist': out[81], - '_impl.contact__efc_address': out[82], - '_impl.contact__frame': out[83], - '_impl.contact__friction': out[84], - '_impl.contact__geom': out[85], - '_impl.contact__geomcollisionid': out[86], - '_impl.contact__includemargin': out[87], - '_impl.contact__pos': out[88], - '_impl.contact__solimp': out[89], - '_impl.contact__solref': out[90], - '_impl.contact__solreffriction': out[91], - '_impl.contact__type': out[92], - '_impl.contact__worldid': out[93], - '_impl.efc__D': out[94], - '_impl.efc__J': out[95], - '_impl.efc__Jaref': out[96], - '_impl.efc__Ma': out[97], - '_impl.efc__Mgrad': out[98], - '_impl.efc__alpha': out[99], - '_impl.efc__aref': out[100], - '_impl.efc__beta': out[101], - '_impl.efc__cost': out[102], - '_impl.efc__done': out[103], - '_impl.efc__force': out[104], - '_impl.efc__frictionloss': out[105], - '_impl.efc__gauss': out[106], - '_impl.efc__grad': out[107], - '_impl.efc__grad_dot': out[108], - '_impl.efc__id': out[109], - '_impl.efc__jv': out[110], - '_impl.efc__margin': out[111], - '_impl.efc__mv': out[112], - '_impl.efc__pos': out[113], - '_impl.efc__prev_Mgrad': out[114], - '_impl.efc__prev_cost': out[115], - '_impl.efc__prev_grad': out[116], - '_impl.efc__quad': out[117], - '_impl.efc__quad_gauss': out[118], - '_impl.efc__search': out[119], - '_impl.efc__search_dot': out[120], - '_impl.efc__state': out[121], - '_impl.efc__type': out[122], - '_impl.efc__vel': out[123], + '_impl.crb': out[14], + 'cvel': out[15], + '_impl.energy': out[16], + '_impl.flexedge_J': out[17], + '_impl.flexedge_length': out[18], + '_impl.flexedge_velocity': out[19], + '_impl.flexvert_xpos': out[20], + 'geom_xmat': out[21], + 'geom_xpos': out[22], + '_impl.light_xdir': out[23], + '_impl.light_xpos': out[24], + '_impl.nacon': out[25], + '_impl.ncollision': out[26], + '_impl.ne': out[27], + '_impl.nefc': out[28], + '_impl.nf': out[29], + '_impl.nl': out[30], + '_impl.qLD': out[31], + '_impl.qLDiagInv': out[32], + '_impl.qM': out[33], + 'qacc': out[34], + 'qacc_smooth': out[35], + 'qacc_warmstart': out[36], + 'qfrc_actuator': out[37], + 'qfrc_bias': out[38], + 'qfrc_constraint': out[39], + '_impl.qfrc_damper': out[40], + 'qfrc_fluid': out[41], + 'qfrc_gravcomp': out[42], + 'qfrc_passive': out[43], + 'qfrc_smooth': out[44], + '_impl.qfrc_spring': out[45], + 'qpos': out[46], + 'qvel': out[47], + 'sensordata': out[48], + 'site_xmat': out[49], + 'site_xpos': out[50], + '_impl.solver_niter': out[51], + '_impl.subtree_angmom': out[52], + 'subtree_com': out[53], + '_impl.subtree_linvel': out[54], + '_impl.ten_J': out[55], + 'ten_length': out[56], + '_impl.ten_velocity': out[57], + '_impl.ten_wrapadr': out[58], + '_impl.ten_wrapnum': out[59], + 'time': out[60], + '_impl.wrap_obj': out[61], + '_impl.wrap_xpos': out[62], + 'xanchor': out[63], + 'xaxis': out[64], + 'ximat': out[65], + 'xipos': out[66], + 'xmat': out[67], + 'xpos': out[68], + 'xquat': out[69], + '_impl.contact__dim': out[70], + '_impl.contact__dist': out[71], + '_impl.contact__efc_address': out[72], + '_impl.contact__frame': out[73], + '_impl.contact__friction': out[74], + '_impl.contact__geom': out[75], + '_impl.contact__geomcollisionid': out[76], + '_impl.contact__includemargin': out[77], + '_impl.contact__pos': out[78], + '_impl.contact__solimp': out[79], + '_impl.contact__solref': out[80], + '_impl.contact__solreffriction': out[81], + '_impl.contact__type': out[82], + '_impl.contact__worldid': out[83], + '_impl.efc__D': out[84], + '_impl.efc__J': out[85], + '_impl.efc__Ma': out[86], + '_impl.efc__aref': out[87], + '_impl.efc__force': out[88], + '_impl.efc__frictionloss': out[89], + '_impl.efc__id': out[90], + '_impl.efc__margin': out[91], + '_impl.efc__pos': out[92], + '_impl.efc__state': out[93], + '_impl.efc__type': out[94], + '_impl.efc__vel': out[95], }) return d diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index ff12107c..ac5f3470 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.py @@ -42,10 +42,13 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _kinematics_shim( # Model nworld: int, + body_branch_start: wp.array(dtype=int), + body_branches: wp.array(dtype=int), body_ipos: wp.array2d(dtype=wp.vec3), body_iquat: wp.array2d(dtype=wp.quat), body_jntadr: wp.array(dtype=int), @@ -55,7 +58,6 @@ def _kinematics_shim( body_pos: wp.array2d(dtype=wp.vec3), body_quat: wp.array2d(dtype=wp.quat), body_rootid: wp.array(dtype=int), - body_tree: tuple[wp.array(dtype=int), ...], body_weldid: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), geom_pos: wp.array2d(dtype=wp.vec3), @@ -64,6 +66,8 @@ def _kinematics_shim( jnt_pos: wp.array2d(dtype=wp.vec3), jnt_qposadr: wp.array(dtype=int), jnt_type: wp.array(dtype=int), + nbody: int, + nbranch: int, ngeom: int, nsite: int, qpos0: wp.array2d(dtype=float), @@ -90,6 +94,8 @@ def _kinematics_shim( _m.opt = _o _d.efc = _e _d.contact = _c + _m.body_branch_start = body_branch_start + _m.body_branches = body_branches _m.body_ipos = body_ipos _m.body_iquat = body_iquat _m.body_jntadr = body_jntadr @@ -99,7 +105,6 @@ def _kinematics_shim( _m.body_pos = body_pos _m.body_quat = body_quat _m.body_rootid = body_rootid - _m.body_tree = body_tree _m.body_weldid = body_weldid _m.geom_bodyid = geom_bodyid _m.geom_pos = geom_pos @@ -108,6 +113,8 @@ def _kinematics_shim( _m.jnt_pos = jnt_pos _m.jnt_qposadr = jnt_qposadr _m.jnt_type = jnt_type + _m.nbody = nbody + _m.nbranch = nbranch _m.ngeom = ngeom _m.nsite = nsite _m.qpos0 = qpos0 @@ -208,6 +215,8 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): ) out = jf( d.qpos.shape[0], + m._impl.body_branch_start, + m._impl.body_branches, m.body_ipos, m.body_iquat, m.body_jntadr, @@ -217,7 +226,6 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): m.body_pos, m.body_quat, m.body_rootid, - m._impl.body_tree, m.body_weldid, m.geom_bodyid, m.geom_pos, @@ -226,6 +234,8 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): m.jnt_pos, m.jnt_qposadr, m.jnt_type, + m.nbody, + m._impl.nbranch, m.ngeom, m.nsite, m.qpos0, diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index 29dc02c3..dac6fa55 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -72,7 +72,7 @@ class BlockDim: energy_vel_kinetic: int euler_dense: int linesearch_iterative: int - mul_m_dense: int + qderiv_actuator_dense: int ray: int segmented_sort: int tendon_velocity: int @@ -93,7 +93,7 @@ class BlockDim: class StatisticWarp(PyTreeNode): """Derived fields from Statistic.""" - meaninertia: float + meaninertia: jax.Array class OptionWarp(PyTreeNode): """Derived fields from Option.""" @@ -104,9 +104,7 @@ class OptionWarp(PyTreeNode): contact_sensor_maxmatch: int graph_conditional: bool graph_mode: GraphMode - has_fluid: bool impratio_invsqrt: jax.Array - is_sparse: bool ls_parallel: bool ls_parallel_min_step: float run_collision_detection: bool @@ -120,8 +118,11 @@ class ModelWarp(PyTreeNode): M_rownnz: np.ndarray actuator_trntype_body_adr: np.ndarray block_dim: BlockDim + body_branch_start: np.ndarray + body_branches: np.ndarray body_fluid_ellipsoid: np.ndarray body_tree: Tuple[np.ndarray, ...] + cam_projection: np.ndarray collision_sensor_adr: np.ndarray dof_tri_col: np.ndarray dof_tri_row: np.ndarray @@ -146,11 +147,16 @@ class ModelWarp(PyTreeNode): flex_vertadr: np.ndarray flex_vertbodyid: np.ndarray flex_vertnum: np.ndarray + flexedge_J_colind: np.ndarray + flexedge_J_rowadr: np.ndarray + flexedge_J_rownnz: np.ndarray flexedge_invweight0: np.ndarray flexedge_length0: np.ndarray geom_pair_type_count: Tuple[int, ...] geom_plugin_index: np.ndarray + has_fluid: bool has_sdf_geom: bool + is_sparse: bool jnt_limited_ball_adr: np.ndarray jnt_limited_slide_hinge_adr: np.ndarray light_active: jax.Array @@ -169,6 +175,7 @@ class ModelWarp(PyTreeNode): mesh_polyvertnum: np.ndarray mocap_bodyid: np.ndarray nacttrnbody: int + nbranch: int nflex: int nflexedge: int nflexelem: int @@ -185,6 +192,7 @@ class ModelWarp(PyTreeNode): nsensorcollision: int nsensorcontact: int nsensortaxel: int + ntree: int nv_pad: int nxn_geom_pair: np.ndarray nxn_geom_pair_filtered: np.ndarray @@ -198,9 +206,9 @@ class ModelWarp(PyTreeNode): qLD_updates: Tuple[np.ndarray, ...] qM_fullm_i: np.ndarray qM_fullm_j: np.ndarray - qM_madr_ij: np.ndarray - qM_mulm_i: np.ndarray - qM_mulm_j: np.ndarray + qM_mulm_col: np.ndarray + qM_mulm_madr: np.ndarray + qM_mulm_rowadr: np.ndarray qM_tiles: Tuple[TileSet, ...] rangefinder_sensor_adr: np.ndarray sensor_acc_adr: np.ndarray @@ -228,6 +236,9 @@ class ModelWarp(PyTreeNode): tendon_jnt_adr: np.ndarray tendon_limited_adr: np.ndarray tendon_site_pair_adr: np.ndarray + tree_bodynum: np.ndarray + tree_dofadr: np.ndarray + tree_dofnum: np.ndarray wrap_geom_adr: np.ndarray wrap_jnt_adr: np.ndarray wrap_pulley_scale: np.ndarray @@ -242,9 +253,6 @@ class DataWarp(PyTreeNode): cfrc_ext: jax.Array cfrc_int: jax.Array cinert: jax.Array - collision_pair: jax.Array - collision_pairid: jax.Array - collision_worldid: jax.Array contact__dim: jax.Array contact__dist: jax.Array contact__efc_address: jax.Array @@ -262,31 +270,13 @@ class DataWarp(PyTreeNode): crb: jax.Array efc__D: jax.Array efc__J: jax.Array - efc__Jaref: jax.Array efc__Ma: jax.Array - efc__Mgrad: jax.Array - efc__alpha: jax.Array efc__aref: jax.Array - efc__beta: jax.Array - efc__cost: jax.Array - efc__done: jax.Array efc__force: jax.Array efc__frictionloss: jax.Array - efc__gauss: jax.Array - efc__grad: jax.Array - efc__grad_dot: jax.Array efc__id: jax.Array - efc__jv: jax.Array efc__margin: jax.Array - efc__mv: jax.Array efc__pos: jax.Array - efc__prev_Mgrad: jax.Array - efc__prev_cost: jax.Array - efc__prev_grad: jax.Array - efc__quad: jax.Array - efc__quad_gauss: jax.Array - efc__search: jax.Array - efc__search_dot: jax.Array efc__state: jax.Array efc__type: jax.Array efc__vel: jax.Array @@ -297,20 +287,15 @@ class DataWarp(PyTreeNode): flexvert_xpos: jax.Array light_xdir: jax.Array light_xpos: jax.Array + naccdmax: int nacon: jax.Array naconmax: int ncollision: jax.Array ne: jax.Array - ne_connect: jax.Array - ne_flex: jax.Array - ne_jnt: jax.Array - ne_ten: jax.Array - ne_weld: jax.Array nefc: jax.Array nf: jax.Array njmax: int nl: jax.Array - nsolving: jax.Array nworld: int qLD: jax.Array qLDiagInv: jax.Array @@ -319,7 +304,6 @@ class DataWarp(PyTreeNode): qfrc_spring: jax.Array solver_niter: jax.Array subtree_angmom: jax.Array - subtree_bodyvel: jax.Array subtree_linvel: jax.Array ten_J: jax.Array ten_velocity: jax.Array @@ -329,9 +313,6 @@ class DataWarp(PyTreeNode): wrap_xpos: jax.Array shape = property(lambda self: self.cacc.shape) DATA_NON_VMAP = { - 'collision_pair', - 'collision_pairid', - 'collision_worldid', 'contact__dim', 'contact__dist', 'contact__efc_address', @@ -346,11 +327,11 @@ DATA_NON_VMAP = { 'contact__solreffriction', 'contact__type', 'contact__worldid', + 'naccdmax', 'nacon', 'naconmax', 'ncollision', 'njmax', - 'nsolving', 'nworld', } @@ -394,9 +375,6 @@ _NDIM = { 'cfrc_ext': 3, 'cfrc_int': 3, 'cinert': 3, - 'collision_pair': 2, - 'collision_pairid': 2, - 'collision_worldid': 1, 'contact__dim': 1, 'contact__dist': 1, 'contact__efc_address': 2, @@ -416,31 +394,13 @@ _NDIM = { 'cvel': 3, 'efc__D': 2, 'efc__J': 3, - 'efc__Jaref': 2, 'efc__Ma': 2, - 'efc__Mgrad': 2, - 'efc__alpha': 1, 'efc__aref': 2, - 'efc__beta': 1, - 'efc__cost': 1, - 'efc__done': 1, 'efc__force': 2, 'efc__frictionloss': 2, - 'efc__gauss': 1, - 'efc__grad': 2, - 'efc__grad_dot': 1, 'efc__id': 2, - 'efc__jv': 2, 'efc__margin': 2, - 'efc__mv': 2, 'efc__pos': 2, - 'efc__prev_Mgrad': 2, - 'efc__prev_cost': 1, - 'efc__prev_grad': 2, - 'efc__quad': 3, - 'efc__quad_gauss': 2, - 'efc__search': 2, - 'efc__search_dot': 1, 'efc__state': 2, 'efc__type': 2, 'efc__vel': 2, @@ -456,20 +416,15 @@ _NDIM = { 'light_xpos': 3, 'mocap_pos': 3, 'mocap_quat': 3, + 'naccdmax': 0, 'nacon': 1, 'naconmax': 0, 'ncollision': 1, 'ne': 1, - 'ne_connect': 1, - 'ne_flex': 1, - 'ne_jnt': 1, - 'ne_ten': 1, - 'ne_weld': 1, 'nefc': 1, 'nf': 1, 'njmax': 0, 'nl': 1, - 'nsolving': 1, 'nworld': 0, 'qLD': 3, 'qLDiagInv': 2, @@ -495,7 +450,6 @@ _NDIM = { 'site_xpos': 3, 'solver_niter': 1, 'subtree_angmom': 3, - 'subtree_bodyvel': 3, 'subtree_com': 3, 'subtree_linvel': 3, 'ten_J': 3, @@ -549,7 +503,7 @@ _NDIM = { 'block_dim__energy_vel_kinetic': 0, 'block_dim__euler_dense': 0, 'block_dim__linesearch_iterative': 0, - 'block_dim__mul_m_dense': 0, + 'block_dim__qderiv_actuator_dense': 0, 'block_dim__ray': 0, 'block_dim__segmented_sort': 0, 'block_dim__tendon_velocity': 0, @@ -557,6 +511,8 @@ _NDIM = { 'block_dim__update_gradient_JTDAJ_sparse': 0, 'block_dim__update_gradient_cholesky': 0, 'block_dim__update_gradient_cholesky_blocked': 0, + 'body_branch_start': 1, + 'body_branches': 1, 'body_conaffinity': 1, 'body_contype': 1, 'body_dofadr': 1, @@ -579,15 +535,17 @@ _NDIM = { 'body_rootid': 1, 'body_subtreemass': 2, 'body_tree': -1, + 'body_treeid': 1, 'body_weldid': 1, 'cam_bodyid': 1, - 'cam_fovy': 1, - 'cam_intrinsic': 2, + 'cam_fovy': 2, + 'cam_intrinsic': 3, 'cam_mat0': 4, 'cam_mode': 1, 'cam_pos': 3, 'cam_pos0': 3, 'cam_poscom0': 3, + 'cam_projection': 1, 'cam_quat': 3, 'cam_resolution': 2, 'cam_sensorsize': 2, @@ -603,6 +561,7 @@ _NDIM = { 'dof_parentid': 1, 'dof_solimp': 3, 'dof_solref': 3, + 'dof_treeid': 1, 'dof_tri_col': 1, 'dof_tri_row': 1, 'eq_active0': 1, @@ -635,6 +594,9 @@ _NDIM = { 'flex_vertadr': 1, 'flex_vertbodyid': 1, 'flex_vertnum': 1, + 'flexedge_J_colind': 1, + 'flexedge_J_rowadr': 1, + 'flexedge_J_rownnz': 1, 'flexedge_invweight0': 1, 'flexedge_length0': 1, 'geom_aabb': 4, @@ -661,12 +623,14 @@ _NDIM = { 'geom_solmix': 2, 'geom_solref': 3, 'geom_type': 1, + 'has_fluid': 0, 'has_sdf_geom': 0, 'hfield_adr': 1, 'hfield_data': 1, 'hfield_ncol': 1, 'hfield_nrow': 1, 'hfield_size': 2, + 'is_sparse': 0, 'jnt_actfrclimited': 1, 'jnt_actfrcrange': 3, 'jnt_actgravcomp': 1, @@ -697,7 +661,7 @@ _NDIM = { 'light_type': 2, 'mapM2M': 1, 'mat_rgba': 3, - 'mat_texid': 2, + 'mat_texid': 3, 'mat_texrepeat': 3, 'mesh_face': 2, 'mesh_faceadr': 1, @@ -724,6 +688,7 @@ _NDIM = { 'na': 0, 'nacttrnbody': 0, 'nbody': 0, + 'nbranch': 0, 'ncam': 0, 'neq': 0, 'nexclude': 0, @@ -765,6 +730,7 @@ _NDIM = { 'nsensortaxel': 0, 'nsite': 0, 'ntendon': 0, + 'ntree': 0, 'nu': 0, 'nv': 0, 'nv_pad': 0, @@ -785,10 +751,8 @@ _NDIM = { 'opt__enableflags': 0, 'opt__graph_conditional': 0, 'opt__gravity': 2, - 'opt__has_fluid': 0, 'opt__impratio_invsqrt': 1, 'opt__integrator': 0, - 'opt__is_sparse': 0, 'opt__iterations': 0, 'opt__ls_iterations': 0, 'opt__ls_parallel': 0, @@ -817,9 +781,9 @@ _NDIM = { 'qLD_updates': -1, 'qM_fullm_i': 1, 'qM_fullm_j': 1, - 'qM_madr_ij': 1, - 'qM_mulm_i': 1, - 'qM_mulm_j': 1, + 'qM_mulm_col': 1, + 'qM_mulm_madr': 1, + 'qM_mulm_rowadr': 1, 'qM_tiles': -1, 'qpos0': 2, 'qpos_spring': 2, @@ -856,7 +820,7 @@ _NDIM = { 'site_quat': 3, 'site_size': 2, 'site_type': 1, - 'stat__meaninertia': 0, + 'stat__meaninertia': 1, 'taxel_sensorid': 1, 'taxel_vertadr': 1, 'ten_wrapadr_site': 1, @@ -883,6 +847,9 @@ _NDIM = { 'tendon_solref_fri': 3, 'tendon_solref_lim': 3, 'tendon_stiffness': 2, + 'tree_bodynum': 1, + 'tree_dofadr': 1, + 'tree_dofnum': 1, 'wrap_geom_adr': 1, 'wrap_jnt_adr': 1, 'wrap_objid': 1, @@ -902,10 +869,8 @@ _NDIM = { 'enableflags': 0, 'graph_conditional': 0, 'gravity': 2, - 'has_fluid': 0, 'impratio_invsqrt': 1, 'integrator': 0, - 'is_sparse': 0, 'iterations': 0, 'ls_iterations': 0, 'ls_parallel': 0, @@ -921,7 +886,7 @@ _NDIM = { 'viscosity': 1, 'wind': 2, }, - 'Statistic': {'meaninertia': 0}, + 'Statistic': {'meaninertia': 1}, } _BATCH_DIM = { 'Data': { @@ -939,9 +904,6 @@ _BATCH_DIM = { 'cfrc_ext': True, 'cfrc_int': True, 'cinert': True, - 'collision_pair': False, - 'collision_pairid': False, - 'collision_worldid': False, 'contact__dim': False, 'contact__dist': False, 'contact__efc_address': False, @@ -961,31 +923,13 @@ _BATCH_DIM = { 'cvel': True, 'efc__D': True, 'efc__J': True, - 'efc__Jaref': True, 'efc__Ma': True, - 'efc__Mgrad': True, - 'efc__alpha': True, 'efc__aref': True, - 'efc__beta': True, - 'efc__cost': True, - 'efc__done': True, 'efc__force': True, 'efc__frictionloss': True, - 'efc__gauss': True, - 'efc__grad': True, - 'efc__grad_dot': True, 'efc__id': True, - 'efc__jv': True, 'efc__margin': True, - 'efc__mv': True, 'efc__pos': True, - 'efc__prev_Mgrad': True, - 'efc__prev_cost': True, - 'efc__prev_grad': True, - 'efc__quad': True, - 'efc__quad_gauss': True, - 'efc__search': True, - 'efc__search_dot': True, 'efc__state': True, 'efc__type': True, 'efc__vel': True, @@ -1001,20 +945,15 @@ _BATCH_DIM = { 'light_xpos': True, 'mocap_pos': True, 'mocap_quat': True, + 'naccdmax': False, 'nacon': False, 'naconmax': False, 'ncollision': False, 'ne': True, - 'ne_connect': True, - 'ne_flex': True, - 'ne_jnt': True, - 'ne_ten': True, - 'ne_weld': True, 'nefc': True, 'nf': True, 'njmax': False, 'nl': True, - 'nsolving': False, 'nworld': False, 'qLD': True, 'qLDiagInv': True, @@ -1040,7 +979,6 @@ _BATCH_DIM = { 'site_xpos': True, 'solver_niter': True, 'subtree_angmom': True, - 'subtree_bodyvel': True, 'subtree_com': True, 'subtree_linvel': True, 'ten_J': True, @@ -1094,7 +1032,7 @@ _BATCH_DIM = { 'block_dim__energy_vel_kinetic': False, 'block_dim__euler_dense': False, 'block_dim__linesearch_iterative': False, - 'block_dim__mul_m_dense': False, + 'block_dim__qderiv_actuator_dense': False, 'block_dim__ray': False, 'block_dim__segmented_sort': False, 'block_dim__tendon_velocity': False, @@ -1102,6 +1040,8 @@ _BATCH_DIM = { 'block_dim__update_gradient_JTDAJ_sparse': False, 'block_dim__update_gradient_cholesky': False, 'block_dim__update_gradient_cholesky_blocked': False, + 'body_branch_start': False, + 'body_branches': False, 'body_conaffinity': False, 'body_contype': False, 'body_dofadr': False, @@ -1124,15 +1064,17 @@ _BATCH_DIM = { 'body_rootid': False, 'body_subtreemass': True, 'body_tree': False, + 'body_treeid': False, 'body_weldid': False, 'cam_bodyid': False, - 'cam_fovy': False, - 'cam_intrinsic': False, + 'cam_fovy': True, + 'cam_intrinsic': True, 'cam_mat0': True, 'cam_mode': False, 'cam_pos': True, 'cam_pos0': True, 'cam_poscom0': True, + 'cam_projection': False, 'cam_quat': True, 'cam_resolution': False, 'cam_sensorsize': False, @@ -1148,6 +1090,7 @@ _BATCH_DIM = { 'dof_parentid': False, 'dof_solimp': True, 'dof_solref': True, + 'dof_treeid': False, 'dof_tri_col': False, 'dof_tri_row': False, 'eq_active0': False, @@ -1180,6 +1123,9 @@ _BATCH_DIM = { 'flex_vertadr': False, 'flex_vertbodyid': False, 'flex_vertnum': False, + 'flexedge_J_colind': False, + 'flexedge_J_rowadr': False, + 'flexedge_J_rownnz': False, 'flexedge_invweight0': False, 'flexedge_length0': False, 'geom_aabb': True, @@ -1206,12 +1152,14 @@ _BATCH_DIM = { 'geom_solmix': True, 'geom_solref': True, 'geom_type': False, + 'has_fluid': False, 'has_sdf_geom': False, 'hfield_adr': False, 'hfield_data': False, 'hfield_ncol': False, 'hfield_nrow': False, 'hfield_size': False, + 'is_sparse': False, 'jnt_actfrclimited': False, 'jnt_actfrcrange': True, 'jnt_actgravcomp': False, @@ -1242,7 +1190,7 @@ _BATCH_DIM = { 'light_type': True, 'mapM2M': False, 'mat_rgba': True, - 'mat_texid': False, + 'mat_texid': True, 'mat_texrepeat': True, 'mesh_face': False, 'mesh_faceadr': False, @@ -1269,6 +1217,7 @@ _BATCH_DIM = { 'na': False, 'nacttrnbody': False, 'nbody': False, + 'nbranch': False, 'ncam': False, 'neq': False, 'nexclude': False, @@ -1310,6 +1259,7 @@ _BATCH_DIM = { 'nsensortaxel': False, 'nsite': False, 'ntendon': False, + 'ntree': False, 'nu': False, 'nv': False, 'nv_pad': False, @@ -1330,10 +1280,8 @@ _BATCH_DIM = { 'opt__enableflags': False, 'opt__graph_conditional': False, 'opt__gravity': True, - 'opt__has_fluid': False, 'opt__impratio_invsqrt': True, 'opt__integrator': False, - 'opt__is_sparse': False, 'opt__iterations': False, 'opt__ls_iterations': False, 'opt__ls_parallel': False, @@ -1362,9 +1310,9 @@ _BATCH_DIM = { 'qLD_updates': False, 'qM_fullm_i': False, 'qM_fullm_j': False, - 'qM_madr_ij': False, - 'qM_mulm_i': False, - 'qM_mulm_j': False, + 'qM_mulm_col': False, + 'qM_mulm_madr': False, + 'qM_mulm_rowadr': False, 'qM_tiles': False, 'qpos0': True, 'qpos_spring': True, @@ -1401,7 +1349,7 @@ _BATCH_DIM = { 'site_quat': True, 'site_size': False, 'site_type': False, - 'stat__meaninertia': False, + 'stat__meaninertia': True, 'taxel_sensorid': False, 'taxel_vertadr': False, 'ten_wrapadr_site': False, @@ -1428,6 +1376,9 @@ _BATCH_DIM = { 'tendon_solref_fri': True, 'tendon_solref_lim': True, 'tendon_stiffness': True, + 'tree_bodynum': False, + 'tree_dofadr': False, + 'tree_dofnum': False, 'wrap_geom_adr': False, 'wrap_jnt_adr': False, 'wrap_objid': False, @@ -1447,10 +1398,8 @@ _BATCH_DIM = { 'enableflags': False, 'graph_conditional': False, 'gravity': True, - 'has_fluid': False, 'impratio_invsqrt': True, 'integrator': False, - 'is_sparse': False, 'iterations': False, 'ls_iterations': False, 'ls_parallel': False, @@ -1466,5 +1415,5 @@ _BATCH_DIM = { 'viscosity': True, 'wind': True, }, - 'Statistic': {'meaninertia': False}, + 'Statistic': {'meaninertia': True}, }