diff --git a/mjx/mujoco/mjx/_src/io_test.py b/mjx/mujoco/mjx/_src/io_test.py index 4c52cdcd..a52867ba 100644 --- a/mjx/mujoco/mjx/_src/io_test.py +++ b/mjx/mujoco/mjx/_src/io_test.py @@ -434,9 +434,9 @@ class DataIOTest(parameterized.TestCase): if not mjx_io.has_cuda_gpu_device(): self.skipTest('No CUDA GPU device.') m = mujoco.MjModel.from_xml_string(_MULTIPLE_CONVEX_OBJECTS) - d = mjx.make_data(m, impl='warp', nconmax=9, njmax=11) + d = mjx.make_data(m, impl='warp', nconmax=9, njmax=23) self.assertEqual(d._impl.contact__dist.shape[0], 9) - self.assertEqual(d._impl.efc__J.shape[0], 11) + self.assertEqual(d._impl.efc__pos.shape[0], 23) @parameterized.parameters('jax', 'c') def test_put_data(self, impl: str): diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 14004150..70152333 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -754,7 +754,7 @@ class Model(PyTreeNode): geom_solref: jax.Array geom_solimp: jax.Array geom_size: jax.Array - geom_aabb: np.ndarray + geom_aabb: jax.Array geom_rbound: jax.Array geom_pos: jax.Array geom_quat: 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 82e53459..4b9c45da 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/__init__.py @@ -65,7 +65,9 @@ from mujoco.mjx.third_party.mujoco_warp._src.smooth import tendon as tendon from mujoco.mjx.third_party.mujoco_warp._src.smooth import transmission as transmission 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 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 from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType as BiasType from mujoco.mjx.third_party.mujoco_warp._src.types import BroadphaseFilter as BroadphaseFilter @@ -82,5 +84,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import IntegratorType as Inte 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 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 from mujoco.mjx.third_party.mujoco_warp._src.types import TrnType as TrnType 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 b4b629c6..32ea187c 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/benchmark.py @@ -45,28 +45,46 @@ def _sum(stack1, stack2): @wp.kernel def ctrl_noise( # Model: + opt_timestep: wp.array(dtype=float), actuator_ctrllimited: wp.array(dtype=bool), actuator_ctrlrange: wp.array2d(dtype=wp.vec2), + # Data in: + ctrl_in: wp.array2d(dtype=float), # In: ctrl_center: wp.array1d(dtype=float), step: int, - ctrlnoise: float, + ctrlnoisestd: float, + ctrlnoiserate: float, # Data out: ctrl_out: wp.array2d(dtype=float), ): worldid, actid = wp.tid() - center = 0.0 - radius = 1.0 + # convert rate and scale to discrete time (Ornstein-Uhlenbeck) + rate = wp.exp(-opt_timestep[0] / ctrlnoiserate) + scale = ctrlnoisestd * wp.sqrt(1.0 - rate * rate) + + midpoint = 0.0 + halfrange = 1.0 ctrlrange = actuator_ctrlrange[0, actid] + is_limited = actuator_ctrllimited[actid] + if is_limited: + midpoint = 0.5 * (ctrlrange[1] + ctrlrange[0]) + halfrange = 0.5 * (ctrlrange[1] - ctrlrange[0]) if ctrl_center.shape[0] > 0: - center = ctrl_center[actid] - elif actuator_ctrllimited[actid]: - center = (ctrlrange[1] + ctrlrange[0]) / 2.0 - radius = (ctrlrange[1] - ctrlrange[0]) / 2.0 - radius *= ctrlnoise - noise = 2.0 * halton((step + 1) * (worldid + 1), actid + 2) - 1.0 - ctrl_out[worldid, actid] = center + radius * noise + midpoint = ctrl_center[actid] + + # exponential convergence to midpoint at ctrlnoiserate + ctrl = rate * ctrl_in[worldid, actid] + (1.0 - rate) * midpoint + + # add noise + ctrl += scale * halfrange * (2.0 * halton((step + 1) * (worldid + 1), actid + 2) - 1.0) + + # clip to range if limited + if is_limited: + ctrl = wp.clamp(ctrl, ctrlrange[0], ctrlrange[1]) + + ctrl_out[worldid, actid] = ctrl def benchmark( @@ -82,28 +100,24 @@ def benchmark( """Benchmark a function of Model and Data. Args: - fn (Callable[[Model, Data], None]): Function to benchmark. - m (Model): The model containing kinematic and dynamic information (device). - d (Data): The data object containing the current state and output information (device). - nstep (int): Number of timesteps. - ctrls (list, optional): control sequence to apply during benchmarking. - Default is None. - event_trace (bool, optional): If True, time routines decorated with @event_scope. - Default is False. - measure_alloc (bool, optional): If True, record number of contacts and constraints. - Default is False. - measure_solver_niter (bool, False): If True, record the number of solver iterations. - Default is False. - Returns: - float: Time to JIT fn. - float: Total time to run the benchmark. - dict: Trace. - list: Number of contacts. - list: Number of constraints. - list: Number of solver iterations. - int: Number of converged worlds. - """ + fn: Function to benchmark. + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output information (device). + nstep: Number of timesteps. + ctrls: Control sequence to apply during benchmarking. + 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. + Returns: + - Time to JIT fn. + - Total time to run the benchmark. + - Trace. + - Number of contacts. + - Number of constraints. + - Number of solver iterations. + - Number of converged worlds. + """ trace = {} nacon, nefc, solver_niter = [], [], [] center = wp.array([], dtype=wp.float32) @@ -126,7 +140,7 @@ def benchmark( wp.launch( ctrl_noise, dim=(d.nworld, m.nu), - inputs=[m.actuator_ctrllimited, m.actuator_ctrlrange, center, i, 0.01], + inputs=[m.opt.timestep, m.actuator_ctrllimited, m.actuator_ctrlrange, d.ctrl, center, i, 0.01, 0.1], outputs=[d.ctrl], ) wp.synchronize() diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py index a080679b..4a25dd21 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/block_cholesky.py @@ -29,11 +29,10 @@ def create_blocked_cholesky_func(block_size: int): # Out: L: wp.array(dtype=float, ndim=2), ): - """ - Computes the Cholesky factorization of a symmetric positive definite matrix A in blocks. + """Computes the Cholesky factorization of a symmetric positive definite matrix A in blocks. + It returns a lower-triangular matrix L such that A = L L^T. """ - num_threads_per_block = wp.block_dim() # Round up active_matrix_size to next multiple of block_size @@ -120,11 +119,11 @@ def create_blocked_cholesky_solve_func(block_size: int): # Out: x: wp.array(dtype=float, ndim=2), ): - """ - Solves A x = b given the Cholesky factor L (A = L L^T) using - blocked forward and backward substitution. - """ + """Block Cholesky factorization and solve. + Solves A x = b given the Cholesky factor L (A = L L^T) using blocked forward and backward + substitution. + """ num_threads_per_block = wp.block_dim() # Round up active_matrix_size to next multiple of block_size 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 22a06c84..34bbb960 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 @@ -16,6 +16,7 @@ import warp as wp from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd +from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import multicontact from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import epa_legacy from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import gjk_legacy from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk_legacy import multicontact_legacy @@ -26,6 +27,8 @@ from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame from mujoco.mjx.third_party.mujoco_warp._src.math import upper_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 Data from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType @@ -88,6 +91,7 @@ def ccd_kernel_builder( epa_exact_neg_distance: bool, depth_extension: float, is_hfield: bool, + use_multiccd: bool, ): @wp.func def eval_ccd_write_contact( @@ -96,6 +100,7 @@ def ccd_kernel_builder( geom_type: wp.array(dtype=int), # Data in: naconmax_in: int, + # In: epa_vert_in: wp.array2d(dtype=wp.vec3), epa_vert1_in: wp.array2d(dtype=wp.vec3), epa_vert2_in: wp.array2d(dtype=wp.vec3), @@ -118,7 +123,6 @@ 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), - # In: geom1: Geom, geom2: Geom, geoms: wp.vec2i, @@ -134,8 +138,8 @@ def ccd_kernel_builder( x1: wp.vec3, x2: wp.vec3, count: int, + pairid: wp.vec2i, # 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), @@ -147,6 +151,9 @@ def ccd_kernel_builder( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ) -> int: # TODO(kbayes): remove legacy GJK once multicontact can be enabled if wp.static(legacy_gjk): @@ -176,12 +183,18 @@ def ccd_kernel_builder( frame = make_frame(normal) else: points = mat3c() + witness1 = mat3c() + witness2 = mat3c() geom1.margin = margin geom2.margin = margin - dist, ncontact, witness1, witness2 = ccd( - False, # ignored for box-box, multiccd always on - opt_ccd_tolerance[worldid], - 0.0, + if pairid[1] >= 0: + # if collision sensor, set large cutoff to work with various sensor cutoff values + cutoff = 1.0e32 + else: + cutoff = 0.0 + dist, ncontact, w1, w2, idx = ccd( + opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]], + cutoff, ccd_iterations, geom1, geom2, @@ -200,29 +213,60 @@ def ccd_kernel_builder( epa_index_in[tid], epa_map_in[tid], epa_horizon_in[tid], - multiccd_polygon_in[tid], - multiccd_clipped_in[tid], - multiccd_pnormal_in[tid], - multiccd_pdist_in[tid], - multiccd_idx1_in[tid], - multiccd_idx2_in[tid], - multiccd_n1_in[tid], - multiccd_n2_in[tid], - multiccd_endvert_in[tid], - multiccd_face1_in[tid], - multiccd_face2_in[tid], ) - if dist >= 0.0: + + if dist >= 0.0 and pairid[1] == -1: return 0 + 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], + w1, + w2, + geom1, + geom2, + geomtype1, + geomtype2, + ) + for i in range(ncontact): points[i] = 0.5 * (witness1[i] + witness2[i]) normal = witness1[0] - witness2[0] frame = make_frame(normal) + # flip if collision sensor + if pairid[1] >= 0: + frame *= -1.0 + geoms = wp.vec2i(geoms[1], geoms[0]) + for i in range(ncontact): write_contact( naconmax_in, + i, dist, points[i], frame, @@ -234,8 +278,8 @@ def ccd_kernel_builder( solreffriction, solimp, geoms, + pairid, worldid, - nacon_out, contact_dist_out, contact_pos_out, contact_frame_out, @@ -247,6 +291,9 @@ def ccd_kernel_builder( contact_dim_out, contact_geom_out, contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, ) if count + (i + 1) >= MJ_MAXCONPAIR: return i + 1 @@ -266,7 +313,7 @@ def ccd_kernel_builder( geom_solref: wp.array2d(dtype=wp.vec2), geom_solimp: wp.array2d(dtype=vec5), geom_size: wp.array2d(dtype=wp.vec3), - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_rbound: wp.array2d(dtype=float), geom_friction: wp.array2d(dtype=wp.vec3), geom_margin: wp.array2d(dtype=float), @@ -302,9 +349,10 @@ def ccd_kernel_builder( geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), collision_pair_in: wp.array(dtype=wp.vec2i), - collision_pairid_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=wp.vec2i), collision_worldid_in: wp.array(dtype=int), ncollision_in: wp.array(dtype=int), + # In: epa_vert_in: wp.array2d(dtype=wp.vec3), epa_vert1_in: wp.array2d(dtype=wp.vec3), epa_vert2_in: wp.array2d(dtype=wp.vec3), @@ -340,6 +388,8 @@ def ccd_kernel_builder( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), ): tid = wp.tid() if tid >= ncollision_in[0]: @@ -384,15 +434,17 @@ def ccd_kernel_builder( worldid, ) + geom_size_id = worldid % geom_size.shape[0] + geom1_dataid = geom_dataid[g1] geom1 = geom( geomtype1, geom1_dataid, - geom_size[worldid, g1], + geom_size[geom_size_id, g1], mesh_vertadr, mesh_vertnum, - mesh_vert, mesh_graphadr, + mesh_vert, mesh_graph, mesh_polynum, mesh_polyadr, @@ -411,11 +463,11 @@ def ccd_kernel_builder( geom2 = geom( geomtype2, geom2_dataid, - geom_size[worldid, g2], + geom_size[geom_size_id, g2], mesh_vertadr, mesh_vertnum, - mesh_vert, mesh_graphadr, + mesh_vert, mesh_graph, mesh_polynum, mesh_polyadr, @@ -433,9 +485,9 @@ def ccd_kernel_builder( # see MuJoCo mjc_ConvexHField if wp.static(is_hfield): # height field subgrid - nrow = hfield_nrow[g1] - ncol = hfield_ncol[g1] - size = hfield_size[g1] + nrow = hfield_nrow[geom1_dataid] + ncol = hfield_ncol[geom1_dataid] + size = hfield_size[geom1_dataid] # subgrid x_scale = 0.5 * float(ncol - 1) / size[0] @@ -541,7 +593,7 @@ def ccd_kernel_builder( x1, geom2.pos, count, - nacon_out, + collision_pairid_in[tid], contact_dist_out, contact_pos_out, contact_frame_out, @@ -553,6 +605,9 @@ def ccd_kernel_builder( contact_dim_out, contact_geom_out, contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, ) count += ncontact if count >= MJ_MAXCONPAIR: @@ -599,7 +654,7 @@ def ccd_kernel_builder( geom1.pos, geom2.pos, 0, - nacon_out, + collision_pairid_in[tid], contact_dist_out, contact_pos_out, contact_frame_out, @@ -611,6 +666,9 @@ def ccd_kernel_builder( contact_dim_out, contact_geom_out, contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, ) return ccd_kernel @@ -625,19 +683,73 @@ def convex_narrowphase(m: Model, d: Data): determine the distance between shapes and the Expanding Polytope Algorithm (EPA) to find the penetration depth and contact normal for colliding pairs. - The convex geom types handled by this function are SPHERE, CAPSULE, ELLIPSOID, CYLINDER, - BOX, MESH, HFIELD. + The convex geom types handled by this function are `SPHERE`, `CAPSULE`, `ELLIPSOID`, `CYLINDER`, + `BOX`, `MESH`, `HFIELD`. To optimize performance, this function dynamically builds and launches a specialized kernel for each type of convex collision pair present in the model, avoiding unnecessary computations for non-existent pair types. """ + # TODO(team): fix early return? + if not any(m.geom_pair_type_count[upper_trid_index(len(GeomType), g[0].value, g[1].value)] for g in _CONVEX_COLLISION_PAIRS): + return + + # set to true to enable multiccd + use_multiccd = False + nmaxpolygon = m.nmaxpolygon if use_multiccd else 0 + nmaxmeshdeg = m.nmaxmeshdeg if use_multiccd else 0 + + # epa_vert: vertices in EPA polytope in Minkowski space + epa_vert = wp.empty(shape=(d.naconmax, 5 + m.opt.ccd_iterations), dtype=wp.vec3) + # epa_vert1: vertices in EPA polytope in geom 1 space + epa_vert1 = wp.empty(shape=(d.naconmax, 5 + m.opt.ccd_iterations), dtype=wp.vec3) + # epa_vert2: vertices in EPA polytope in geom 2 space + epa_vert2 = wp.empty(shape=(d.naconmax, 5 + m.opt.ccd_iterations), dtype=wp.vec3) + # epa_vert_index1: vertex indices in EPA polytope for geom 1 + epa_vert_index1 = wp.empty(shape=(d.naconmax, 5 + m.opt.ccd_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 + m.opt.ccd_iterations), dtype=int) + # epa_face: faces of polytope represented by three indices + epa_face = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * m.opt.ccd_iterations), dtype=wp.vec3i) + # epa_pr: projection of origin on polytope faces + epa_pr = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * m.opt.ccd_iterations), dtype=wp.vec3) + # epa_norm2: epa_pr * epa_pr + epa_norm2 = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * m.opt.ccd_iterations), dtype=float) + # epa_index: index of face in polytope map + epa_index = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * m.opt.ccd_iterations), dtype=int) + # epa_map: status of faces in polytope + epa_map = wp.empty(shape=(d.naconmax, 6 + MJ_MAX_EPAFACES * m.opt.ccd_iterations), dtype=int) + # epa_horizon: index pair (i j) of edges on horizon + epa_horizon = wp.empty(shape=(d.naconmax, 2 * MJ_MAX_EPAHORIZON), dtype=int) + # multiccd_polygon: clipped contact surface + multiccd_polygon = wp.empty(shape=(d.naconmax, 2 * nmaxpolygon), dtype=wp.vec3) + # multiccd_clipped: clipped contact surface (intermediate) + multiccd_clipped = wp.empty(shape=(d.naconmax, 2 * nmaxpolygon), dtype=wp.vec3) + # multiccd_pnormal: plane normal of clipping polygon + multiccd_pnormal = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + # multiccd_pdist: plane distance of clipping polygon + multiccd_pdist = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=float) + # multiccd_idx1: list of normal index candidates for Geom 1 + multiccd_idx1 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=int) + # multiccd_idx2: list of normal index candidates for Geom 2 + multiccd_idx2 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=int) + # multiccd_n1: list of normal candidates for Geom 1 + multiccd_n1 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + # multiccd_n2: list of normal candidates for Geom 1 + multiccd_n2 = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + # multiccd_endvert: list of edge vertices candidates + multiccd_endvert = wp.empty(shape=(d.naconmax, nmaxmeshdeg), dtype=wp.vec3) + # multiccd_face1: contact face + multiccd_face1 = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + # multiccd_face2: contact face + multiccd_face2 = wp.empty(shape=(d.naconmax, nmaxpolygon), dtype=wp.vec3) + for geom_pair in _CONVEX_COLLISION_PAIRS: g1 = geom_pair[0].value g2 = geom_pair[1].value if m.geom_pair_type_count[upper_trid_index(len(GeomType), g1, g2)]: wp.launch( - ccd_kernel_builder(m.opt.legacy_gjk, g1, g2, m.opt.ccd_iterations, True, 1e9, g1 == GeomType.HFIELD), + ccd_kernel_builder(m.opt.legacy_gjk, g1, g2, m.opt.ccd_iterations, True, 1e9, g1 == GeomType.HFIELD, use_multiccd), dim=d.naconmax, inputs=[ m.opt.ccd_tolerance, @@ -687,28 +799,28 @@ def convex_narrowphase(m: Model, d: Data): d.collision_pairid, d.collision_worldid, d.ncollision, - d.epa_vert, - d.epa_vert1, - d.epa_vert2, - d.epa_vert_index1, - d.epa_vert_index2, - d.epa_face, - d.epa_pr, - d.epa_norm2, - d.epa_index, - d.epa_map, - d.epa_horizon, - d.multiccd_polygon, - d.multiccd_clipped, - d.multiccd_pnormal, - d.multiccd_pdist, - d.multiccd_idx1, - d.multiccd_idx2, - d.multiccd_n1, - d.multiccd_n2, - d.multiccd_endvert, - d.multiccd_face1, - d.multiccd_face2, + epa_vert, + epa_vert1, + epa_vert2, + epa_vert_index1, + epa_vert_index2, + epa_face, + epa_pr, + epa_norm2, + epa_index, + epa_map, + epa_horizon, + multiccd_polygon, + multiccd_clipped, + multiccd_pnormal, + multiccd_pdist, + multiccd_idx1, + multiccd_idx2, + multiccd_n1, + multiccd_n2, + multiccd_endvert, + multiccd_face1, + multiccd_face2, ], outputs=[ d.nacon, @@ -723,5 +835,7 @@ def convex_narrowphase(m: Model, d: Data): d.contact.dim, d.contact.geom, d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, ], ) 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 24891934..ef314275 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 @@ -26,7 +26,6 @@ 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 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.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope @@ -233,12 +232,11 @@ def _obb_filter( return True -@cache_kernel -def _broadphase_filter(opt_broadphase_filter: int): +def _broadphase_filter(m: Model): @wp.func def func( # Model: - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_rbound: wp.array2d(dtype=float), geom_margin: wp.array2d(dtype=float), # Data in: @@ -254,30 +252,28 @@ def _broadphase_filter(opt_broadphase_filter: int): # 4: aabb # 8: obb - center1 = geom_aabb[geom1, 0] - center2 = geom_aabb[geom2, 0] - size1 = geom_aabb[geom1, 1] - size2 = geom_aabb[geom2, 1] - rbound1 = geom_rbound[worldid, geom1] - rbound2 = geom_rbound[worldid, geom2] - margin1 = geom_margin[worldid, geom1] - margin2 = geom_margin[worldid, geom2] - xpos1 = geom_xpos_in[worldid, geom1] - xpos2 = geom_xpos_in[worldid, geom2] - xmat1 = geom_xmat_in[worldid, geom1] - xmat2 = geom_xmat_in[worldid, geom2] + aabb_id = worldid % geom_aabb.shape[0] if wp.static(m.geom_aabb.shape[0] > 1) else 0 + center1, center2 = geom_aabb[aabb_id, geom1, 0], geom_aabb[aabb_id, geom2, 0] + size1, size2 = geom_aabb[aabb_id, geom1, 1], geom_aabb[aabb_id, geom2, 1] + + rbound_id = worldid % geom_rbound.shape[0] if wp.static(m.geom_rbound.shape[0] > 1) else 0 + rbound1, rbound2 = geom_rbound[rbound_id, geom1], geom_rbound[rbound_id, geom2] + margin_id = worldid % geom_margin.shape[0] if wp.static(m.geom_margin.shape[0] > 1) else 0 + margin1, margin2 = geom_margin[margin_id, geom1], geom_margin[margin_id, geom2] + xpos1, xpos2 = geom_xpos_in[worldid, geom1], geom_xpos_in[worldid, geom2] + xmat1, xmat2 = geom_xmat_in[worldid, geom1], geom_xmat_in[worldid, geom2] if rbound1 == 0.0 or rbound2 == 0.0: - if wp.static(opt_broadphase_filter & BroadphaseFilter.PLANE): + if wp.static(m.opt.broadphase_filter & BroadphaseFilter.PLANE): return _plane_filter(rbound1, rbound2, margin1, margin2, xpos1, xpos2, xmat1, xmat2) else: - if wp.static(opt_broadphase_filter & BroadphaseFilter.SPHERE): + if wp.static(m.opt.broadphase_filter & BroadphaseFilter.SPHERE): if not _sphere_filter(rbound1, rbound2, margin1, margin2, xpos1, xpos2): return False - if wp.static(opt_broadphase_filter & BroadphaseFilter.AABB): + if wp.static(m.opt.broadphase_filter & BroadphaseFilter.AABB): if not _aabb_filter(center1, center2, size1, size2, margin1, margin2, xpos1, xpos2, xmat1, xmat2): return False - if wp.static(opt_broadphase_filter & BroadphaseFilter.OBB): + if wp.static(m.opt.broadphase_filter & BroadphaseFilter.OBB): if not _obb_filter(center1, center2, size1, size2, margin1, margin2, xpos1, xpos2, xmat1, xmat2): return False @@ -290,7 +286,7 @@ def _broadphase_filter(opt_broadphase_filter: int): def _add_geom_pair( # Model: geom_type: wp.array(dtype=int), - nxn_pairid: wp.array(dtype=int), + nxn_pairid: wp.array(dtype=wp.vec2i), # Data in: naconmax_in: int, # In: @@ -300,7 +296,7 @@ def _add_geom_pair( nxnid: int, # Data out: collision_pair_out: wp.array(dtype=wp.vec2i), - collision_pairid_out: wp.array(dtype=int), + collision_pairid_out: wp.array(dtype=wp.vec2i), collision_worldid_out: wp.array(dtype=int), ncollision_out: wp.array(dtype=int), ): @@ -334,64 +330,76 @@ def _binary_search(values: wp.array(dtype=Any), value: Any, lower: int, upper: i return upper -@wp.kernel -def _sap_project( - # Model: - geom_rbound: wp.array2d(dtype=float), - geom_margin: wp.array2d(dtype=float), - # Data in: - geom_xpos_in: wp.array2d(dtype=wp.vec3), - # In: - direction_in: wp.vec3, - # Data out: - sap_projection_lower_out: wp.array2d(dtype=float), # kernel_analyzer: ignore - sap_projection_upper_out: wp.array2d(dtype=float), - sap_sort_index_out: wp.array2d(dtype=int), # kernel_analyzer: ignore -): - worldid, geomid = wp.tid() +def _sap_project(opt_broadphase: int): + @nested_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), + # In: + direction_in: wp.vec3, + # Out: + projection_lower_out: wp.array2d(dtype=float), + projection_upper_out: wp.array2d(dtype=float), + sort_index_out: wp.array2d(dtype=int), + segmented_index_out: wp.array(dtype=int), + ): + worldid, geomid = wp.tid() - xpos = geom_xpos_in[worldid, geomid] - rbound = geom_rbound[worldid, geomid] + xpos = geom_xpos_in[worldid, geomid] + rbound = geom_rbound[worldid % geom_rbound.shape[0], geomid] - if rbound == 0.0: - # geom is a plane - rbound = MJ_MAXVAL + if rbound == 0.0: + # geom is a plane + rbound = MJ_MAXVAL - radius = rbound + geom_margin[worldid, geomid] - center = wp.dot(direction_in, xpos) + radius = rbound + geom_margin[worldid % geom_margin.shape[0], geomid] + center = wp.dot(direction_in, xpos) - sap_sort_index_out[worldid, geomid] = geomid - if not wp.isnan(center): - sap_projection_lower_out[worldid, geomid] = center - radius - sap_projection_upper_out[worldid, geomid] = center + radius - else: - sap_projection_lower_out[worldid, geomid] = MJ_MAXVAL - sap_projection_upper_out[worldid, geomid] = MJ_MAXVAL + sort_index_out[worldid, geomid] = geomid + if not wp.isnan(center): + projection_lower_out[worldid, geomid] = center - radius + projection_upper_out[worldid, geomid] = center + radius + else: + projection_lower_out[worldid, geomid] = MJ_MAXVAL + projection_upper_out[worldid, geomid] = MJ_MAXVAL + + if wp.static(opt_broadphase == BroadphaseType.SAP_SEGMENTED): + if geomid == 0: + segmented_index_out[worldid] = worldid * ngeom + if worldid == nworld_in - 1: + segmented_index_out[nworld_in] = nworld_in * ngeom + + return sap_project @wp.kernel def _sap_range( # Model: ngeom: int, - # Data in: - sap_projection_lower_in: wp.array2d(dtype=float), # kernel_analyzer: ignore - sap_projection_upper_in: wp.array2d(dtype=float), - sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore - # Data out: - sap_range_out: wp.array2d(dtype=int), + # In: + projection_lower_in: wp.array2d(dtype=float), + projection_upper_in: wp.array2d(dtype=float), + sort_index_in: wp.array2d(dtype=int), + # Out: + range_out: wp.array2d(dtype=int), ): worldid, geomid = wp.tid() # current bounding geom - idx = sap_sort_index_in[worldid, geomid] + idx = sort_index_in[worldid, geomid] - upper = sap_projection_upper_in[worldid, idx] + upper = projection_upper_in[worldid, idx] - limit = _binary_search(sap_projection_lower_in[worldid], upper, geomid + 1, ngeom) + limit = _binary_search(projection_lower_in[worldid], upper, geomid + 1, ngeom) limit = wp.min(ngeom - 1, limit) # range of geoms for the sweep and prune process - sap_range_out[worldid, geomid] = limit - geomid + range_out[worldid, geomid] = limit - geomid @cache_kernel @@ -401,45 +409,45 @@ def _sap_broadphase(broadphase_filter): # Model: ngeom: int, geom_type: wp.array(dtype=int), - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_rbound: wp.array2d(dtype=float), geom_margin: wp.array2d(dtype=float), - nxn_pairid: wp.array(dtype=int), + 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), - sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore - sap_cumulative_sum_in: wp.array(dtype=int), # kernel_analyzer: ignore # In: + sort_index_in: wp.array2d(dtype=int), + cumulative_sum_in: wp.array(dtype=int), nsweep_in: int, # Data out: collision_pair_out: wp.array(dtype=wp.vec2i), - collision_pairid_out: wp.array(dtype=int), + collision_pairid_out: wp.array(dtype=wp.vec2i), collision_worldid_out: wp.array(dtype=int), ncollision_out: wp.array(dtype=int), ): worldgeomid = wp.tid() nworldgeom = nworld_in * ngeom - nworkpackages = sap_cumulative_sum_in[nworldgeom - 1] + nworkpackages = cumulative_sum_in[nworldgeom - 1] while worldgeomid < nworkpackages: # binary search to find current and next geom pair indices - i = _binary_search(sap_cumulative_sum_in, worldgeomid, 0, nworldgeom) + i = _binary_search(cumulative_sum_in, worldgeomid, 0, nworldgeom) j = i + worldgeomid + 1 if i > 0: - j -= sap_cumulative_sum_in[i - 1] + j -= cumulative_sum_in[i - 1] worldid = i // ngeom i = i % ngeom j = j % ngeom # get geom indices and swap if necessary - geom1 = sap_sort_index_in[worldid, i] - geom2 = sap_sort_index_in[worldid, j] + geom1 = sort_index_in[worldid, i] + geom2 = sort_index_in[worldid, j] # find linear index of (geom1, geom2) in upper triangular nxn_pairid if geom2 < geom1: @@ -448,10 +456,14 @@ def _sap_broadphase(broadphase_filter): idx = upper_tri_index(ngeom, geom1, geom2) worldgeomid += nsweep_in - if nxn_pairid[idx] < -1: + pairid = nxn_pairid[idx] + if pairid[0] < -1 and pairid[1] < 0: continue - if broadphase_filter(geom_aabb, geom_rbound, geom_margin, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid): + if ( + broadphase_filter(geom_aabb, geom_rbound, geom_margin, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid) + or pairid[1] >= 0 + ): _add_geom_pair( geom_type, nxn_pairid, @@ -472,22 +484,25 @@ def _sap_broadphase(broadphase_filter): def _segmented_sort(tile_size: int): @wp.kernel def segmented_sort( - # Data in: - sap_projection_lower_in: wp.array2d(dtype=float), # kernel_analyzer: ignore - sap_sort_index_in: wp.array2d(dtype=int), # kernel_analyzer: ignore + # In: + projection_lower_in: wp.array2d(dtype=float), + sort_index_in: wp.array2d(dtype=int), + # Out: + projection_lower_out: wp.array2d(dtype=float), + sort_index_out: wp.array2d(dtype=int), ): worldid = wp.tid() # Load input into shared memory - keys = wp.tile_load(sap_projection_lower_in[worldid], shape=tile_size, storage="shared") - values = wp.tile_load(sap_sort_index_in[worldid], shape=tile_size, storage="shared") + keys = wp.tile_load(projection_lower_in[worldid], shape=tile_size, storage="shared") + values = wp.tile_load(sort_index_in[worldid], shape=tile_size, storage="shared") # Perform in-place sorting wp.tile_sort(keys, values) # Store sorted shared memory into output arrays - wp.tile_store(sap_projection_lower_in[worldid], keys) - wp.tile_store(sap_sort_index_in[worldid], values) + wp.tile_store(projection_lower_out[worldid], keys) + wp.tile_store(sort_index_out[worldid], values) return segmented_sort @@ -505,11 +520,11 @@ def sap_broadphase(m: Model, d: Data): bounding sphere check is performed. If this check passes, the pair is added to the collision arrays in `d` for the narrowphase stage. - Two sorting strategies are supported, controlled by `m.opt.broadphase`: + Two sorting strategies are supported, controlled by `m.opt.broadphase` + - `SAP_TILE`: Uses a tile-based sort. - `SAP_SEGMENTED`: Uses a segmented sort. """ - nworldgeom = d.nworld * m.ngeom # TODO(team): direction @@ -518,58 +533,52 @@ def sap_broadphase(m: Model, d: Data): direction = wp.vec3(0.5935, 0.7790, 0.1235) direction = wp.normalize(direction) + projection_lower = wp.empty((d.nworld, m.ngeom, 2), dtype=float) + projection_upper = wp.empty((d.nworld, m.ngeom), dtype=float) + sort_index = wp.empty((d.nworld, m.ngeom, 2), dtype=int) + range_ = wp.empty((d.nworld, m.ngeom), dtype=int) + cumulative_sum = wp.empty((d.nworld, m.ngeom), dtype=int) + segmented_index = wp.empty(d.nworld + 1 if m.opt.broadphase == BroadphaseType.SAP_SEGMENTED else 0, dtype=int) + wp.launch( - kernel=_sap_project, + kernel=_sap_project(m.opt.broadphase), dim=(d.nworld, m.ngeom), - inputs=[ - m.geom_rbound, - m.geom_margin, - d.geom_xpos, - direction, - ], + inputs=[m.ngeom, m.geom_rbound, m.geom_margin, d.nworld, d.geom_xpos, direction], outputs=[ - d.sap_projection_lower.reshape((-1, m.ngeom)), - d.sap_projection_upper, - d.sap_sort_index.reshape((-1, m.ngeom)), + projection_lower.reshape((-1, m.ngeom)), + projection_upper, + sort_index.reshape((-1, m.ngeom)), + segmented_index, ], ) if m.opt.broadphase == BroadphaseType.SAP_TILE: wp.launch_tiled( kernel=_segmented_sort(m.ngeom), - dim=(d.nworld), - inputs=[d.sap_projection_lower.reshape((-1, m.ngeom)), d.sap_sort_index.reshape((-1, m.ngeom))], + dim=d.nworld, + inputs=[projection_lower.reshape((-1, m.ngeom)), sort_index.reshape((-1, m.ngeom))], + outputs=[projection_lower.reshape((-1, m.ngeom)), sort_index.reshape((-1, m.ngeom))], block_dim=m.block_dim.segmented_sort, ) else: wp.utils.segmented_sort_pairs( - d.sap_projection_lower.reshape((-1, m.ngeom)), - d.sap_sort_index.reshape((-1, m.ngeom)), - nworldgeom, - d.sap_segment_index.reshape(-1), + projection_lower.reshape((-1, m.ngeom)), sort_index.reshape((-1, m.ngeom)), nworldgeom, segmented_index ) wp.launch( kernel=_sap_range, dim=(d.nworld, m.ngeom), - inputs=[ - m.ngeom, - d.sap_projection_lower.reshape((-1, m.ngeom)), - d.sap_projection_upper, - d.sap_sort_index.reshape((-1, m.ngeom)), - ], - outputs=[ - d.sap_range, - ], + inputs=[m.ngeom, projection_lower.reshape((-1, m.ngeom)), projection_upper, sort_index.reshape((-1, m.ngeom))], + outputs=[range_], ) # scan is used for load balancing among the threads - wp.utils.array_scan(d.sap_range.reshape(-1), d.sap_cumulative_sum.reshape(-1), True) + wp.utils.array_scan(range_.reshape(-1), cumulative_sum.reshape(-1), True) # estimate number of overlap checks # assumes each geom has 5 other geoms (batched over all worlds) nsweep = 5 * nworldgeom - broadphase_filter = _broadphase_filter(m.opt.broadphase_filter) + broadphase_filter = _broadphase_filter(m) wp.launch( kernel=_sap_broadphase(broadphase_filter), dim=nsweep, @@ -584,16 +593,11 @@ def sap_broadphase(m: Model, d: Data): d.naconmax, d.geom_xpos, d.geom_xmat, - d.sap_sort_index.reshape((-1, m.ngeom)), - d.sap_cumulative_sum.reshape(-1), + 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.collision_pair, d.collision_pairid, d.collision_worldid, d.ncollision], ) @@ -603,18 +607,18 @@ def _nxn_broadphase(broadphase_filter): def kernel( # Model: geom_type: wp.array(dtype=int), - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_rbound: wp.array2d(dtype=float), geom_margin: wp.array2d(dtype=float), nxn_geom_pair: wp.array(dtype=wp.vec2i), - nxn_pairid: wp.array(dtype=int), + 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), # Data out: collision_pair_out: wp.array(dtype=wp.vec2i), - collision_pairid_out: wp.array(dtype=int), + collision_pairid_out: wp.array(dtype=wp.vec2i), collision_worldid_out: wp.array(dtype=int), ncollision_out: wp.array(dtype=int), ): @@ -624,7 +628,10 @@ def _nxn_broadphase(broadphase_filter): geom1 = geom[0] geom2 = geom[1] - if broadphase_filter(geom_aabb, geom_rbound, geom_margin, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid): + if ( + broadphase_filter(geom_aabb, geom_rbound, geom_margin, geom_xpos_in, geom_xmat_in, geom1, geom2, worldid) + or nxn_pairid[elementid][1] >= 0 + ): _add_geom_pair( geom_type, nxn_pairid, @@ -656,8 +663,7 @@ def nxn_broadphase(m: Model, d: Data): The initial list of pairs is filtered at model creation time to exclude pairs based on `contype`/`conaffinity`, parent-child relationships, and explicit `` tags. """ - - broadphase_filter = _broadphase_filter(m.opt.broadphase_filter) + broadphase_filter = _broadphase_filter(m) wp.launch( _nxn_broadphase(broadphase_filter), dim=(d.nworld, m.nxn_geom_pair_filtered.shape[0]), @@ -709,7 +715,6 @@ 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]) @@ -721,7 +726,4 @@ def collision(m: Model, d: Data): else: sap_broadphase(m, d) - if m.opt.graph_conditional: - wp.capture_if(condition=d.ncollision, on_true=_narrowphase, m=m, d=d) - else: - _narrowphase(m, d) + _narrowphase(m, d) 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 892a7115..bd0b29bb 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 @@ -924,7 +924,7 @@ def _polytope2( geomtype1: int, geomtype2: int, ) -> Tuple[Polytope, GJKResult]: - """Create polytope for EPA given a 1-simplex from GJK""" + """Create polytope for EPA given a 1-simplex from GJK.""" diff = simplex[1] - simplex[0] # find component with smallest magnitude (so cross product is largest) @@ -1023,7 +1023,7 @@ def _polytope3( geomtype1: int, geomtype2: int, ) -> Polytope: - """Create polytope for EPA given a 2-simplex from GJK""" + """Create polytope for EPA given a 2-simplex from GJK.""" # get normals in both directions n = wp.cross(simplex[1] - simplex[0], simplex[2] - simplex[0]) if wp.norm_l2(n) < MJ_MINVAL: @@ -1123,7 +1123,7 @@ def _polytope4( geomtype1: int, geomtype2: int, ) -> Tuple[Polytope, GJKResult]: - """Create polytope for EPA given a 3-simplex from GJK""" + """Create polytope for EPA given a 3-simplex from GJK.""" pt.vert[0] = simplex[0] pt.vert[1] = simplex[1] pt.vert[2] = simplex[2] @@ -1906,7 +1906,7 @@ def _set_edge( # recover multiple contacts from EPA polytope @wp.func -def _multicontact( +def multicontact( # In: polygon: wp.array(dtype=wp.vec3), clipped: wp.array(dtype=wp.vec3), @@ -1919,7 +1919,10 @@ def _multicontact( endvert: wp.array(dtype=wp.vec3), face1: wp.array(dtype=wp.vec3), face2: wp.array(dtype=wp.vec3), - pt: Polytope, + 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), face: wp.vec3i, x1: wp.vec3, x2: wp.vec3, @@ -1953,8 +1956,8 @@ def _multicontact( polymap = geom2.mesh_polymap # get dimensions of features of geoms 1 and 2 - nface1, feature_index1, feature_vertex1 = _feature_dim(face, pt.vert_index1, pt.vert1) - nface2, feature_index2, feature_vertex2 = _feature_dim(face, pt.vert_index2, pt.vert2) + 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) dir = x2 - x1 dir_neg = -dir @@ -2070,7 +2073,7 @@ def _multicontact( # recover geom1 matching edge or face if is_edge_contact_geom1: - nface1 = _set_edge(pt.vert1, endvert, face[0], i, face1) + nface1 = _set_edge(epa_vert1, endvert, face[0], i, face1) else: ind = wp.where(is_edge_contact_geom2, idx1[j], idx1[i]) if geomtype1 == GeomType.BOX: @@ -2091,7 +2094,7 @@ def _multicontact( # recover geom2 matching edge or face if is_edge_contact_geom2: - nface2 = _set_edge(pt.vert2, endvert, face[0], i, face2) + nface2 = _set_edge(epa_vert2, endvert, face[0], i, face2) else: if geomtype2 == GeomType.BOX: nface2 = _box_face(geom2.rot, geom2.pos, geom2.size, idx2[j], face2) @@ -2147,7 +2150,6 @@ def _inflate(dist: float, x1: wp.vec3, x2: wp.vec3, margin1: float, margin2: flo @wp.func def ccd( # In: - multiccd: bool, tolerance: float, cutoff: float, ccd_iterations: int, @@ -2168,21 +2170,8 @@ def ccd( face_index: wp.array(dtype=int), face_map: wp.array(dtype=int), horizon: wp.array(dtype=int), - polygon: wp.array(dtype=wp.vec3), - clipped: wp.array(dtype=wp.vec3), - plane_normal: wp.array(dtype=wp.vec3), - plane_dist: wp.array(dtype=float), - idx1: wp.array(dtype=int), - idx2: wp.array(dtype=int), - n1: wp.array(dtype=wp.vec3), - n2: wp.array(dtype=wp.vec3), - endvert: wp.array(dtype=wp.vec3), - face1: wp.array(dtype=wp.vec3), - face2: wp.array(dtype=wp.vec3), -) -> Tuple[float, int, mat3c, mat3c]: +) -> Tuple[float, int, wp.vec3, wp.vec3, int]: """General convex collision detection via GJK/EPA.""" - witness1 = mat3c() - witness2 = mat3c() full_margin1 = 0.0 full_margin2 = 0.0 size1 = 0.0 @@ -2211,13 +2200,9 @@ def ccd( # shallow penetration, inflate contact if result.dist > tolerance: if result.dist == FLOAT_MAX: - witness1[0] = result.x1 - witness2[0] = result.x2 - return result.dist, 1, witness1, witness2 + return result.dist, 1, result.x1, result.x2, -1 dist, x1, x2 = _inflate(result.dist, result.x1, result.x2, full_margin1, full_margin2) - witness1[0] = x1 - witness2[0] = x2 - return dist, 1, witness1, witness2 + return dist, 1, x1, x2, -1 # deep penetration, reset initial conditions and rerun GJK + EPA geom1.margin = full_margin1 - size1 @@ -2230,9 +2215,7 @@ def ccd( # no penetration depth to recover if result.dist > tolerance or result.dim < 2: - witness1[0] = result.x1 - witness2[0] = result.x2 - return result.dist, 1, witness1, witness2 + return result.dist, 1, result.x1, result.x2, -1 pt = Polytope() pt.nface = 0 @@ -2312,47 +2295,9 @@ def ccd( # origin on boundary (objects are not considered penetrating) if pt.status: - witness1[0] = result.x1 - witness2[0] = result.x2 - return result.dist, 1, witness1, witness2 + return result.dist, 1, result.x1, result.x2, -1 dist, x1, x2, idx = _epa(tolerance, ccd_iterations, pt, geom1, geom2, geomtype1, geomtype2, is_discrete) if idx == -1: - return FLOAT_MAX, 0, witness1, witness2 - - # multiccd is always on for box-box collisions - if geomtype1 == GeomType.BOX and geomtype2 == GeomType.BOX: - multiccd = True - - if ( - multiccd - and (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)) - ): - num, w1, w2 = _multicontact( - polygon, - clipped, - plane_normal, - plane_dist, - idx1, - idx2, - n1, - n2, - endvert, - face1, - face2, - pt, - pt.face[idx], - x1, - x2, - geom1, - geom2, - geomtype1, - geomtype2, - ) - if num > 0: - return dist, num, w1, w2 - witness1[0] = x1 - witness2[0] = x2 - return dist, 1, witness1, witness2 + return FLOAT_MAX, 0, wp.vec3(), wp.vec3(), -1 + return dist, 1, x1, x2, idx diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py index b2eed4ab..799b4414 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py @@ -24,7 +24,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL def hfield_filter( # Model: geom_dataid: wp.array(dtype=int), - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_rbound: wp.array2d(dtype=float), geom_margin: wp.array2d(dtype=float), hfield_size: wp.array(dtype=wp.vec4), @@ -43,17 +43,20 @@ def hfield_filter( # height field info hfdataid = geom_dataid[g1] size1 = hfield_size[hfdataid] + + # geom info + rbound_id = worldid % geom_rbound.shape[0] + margin_id = worldid % geom_margin.shape[0] + pos1 = geom_xpos_in[worldid, g1] mat1 = geom_xmat_in[worldid, g1] mat1T = wp.transpose(mat1) - - # geom info pos2 = geom_xpos_in[worldid, g2] pos = mat1T @ (pos2 - pos1) - r2 = geom_rbound[worldid, g2] + r2 = geom_rbound[rbound_id, g2] # TODO(team): margin? - margin = wp.max(geom_margin[worldid, g1], geom_margin[worldid, g2]) + margin = wp.max(geom_margin[margin_id, g1], geom_margin[margin_id, g2]) # box-sphere test: horizontal plane for i in range(2): @@ -78,8 +81,9 @@ def hfield_filter( ymin = MJ_MAXVAL zmin = MJ_MAXVAL - center2 = geom_aabb[g2, 0] - size2 = geom_aabb[g2, 1] + aabb_id = worldid % geom_aabb.shape[0] + center2 = geom_aabb[aabb_id, g2, 0] + size2 = geom_aabb[aabb_id, g2, 1] pos += mat1T @ center2 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 778e0c85..fef22032 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 @@ -34,6 +34,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.math import safe_div from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINMU from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL +from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import Model @@ -44,8 +45,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_k wp.set_module_options({"enable_backward": False}) -_HUGE_VAL = 1e6 - class mat43f(wp.types.matrix(shape=(4, 3), dtype=wp.float32)): pass @@ -88,8 +87,8 @@ def geom( geom_size: wp.vec3, 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), @@ -146,16 +145,16 @@ def plane_convex(plane_normal: wp.vec3, plane_pos: wp.vec3, convex: Geom) -> Tup """Core contact geometry calculation for plane-convex collision. Args: - plane_normal: Normal vector of the plane - plane_pos: Position point on the plane - convex: Convex geometry object containing position, rotation, and mesh data + plane_normal: Normal vector of the plane. + plane_pos: Position point on the plane. + convex: Convex geometry object containing position, rotation, and mesh data. Returns: - Tuple containing: - contact_dist: Vector of contact distances (wp.inf for unpopulated contacts) - contact_pos: Matrix of contact positions (one per row) - contact_normal: Matrix of contact normal vectors (one per row) + - Vector of contact distances (wp.inf 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_pos = mat43f() @@ -375,6 +374,7 @@ def write_contact( # Data in: naconmax_in: int, # In: + id_: int, dist_in: float, pos_in: wp.vec3, frame_in: wp.mat33, @@ -386,9 +386,9 @@ def write_contact( solreffriction_in: wp.vec2, solimp_in: vec5, geoms_in: wp.vec2i, + pairid_in: wp.vec2i, worldid_in: 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), @@ -400,22 +400,40 @@ def write_contact( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): - if dist_in - margin_in < 0.0: - cid = wp.atomic_add(nacon_out, 0, 1) - if cid < naconmax_in: - contact_dist_out[cid] = dist_in - contact_pos_out[cid] = pos_in - contact_frame_out[cid] = frame_in - contact_geom_out[cid] = geoms_in - contact_worldid_out[cid] = worldid_in - includemargin = margin_in - gap_in - contact_includemargin_out[cid] = includemargin - contact_dim_out[cid] = condim_in - contact_friction_out[cid] = friction_in - contact_solref_out[cid] = solref_in - contact_solreffriction_out[cid] = solreffriction_in - contact_solimp_out[cid] = solimp_in + 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 + + contact_type = 0 + + if pairid_in[0] >= -1 and active: + contact_type |= ContactType.CONSTRAINT + + if pairid_in[1] >= 0: + contact_type |= ContactType.SENSOR + + cid = wp.atomic_add(nacon_out, 0, 1) + if cid < naconmax_in: + contact_dist_out[cid] = dist_in + contact_pos_out[cid] = pos_in + contact_frame_out[cid] = frame_in + contact_geom_out[cid] = geoms_in + contact_worldid_out[cid] = worldid_in + includemargin = margin_in - gap_in + contact_includemargin_out[cid] = includemargin + contact_dim_out[cid] = condim_in + contact_friction_out[cid] = friction_in + contact_solref_out[cid] = solref_in + contact_solreffriction_out[cid] = solreffriction_in + contact_solimp_out[cid] = solimp_in + contact_type_out[cid] = contact_type + contact_geomcollisionid_out[cid] = id_ @wp.func @@ -438,13 +456,16 @@ def contact_params( pair_friction: wp.array2d(dtype=vec5), # Data in: collision_pair_in: wp.array(dtype=wp.vec2i), - collision_pairid_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=wp.vec2i), # In: cid: int, worldid: int, ): geoms = collision_pair_in[cid] - pairid = collision_pairid_in[cid] + pairid = collision_pairid_in[cid][0] + + # TODO(team): early return if collision sensor but no contact + # (ie, pairid[0] < -1 and pairid[1] < 0) if pairid > -1: margin = pair_margin[worldid, pairid] @@ -457,9 +478,15 @@ def contact_params( else: g1 = geoms[0] g2 = geoms[1] + solmix_id = worldid % geom_solmix.shape[0] + friction_id = worldid % geom_friction.shape[0] + solref_id = worldid % geom_solref.shape[0] + solimp_id = worldid % geom_solimp.shape[0] + margin_id = worldid % geom_margin.shape[0] + gap_id = worldid % geom_gap.shape[0] - solmix1 = geom_solmix[worldid, g1] - solmix2 = geom_solmix[worldid, g2] + solmix1 = geom_solmix[solmix_id, g1] + solmix2 = geom_solmix[solmix_id, g2] condim1 = geom_condim[g1] condim2 = geom_condim[g2] @@ -471,18 +498,18 @@ def contact_params( if p1 > p2: mix = 1.0 condim = condim1 - max_geom_friction = geom_friction[worldid, g1] + max_geom_friction = geom_friction[friction_id, g1] elif p2 > p1: mix = 0.0 condim = condim2 - max_geom_friction = geom_friction[worldid, g2] + max_geom_friction = geom_friction[friction_id, g2] else: mix = safe_div(solmix1, solmix1 + solmix2) mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 < MJ_MINVAL), 0.5, mix) mix = wp.where((solmix1 < MJ_MINVAL) and (solmix2 >= MJ_MINVAL), 0.0, mix) mix = wp.where((solmix1 >= MJ_MINVAL) and (solmix2 < MJ_MINVAL), 1.0, mix) condim = wp.max(condim1, condim2) - max_geom_friction = wp.max(geom_friction[worldid, g1], geom_friction[worldid, g2]) + max_geom_friction = wp.max(geom_friction[friction_id, g1], geom_friction[friction_id, g2]) friction = vec5( wp.max(MJ_MINMU, max_geom_friction[0]), @@ -492,18 +519,16 @@ def contact_params( wp.max(MJ_MINMU, max_geom_friction[2]), ) - if geom_solref[worldid, g1][0] > 0.0 and geom_solref[worldid, g2][0] > 0.0: - solref = mix * geom_solref[worldid, g1] + (1.0 - mix) * geom_solref[worldid, g2] + if geom_solref[solref_id, g1][0] > 0.0 and geom_solref[solref_id, g2][0] > 0.0: + solref = mix * geom_solref[solref_id, g1] + (1.0 - mix) * geom_solref[solref_id, g2] else: - solref = wp.min(geom_solref[worldid, g1], geom_solref[worldid, g2]) + solref = wp.min(geom_solref[solref_id, g1], geom_solref[solref_id, g2]) solreffriction = wp.vec2(0.0, 0.0) - - solimp = mix * geom_solimp[worldid, g1] + (1.0 - mix) * geom_solimp[worldid, g2] - + solimp = mix * geom_solimp[solimp_id, g1] + (1.0 - mix) * geom_solimp[solimp_id, g2] # geom priority is ignored - margin = wp.max(geom_margin[worldid, g1], geom_margin[worldid, g2]) - gap = wp.max(geom_gap[worldid, g1], geom_gap[worldid, g2]) + margin = wp.max(geom_margin[margin_id, g1], geom_margin[margin_id, g2]) + gap = wp.max(geom_gap[gap_id, g1], geom_gap[gap_id, g2]) return geoms, margin, gap, condim, friction, solref, solreffriction, solimp @@ -524,8 +549,8 @@ def plane_sphere_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -537,38 +562,45 @@ def plane_sphere_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contact between a sphere and a plane.""" - dist, pos = plane_sphere(plane.normal, plane.pos, sphere.pos, sphere.size[0]) + normal = plane.normal + dist, pos = plane_sphere(normal, plane.pos, sphere.pos, sphere.size[0]) - if dist - margin < 0: - write_contact( - naconmax_in, - dist, - pos, - make_frame(plane.normal), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + 0, + dist, + pos, + make_frame(normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -587,8 +619,8 @@ def sphere_sphere_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -600,38 +632,44 @@ def sphere_sphere_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contact between two spheres.""" dist, pos, normal = sphere_sphere(sphere1.pos, sphere1.size[0], sphere2.pos, sphere2.size[0]) - if dist - margin < 0: - write_contact( - naconmax_in, - dist, - pos, - make_frame(normal), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + 0, + dist, + pos, + make_frame(normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -650,8 +688,8 @@ def sphere_capsule_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -663,6 +701,9 @@ def sphere_capsule_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates one contact between a sphere and a capsule.""" # capsule axis @@ -670,34 +711,37 @@ def sphere_capsule_wrapper( dist, pos, normal = sphere_capsule(sphere.pos, sphere.size[0], cap.pos, axis, cap.size[0], cap.size[1]) - if dist - margin < 0: - write_contact( - naconmax_in, - dist, - pos, - make_frame(normal), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + 0, + dist, + pos, + make_frame(normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -716,8 +760,8 @@ def capsule_capsule_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -729,6 +773,9 @@ def capsule_capsule_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between two capsules.""" # capsule axes @@ -746,34 +793,37 @@ def capsule_capsule_wrapper( cap2.size[1], # half_length2 ) - if dist - margin < 0: - write_contact( - naconmax_in, - dist, - pos, - make_frame(normal), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + 0, + dist, + pos, + make_frame(normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -792,8 +842,8 @@ def plane_capsule_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -805,6 +855,9 @@ def plane_capsule_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between a capsule and a plane.""" # capsule axis @@ -820,35 +873,37 @@ def plane_capsule_wrapper( ) for i in range(2): - disti = dist[i] - if disti - margin < 0.0: - write_contact( - naconmax_in, - disti, - pos[i], - frame, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + i, + dist[i], + pos[i], + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -867,8 +922,8 @@ def plane_ellipsoid_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -880,38 +935,44 @@ def plane_ellipsoid_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between an ellipsoid and a plane.""" dist, pos, normal = plane_ellipsoid(plane.normal, plane.pos, ellipsoid.pos, ellipsoid.rot, ellipsoid.size) - if dist - margin < 0: - write_contact( - naconmax_in, - dist, - pos, - make_frame(normal), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + 0, + dist, + pos, + make_frame(normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -930,8 +991,8 @@ def plane_box_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -943,44 +1004,46 @@ def plane_box_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between a box and a plane.""" dist, pos, normal = plane_box(plane.normal, plane.pos, box.pos, box.rot, box.size) frame = make_frame(normal) for i in range(4): - disti = dist[i] - if disti - margin < 0.0: - write_contact( - naconmax_in, - disti, - pos[i], - frame, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) - - -_HUGE_VAL = 1e6 + write_contact( + naconmax_in, + i, + dist[i], + pos[i], + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -999,8 +1062,8 @@ def plane_convex_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -1012,41 +1075,46 @@ def plane_convex_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between a plane and a convex object.""" dist, pos, normal = plane_convex(plane.normal, plane.pos, convex) frame = make_frame(normal) for i in range(4): - disti = dist[i] - if disti - margin < 0.0: - write_contact( - naconmax_in, - disti, - pos[i], - frame, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + i, + dist[i], + pos[i], + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -1065,8 +1133,8 @@ def sphere_cylinder_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -1078,6 +1146,9 @@ def sphere_cylinder_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between a sphere and a cylinder.""" # cylinder axis @@ -1092,34 +1163,37 @@ def sphere_cylinder_wrapper( cylinder.size[1], # cylinder half_height ) - if dist - margin < 0.0: - write_contact( - naconmax_in, - dist, - pos, - make_frame(normal), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + 0, + dist, + pos, + make_frame(normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -1138,8 +1212,8 @@ def plane_cylinder_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -1151,6 +1225,9 @@ def plane_cylinder_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between a cylinder and a plane.""" # cylinder axis @@ -1167,35 +1244,37 @@ def plane_cylinder_wrapper( frame = make_frame(normal) for i in range(4): - disti = dist[i] - if disti - margin < 0.0: - write_contact( - naconmax_in, - disti, - pos[i], - frame, - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + i, + dist[i], + pos[i], + frame, + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -1214,8 +1293,8 @@ def sphere_box_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -1227,37 +1306,43 @@ def sphere_box_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): dist, pos, normal = sphere_box(sphere.pos, sphere.size[0], box.pos, box.rot, box.size) - if dist - margin < 0.0: - write_contact( - naconmax_in, - dist, - pos, - make_frame(normal), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + 0, + dist, + pos, + make_frame(normal), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -1276,8 +1361,8 @@ def capsule_box_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -1289,6 +1374,9 @@ def capsule_box_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between a capsule and a box.""" # Extract capsule axis @@ -1307,35 +1395,37 @@ def capsule_box_wrapper( # Loop over the contacts and write them for i in range(2): - disti = dist[i] - if disti - margin < 0.0: - write_contact( - naconmax_in, - disti, - pos[i], - make_frame(normal[i]), - margin, - gap, - condim, - friction, - solref, - solreffriction, - solimp, - geoms, - worldid, - nacon_out, - contact_dist_out, - contact_pos_out, - contact_frame_out, - contact_includemargin_out, - contact_friction_out, - contact_solref_out, - contact_solreffriction_out, - contact_solimp_out, - contact_dim_out, - contact_geom_out, - contact_worldid_out, - ) + write_contact( + naconmax_in, + i, + dist[i], + pos[i], + make_frame(normal[i]), + margin, + gap, + condim, + friction, + solref, + solreffriction, + solimp, + geoms, + pairid, + worldid, + contact_dist_out, + contact_pos_out, + contact_frame_out, + contact_includemargin_out, + contact_friction_out, + contact_solref_out, + contact_solreffriction_out, + contact_solimp_out, + contact_dim_out, + contact_geom_out, + contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, + ) @wp.func @@ -1354,8 +1444,8 @@ def box_box_wrapper( solreffriction: wp.vec2, solimp: vec5, geoms: wp.vec2i, + pairid: wp.vec2i, # 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), @@ -1367,6 +1457,9 @@ def box_box_wrapper( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): """Calculates contacts between two boxes.""" # Call the core function to get contact geometry @@ -1377,14 +1470,13 @@ def box_box_wrapper( box2.pos, box2.rot, box2.size, + margin, ) for i in range(8): - if dist[i] - margin >= 0.0: - continue - write_contact( naconmax_in, + i, dist[i], pos[i], make_frame(normal[i]), @@ -1396,8 +1488,8 @@ def box_box_wrapper( solreffriction, solimp, geoms, + pairid, worldid, - nacon_out, contact_dist_out, contact_pos_out, contact_frame_out, @@ -1409,6 +1501,9 @@ def box_box_wrapper( contact_dim_out, contact_geom_out, contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, ) @@ -1466,15 +1561,10 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_ 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), @@ -1485,6 +1575,11 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_ 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), @@ -1493,15 +1588,14 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_ 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, collision_pair_in: wp.array(dtype=wp.vec2i), - collision_pairid_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=wp.vec2i), collision_worldid_in: wp.array(dtype=int), ncollision_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), @@ -1513,6 +1607,9 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_ contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): tid = wp.tid() @@ -1551,14 +1648,15 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_ ) geom1_dataid = geom_dataid[g1] + geom1 = geom( type1, geom1_dataid, - geom_size[worldid, g1], + geom_size[worldid % geom_size.shape[0], g1], mesh_vertadr, mesh_vertnum, - mesh_vert, mesh_graphadr, + mesh_vert, mesh_graph, mesh_polynum, mesh_polyadr, @@ -1577,11 +1675,11 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_ geom2 = geom( type2, geom2_dataid, - geom_size[worldid, g2], + geom_size[worldid % geom_size.shape[0], g2], mesh_vertadr, mesh_vertnum, - mesh_vert, mesh_graphadr, + mesh_vert, mesh_graph, mesh_polynum, mesh_polyadr, @@ -1614,7 +1712,7 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_ solreffriction, solimp, geoms, - nacon_out, + collision_pairid_in[tid], contact_dist_out, contact_pos_out, contact_frame_out, @@ -1626,6 +1724,9 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_ contact_dim_out, contact_geom_out, contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, ) return _primitive_narrowphase @@ -1652,9 +1753,9 @@ def primitive_narrowphase(m: Model, d: Data): identified during the broadphase stage. It computes detailed contact information such as distance, position, and frame, and populates the `d.contact` array. - The primitive geom types handled are PLANE, SPHERE, CAPSULE, CYLINDER, BOX. + The primitive geom types: `PLANE`, `SPHERE`, `CAPSULE`, `CYLINDER`, and `BOX`. - It also handles collisions between planes and convex hulls. + Additionally, collisions between planes and convex hulls. To improve performance, it dynamically builds and launches a kernel tailored to the specific primitive collision types present in the model, avoiding @@ -1677,15 +1778,10 @@ def primitive_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, @@ -1696,6 +1792,11 @@ def primitive_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, @@ -1703,16 +1804,15 @@ def primitive_narrowphase(m: Model, d: Data): m.pair_margin, m.pair_gap, m.pair_friction, - d.naconmax, d.geom_xpos, d.geom_xmat, + d.naconmax, d.collision_pair, d.collision_pairid, d.collision_worldid, d.ncollision, ], outputs=[ - d.nacon, d.contact.dist, d.contact.pos, d.contact.frame, @@ -1724,5 +1824,8 @@ def primitive_narrowphase(m: Model, d: Data): d.contact.dim, d.contact.geom, d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + d.nacon, ], ) 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 f20df572..499bd806 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 @@ -55,7 +55,6 @@ def closest_segment_point_and_dist(a: wp.vec3, b: wp.vec3, pt: wp.vec3) -> Tuple @wp.func def closest_segment_to_segment_points(a0: wp.vec3, a1: wp.vec3, b0: wp.vec3, b1: wp.vec3) -> Tuple[wp.vec3, wp.vec3]: """Returns closest points between two line segments.""" - dir_a, len_a = normalize_with_norm(a1 - a0) dir_b, len_b = normalize_with_norm(b1 - b0) @@ -122,16 +121,15 @@ def sphere_sphere( """Sphere-sphere collision calculation. Args: - pos1: Center position of the first sphere - radius1: Radius of the first sphere - pos2: Center position of the second sphere - radius2: Radius of the second sphere + pos1: Center position of the first sphere. + radius1: Radius of the first sphere. + pos2: Center position of the second sphere. + radius2: Radius of the second sphere. Returns: - Tuple containing: - dist: Distance between sphere surfaces (negative if overlapping) - pos: Contact position - n: Contact normal vector + - Distance between sphere surfaces (negative if overlapping). + - Contact position. + - Contact normal vector. """ dir = pos2 - pos1 dist = wp.length(dir) @@ -157,20 +155,18 @@ def sphere_capsule( """Core contact geometry calculation for sphere-capsule collision. Args: - sphere_pos: Center position of the sphere - sphere_radius: Radius of the sphere - capsule_pos: Center position of the capsule - capsule_axis: Axis direction of the capsule - capsule_radius: Radius of the capsule - capsule_half_length: Half length of the capsule + sphere_pos: Center position of the sphere. + sphere_radius: Radius of the sphere. + capsule_pos: Center position of the capsule. + capsule_axis: Axis direction of the capsule. + capsule_radius: Radius of the capsule. + capsule_half_length: Half length of the capsule. Returns: - Tuple containing: - contact_dist: Vector of contact distances - contact_pos: Matrix of contact positions (one per row) - contact_normals: Matrix of contact normal vectors (one per row) + - Vector of contact distances. + - Matrix of contact positions (one per row). + - Matrix of contact normal vectors (one per row). """ - # Calculate capsule segment segment = capsule_axis * capsule_half_length @@ -196,22 +192,20 @@ def capsule_capsule( """Core contact geometry calculation for capsule-capsule collision. Args: - cap1_pos: Center position of the first capsule - cap1_axis: Axis direction of the first capsule - cap1_radius: Radius of the first capsule - cap1_half_length: Half length of the first capsule - cap2_pos: Center position of the second capsule - cap2_axis: Axis direction of the second capsule - cap2_radius: Radius of the second capsule - cap2_half_length: Half length of the second capsule + cap1_pos: Center position of the first capsule. + cap1_axis: Axis direction of the first capsule. + cap1_radius: Radius of the first capsule. + cap1_half_length: Half length of the first capsule. + cap2_pos: Center position of the second capsule. + cap2_axis: Axis direction of the second capsule. + cap2_radius: Radius of the second capsule. + cap2_half_length: Half length of the second capsule. Returns: - Tuple containing: - contact_dist: Vector of contact distances - contact_pos: Matrix of contact positions (one per row) - contact_normals: Matrix of contact normal vectors (one per row) + - Vector of contact distances. + - Matrix of contact positions (one per row). + - Matrix of contact normal vectors (one per row). """ - # TODO(team): parallel axes case # Calculate capsule segments @@ -243,20 +237,18 @@ def plane_capsule( """Core contact geometry calculation for plane-capsule collision. Args: - plane_normal: Normal vector of the plane - plane_pos: Position point on the plane - capsule_pos: Center position of the capsule - capsule_axis: Axis direction of the capsule - capsule_radius: Radius of the capsule - capsule_half_length: Half length of the capsule + plane_normal: Normal vector of the plane. + plane_pos: Position point on the plane. + capsule_pos: Center position of the capsule. + capsule_axis: Axis direction of the capsule. + capsule_radius: Radius of the capsule. + capsule_half_length: Half length of the capsule. Returns: - Tuple containing: - contact_dist: Vector of contact distances - contact_pos: Matrix of contact positions (one per row) - contact_frame: Contact frame for both contacts + - Vector of contact distances. + - Matrix of contact positions (one per row). + - Contact frame for both contacts. """ - n = plane_normal axis = capsule_axis @@ -297,17 +289,16 @@ def plane_ellipsoid( """Core contact geometry calculation for plane-ellipsoid collision. Args: - plane_normal: Normal vector of the plane - plane_pos: Position point on the plane - ellipsoid_pos: Center position of the ellipsoid - ellipsoid_rot: Rotation matrix of the ellipsoid - ellipsoid_size: Size (radii) of the ellipsoid along each axis + plane_normal: Normal vector of the plane. + plane_pos: Position point on the plane. + ellipsoid_pos: Center position of the ellipsoid. + ellipsoid_rot: Rotation matrix of the ellipsoid. + ellipsoid_size: Size (radii) of the ellipsoid along each axis. Returns: - Tuple containing: - contact_dist: Vector of contact distances - contact_pos: Matrix of contact positions (one per row) - contact_normals: Matrix of contact normal vectors (one per row) + - Vector of contact distances. + - Matrix of contact positions (one per row). + - Matrix of contact normal vectors (one per row). """ sphere_support = -wp.normalize(wp.cw_mul(wp.transpose(ellipsoid_rot) @ plane_normal, ellipsoid_size)) pos = ellipsoid_pos + ellipsoid_rot @ wp.cw_mul(sphere_support, ellipsoid_size) @@ -329,19 +320,17 @@ def plane_box( """Core contact geometry calculation for plane-box collision. Args: - plane_normal: Normal vector of the plane - plane_pos: Position point on the plane - box_pos: Center position of the box - box_rot: Rotation matrix of the box - box_size: Half-extents of the box along each axis + plane_normal: Normal vector of the plane. + plane_pos: Position point on the plane. + box_pos: Center position of the box. + box_rot: Rotation matrix of the box. + box_size: Half-extents of the box along each axis. Returns: - Tuple containing: - contact_dist: Vector of contact distances (wp.inf for unpopulated contacts) - contact_pos: Matrix of contact positions (one per row) - contact_normal: contact normal vector + - Vector of contact distances (wp.inf for unpopulated contacts). + - Matrix of contact positions (one per row). + - Contact normal vector. """ - corner = wp.vec3() center_dist = wp.dot(box_pos - plane_pos, plane_normal) @@ -389,18 +378,17 @@ def sphere_cylinder( """Core contact geometry calculation for sphere-cylinder collision. Args: - sphere_pos: Center position of the sphere - sphere_radius: Radius of the sphere - cylinder_pos: Center position of the cylinder - cylinder_axis: Axis direction of the cylinder - cylinder_radius: Radius of the cylinder - cylinder_half_height: Half height of the cylinder + sphere_pos: Center position of the sphere. + sphere_radius: Radius of the sphere. + cylinder_pos: Center position of the cylinder. + cylinder_axis: Axis direction of the cylinder. + cylinder_radius: Radius of the cylinder. + cylinder_half_height: Half height of the cylinder. Returns: - Tuple containing: - contact_dist: Vector of contact distances - contact_pos: Matrix of contact positions (one per row) - contact_normals: Matrix of contact normal vectors (one per row) + - Vector of contact distances. + - Matrix of contact positions (one per row). + - Matrix of contact normal vectors (one per row). """ vec = sphere_pos - cylinder_pos x = wp.dot(vec, cylinder_axis) @@ -462,20 +450,18 @@ def plane_cylinder( """Core contact geometry calculation for plane-cylinder collision. Args: - plane_normal: Normal vector of the plane - plane_pos: Position point on the plane - cylinder_center: Center position of the cylinder - cylinder_axis: Axis direction of the cylinder - cylinder_radius: Radius of the cylinder - cylinder_half_height: Half height of the cylinder + plane_normal: Normal vector of the plane. + plane_pos: Position point on the plane. + cylinder_center: Center position of the cylinder. + cylinder_axis: Axis direction of the cylinder. + cylinder_radius: Radius of the cylinder. + cylinder_half_height: Half height of the cylinder. Returns: - Tuple containing: - contact_dist: Vector of contact distances - contact_pos: Matrix of contact positions (one per row) - contact_normals: Matrix of contact normal vectors (one per row) + - Vector of contact distances. + - Matrix of contact positions (one per row). + - Matrix of contact normal vectors (one per row). """ - # Initialize output matrices contact_dist = wp.vec4(wp.inf) contact_pos = mat43f() @@ -589,24 +575,25 @@ def box_box( box2_pos: wp.vec3, box2_rot: wp.mat33, box2_size: wp.vec3, + margin: float = 0.0, # kernel_analyzer: off ) -> Tuple[vec8f, mat83f, mat83f]: """Core contact geometry calculation for box-box collision. Args: - box1_pos: Center position of the first box - box1_rot: Rotation matrix of the first box - box1_size: Half-extents of the first box along each axis - box2_pos: Center position of the second box - box2_rot: Rotation matrix of the second box - box2_size: Half-extents of the second box along each axis + box1_pos: Center position of the first box. + box1_rot: Rotation matrix of the first box. + box1_size: Half-extents of the first box along each axis. + box2_pos: Center position of the second box. + box2_rot: Rotation matrix of the second box. + box2_size: Half-extents of the second box along each axis. + margin: Distance threshold for early contact generation (default: 0.0). + When positive, contacts are generated before boxes overlap. Returns: - Tuple containing: - contact_dist: Vector of contact distances (wp.inf for unpopulated contacts) - contact_pos: Matrix of contact positions (one per row) - contact_normals: Matrix of contact normal vectors (one per row) + - Vector of contact distances (wp.inf 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): @@ -630,7 +617,7 @@ def box_box( # Compute axis of maximum separation s_sum_3 = 3.0 * (box1_size + box2_size) - separation = wp.float32(s_sum_3[0] + s_sum_3[1] + s_sum_3[2]) + separation = wp.float32(margin + s_sum_3[0] + s_sum_3[1] + s_sum_3[2]) axis_code = wp.int32(-1) # First test: consider boxes' face normals @@ -639,7 +626,7 @@ def box_box( c2 = -wp.abs(pos12[i]) + box2_size[i] + plen1[i] - if c1 < 0.0 or c2 < 0.0: + if c1 < -margin or c2 < -margin: return contact_dist, contact_pos, contact_normals if c1 < separation: @@ -684,7 +671,7 @@ def box_box( c3 -= wp.abs(box_dist) # Early exit: no collision if separated along this axis - if c3 < 0.0: + if c3 < -margin: return contact_dist, contact_pos, contact_normals # Track minimum separation and which edge-edge pair it occurs on @@ -811,7 +798,7 @@ def box_box( n = wp.int32(0) for i in range(m): - if points[i][2] > 0.0: + if points[i][2] > margin: continue if i != n: points[n] = points[i] @@ -925,7 +912,7 @@ def box_box( c2 = lc + ld * c1 if wp.abs(c2) > s[1 - q]: continue - if (lua[2] + lub[2] * c1) * innorm > 0.0: + if (lua[2] + lub[2] * c1) * innorm > margin: continue points[n] = lua * 0.5 + c1 * lub * 0.5 @@ -969,7 +956,7 @@ def box_box( vtmp2 = points[n] - vtmp tc1 = wp.length_sq(vtmp2) - if vtmp[2] > 0 and tc1 > 0.0: + if vtmp[2] > 0 and tc1 > margin * margin: continue points[n] = 0.5 * (points[n] + vtmp) @@ -1000,7 +987,7 @@ def box_box( c1 += pu[i, 2] * innorm * pu[i, 2] * innorm - if pu[i, 2] > 0 and c1 > 0.0: + if pu[i, 2] > 0 and c1 > margin * margin: continue tmp_p = wp.vec3(pu[i, 0], pu[i, 1], 0.0) @@ -1047,19 +1034,17 @@ def sphere_box( """Core contact geometry calculation for sphere-box collision. Args: - sphere_pos: Center position of the sphere - sphere_radius: Radius of the sphere - box_pos: Center position of the box - box_rot: Rotation matrix of the box - box_size: Half-extents of the box along each axis + sphere_pos: Center position of the sphere. + sphere_radius: Radius of the sphere. + box_pos: Center position of the box. + box_rot: Rotation matrix of the box. + box_size: Half-extents of the box along each axis. Returns: - Tuple containing: - contact_dist: Vector of contact distances - contact_pos: contact positions - contact_normal: contact normal vectors + - Vector of contact distances. + - Contact positions. + - Contact normal vectors. """ - center = wp.transpose(box_rot) @ (sphere_pos - box_pos) clamped = wp.max(-box_size, wp.min(box_size, center)) @@ -1106,21 +1091,19 @@ def capsule_box( """Core contact geometry calculation for capsule-box collision. Args: - capsule_pos: Center position of the capsule - capsule_axis: Axis direction of the capsule - capsule_radius: Radius of the capsule - capsule_half_length: Half length of the capsule - box_pos: Center position of the box - box_rot: Rotation matrix of the box - box_size: Half-extents of the box along each axis + capsule_pos: Center position of the capsule. + capsule_axis: Axis direction of the capsule. + capsule_radius: Radius of the capsule. + capsule_half_length: Half length of the capsule. + box_pos: Center position of the box. + box_rot: Rotation matrix of the box. + box_size: Half-extents of the box along each axis. Returns: - Tuple containing: - contact_dist: Vector of contact distances (wp.inf for unpopulated contacts) - contact_pos: Matrix of contact positions (one per row) - contact_normals: Matrix of contact normal vectors (one per row) + - Vector of contact distances (wp.inf for unpopulated contacts). + - Matrix of contact positions (one per row). + - Matrix of contact normal vectors (one per row). """ - # Based on the mjc implementation boxmatT = wp.transpose(box_rot) pos = boxmatT @ (capsule_pos - box_pos) 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 fdd988ac..15ffdb41 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 @@ -17,7 +17,6 @@ from typing import Tuple import warp as wp -from mujoco.mjx.third_party.mujoco_warp._src import math from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact @@ -78,8 +77,8 @@ class MeshData: @wp.func def get_sdf_params( # Model: - oct_aabb: wp.array2d(dtype=wp.vec3), oct_child: wp.array(dtype=vec8i), + oct_aabb: wp.array2d(dtype=wp.vec3), oct_coeff: wp.array(dtype=vec8f), plugin: wp.array(dtype=int), plugin_attr: wp.array(dtype=wp.vec3f), @@ -225,7 +224,7 @@ def user_sdf_grad(p: wp.vec3, attr: wp.vec3, sdf_type: int) -> wp.vec3: @wp.func def find_oct( - oct_aabb: wp.array2d(dtype=wp.vec3), oct_child: wp.array(dtype=vec8i), p: wp.vec3, grad: bool + oct_child: wp.array(dtype=vec8i), oct_aabb: wp.array2d(dtype=wp.vec3), p: wp.vec3, grad: bool ) -> Tuple[int, Tuple[vec8f, vec8f, vec8f]]: stack = int(0) niter = int(100) @@ -331,7 +330,7 @@ def box_project(center: wp.vec3, half_size: wp.vec3, xyz: wp.vec3) -> Tuple[floa @wp.func def sample_volume_sdf(xyz: wp.vec3, volume_data: VolumeData) -> float: dist0, point = box_project(volume_data.center, volume_data.half_size, xyz) - node, weights = find_oct(volume_data.oct_aabb, volume_data.oct_child, point, grad=False) + node, weights = find_oct(volume_data.oct_child, volume_data.oct_aabb, point, grad=False) return dist0 + wp.dot(weights[0], volume_data.oct_coeff[node]) @@ -348,7 +347,7 @@ def sample_volume_grad(xyz: wp.vec3, volume_data: VolumeData) -> wp.vec3: grad_y = (sample_volume_sdf(xyz + dy, volume_data) - f) / h grad_z = (sample_volume_sdf(xyz + dz, volume_data) - f) / h return wp.vec3(grad_x, grad_y, grad_z) - node, weights = find_oct(volume_data.oct_aabb, volume_data.oct_child, point, grad=True) + node, weights = find_oct(volume_data.oct_child, volume_data.oct_aabb, point, grad=True) grad_x = wp.dot(weights[0], volume_data.oct_coeff[node]) grad_y = wp.dot(weights[1], volume_data.oct_coeff[node]) grad_z = wp.dot(weights[2], volume_data.oct_coeff[node]) @@ -371,8 +370,8 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: Volume dist = ray_mesh( mesh_data.nmeshface, mesh_data.mesh_vertadr, - mesh_data.mesh_vert, mesh_data.mesh_faceadr, + mesh_data.mesh_vert, mesh_data.mesh_face, mesh_data.data_id, mesh_data.pos, @@ -384,8 +383,8 @@ def sdf(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: Volume return -ray_mesh( mesh_data.nmeshface, mesh_data.mesh_vertadr, - mesh_data.mesh_vert, mesh_data.mesh_faceadr, + mesh_data.mesh_vert, mesh_data.mesh_face, mesh_data.data_id, mesh_data.pos, @@ -420,8 +419,8 @@ def sdf_grad(type: int, p: wp.vec3, attr: wp.vec3, sdf_type: int, volume_data: V dist = ray_mesh( mesh_data.nmeshface, mesh_data.mesh_vertadr, - mesh_data.mesh_vert, mesh_data.mesh_faceadr, + mesh_data.mesh_vert, mesh_data.mesh_face, mesh_data.data_id, mesh_data.pos, @@ -620,6 +619,9 @@ def gradient_descent( def _sdf_narrowphase( # Model: nmeshface: int, + oct_child: wp.array(dtype=vec8i), + oct_aabb: wp.array2d(dtype=wp.vec3), + oct_coeff: wp.array(dtype=vec8f), geom_type: wp.array(dtype=int), geom_condim: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), @@ -628,21 +630,16 @@ def _sdf_narrowphase( geom_solref: wp.array2d(dtype=wp.vec2), geom_solimp: wp.array2d(dtype=vec5), geom_size: wp.array2d(dtype=wp.vec3), - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_friction: wp.array2d(dtype=wp.vec3), geom_margin: wp.array2d(dtype=float), geom_gap: wp.array2d(dtype=float), - 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_faceadr: wp.array(dtype=int), - mesh_face: wp.array(dtype=wp.vec3i), mesh_graphadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_face: wp.array(dtype=wp.vec3i), mesh_graph: wp.array(dtype=int), mesh_polynum: wp.array(dtype=int), mesh_polyadr: wp.array(dtype=int), @@ -653,9 +650,11 @@ def _sdf_narrowphase( mesh_polymapadr: wp.array(dtype=int), mesh_polymapnum: wp.array(dtype=int), mesh_polymap: wp.array(dtype=int), - oct_aabb: wp.array2d(dtype=wp.vec3), - oct_child: wp.array(dtype=vec8i), - oct_coeff: wp.array(dtype=vec8f), + 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), @@ -663,23 +662,21 @@ def _sdf_narrowphase( pair_margin: wp.array2d(dtype=float), pair_gap: wp.array2d(dtype=float), pair_friction: wp.array2d(dtype=vec5), - # In: plugin: wp.array(dtype=int), plugin_attr: wp.array(dtype=wp.vec3f), geom_plugin_index: wp.array(dtype=int), # Data in: - naconmax_in: int, geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), + naconmax_in: int, collision_pair_in: wp.array(dtype=wp.vec2i), - collision_pairid_in: wp.array(dtype=int), + 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: - 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), @@ -691,6 +688,9 @@ def _sdf_narrowphase( contact_dim_out: wp.array(dtype=int), contact_geom_out: wp.array(dtype=wp.vec2i), contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + nacon_out: wp.array(dtype=int), ): i, contact_tid = wp.tid() if i >= sdf_initpoints: @@ -724,18 +724,21 @@ def _sdf_narrowphase( contact_tid, worldid, ) + + geom_size_id = worldid % geom_size.shape[0] + aabb_id = worldid % geom_aabb.shape[0] + g1 = geoms[0] type1 = geom_type[g1] - geom1_dataid = geom_dataid[g1] geom1 = geom( type1, geom1_dataid, - geom_size[worldid, g1], + geom_size[geom_size_id, g1], mesh_vertadr, mesh_vertnum, - mesh_vert, mesh_graphadr, + mesh_vert, mesh_graph, mesh_polynum, mesh_polyadr, @@ -754,11 +757,11 @@ def _sdf_narrowphase( geom2 = geom( type2, geom2_dataid, - geom_size[worldid, g2], + geom_size[geom_size_id, g2], mesh_vertadr, mesh_vertnum, - mesh_vert, mesh_graphadr, + mesh_vert, mesh_graph, mesh_polynum, mesh_polyadr, @@ -777,12 +780,12 @@ def _sdf_narrowphase( g1_to_g2_rot = wp.transpose(geom1.rot) * geom2.rot g1_to_g2_pos = wp.transpose(geom1.rot) * (geom2.pos - geom1.pos) - aabb_pos = geom_aabb[g1, 0] - aabb_size = geom_aabb[g1, 1] + aabb_pos = geom_aabb[aabb_id, g1, 0] + aabb_size = geom_aabb[aabb_id, g1, 1] identity = wp.identity(3, dtype=float) aabb1 = transform_aabb(aabb_pos, aabb_size, wp.vec3(0.0), identity) - aabb_pos = geom_aabb[g2, 0] - aabb_size = geom_aabb[g2, 1] + aabb_pos = geom_aabb[aabb_id, g2, 0] + aabb_size = geom_aabb[aabb_id, g2, 1] aabb2 = transform_aabb(aabb_pos, aabb_size, g1_to_g2_pos, g1_to_g2_rot) aabb_intersection = AABB() aabb_intersection.min = wp.max(aabb1.min, aabb2.min) @@ -794,11 +797,11 @@ def _sdf_narrowphase( rot1 = geom1.rot attr1, g1_plugin_id, volume_data1, mesh_data1 = get_sdf_params( - oct_aabb, oct_child, oct_coeff, plugin, plugin_attr, type1, geom1.size, g1_plugin, geom_dataid[g1] + oct_child, oct_aabb, oct_coeff, plugin, plugin_attr, type1, geom1.size, g1_plugin, geom_dataid[g1] ) attr2, g2_plugin_id, volume_data2, mesh_data2 = get_sdf_params( - oct_aabb, oct_child, oct_coeff, plugin, plugin_attr, type2, geom2.size, g2_plugin, geom_dataid[g2] + oct_child, oct_aabb, oct_coeff, plugin, plugin_attr, type2, geom2.size, g2_plugin, geom_dataid[g2] ) mesh_data1.nmeshface = nmeshface @@ -851,6 +854,7 @@ def _sdf_narrowphase( ) write_contact( naconmax_in, + 0, dist, pos, make_frame(n), @@ -862,8 +866,8 @@ def _sdf_narrowphase( solreffriction, solimp, geoms, + collision_pairid_in[contact_tid], worldid, - nacon_out, contact_dist_out, contact_pos_out, contact_frame_out, @@ -875,6 +879,9 @@ def _sdf_narrowphase( contact_dim_out, contact_geom_out, contact_worldid_out, + contact_type_out, + contact_geomcollisionid_out, + nacon_out, ) @@ -885,6 +892,9 @@ def sdf_narrowphase(m: Model, d: Data): dim=(m.opt.sdf_initpoints, d.naconmax), inputs=[ m.nmeshface, + m.oct_child, + m.oct_aabb, + m.oct_coeff, m.geom_type, m.geom_condim, m.geom_dataid, @@ -897,17 +907,12 @@ def sdf_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_faceadr, - m.mesh_face, m.mesh_graphadr, + m.mesh_vert, + m.mesh_face, m.mesh_graph, m.mesh_polynum, m.mesh_polyadr, @@ -918,9 +923,11 @@ def sdf_narrowphase(m: Model, d: Data): m.mesh_polymapadr, m.mesh_polymapnum, m.mesh_polymap, - m.oct_aabb, - m.oct_child, - m.oct_coeff, + m.hfield_size, + m.hfield_nrow, + m.hfield_ncol, + m.hfield_adr, + m.hfield_data, m.pair_dim, m.pair_solref, m.pair_solreffriction, @@ -931,9 +938,9 @@ def sdf_narrowphase(m: Model, d: Data): m.plugin, m.plugin_attr, m.geom_plugin_index, - d.naconmax, d.geom_xpos, d.geom_xmat, + d.naconmax, d.collision_pair, d.collision_pairid, d.collision_worldid, @@ -942,7 +949,6 @@ def sdf_narrowphase(m: Model, d: Data): m.opt.sdf_iterations, ], outputs=[ - d.nacon, d.contact.dist, d.contact.pos, d.contact.frame, @@ -954,5 +960,8 @@ def sdf_narrowphase(m: Model, d: Data): d.contact.dim, d.contact.geom, d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + d.nacon, ], ) diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py index ae9f8523..f0fc7a30 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/constraint.py @@ -19,6 +19,7 @@ from mujoco.mjx.third_party.mujoco_warp._src import math from mujoco.mjx.third_party.mujoco_warp._src import support from mujoco.mjx.third_party.mujoco_warp._src import types from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType +from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType from mujoco.mjx.third_party.mujoco_warp._src.types import vec5 from mujoco.mjx.third_party.mujoco_warp._src.types import vec11 from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope @@ -27,16 +28,16 @@ wp.set_module_options({"enable_backward": False}) @wp.kernel -def zero_constraint_counts( +def _zero_constraint_counts( # Data out: ne_out: wp.array(dtype=int), + 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), - nf_out: wp.array(dtype=int), - nl_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), ): worldid = wp.tid() @@ -141,7 +142,6 @@ def _efc_equality_connect( eq_data: wp.array2d(dtype=vec11), eq_connect_adr: wp.array(dtype=int), # Data in: - njmax_in: int, qvel_in: wp.array2d(dtype=float), eq_active_in: wp.array2d(dtype=bool), xpos_in: wp.array2d(dtype=wp.vec3), @@ -149,10 +149,10 @@ def _efc_equality_connect( site_xpos_in: wp.array2d(dtype=wp.vec3), subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, # In: refsafe_in: int, # Data out: - ne_connect_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), @@ -163,9 +163,9 @@ 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() eqid = eq_connect_adr[eqconnectid] @@ -178,7 +178,7 @@ def _efc_equality_connect( if efcid + 3 >= njmax_in: return - data = eq_data[worldid, eqid] + data = eq_data[worldid % eq_data.shape[0], eqid] anchor1 = wp.vec3f(data[0], data[1], data[2]) anchor2 = wp.vec3f(data[3], data[4], data[5]) @@ -231,18 +231,20 @@ def _efc_equality_connect( efc_J_out[worldid, efcid + 2, dofid] = j1mj2[2] Jqvel += j1mj2 * qvel_in[worldid, dofid] - invweight = body_invweight0[worldid, body1id][0] + body_invweight0[worldid, body2id][0] + body_invweight0_id = worldid % body_invweight0.shape[0] + invweight = body_invweight0[body_invweight0_id, body1id][0] + body_invweight0[body_invweight0_id, body2id][0] pos_imp = wp.length(pos) - solref = eq_solref[worldid, eqid] - solimp = eq_solimp[worldid, eqid] + solref = eq_solref[worldid % eq_solref.shape[0], eqid] + solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] for i in range(3): efcidi = efcid + i _update_efc_row( worldid, - opt_timestep[worldid], + timestep, refsafe_in, efcidi, pos[i], @@ -282,14 +284,13 @@ def _efc_equality_joint( eq_data: wp.array2d(dtype=vec11), eq_jnt_adr: wp.array(dtype=int), # Data in: - njmax_in: int, qpos_in: wp.array2d(dtype=float), qvel_in: wp.array2d(dtype=float), eq_active_in: wp.array2d(dtype=bool), + njmax_in: int, # In: refsafe_in: int, # Data out: - ne_jnt_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), @@ -300,6 +301,7 @@ 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] @@ -318,10 +320,12 @@ def _efc_equality_joint( jntid_1 = eq_obj1id[eqid] jntid_2 = eq_obj2id[eqid] - data = eq_data[worldid, eqid] + data = eq_data[worldid % eq_data.shape[0], eqid] dofadr1 = jnt_dofadr[jntid_1] qposadr1 = jnt_qposadr[jntid_1] efc_J_out[worldid, efcid, dofadr1] = 1.0 + qpos0_id = worldid % qpos0.shape[0] + dof_invweight0_id = worldid % dof_invweight0.shape[0] if jntid_2 > -1: # Two joint constraint @@ -333,28 +337,28 @@ def _efc_equality_joint( rhs = data[0] + dif * (data[1] + dif * (data[2] + dif * (data[3] + dif * data[4]))) deriv_2 = data[1] + dif * (2.0 * data[2] + dif * (3.0 * data[3] + dif * 4.0 * data[4])) - pos = qpos_in[worldid, qposadr1] - qpos0[worldid, qposadr1] - rhs + pos = qpos_in[worldid, qposadr1] - qpos0[qpos0_id, qposadr1] - rhs Jqvel = qvel_in[worldid, dofadr1] - qvel_in[worldid, dofadr2] * deriv_2 - invweight = dof_invweight0[worldid, dofadr1] + dof_invweight0[worldid, dofadr2] + invweight = dof_invweight0[dof_invweight0_id, dofadr1] + dof_invweight0[dof_invweight0_id, dofadr2] efc_J_out[worldid, efcid, dofadr2] = -deriv_2 else: # Single joint constraint - pos = qpos_in[worldid, qposadr1] - qpos0[worldid, qposadr1] - data[0] + pos = qpos_in[worldid, qposadr1] - qpos0[qpos0_id, qposadr1] - data[0] Jqvel = qvel_in[worldid, dofadr1] - invweight = dof_invweight0[worldid, dofadr1] + invweight = dof_invweight0[dof_invweight0_id, dofadr1] # Update constraint parameters _update_efc_row( worldid, - opt_timestep[worldid], + opt_timestep[worldid % opt_timestep.shape[0]], refsafe_in, efcid, pos, pos, invweight, - eq_solref[worldid, eqid], - eq_solimp[worldid, eqid], + eq_solref[worldid % eq_solref.shape[0], eqid], + eq_solimp[worldid % eq_solimp.shape[0], eqid], 0.0, Jqvel, 0.0, @@ -381,19 +385,18 @@ def _efc_equality_tendon( eq_solref: wp.array2d(dtype=wp.vec2), eq_solimp: wp.array2d(dtype=vec5), eq_data: wp.array2d(dtype=vec11), - eq_ten_adr: wp.array(dtype=int), tendon_length0: wp.array2d(dtype=float), tendon_invweight0: wp.array2d(dtype=float), + eq_ten_adr: wp.array(dtype=int), # Data in: - njmax_in: int, qvel_in: wp.array2d(dtype=float), eq_active_in: wp.array2d(dtype=bool), - ten_length_in: wp.array2d(dtype=float), ten_J_in: wp.array3d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + njmax_in: int, # In: refsafe_in: int, # Data out: - ne_ten_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), @@ -404,6 +407,7 @@ 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] @@ -420,16 +424,18 @@ def _efc_equality_tendon( obj1id = eq_obj1id[eqid] obj2id = eq_obj2id[eqid] - data = eq_data[worldid, eqid] - solref = eq_solref[worldid, eqid] - solimp = eq_solimp[worldid, eqid] - pos1 = ten_length_in[worldid, obj1id] - tendon_length0[worldid, obj1id] + data = eq_data[worldid % eq_data.shape[0], eqid] + solref = eq_solref[worldid % eq_solref.shape[0], eqid] + solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] + tendon_length0_id = worldid % tendon_length0.shape[0] + tendon_invweight0_id = worldid % tendon_invweight0.shape[0] + pos1 = ten_length_in[worldid, obj1id] - tendon_length0[tendon_length0_id, obj1id] jac1 = ten_J_in[worldid, obj1id] if obj2id > -1: - invweight = tendon_invweight0[worldid, obj1id] + tendon_invweight0[worldid, obj2id] + invweight = tendon_invweight0[tendon_invweight0_id, obj1id] + tendon_invweight0[tendon_invweight0_id, obj2id] - pos2 = ten_length_in[worldid, obj2id] - tendon_length0[worldid, obj2id] + pos2 = ten_length_in[worldid, obj2id] - tendon_length0[tendon_length0_id, obj2id] jac2 = ten_J_in[worldid, obj2id] dif = pos2 @@ -440,7 +446,7 @@ def _efc_equality_tendon( pos = pos1 - (data[0] + data[1] * dif + data[2] * dif2 + data[3] * dif3 + data[4] * dif4) deriv = data[1] + 2.0 * data[2] * dif + 3.0 * data[3] * dif2 + 4.0 * data[4] * dif3 else: - invweight = tendon_invweight0[worldid, obj1id] + invweight = tendon_invweight0[tendon_invweight0_id, obj1id] pos = pos1 - data[0] deriv = 0.0 @@ -455,7 +461,7 @@ def _efc_equality_tendon( _update_efc_row( worldid, - opt_timestep[worldid], + opt_timestep[worldid % opt_timestep.shape[0]], refsafe_in, efcid, pos, @@ -484,13 +490,13 @@ def _efc_friction_dof( # Model: nv: int, opt_timestep: wp.array(dtype=float), - dof_invweight0: wp.array2d(dtype=float), - dof_frictionloss: wp.array2d(dtype=float), - dof_solimp: wp.array2d(dtype=vec5), dof_solref: wp.array2d(dtype=wp.vec2), + dof_solimp: wp.array2d(dtype=vec5), + dof_frictionloss: wp.array2d(dtype=float), + dof_invweight0: wp.array2d(dtype=float), # Data in: - njmax_in: int, qvel_in: wp.array2d(dtype=float), + njmax_in: int, # In: refsafe_in: int, # Data out: @@ -508,7 +514,9 @@ def _efc_friction_dof( ): worldid, dofid = wp.tid() - if dof_frictionloss[worldid, dofid] <= 0.0: + dof_frictionloss_id = worldid % dof_frictionloss.shape[0] + + if dof_frictionloss[dof_frictionloss_id, dofid] <= 0.0: return wp.atomic_add(nf_out, worldid, 1) @@ -523,19 +531,22 @@ def _efc_friction_dof( efc_J_out[worldid, efcid, dofid] = 1.0 Jqvel = qvel_in[worldid, dofid] + dof_invweight0_id = worldid % dof_invweight0.shape[0] + dof_solref_id = worldid % dof_solref.shape[0] + dof_solimp_id = worldid % dof_solimp.shape[0] _update_efc_row( worldid, - opt_timestep[worldid], + opt_timestep[worldid % opt_timestep.shape[0]], refsafe_in, efcid, 0.0, 0.0, - dof_invweight0[worldid, dofid], - dof_solref[worldid, dofid], - dof_solimp[worldid, dofid], + dof_invweight0[dof_invweight0_id, dofid], + dof_solref[dof_solref_id, dofid], + dof_solimp[dof_solimp_id, dofid], 0.0, Jqvel, - dof_frictionloss[worldid, dofid], + dof_frictionloss[dof_frictionloss_id, dofid], ConstraintType.FRICTION_DOF, dofid, efc_type_out, @@ -559,9 +570,9 @@ def _efc_friction_tendon( tendon_frictionloss: wp.array2d(dtype=float), tendon_invweight0: wp.array2d(dtype=float), # Data in: - njmax_in: int, qvel_in: wp.array2d(dtype=float), ten_J_in: wp.array3d(dtype=float), + njmax_in: int, # In: refsafe_in: int, # Data out: @@ -579,7 +590,9 @@ def _efc_friction_tendon( ): worldid, tenid = wp.tid() - frictionloss = tendon_frictionloss[worldid, tenid] + tendon_frictionloss_id = worldid % tendon_frictionloss.shape[0] + + frictionloss = tendon_frictionloss[tendon_frictionloss_id, tenid] if frictionloss <= 0.0: return @@ -597,16 +610,19 @@ def _efc_friction_tendon( efc_J_out[worldid, efcid, i] = J Jqvel += J * qvel_in[worldid, i] + tendon_invweight0_id = worldid % tendon_invweight0.shape[0] + tendon_solref_fri_id = worldid % tendon_solref_fri.shape[0] + tendon_solimp_fri_id = worldid % tendon_solimp_fri.shape[0] _update_efc_row( worldid, - opt_timestep[worldid], + opt_timestep[worldid % opt_timestep.shape[0]], refsafe_in, efcid, 0.0, 0.0, - tendon_invweight0[worldid, tenid], - tendon_solref_fri[worldid, tenid], - tendon_solimp_fri[worldid, tenid], + tendon_invweight0[tendon_invweight0_id, tenid], + tendon_solref_fri[tendon_solref_fri_id, tenid], + tendon_solimp_fri[tendon_solimp_fri_id, tenid], 0.0, Jqvel, frictionloss, @@ -643,7 +659,6 @@ def _efc_equality_weld( eq_data: wp.array2d(dtype=vec11), eq_wld_adr: wp.array(dtype=int), # Data in: - njmax_in: int, qvel_in: wp.array2d(dtype=float), eq_active_in: wp.array2d(dtype=bool), xpos_in: wp.array2d(dtype=wp.vec3), @@ -652,10 +667,10 @@ def _efc_equality_weld( site_xpos_in: wp.array2d(dtype=wp.vec3), subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, # In: refsafe_in: int, # Data out: - ne_weld_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), @@ -666,6 +681,7 @@ 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] @@ -684,7 +700,7 @@ def _efc_equality_weld( obj1id = eq_obj1id[eqid] obj2id = eq_obj2id[eqid] - data = eq_data[worldid, eqid] + data = eq_data[worldid % eq_data.shape[0], eqid] anchor1 = wp.vec3(data[0], data[1], data[2]) anchor2 = wp.vec3(data[3], data[4], data[5]) relpose = wp.quat(data[6], data[7], data[8], data[9]) @@ -697,8 +713,9 @@ def _efc_equality_weld( pos1 = site_xpos_in[worldid, obj1id] pos2 = site_xpos_in[worldid, obj2id] - quat = math.mul_quat(xquat_in[worldid, body1id], site_quat[worldid, obj1id]) - quat1 = math.quat_inv(math.mul_quat(xquat_in[worldid, body2id], site_quat[worldid, obj2id])) + site_quat_id = worldid % site_quat.shape[0] + quat = math.mul_quat(xquat_in[worldid, body1id], site_quat[site_quat_id, obj1id]) + quat1 = math.quat_inv(math.mul_quat(xquat_in[worldid, body2id], site_quat[site_quat_id, obj2id])) else: body1id = obj1id @@ -757,14 +774,15 @@ def _efc_equality_weld( crotq = math.mul_quat(quat1, quat) # copy axis components crot = wp.vec3(crotq[1], crotq[2], crotq[3]) * torquescale - invweight_t = body_invweight0[worldid, body1id][0] + body_invweight0[worldid, body2id][0] + body_invweight0_id = worldid % body_invweight0.shape[0] + invweight_t = body_invweight0[body_invweight0_id, body1id][0] + body_invweight0[body_invweight0_id, body2id][0] pos_imp = wp.sqrt(wp.length_sq(cpos) + wp.length_sq(crot)) - solref = eq_solref[worldid, eqid] - solimp = eq_solimp[worldid, eqid] + solref = eq_solref[worldid % eq_solref.shape[0], eqid] + solimp = eq_solimp[worldid % eq_solimp.shape[0], eqid] - timestep = opt_timestep[worldid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] for i in range(3): _update_efc_row( @@ -792,7 +810,7 @@ def _efc_equality_weld( efc_frictionloss_out, ) - invweight_r = body_invweight0[worldid, body1id][1] + body_invweight0[worldid, body2id][1] + invweight_r = body_invweight0[body_invweight0_id, body1id][1] + body_invweight0[body_invweight0_id, body2id][1] for i in range(3): _update_efc_row( @@ -832,12 +850,12 @@ def _efc_limit_slide_hinge( jnt_solimp: wp.array2d(dtype=vec5), jnt_range: wp.array2d(dtype=wp.vec2), jnt_margin: wp.array2d(dtype=float), - jnt_limited_slide_hinge_adr: wp.array(dtype=int), dof_invweight0: wp.array2d(dtype=float), + jnt_limited_slide_hinge_adr: wp.array(dtype=int), # Data in: - njmax_in: int, qpos_in: wp.array2d(dtype=float), qvel_in: wp.array2d(dtype=float), + njmax_in: int, # In: refsafe_in: int, # Data out: @@ -855,10 +873,12 @@ def _efc_limit_slide_hinge( ): worldid, jntlimitedid = wp.tid() jntid = jnt_limited_slide_hinge_adr[jntlimitedid] - jntrange = jnt_range[worldid, jntid] + jnt_range_id = worldid % jnt_range.shape[0] + jntrange = jnt_range[jnt_range_id, jntid] qpos = qpos_in[worldid, jnt_qposadr[jntid]] - jntmargin = jnt_margin[worldid, jntid] + jnt_margin_id = worldid % jnt_margin.shape[0] + jntmargin = jnt_margin[jnt_margin_id, jntid] dist_min, dist_max = qpos - jntrange[0], jntrange[1] - qpos pos = wp.min(dist_min, dist_max) - jntmargin active = pos < 0 @@ -879,16 +899,19 @@ def _efc_limit_slide_hinge( efc_J_out[worldid, efcid, dofadr] = J Jqvel = J * qvel_in[worldid, dofadr] + dof_invweight0_id = worldid % dof_invweight0.shape[0] + jnt_solref_id = worldid % jnt_solref.shape[0] + jnt_solimp_id = worldid % jnt_solimp.shape[0] _update_efc_row( worldid, - opt_timestep[worldid], + opt_timestep[worldid % opt_timestep.shape[0]], refsafe_in, efcid, pos, pos, - dof_invweight0[worldid, dofadr], - jnt_solref[worldid, jntid], - jnt_solimp[worldid, jntid], + dof_invweight0[dof_invweight0_id, dofadr], + jnt_solref[jnt_solref_id, jntid], + jnt_solimp[jnt_solimp_id, jntid], jntmargin, Jqvel, 0.0, @@ -916,12 +939,12 @@ def _efc_limit_ball( jnt_solimp: wp.array2d(dtype=vec5), jnt_range: wp.array2d(dtype=wp.vec2), jnt_margin: wp.array2d(dtype=float), - jnt_limited_ball_adr: wp.array(dtype=int), dof_invweight0: wp.array2d(dtype=float), + jnt_limited_ball_adr: wp.array(dtype=int), # Data in: - njmax_in: int, qpos_in: wp.array2d(dtype=float), qvel_in: wp.array2d(dtype=float), + njmax_in: int, # In: refsafe_in: int, # Data out: @@ -945,9 +968,11 @@ def _efc_limit_ball( jnt_quat = wp.quat(qpos[qposadr + 0], qpos[qposadr + 1], qpos[qposadr + 2], qpos[qposadr + 3]) jnt_quat = wp.normalize(jnt_quat) axis_angle = math.quat_to_vel(jnt_quat) - jntrange = jnt_range[worldid, jntid] + jnt_range_id = worldid % jnt_range.shape[0] + jntrange = jnt_range[jnt_range_id, jntid] axis, angle = math.normalize_with_norm(axis_angle) - jntmargin = jnt_margin[worldid, jntid] + jnt_margin_id = worldid % jnt_margin.shape[0] + jntmargin = jnt_margin[jnt_margin_id, jntid] pos = wp.max(jntrange[0], jntrange[1]) - angle - jntmargin active = pos < 0 @@ -972,16 +997,19 @@ def _efc_limit_ball( Jqvel -= axis[1] * qvel_in[worldid, dofadr + 1] Jqvel -= axis[2] * qvel_in[worldid, dofadr + 2] + dof_invweight0_id = worldid % dof_invweight0.shape[0] + jnt_solref_id = worldid % jnt_solref.shape[0] + jnt_solimp_id = worldid % jnt_solimp.shape[0] _update_efc_row( worldid, - opt_timestep[worldid], + opt_timestep[worldid % opt_timestep.shape[0]], refsafe_in, efcid, pos, pos, - dof_invweight0[worldid, dofadr], - jnt_solref[worldid, jntid], - jnt_solimp[worldid, jntid], + dof_invweight0[dof_invweight0_id, dofadr], + jnt_solref[jnt_solref_id, jntid], + jnt_solimp[jnt_solimp_id, jntid], jntmargin, Jqvel, 0.0, @@ -1006,19 +1034,19 @@ def _efc_limit_tendon( jnt_dofadr: wp.array(dtype=int), tendon_adr: wp.array(dtype=int), tendon_num: wp.array(dtype=int), - tendon_limited_adr: wp.array(dtype=int), tendon_solref_lim: wp.array2d(dtype=wp.vec2), tendon_solimp_lim: wp.array2d(dtype=vec5), tendon_range: wp.array2d(dtype=wp.vec2), tendon_margin: wp.array2d(dtype=float), tendon_invweight0: wp.array2d(dtype=float), - wrap_objid: wp.array(dtype=int), wrap_type: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), + tendon_limited_adr: wp.array(dtype=int), # Data in: - njmax_in: int, qvel_in: wp.array2d(dtype=float), - ten_length_in: wp.array2d(dtype=float), ten_J_in: wp.array3d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + njmax_in: int, # In: refsafe_in: int, # Data out: @@ -1037,10 +1065,12 @@ def _efc_limit_tendon( worldid, tenlimitedid = wp.tid() tenid = tendon_limited_adr[tenlimitedid] - tenrange = tendon_range[worldid, tenid] + tendon_range_id = worldid % tendon_range.shape[0] + tenrange = tendon_range[tendon_range_id, tenid] length = ten_length_in[worldid, tenid] dist_min, dist_max = length - tenrange[0], tenrange[1] - length - tenmargin = tendon_margin[worldid, tenid] + tendon_margin_id = worldid % tendon_margin.shape[0] + tenmargin = tendon_margin[tendon_margin_id, tenid] pos = wp.min(dist_min, dist_max) - tenmargin active = pos < 0 @@ -1071,16 +1101,19 @@ def _efc_limit_tendon( efc_J_out[worldid, efcid, i] = J Jqvel += J * qvel_in[worldid, i] + tendon_invweight0_id = worldid % tendon_invweight0.shape[0] + tendon_solref_lim_id = worldid % tendon_solref_lim.shape[0] + tendon_solimp_lim_id = worldid % tendon_solimp_lim.shape[0] _update_efc_row( worldid, - opt_timestep[worldid], + opt_timestep[worldid % opt_timestep.shape[0]], refsafe_in, efcid, pos, pos, - tendon_invweight0[worldid, tenid], - tendon_solref_lim[worldid, tenid], - tendon_solimp_lim[worldid, tenid], + tendon_invweight0[tendon_invweight0_id, tenid], + tendon_solref_lim[tendon_solref_lim_id, tenid], + tendon_solimp_lim[tendon_solimp_lim_id, tenid], tenmargin, Jqvel, 0.0, @@ -1109,11 +1142,11 @@ def _efc_contact_pyramidal( dof_bodyid: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), # Data in: - njmax_in: int, - nacon_in: wp.array(dtype=int), qvel_in: wp.array2d(dtype=float), subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, + nacon_in: wp.array(dtype=int), # In: refsafe_in: int, dist_in: wp.array(dtype=float), @@ -1126,6 +1159,7 @@ def _efc_contact_pyramidal( friction_in: wp.array(dtype=vec5), solref_in: wp.array(dtype=wp.vec2), solimp_in: wp.array(dtype=vec5), + type_in: wp.array(dtype=int), # Data out: nefc_out: wp.array(dtype=int), contact_efc_address_out: wp.array2d(dtype=int), @@ -1144,6 +1178,9 @@ def _efc_contact_pyramidal( if conid >= nacon_in[0]: return + if not type_in[conid] & ContactType.CONSTRAINT: + return + condim = condim_in[conid] if condim == 1 and dimid > 0: @@ -1163,8 +1200,9 @@ def _efc_contact_pyramidal( contact_efc_address_out[conid, dimid] = -1 return - timestep = opt_timestep[worldid] - impratio = opt_impratio[worldid] + opt_timestep_id = worldid % opt_timestep.shape[0] + timestep = opt_timestep[opt_timestep_id] + impratio = opt_impratio[opt_timestep_id] contact_efc_address_out[conid, dimid] = efcid geom = geom_in[conid] @@ -1175,7 +1213,8 @@ def _efc_contact_pyramidal( frame = frame_in[conid] # pyramidal has common invweight across all edges - invweight = body_invweight0[worldid, body1][0] + body_invweight0[worldid, body2][0] + body_invweight0_id = worldid % body_invweight0.shape[0] + invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] if condim > 1: dimid2 = dimid / 2 + 1 @@ -1274,11 +1313,11 @@ def _efc_contact_elliptic( dof_bodyid: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), # Data in: - njmax_in: int, - nacon_in: wp.array(dtype=int), qvel_in: wp.array2d(dtype=float), subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, + nacon_in: wp.array(dtype=int), # In: refsafe_in: int, dist_in: wp.array(dtype=float), @@ -1292,6 +1331,7 @@ def _efc_contact_elliptic( solref_in: wp.array(dtype=wp.vec2), solreffriction_in: wp.array(dtype=wp.vec2), solimp_in: wp.array(dtype=vec5), + type_in: wp.array(dtype=int), # Data out: nefc_out: wp.array(dtype=int), contact_efc_address_out: wp.array2d(dtype=int), @@ -1310,6 +1350,9 @@ def _efc_contact_elliptic( if conid >= nacon_in[0]: return + if not type_in[conid] & ContactType.CONSTRAINT: + return + condim = condim_in[conid] if dimid > condim - 1: @@ -1327,8 +1370,9 @@ def _efc_contact_elliptic( contact_efc_address_out[conid, dimid] = -1 return - timestep = opt_timestep[worldid] - impratio = opt_impratio[worldid] + opt_timestep_id = worldid % opt_timestep.shape[0] + timestep = opt_timestep[opt_timestep_id] + impratio = opt_impratio[opt_timestep_id] contact_efc_address_out[conid, dimid] = efcid geom = geom_in[conid] @@ -1375,7 +1419,8 @@ def _efc_contact_elliptic( efc_J_out[worldid, efcid, i] = J Jqvel += J * qvel_in[worldid, i] - invweight = body_invweight0[worldid, body1][0] + body_invweight0[worldid, body2][0] + body_invweight0_id = worldid % body_invweight0.shape[0] + invweight = body_invweight0[body_invweight0_id, body1][0] + body_invweight0[body_invweight0_id, body2][0] ref = solref_in[conid] pos_aref = pos @@ -1448,11 +1493,10 @@ def _num_equality( @event_scope def make_constraint(m: types.Model, d: types.Data): """Creates constraint jacobians and other supporting data.""" - wp.launch( - zero_constraint_counts, + _zero_constraint_counts, dim=d.nworld, - inputs=[d.ne, d.ne_connect, d.ne_weld, d.ne_jnt, d.ne_ten, d.nf, d.nl, d.nefc], + inputs=[d.ne, d.nf, d.nl, d.nefc, d.ne_connect, d.ne_weld, d.ne_jnt, d.ne_ten], ) if not (m.opt.disableflags & types.DisableBit.CONSTRAINT): @@ -1478,7 +1522,6 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solimp, m.eq_data, m.eq_connect_adr, - d.njmax, d.qvel, d.eq_active, d.xpos, @@ -1486,10 +1529,10 @@ def make_constraint(m: types.Model, d: types.Data): d.site_xpos, d.subtree_com, d.cdof, + d.njmax, refsafe, ], outputs=[ - d.ne_connect, d.nefc, d.efc.type, d.efc.id, @@ -1500,6 +1543,7 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, + d.ne_connect, ], ) wp.launch( @@ -1522,7 +1566,6 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solimp, m.eq_data, m.eq_wld_adr, - d.njmax, d.qvel, d.eq_active, d.xpos, @@ -1531,10 +1574,10 @@ def make_constraint(m: types.Model, d: types.Data): d.site_xpos, d.subtree_com, d.cdof, + d.njmax, refsafe, ], outputs=[ - d.ne_weld, d.nefc, d.efc.type, d.efc.id, @@ -1545,6 +1588,7 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, + d.ne_weld, ], ) wp.launch( @@ -1563,14 +1607,13 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solimp, m.eq_data, m.eq_jnt_adr, - d.njmax, d.qpos, d.qvel, d.eq_active, + d.njmax, refsafe, ], outputs=[ - d.ne_jnt, d.nefc, d.efc.type, d.efc.id, @@ -1581,6 +1624,7 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, + d.ne_jnt, ], ) wp.launch( @@ -1594,18 +1638,17 @@ def make_constraint(m: types.Model, d: types.Data): m.eq_solref, m.eq_solimp, m.eq_data, - m.eq_ten_adr, m.tendon_length0, m.tendon_invweight0, - d.njmax, + m.eq_ten_adr, d.qvel, d.eq_active, - d.ten_length, d.ten_J, + d.ten_length, + d.njmax, refsafe, ], outputs=[ - d.ne_ten, d.nefc, d.efc.type, d.efc.id, @@ -1616,6 +1659,7 @@ def make_constraint(m: types.Model, d: types.Data): d.efc.vel, d.efc.aref, d.efc.frictionloss, + d.ne_ten, ], ) @@ -1633,12 +1677,12 @@ def make_constraint(m: types.Model, d: types.Data): inputs=[ m.nv, m.opt.timestep, - m.dof_invweight0, - m.dof_frictionloss, - m.dof_solimp, m.dof_solref, - d.njmax, + m.dof_solimp, + m.dof_frictionloss, + m.dof_invweight0, d.qvel, + d.njmax, refsafe, ], outputs=[ @@ -1666,9 +1710,9 @@ def make_constraint(m: types.Model, d: types.Data): m.tendon_solimp_fri, m.tendon_frictionloss, m.tendon_invweight0, - d.njmax, d.qvel, d.ten_J, + d.njmax, refsafe, ], outputs=[ @@ -1700,11 +1744,11 @@ def make_constraint(m: types.Model, d: types.Data): m.jnt_solimp, m.jnt_range, m.jnt_margin, - m.jnt_limited_ball_adr, m.dof_invweight0, - d.njmax, + m.jnt_limited_ball_adr, d.qpos, d.qvel, + d.njmax, refsafe, ], outputs=[ @@ -1734,11 +1778,11 @@ def make_constraint(m: types.Model, d: types.Data): m.jnt_solimp, m.jnt_range, m.jnt_margin, - m.jnt_limited_slide_hinge_adr, m.dof_invweight0, - d.njmax, + m.jnt_limited_slide_hinge_adr, d.qpos, d.qvel, + d.njmax, refsafe, ], outputs=[ @@ -1765,18 +1809,18 @@ def make_constraint(m: types.Model, d: types.Data): m.jnt_dofadr, m.tendon_adr, m.tendon_num, - m.tendon_limited_adr, m.tendon_solref_lim, m.tendon_solimp_lim, m.tendon_range, m.tendon_margin, m.tendon_invweight0, - m.wrap_objid, m.wrap_type, - d.njmax, + m.wrap_objid, + m.tendon_limited_adr, d.qvel, - d.ten_length, d.ten_J, + d.ten_length, + d.njmax, refsafe, ], outputs=[ @@ -1809,11 +1853,11 @@ def make_constraint(m: types.Model, d: types.Data): m.body_invweight0, m.dof_bodyid, m.geom_bodyid, - d.njmax, - d.nacon, d.qvel, d.subtree_com, d.cdof, + d.njmax, + d.nacon, refsafe, d.contact.dist, d.contact.dim, @@ -1825,6 +1869,7 @@ def make_constraint(m: types.Model, d: types.Data): d.contact.friction, d.contact.solref, d.contact.solimp, + d.contact.type, ], outputs=[ d.nefc, @@ -1853,11 +1898,11 @@ def make_constraint(m: types.Model, d: types.Data): m.body_invweight0, m.dof_bodyid, m.geom_bodyid, - d.njmax, - d.nacon, d.qvel, d.subtree_com, d.cdof, + d.njmax, + d.nacon, refsafe, d.contact.dist, d.contact.dim, @@ -1870,6 +1915,7 @@ def make_constraint(m: types.Model, d: types.Data): d.contact.solref, d.contact.solreffriction, d.contact.solimp, + d.contact.type, ], outputs=[ d.nefc, 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 d1023013..13a5dbf2 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/derivative.py @@ -15,7 +15,6 @@ import warp as wp -from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m from mujoco.mjx.third_party.mujoco_warp._src.types import BiasType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit @@ -52,8 +51,8 @@ def _qderiv_actuator_passive( # In: qMi: wp.array(dtype=int), qMj: wp.array(dtype=int), - # Data out: - qM_integration_out: wp.array3d(dtype=float), + # Out: + qDeriv_out: wp.array3d(dtype=float), ): worldid, elemid = wp.tid() @@ -62,14 +61,17 @@ def _qderiv_actuator_passive( qderiv = float(0.0) if not opt_disableflags & DisableBit.ACTUATION: + actuator_gainprm_id = worldid % actuator_gainprm.shape[0] + actuator_biasprm_id = worldid % actuator_biasprm.shape[0] + for actid in range(nu): if actuator_gaintype[actid] == GainType.AFFINE: - gain = actuator_gainprm[worldid, actid][2] + gain = actuator_gainprm[actuator_gainprm_id, actid][2] else: gain = 0.0 if actuator_biastype[actid] == BiasType.AFFINE: - bias = actuator_biasprm[worldid, actid][2] + bias = actuator_biasprm[actuator_biasprm_id, actid][2] else: bias = 0.0 @@ -85,17 +87,17 @@ def _qderiv_actuator_passive( # TODO(team): fluid model derivative if not opt_disableflags & DisableBit.DAMPER and dofiid == dofjid: - qderiv -= dof_damping[worldid, dofiid] + qderiv -= dof_damping[worldid % dof_damping.shape[0], dofiid] - qderiv *= opt_timestep[worldid] + qderiv *= opt_timestep[worldid % opt_timestep.shape[0]] if opt_is_sparse: - qM_integration_out[worldid, 0, elemid] = qM_in[worldid, 0, elemid] - qderiv + qDeriv_out[worldid, 0, elemid] = qM_in[worldid, 0, elemid] - qderiv else: qM = qM_in[worldid, dofiid, dofjid] - qderiv - qM_integration_out[worldid, dofiid, dofjid] = qM + qDeriv_out[worldid, dofiid, dofjid] = qM if dofiid != dofjid: - qM_integration_out[worldid, dofjid, dofiid] = qM + qDeriv_out[worldid, dofjid, dofiid] = qM # TODO(team): improve performance with tile operations? @@ -111,35 +113,36 @@ def _qderiv_tendon_damping( # In: qMi: wp.array(dtype=int), qMj: wp.array(dtype=int), - # Data out: - qM_integration_out: wp.array3d(dtype=float), + # Out: + qDeriv_out: wp.array3d(dtype=float), ): worldid, elemid = wp.tid() dofiid = qMi[elemid] dofjid = qMj[elemid] qderiv = float(0.0) + tendon_damping_id = worldid % tendon_damping.shape[0] for tenid in range(ntendon): - qderiv -= ten_J_in[worldid, tenid, dofiid] * ten_J_in[worldid, tenid, dofjid] * tendon_damping[worldid, tenid] - qderiv *= opt_timestep[worldid] + qderiv -= ten_J_in[worldid, tenid, dofiid] * ten_J_in[worldid, tenid, dofjid] * tendon_damping[tendon_damping_id, tenid] + + qderiv *= opt_timestep[worldid % opt_timestep.shape[0]] if opt_is_sparse: - qM_integration_out[worldid, 0, elemid] -= qderiv + qDeriv_out[worldid, 0, elemid] -= qderiv else: - qM_integration_out[worldid, dofiid, dofjid] -= qderiv + qDeriv_out[worldid, dofiid, dofjid] -= qderiv if dofiid != dofjid: - qM_integration_out[worldid, dofjid, dofiid] -= qderiv + qDeriv_out[worldid, dofjid, dofiid] -= qderiv @event_scope -def deriv_smooth_vel(m: Model, d: Data, flg_forward: bool = True): +def deriv_smooth_vel(m: Model, d: Data, qDeriv: wp.array2d(dtype=float)): """Analytical derivative of smooth forces w.r.t. velocities. Args: - m (Model): The model containing kinematic and dynamic information (device). - d (Data): The data object containing the current state and output arrays (device). - flg_forward (bool, optional): If True forward dynamics else inverse dynamics routine. - Default is True. + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + qDeriv: Analytical derivative of smooth forces w.r.t. velocity. """ qMi = m.qM_fullm_i if m.opt.is_sparse else m.dof_tri_row qMj = m.qM_fullm_j if m.opt.is_sparse else m.dof_tri_col @@ -168,24 +171,18 @@ def deriv_smooth_vel(m: Model, d: Data, flg_forward: bool = True): qMi, qMj, ], - outputs=[d.qM_integration], + outputs=[qDeriv], ) else: # TODO(team): directly utilize qM for these settings - wp.copy(d.qM_integration, d.qM) + wp.copy(qDeriv, d.qM) if not m.opt.disableflags & DisableBit.DAMPER: 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], - outputs=[d.qM_integration], + outputs=[qDeriv], ) - if flg_forward: - wp.copy(d.qfrc_integration, d.efc.Ma) - else: - # qfrc = qM @ qacc - mul_m(m, d, d.qfrc_integration, d.qacc, d.inverse_mul_m_skip, d.qM_integration) - # TODO(team): rne derivative 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 86d2641f..08d26fd6 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py @@ -71,7 +71,7 @@ def _next_position( qpos_out: wp.array2d(dtype=float), ): worldid, jntid = wp.tid() - timestep = opt_timestep[worldid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] jnttype = jnt_type[jntid] qpos_adr = jnt_qposadr[jntid] @@ -137,7 +137,7 @@ def _next_velocity( qvel_out: wp.array2d(dtype=float), ): worldid, dofid = wp.tid() - timestep = opt_timestep[worldid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] qvel_out[worldid, dofid] = qvel_in[worldid, dofid] + qacc_scale_in * qacc_in[worldid, dofid] * timestep @@ -188,11 +188,14 @@ def _next_activation( act_out: wp.array2d(dtype=float), ): worldid, actid = wp.tid() + opt_timestep_id = worldid % opt_timestep.shape[0] + actuator_dynprm_id = worldid % actuator_dynprm.shape[0] + actuator_actrange_id = worldid % actuator_actrange.shape[0] act = _next_act( - opt_timestep[worldid], + opt_timestep[opt_timestep_id], actuator_dyntype[actid], - actuator_dynprm[worldid, actid], - actuator_actrange[worldid, actid], + actuator_dynprm[actuator_dynprm_id, actid], + actuator_actrange[actuator_actrange_id, actid], act_in[worldid, actid], act_dot_in[worldid, actid], act_dot_scale, @@ -206,18 +209,18 @@ def _next_time( # Model: opt_timestep: wp.array(dtype=float), # Data in: + nefc_in: wp.array(dtype=int), + time_in: wp.array(dtype=float), nworld_in: int, naconmax_in: int, njmax_in: int, nacon_in: wp.array(dtype=int), - nefc_in: wp.array(dtype=int), - time_in: wp.array(dtype=float), ncollision_in: wp.array(dtype=int), # Data out: time_out: wp.array(dtype=float), ): worldid = wp.tid() - time_out[worldid] = time_in[worldid] + opt_timestep[worldid] + time_out[worldid] = time_in[worldid] + opt_timestep[worldid % opt_timestep.shape[0]] nefc = nefc_in[worldid] if nefc > njmax_in: @@ -227,16 +230,15 @@ def _next_time( ncollision = ncollision_in[0] if ncollision > naconmax_in: nconmax = int(wp.ceil(float(ncollision) / float(nworld_in))) - wp.printf("ncollision overflow - please increase nconmax to %u\n", nconmax) + wp.printf("broadphase overflow - please increase nconmax to %u or naconmax to %u\n", nconmax, ncollision) if nacon_in[0] > naconmax_in: - nconmax = int(wp.ceil(float(ncollision) / float(nworld_in))) - wp.printf("nacon overflow - please increase nconmax to %u\n", nconmax) + nconmax = int(wp.ceil(float(nacon_in[0]) / float(nworld_in))) + wp.printf("narrowphase overflow - please increase nconmax to %u or naconmax to %u\n", nconmax, nacon_in[0]) def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None): """Advance state and time given activation derivatives and acceleration.""" - # TODO(team): can we assume static timesteps? # advance activations @@ -302,12 +304,12 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None) dim=(d.nworld,), inputs=[ m.opt.timestep, + d.nefc, + d.time, d.nworld, d.naconmax, d.njmax, d.nacon, - d.nefc, - d.time, d.ncollision, ], outputs=[ @@ -324,41 +326,16 @@ def _euler_damp_qfrc_sparse( opt_timestep: wp.array(dtype=float), dof_Madr: wp.array(dtype=int), dof_damping: wp.array2d(dtype=float), - # Data out: + # Out: qM_integration_out: wp.array3d(dtype=float), ): worldid, tid = wp.tid() - timestep = opt_timestep[worldid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] adr = dof_Madr[tid] qM_integration_out[worldid, 0, adr] += timestep * dof_damping[worldid, tid] -def _euler_sparse(m: Model, d: Data): - wp.copy(d.qM_integration, d.qM) - wp.launch( - _euler_damp_qfrc_sparse, - dim=(d.nworld, m.nv), - inputs=[ - m.opt.timestep, - m.dof_Madr, - m.dof_damping, - ], - outputs=[ - d.qM_integration, - ], - ) - smooth.factor_solve_i( - m, - d, - d.qM_integration, - d.qLD_integration, - d.qLDiagInv_integration, - d.qacc_integration, - d.efc.Ma, - ) - - @cache_kernel def _tile_euler_dense(tile: TileSet): @nested_kernel(module="unique", enable_backward=False) @@ -371,23 +348,23 @@ def _tile_euler_dense(tile: TileSet): efc_Ma_in: wp.array2d(dtype=float), # In: adr_in: wp.array(dtype=int), - # Data out: - qacc_integration_out: wp.array2d(dtype=float), + # Out: + qacc_out: wp.array2d(dtype=float), ): worldid, nodeid = wp.tid() - timestep = opt_timestep[worldid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] TILE_SIZE = wp.static(tile.size) dofid = adr_in[nodeid] M_tile = wp.tile_load(qM_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(dofid, dofid)) - damping_tile = wp.tile_load(dof_damping[worldid], shape=(TILE_SIZE,), offset=(dofid,)) + damping_tile = wp.tile_load(dof_damping[worldid % dof_damping.shape[0]], shape=(TILE_SIZE,), offset=(dofid,)) damping_scaled = damping_tile * timestep qm_integration_tile = wp.tile_diag_add(M_tile, damping_scaled) Ma_tile = wp.tile_load(efc_Ma_in[worldid], shape=(TILE_SIZE,), offset=(dofid,)) L_tile = wp.tile_cholesky(qm_integration_tile) - qacc_integration_tile = wp.tile_cholesky_solve(L_tile, Ma_tile) - wp.tile_store(qacc_integration_out[worldid], qacc_integration_tile, offset=(dofid)) + qacc_tile = wp.tile_cholesky_solve(L_tile, Ma_tile) + wp.tile_store(qacc_out[worldid], qacc_tile, offset=(dofid)) return euler_dense @@ -395,32 +372,47 @@ def _tile_euler_dense(tile: TileSet): @event_scope def euler(m: Model, d: Data): """Euler integrator, semi-implicit in velocity.""" - # 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: - _euler_sparse(m, d) + qM = wp.clone(d.qM) + qLD = wp.empty((d.nworld, 1, m.nC), dtype=float) + qLDiagInv = wp.empty((d.nworld, m.nv), dtype=float) + wp.launch( + _euler_damp_qfrc_sparse, + dim=(d.nworld, m.nv), + inputs=[m.opt.timestep, m.dof_Madr, m.dof_damping], + outputs=[qM], + ) + smooth.factor_solve_i(m, d, qM, qLD, qLDiagInv, qacc, d.efc.Ma) else: for tile in m.qM_tiles: 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], - outputs=[d.qacc_integration], + outputs=[qacc], block_dim=m.block_dim.euler_dense, ) - - _advance(m, d, d.qacc_integration) + _advance(m, d, qacc) else: _advance(m, d, d.qacc) -def _rk_perturb_state(m: Model, d: Data, scale: float): +def _rk_perturb_state( + m: Model, + d: Data, + scale: float, + qpos_t0: wp.array2d(dtype=float), + qvel_t0: wp.array2d(dtype=float), + act_t0: Optional[wp.array] = None, +): # position wp.launch( _next_position, dim=(d.nworld, m.njnt), - inputs=[m.opt.timestep, m.jnt_type, m.jnt_qposadr, m.jnt_dofadr, d.qpos_t0, d.qvel, scale], + inputs=[m.opt.timestep, m.jnt_type, m.jnt_qposadr, m.jnt_dofadr, qpos_t0, d.qvel, scale], outputs=[d.qpos], ) @@ -428,16 +420,16 @@ def _rk_perturb_state(m: Model, d: Data, scale: float): wp.launch( _next_velocity, dim=(d.nworld, m.nv), - inputs=[m.opt.timestep, d.qvel_t0, d.qacc, scale], + inputs=[m.opt.timestep, qvel_t0, d.qacc, scale], outputs=[d.qvel], ) # activation - if m.na: + if m.na and act_t0 is not None: wp.launch( _next_activation, dim=(d.nworld, m.na), - inputs=[m.opt.timestep, d.act_t0, d.act_dot, scale, False], + inputs=[m.opt.timestep, act_t0, d.act_dot, scale, False], outputs=[d.act], ) @@ -471,73 +463,94 @@ def _rk_accumulate_activation_velocity( act_dot_out[worldid, actid] += scale * act_dot_in[worldid, actid] -def _rk_accumulate(m: Model, d: Data, scale: float): - """Computes one term of 1/6 k_1 + 1/3 k_2 + 1/3 k_3 + 1/6 k_4""" - +def _rk_accumulate( + m: Model, + d: Data, + scale: float, + qvel_rk: wp.array2d(dtype=float), + qacc_rk: wp.array2d(dtype=float), + act_dot_rk: Optional[wp.array] = None, +): + """Computes one term of 1/6 k_1 + 1/3 k_2 + 1/3 k_3 + 1/6 k_4.""" wp.launch( _rk_accumulate_velocity_acceleration, dim=(d.nworld, m.nv), inputs=[d.qvel, d.qacc, scale], - outputs=[d.qvel_rk, d.qacc_rk], + outputs=[qvel_rk, qacc_rk], ) - if m.na: + if m.na and act_dot_rk is not None: wp.launch( _rk_accumulate_activation_velocity, dim=(d.nworld, m.na), inputs=[d.act_dot, scale], - outputs=[d.act_dot_rk], + outputs=[act_dot_rk], ) @event_scope def rungekutta4(m: Model, d: Data): """Runge-Kutta explicit order 4 integrator.""" - - wp.copy(d.qpos_t0, d.qpos) - wp.copy(d.qvel_t0, d.qvel) - - d.qvel_rk.zero_() - d.qacc_rk.zero_() - d.act_dot_rk.zero_() + qpos_t0 = wp.clone(d.qpos) + qvel_t0 = wp.clone(d.qvel) + qvel_rk = wp.zeros((d.nworld, m.nv), dtype=float) + qacc_rk = wp.zeros((d.nworld, m.nv), dtype=float) if m.na: - wp.copy(d.act_t0, d.act) + act_t0 = wp.clone(d.act) + act_dot_rk = wp.zeros((d.nworld, m.na), dtype=float) + else: + act_t0 = None + act_dot_rk = None A, B = _RK4_A, _RK4_B - _rk_accumulate(m, d, B[0]) + _rk_accumulate(m, d, B[0], qvel_rk, qacc_rk, act_dot_rk) + for i in range(3): a, b = float(A[i][i]), B[i + 1] - _rk_perturb_state(m, d, a) + _rk_perturb_state(m, d, a, qpos_t0, qvel_t0, act_t0) forward(m, d) - _rk_accumulate(m, d, b) + _rk_accumulate(m, d, b, qvel_rk, qacc_rk, act_dot_rk) + + wp.copy(d.qpos, qpos_t0) + wp.copy(d.qvel, qvel_t0) - wp.copy(d.qpos, d.qpos_t0) - wp.copy(d.qvel, d.qvel_t0) if m.na: - wp.copy(d.act, d.act_t0) - wp.copy(d.act_dot, d.act_dot_rk) - _advance(m, d, d.qacc_rk, d.qvel_rk) + wp.copy(d.act, act_t0) + wp.copy(d.act_dot, act_dot_rk) + + _advance(m, d, qacc_rk, qvel_rk) @event_scope def implicit(m: Model, d: Data): """Integrates fully implicit in velocity.""" if ~(m.opt.disableflags | ~(DisableBit.ACTUATION | DisableBit.SPRING | DisableBit.DAMPER)): - derivative.deriv_smooth_vel(m, d) - smooth.factor_solve_i( - m, d, d.qM_integration, d.qLD_integration, d.qLDiagInv_integration, d.qacc_integration, d.qfrc_integration - ) - _advance(m, d, d.qacc_integration) + if m.opt.is_sparse: + qDeriv = wp.empty((d.nworld, 1, m.nM), dtype=float) + qLD = wp.empty((d.nworld, 1, m.nC), dtype=float) + else: + qDeriv = wp.empty((d.nworld, m.nv, m.nv), dtype=float) + qLD = wp.empty((d.nworld, m.nv, m.nv), dtype=float) + qLDiagInv = wp.empty((d.nworld, m.nv), dtype=float) + derivative.deriv_smooth_vel(m, d, qDeriv) + qacc = wp.empty((d.nworld, m.nv), dtype=float) + smooth.factor_solve_i(m, d, qDeriv, qLD, qLDiagInv, qacc, d.efc.Ma) + _advance(m, d, qacc) else: _advance(m, d, d.qacc) @event_scope def fwd_position(m: Model, d: Data, factorize: bool = True): - """Position-dependent computations.""" + """Position-dependent computations. + Args: + m: The model containing kinematic and dynamic information. + d: The data object containing the current state and output arrays. + factorize: Flag to factorize interia matrix. + """ smooth.kinematics(m, d) smooth.com_pos(m, d) smooth.camlight(m, d) @@ -625,7 +638,6 @@ def _tendon_velocity(m: Model, d: Data): @event_scope def fwd_velocity(m: Model, d: Data): """Velocity-dependent computations.""" - _actuator_velocity(m, d) if m.ntendon > 0: @@ -673,10 +685,12 @@ def _actuator_force( ): worldid, uid = wp.tid() + actuator_ctrlrange_id = worldid % actuator_ctrlrange.shape[0] + ctrl = ctrl_in[worldid, uid] if actuator_ctrllimited[uid] and not dsbl_clampctrl: - ctrlrange = actuator_ctrlrange[worldid, uid] + ctrlrange = actuator_ctrlrange[actuator_ctrlrange_id, uid] ctrl = wp.clamp(ctrl, ctrlrange[0], ctrlrange[1]) ctrl_act = ctrl @@ -684,11 +698,11 @@ def _actuator_force( if na and act_first >= 0: act_last = act_first + actuator_actnum[uid] - 1 dyntype = actuator_dyntype[uid] + dynprm = actuator_dynprm[worldid % actuator_dynprm.shape[0], uid] if dyntype == DynType.INTEGRATOR: act_dot = ctrl elif dyntype == DynType.FILTER or dyntype == DynType.FILTEREXACT: - dynprm = actuator_dynprm[worldid, uid] act = act_in[worldid, act_last] act_dot = (ctrl - act) / wp.max(dynprm[0], MJ_MINVAL) elif dyntype == DynType.MUSCLE: @@ -701,15 +715,16 @@ def _actuator_force( act_dot_out[worldid, act_last] = act_dot if actuator_actearly[uid]: + opt_timestep_id = worldid % opt_timestep.shape[0] + actuator_actrange_id = worldid % actuator_actrange.shape[0] if dyntype == DynType.INTEGRATOR or dyntype == DynType.NONE: - dynprm = actuator_dynprm[worldid, uid] act = act_in[worldid, act_last] ctrl_act = _next_act( - opt_timestep[worldid], + opt_timestep[opt_timestep_id], dyntype, dynprm, - actuator_actrange[worldid, uid], + actuator_actrange[actuator_actrange_id, uid], act, act_dot, 1.0, @@ -723,7 +738,7 @@ def _actuator_force( # gain gaintype = actuator_gaintype[uid] - gainprm = actuator_gainprm[worldid, uid] + gainprm = actuator_gainprm[worldid % actuator_gainprm.shape[0], uid] gain = 0.0 if gaintype == GainType.FIXED: @@ -737,7 +752,7 @@ def _actuator_force( # bias biastype = actuator_biastype[uid] - biasprm = actuator_biasprm[worldid, uid] + biasprm = actuator_biasprm[worldid % actuator_biasprm.shape[0], uid] bias = 0.0 # BiasType.NONE if biastype == BiasType.AFFINE: @@ -752,7 +767,7 @@ def _actuator_force( # TODO(team): tendon total force clamping if actuator_forcelimited[uid]: - forcerange = actuator_forcerange[worldid, uid] + forcerange = actuator_forcerange[worldid % actuator_forcerange.shape[0], uid] force = wp.clamp(force, forcerange[0], forcerange[1]) actuator_force_out[worldid, uid] = force @@ -765,7 +780,7 @@ def _tendon_actuator_force( actuator_trnid: wp.array(dtype=wp.vec2i), # Data in: actuator_force_in: wp.array2d(dtype=float), - # Data out: + # Out: ten_actfrc_out: wp.array2d(dtype=float), ): worldid, actid = wp.tid() @@ -779,11 +794,11 @@ def _tendon_actuator_force( @wp.kernel def _tendon_actuator_force_clamp( # Model: - actuator_trntype: wp.array(dtype=int), - actuator_trnid: wp.array(dtype=wp.vec2i), tendon_actfrclimited: wp.array(dtype=bool), tendon_actfrcrange: wp.array2d(dtype=wp.vec2), - # Data in: + actuator_trntype: wp.array(dtype=int), + actuator_trnid: wp.array(dtype=wp.vec2i), + # In: ten_actfrc_in: wp.array2d(dtype=float), # Data out: actuator_force_out: wp.array2d(dtype=float), @@ -794,7 +809,7 @@ def _tendon_actuator_force_clamp( tenid = actuator_trnid[actid][0] if tendon_actfrclimited[tenid]: ten_actfrc = ten_actfrc_in[worldid, tenid] - actfrcrange = tendon_actfrcrange[worldid, tenid] + actfrcrange = tendon_actfrcrange[worldid % tendon_actfrcrange.shape[0], tenid] if ten_actfrc < actfrcrange[0]: actuator_force_out[worldid, actid] *= actfrcrange[0] / ten_actfrc @@ -808,8 +823,8 @@ def _qfrc_actuator( nu: int, ngravcomp: int, jnt_actfrclimited: wp.array(dtype=bool), - jnt_actfrcrange: wp.array2d(dtype=wp.vec2), jnt_actgravcomp: wp.array(dtype=int), + jnt_actfrcrange: wp.array2d(dtype=wp.vec2), dof_jntid: wp.array(dtype=int), # Data in: actuator_moment_in: wp.array3d(dtype=float), @@ -831,7 +846,7 @@ def _qfrc_actuator( qfrc += qfrc_gravcomp_in[worldid, dofid] if jnt_actfrclimited[jntid]: - frcrange = jnt_actfrcrange[worldid, jntid] + frcrange = jnt_actfrcrange[worldid % jnt_actfrcrange.shape[0], jntid] qfrc = wp.clamp(qfrc, frcrange[0], frcrange[1]) qfrc_actuator_out[worldid, dofid] = qfrc @@ -878,29 +893,19 @@ def fwd_actuation(m: Model, d: Data): ) if m.ntendon: - d.ten_actfrc.zero_() - + # total actuator force at tendon + ten_actfrc = wp.zeros((d.nworld, m.ntendon), dtype=float) wp.launch( _tendon_actuator_force, dim=(d.nworld, m.nu), - inputs=[ - m.actuator_trntype, - m.actuator_trnid, - d.actuator_force, - ], - outputs=[d.ten_actfrc], + inputs=[m.actuator_trntype, m.actuator_trnid, d.actuator_force], + outputs=[ten_actfrc], ) wp.launch( _tendon_actuator_force_clamp, dim=(d.nworld, m.nu), - inputs=[ - m.actuator_trntype, - m.actuator_trnid, - m.tendon_actfrclimited, - m.tendon_actfrcrange, - d.ten_actfrc, - ], + inputs=[m.tendon_actfrclimited, m.tendon_actfrcrange, m.actuator_trntype, m.actuator_trnid, ten_actfrc], outputs=[d.actuator_force], ) @@ -911,8 +916,8 @@ def fwd_actuation(m: Model, d: Data): m.nu, m.ngravcomp, m.jnt_actfrclimited, - m.jnt_actfrcrange, m.jnt_actgravcomp, + m.jnt_actfrcrange, m.dof_jntid, d.actuator_moment, d.qfrc_gravcomp, @@ -943,8 +948,13 @@ def _qfrc_smooth( @event_scope def fwd_acceleration(m: Model, d: Data, factorize: bool = False): - """Add up all non-constraint forces, compute qacc_smooth.""" + """Add up all non-constraint forces, compute qacc_smooth. + Args: + m: The model containing kinematic and dynamic information. + d: The data object containing the current state and output arrays. + factorize: Flag to factorize inertia matrix. + """ wp.launch( _qfrc_smooth, dim=(d.nworld, m.nv), 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 8f98bfb3..00d58af7 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/inverse.py @@ -21,6 +21,7 @@ from mujoco.mjx.third_party.mujoco_warp._src import sensor from mujoco.mjx.third_party.mujoco_warp._src import smooth from mujoco.mjx.third_party.mujoco_warp._src import solver from mujoco.mjx.third_party.mujoco_warp._src import support +from mujoco.mjx.third_party.mujoco_warp._src.support import mul_m 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 EnableBit @@ -41,8 +42,8 @@ def _qfrc_eulerdamp( qfrc_out: wp.array2d(dtype=float), ): worldid, dofid = wp.tid() - timestep = opt_timestep[worldid] - qfrc_out[worldid, dofid] += timestep * dof_damping[worldid, dofid] * qacc_in[worldid, dofid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] + qfrc_out[worldid, dofid] += timestep * dof_damping[worldid % dof_damping.shape[0], dofid] * qacc_in[worldid, dofid] @wp.kernel @@ -66,8 +67,15 @@ def _qfrc_inverse( qfrc_inverse_out[worldid, dofid] = qfrc_inverse -def discrete_acc(m: Model, d: Data, qacc: wp.array2d(dtype=float), qfrc: wp.array2d(dtype=float)): - """Convert discrete-time qacc to continuous-time qacc.""" +def discrete_acc(m: Model, d: Data, qacc: wp.array2d(dtype=float)): + """Convert discrete-time qacc to continuous-time qacc. + + Args: + m: The model containing kinematic and dynamic information. + d: The data object containing the current state and output arrays. + qacc: Acceleration. + """ + qfrc = wp.empty((d.nworld, m.nv), dtype=float) if m.opt.integrator == IntegratorType.RK4: raise NotImplementedError("discrete inverse dynamics is not supported by RK4 integrator") @@ -81,7 +89,7 @@ def discrete_acc(m: Model, d: Data, qacc: wp.array2d(dtype=float), qfrc: wp.arra # set qfrc = (d.qM + m.opt.timestep * diag(m.dof_damping)) * d.qacc # d.qM @ d.qacc - support.mul_m(m, d, qfrc, d.qacc, d.inverse_mul_m_skip) + support.mul_m(m, d, qfrc, d.qacc) # qfrc += m.opt.timestep * m.dof_damping * d.qacc wp.launch( @@ -91,7 +99,12 @@ def discrete_acc(m: Model, d: Data, qacc: wp.array2d(dtype=float), qfrc: wp.arra outputs=[qfrc], ) elif m.opt.integrator == IntegratorType.IMPLICITFAST: - derivative.deriv_smooth_vel(m, d, flg_forward=False) + if m.opt.is_sparse: + qDeriv = wp.empty((d.nworld, 1, m.nM), dtype=float) + else: + qDeriv = wp.empty((d.nworld, m.nv, m.nv), dtype=float) + derivative.deriv_smooth_vel(m, d, qDeriv) + mul_m(m, d, qfrc, d.qacc, M=qDeriv) smooth.factor_solve_i(m, d, d.qM, d.qLD, d.qLDiagInv, qacc, qfrc) else: raise NotImplementedError(f"integrator {m.opt.integrator} not implemented.") @@ -102,7 +115,6 @@ def discrete_acc(m: Model, d: Data, qacc: wp.array2d(dtype=float), qfrc: wp.arra def inv_constraint(m: Model, d: Data): """Inverse constraint solver.""" - # no constraints if d.njmax == 0: d.qfrc_constraint.zero_() @@ -122,15 +134,15 @@ def inverse(m: Model, d: Data): invdiscrete = m.opt.enableflags & EnableBit.INVDISCRETE if invdiscrete: # save discrete-time qacc and compute continuous-time qacc - wp.copy(d.qacc_discrete, d.qacc) - discrete_acc(m, d, d.qacc, d.qfrc_integration) + qacc_discrete = wp.clone(d.qacc) + discrete_acc(m, d, d.qacc) inv_constraint(m, d) smooth.rne(m, d) smooth.tendon_bias(m, d, d.qfrc_bias) sensor.sensor_acc(m, d) - support.mul_m(m, d, d.qfrc_inverse, d.qacc, d.inverse_mul_m_skip) + support.mul_m(m, d, d.qfrc_inverse, d.qacc) wp.launch( _qfrc_inverse, @@ -146,4 +158,4 @@ def inverse(m: Model, d: Data): if invdiscrete: # restore discrete-time qacc - wp.copy(d.qacc, d.qacc_discrete) + wp.copy(d.qacc, qacc_discrete) 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 01474c39..db0aafbe 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py @@ -22,36 +22,64 @@ import warp as wp from mujoco.mjx.third_party.mujoco_warp._src import math from mujoco.mjx.third_party.mujoco_warp._src import types from mujoco.mjx.third_party.mujoco_warp._src import warp_util - -# max number of worlds supported -MAX_WORLDS = 2**24 +from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel # tolerance override for float32 _TOLERANCE_F32 = 1.0e-6 -def _max_meshdegree(mjm: mujoco.MjModel) -> int: - if mjm.mesh_polyvertnum.size == 0: - return 4 - return max(3, mjm.mesh_polymapnum.max()) +def _compute_nmaxpolygon(mjm: mujoco.MjModel, geom_pair_type_count: tuple[int, ...]) -> int: + """Compute nmaxpolygon given the geom pairs for the model.""" + nboxbox = geom_pair_type_count[math.upper_trid_index(len(types.GeomType), types.GeomType.BOX.value, types.GeomType.BOX.value)] + nboxmesh = geom_pair_type_count[ + math.upper_trid_index(len(types.GeomType), types.GeomType.BOX.value, types.GeomType.MESH.value) + ] + nmeshmesh = geom_pair_type_count[ + math.upper_trid_index(len(types.GeomType), types.GeomType.MESH.value, types.GeomType.MESH.value) + ] + + # need at least 4 (square sides) if there's a box collision needing multiccd + # TODO(kbayes): remove nboxbox or enable ccd for box-box collisions + box_factor = 4 if nboxbox + nboxmesh > 0 else 0 + + # possibly need to allocate more memory if there's meshes + if nmeshmesh + nboxmesh > 0: + return np.append(mjm.mesh_polyvertnum, box_factor).max() + return box_factor -def _max_npolygon(mjm: mujoco.MjModel) -> int: - if mjm.mesh_polyvertnum.size == 0: - return 4 - return max(4, mjm.mesh_polyvertnum.max()) +def _compute_nmaxmeshdeg(mjm: mujoco.MjModel, geom_pair_type_count: tuple[int, ...]) -> int: + """Compute nmaxmeshdeg given the geom pairs for the model.""" + nboxbox = geom_pair_type_count[math.upper_trid_index(len(types.GeomType), types.GeomType.BOX.value, types.GeomType.BOX.value)] + nboxmesh = geom_pair_type_count[ + math.upper_trid_index(len(types.GeomType), types.GeomType.BOX.value, types.GeomType.MESH.value) + ] + nmeshmesh = geom_pair_type_count[ + math.upper_trid_index(len(types.GeomType), types.GeomType.MESH.value, types.GeomType.MESH.value) + ] + + # need at least 3 (3 edges per vertex) if there's a box collision needing multiccd + # TODO(kbayes): remove nboxbox or enable ccd for box-box collisions + box_factor = 3 if nboxbox + nboxmesh > 0 else 0 + + # possibly need to allocate more memory if there's meshes + if nmeshmesh + nboxmesh > 0: + return np.append(mjm.mesh_polymapnum, box_factor).max() + return box_factor def put_model(mjm: mujoco.MjModel) -> types.Model: - """ - Creates a model on device. + """Creates a model on device. Args: - mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host). + mjm: The model containing kinematic and dynamic information (host). Returns: - Model: The model containing kinematic and dynamic information (device). + The model containing kinematic and dynamic information (device). """ + # check for compatible cuda toolkit and driver versions + warp_util.check_toolkit_driver() + # check supported features for field, field_types, field_str in ( (mjm.actuator_trntype, types.TrnType, "Actuator transmission type"), @@ -134,9 +162,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: is_sparse = mujoco.mj_isSparse(mjm) - # calculate some fields that cannot be easily computed inline - nlsp = mjm.opt.ls_iterations # TODO(team): how to set nlsp? - # dof lower triangle row and column indices (used in solver) dof_tri_row, dof_tri_col = np.tril_indices(mjm.nv) @@ -197,12 +222,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: qM_tiles = tuple(types.TileSet(adr=wp.array(tiles[sz], dtype=int), size=sz) for sz in sorted(tiles.keys())) - # subtree_mass is a precalculated array used in smooth - subtree_mass = np.copy(mjm.body_mass) - # TODO(team): should this be [mjm.nbody - 1, 0) ? - for i in range(mjm.nbody - 1, -1, -1): - subtree_mass[mjm.body_parentid[i]] += subtree_mass[i] - # actuator_moment tiles are grouped by dof size and number of actuators tree_id = np.arange(len(tile_corners), dtype=np.int32) num_trees = int(np.max(tree_id)) if len(tree_id) > 0 else 0 @@ -336,8 +355,8 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: mask = np.array((contype1 & conaffinity2) | (contype2 & conaffinity1), dtype=bool) exclude = np.isin((bodyid1 << 16) + bodyid2, mjm.exclude_signature) - nxn_pairid = -1 * np.ones(len(geom1), dtype=int) - nxn_pairid[~(mask & ~self_collision & ~parent_child_collision & ~exclude)] = -2 + nxn_pairid_contact = -1 * np.ones(len(geom1), dtype=int) + nxn_pairid_contact[~(mask & ~self_collision & ~parent_child_collision & ~exclude)] = -2 # contact pairs for i in range(mjm.npair): @@ -349,25 +368,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: else: pairid = np.int32(math.upper_tri_index(mjm.ngeom, int(pair_geom1), int(pair_geom2))) - nxn_pairid[pairid] = i - - include = nxn_pairid > -2 - nxn_pairid_filtered = nxn_pairid[include] - nxn_geom_pair_filtered = nxn_geom_pair[include] - - # count contact pair types - geom_type_pair_count = np.bincount( - [ - math.upper_trid_index(len(types.GeomType), int(mjm.geom_type[geom1[i]]), int(mjm.geom_type[geom2[i]])) - for i in np.arange(len(geom1)) - if nxn_pairid[i] > -2 - ], - minlength=len(types.GeomType) * (len(types.GeomType) + 1) // 2, - ) - - # Disable collisions if there are no potentially colliding pairs - if np.sum(geom_type_pair_count) == 0: - mjm.opt.disableflags |= types.DisableBit.CONTACT + nxn_pairid_contact[pairid] = i def create_nmodel_batched_array(mjm_array, dtype, expand_dim=True): array = wp.array(mjm_array, dtype=dtype) @@ -375,11 +376,11 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: array._is_batched = True if not expand_dim: array.strides = (0,) + array.strides[1:] - array.shape = (MAX_WORLDS,) + array.shape[1:] + array.shape = (1,) + array.shape[1:] return array array.strides = (0,) + array.strides array.ndim += 1 - array.shape = (MAX_WORLDS,) + array.shape + array.shape = (1,) + array.shape return array # rangefinder @@ -391,13 +392,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: # contact sensor sensor_adr_to_contact_adr = np.clip(np.cumsum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT) - 1, a_min=0, a_max=None) - if nxn_geom_pair_filtered.shape[0] < 250_000: - broadphase = types.BroadphaseType.NXN - elif mjm.ngeom < 1000: - broadphase = types.BroadphaseType.SAP_TILE - else: - broadphase = types.BroadphaseType.SAP_SEGMENTED - condim = np.concatenate((mjm.geom_condim, mjm.pair_dim)) condim_max = np.max(condim) if len(condim) > 0 else 0 @@ -411,47 +405,98 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: if is_collision_sensor.any(): - def _collision_sensor_check(sensor_type, sensor_id, geom_type, err_msg): - for type_, id_ in zip(sensor_type, sensor_id): - if type_ == mujoco.mjtObj.mjOBJ_BODY: - geomnum = mjm.body_geomnum[id_] - geomadr = mjm.body_geomadr[id_] - for geomid in range(geomadr, geomadr + geomnum): - if mjm.geom_type[geomid] == geom_type: - raise NotImplementedError(err_msg) - elif type_ == mujoco.mjtObj.mjOBJ_GEOM: - if mjm.geom_type[id_] == geom_type: - raise NotImplementedError(err_msg) + def not_implemented(objtype, objid, geomtype): + if objtype == mujoco.mjtObj.mjOBJ_BODY: + geomnum = mjm.body_geomnum[objid] + geomadr = mjm.body_geomadr[objid] + for geomid in range(geomadr, geomadr + geomnum): + if mjm.geom_type[geomid] == geomtype: + return True + elif objtype == mujoco.mjtObj.mjOBJ_GEOM: + if mjm.geom_type[objid] == geomtype: + return True + return False - sensor_collision_objtype = mjm.sensor_objtype[is_collision_sensor] - sensor_collision_objid = mjm.sensor_objid[is_collision_sensor] - sensor_collision_reftype = mjm.sensor_reftype[is_collision_sensor] - sensor_collision_refid = mjm.sensor_refid[is_collision_sensor] + for geoms in [ + (types.GeomType.BOX, types.GeomType.BOX), + (types.GeomType.CAPSULE, types.GeomType.BOX), + (types.GeomType.CYLINDER, types.GeomType.BOX), + (types.GeomType.PLANE, types.GeomType.BOX), + ]: + 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, geoms[0]) and not_implemented(reftype, refid, geoms[1]): + raise NotImplementedError(f"Collision sensors with {geoms[0]} and {geoms[1]} are not implemented.") - _collision_sensor_check( - sensor_collision_objtype, - sensor_collision_objid, - mujoco.mjtGeom.mjGEOM_PLANE, - "Collision sensors with planes are not implemented.", - ) - _collision_sensor_check( - sensor_collision_reftype, - sensor_collision_refid, - mujoco.mjtGeom.mjGEOM_PLANE, - "Collision sensors with planes are not implemented.", - ) - _collision_sensor_check( - sensor_collision_objtype, - sensor_collision_objid, - mujoco.mjtGeom.mjGEOM_HFIELD, - "Collision sensors with height fields are not implemented.", - ) - _collision_sensor_check( - sensor_collision_reftype, - sensor_collision_refid, - mujoco.mjtGeom.mjGEOM_HFIELD, - "Collision sensors with height fields are not implemented.", - ) + nxn_pairid_collision = -1 * np.ones(len(geom1), dtype=int) + pairids = [] + collision_geom_adr = [0] + sensor_collision_start_adr = [] + for i in range(sensor_collision_adr.size): + sensorid = sensor_collision_adr[i] + objtype = mjm.sensor_objtype[sensorid] + objid = mjm.sensor_objid[sensorid] + reftype = mjm.sensor_reftype[sensorid] + refid = mjm.sensor_refid[sensorid] + + # get lists of geoms to collide + if objtype == types.ObjType.BODY: + n1 = mjm.body_geomnum[objid] + id1 = mjm.body_geomadr[objid] + else: + n1 = 1 + id1 = objid + if reftype == types.ObjType.BODY: + n2 = mjm.body_geomnum[refid] + id2 = mjm.body_geomadr[refid] + else: + n2 = 1 + id2 = refid + + # collide all pairs + geomid = 0 + for geom1id in range(id1, id1 + n1): + for geom2id in range(id2, id2 + n2): + if geom2id < geom1id: + pairid = np.int32(math.upper_tri_index(mjm.ngeom, int(geom2id), int(geom1id))) + else: + pairid = np.int32(math.upper_tri_index(mjm.ngeom, int(geom1id), int(geom2id))) + + if pairid in pairids: + sensor_collision_start_adr.append(nxn_pairid_collision[pairid]) + else: + pairids.append(pairid) + adr = collision_geom_adr[-1] + geomid + nxn_pairid_collision[pairid] = adr + sensor_collision_start_adr.append(adr) + + geomid += 1 + if i < sensor_collision_adr.size - 1: + collision_geom_adr.append(collision_geom_adr[-1] + n1 * n2) + + include = (nxn_pairid_contact > -2) | (nxn_pairid_collision >= 0) + nxn_pairid = np.hstack([nxn_pairid_contact.reshape((-1, 1)), nxn_pairid_collision.reshape((-1, 1))]) + nxn_pairid_filtered = nxn_pairid[include] + nxn_geom_pair_filtered = nxn_geom_pair[include] + + # count contact pair types + geom_type_pair_count = np.bincount( + [ + math.upper_trid_index(len(types.GeomType), int(mjm.geom_type[geom1[i]]), int(mjm.geom_type[geom2[i]])) + for i in np.arange(len(geom1)) + if nxn_pairid_contact[i] > -2 or nxn_pairid_collision[i] > -1 + ], + minlength=len(types.GeomType) * (len(types.GeomType) + 1) // 2, + ) + + # disable collisions if there are no potentially colliding pairs + disableflags = mjm.opt.disableflags + if np.sum(geom_type_pair_count) == 0: + disableflags |= types.DisableBit.CONTACT if mjm.geom_fluid.size: geom_fluid_params = mjm.geom_fluid.reshape(mjm.ngeom, mujoco.mjNFLUID) @@ -464,6 +509,19 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: if np.any(active_geom): body_fluid_ellipsoid[mjm.geom_bodyid[active_geom]] = True + if nxn_geom_pair_filtered.shape[0] < 250_000: + broadphase = types.BroadphaseType.NXN + elif mjm.ngeom < 1000: + broadphase = types.BroadphaseType.SAP_TILE + else: + broadphase = types.BroadphaseType.SAP_SEGMENTED + + ls_parallel_id = mujoco.mj_name2id(mjm, mujoco.mjtObj.mjOBJ_NUMERIC, "ls_parallel") + if (ls_parallel_id > -1) and (mjm.numeric_data[mjm.numeric_adr[ls_parallel_id]] == 1): + ls_parallel = True + else: + ls_parallel = False + m = types.Model( nq=mjm.nq, nv=mjm.nv, @@ -471,35 +529,35 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: na=mjm.na, nbody=mjm.nbody, njnt=mjm.njnt, + nM=mjm.nM, + nC=mjm.nC, ngeom=mjm.ngeom, nsite=mjm.nsite, ncam=mjm.ncam, nlight=mjm.nlight, - nmat=mjm.nmat, nflex=mjm.nflex, nflexvert=mjm.nflexvert, nflexedge=mjm.nflexedge, nflexelem=mjm.nflexelem, nflexelemdata=mjm.nflexelemdata, - nexclude=mjm.nexclude, - neq=mjm.neq, - nmocap=mjm.nmocap, - ngravcomp=mjm.ngravcomp, - nM=mjm.nM, - nC=mjm.nC, - ntendon=mjm.ntendon, - nwrap=mjm.nwrap, - nsensor=mjm.nsensor, - nsensordata=mjm.nsensordata, - nsensortaxel=sum(mjm.mesh_vertnum[mjm.sensor_objid[mjm.sensor_type == mujoco.mjtSensor.mjSENS_TACTILE]]), nmeshvert=mjm.nmeshvert, nmeshface=mjm.nmeshface, nmeshgraph=mjm.nmeshgraph, nmeshpoly=mjm.nmeshpoly, nmeshpolyvert=mjm.nmeshpolyvert, nmeshpolymap=mjm.nmeshpolymap, - nlsp=nlsp, + nhfield=mjm.nhfield, + nhfielddata=mjm.nhfielddata, + nmat=mjm.nmat, npair=mjm.npair, + nexclude=mjm.nexclude, + neq=mjm.neq, + ntendon=mjm.ntendon, + nwrap=mjm.nwrap, + nsensor=mjm.nsensor, + nmocap=mjm.nmocap, + ngravcomp=mjm.ngravcomp, + nsensordata=mjm.nsensordata, opt=types.Option( timestep=create_nmodel_batched_array(np.array(mjm.opt.timestep), dtype=float, expand_dim=False), tolerance=create_nmodel_batched_array( @@ -518,16 +576,16 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: iterations=mjm.opt.iterations, ls_iterations=mjm.opt.ls_iterations, integrator=mjm.opt.integrator, - disableflags=mjm.opt.disableflags, + disableflags=type(mjm.opt.disableflags)(disableflags), enableflags=mjm.opt.enableflags, impratio=create_nmodel_batched_array(np.array(mjm.opt.impratio), dtype=float, expand_dim=False), is_sparse=bool(is_sparse), - ls_parallel=False, + ls_parallel=ls_parallel, ls_parallel_min_step=1.0e-6, # TODO(team): determine good default setting ccd_iterations=mjm.opt.ccd_iterations, broadphase=int(broadphase), broadphase_filter=int(types.BroadphaseFilter.PLANE | types.BroadphaseFilter.SPHERE | types.BroadphaseFilter.OBB), - graph_conditional=True and warp_util.conditional_graph_supported(), + graph_conditional=True, sdf_initpoints=mjm.opt.sdf_initpoints, sdf_iterations=mjm.opt.sdf_iterations, run_collision_detection=True, @@ -539,23 +597,10 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: ), qpos0=create_nmodel_batched_array(mjm.qpos0, dtype=float), qpos_spring=create_nmodel_batched_array(mjm.qpos_spring, dtype=float), - qM_fullm_i=wp.array(qM_fullm_i, dtype=int), - qM_fullm_j=wp.array(qM_fullm_j, dtype=int), - qM_mulm_i=wp.array(qM_mulm_i, dtype=int), - qM_mulm_j=wp.array(qM_mulm_j, dtype=int), - qM_madr_ij=wp.array(qM_madr_ij, dtype=int), - qLD_updates=qLD_updates, - M_rownnz=wp.array(mjm.M_rownnz, dtype=int), - M_rowadr=wp.array(mjm.M_rowadr, dtype=int), - M_colind=wp.array(mjm.M_colind, dtype=int), - mapM2M=wp.array(mjm.mapM2M, dtype=int), - qM_tiles=qM_tiles, - body_tree=body_tree, body_parentid=wp.array(mjm.body_parentid, dtype=int), body_rootid=wp.array(mjm.body_rootid, dtype=int), body_weldid=wp.array(mjm.body_weldid, dtype=int), body_mocapid=wp.array(mjm.body_mocapid, dtype=int), - mocap_bodyid=wp.array(mocap_bodyid, dtype=int), body_jntnum=wp.array(mjm.body_jntnum, dtype=int), body_jntadr=wp.array(mjm.body_jntadr, dtype=int), body_dofnum=wp.array(mjm.body_dofnum, dtype=int), @@ -568,19 +613,21 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: body_iquat=create_nmodel_batched_array(mjm.body_iquat, dtype=wp.quat), body_mass=create_nmodel_batched_array(mjm.body_mass, dtype=float), body_subtreemass=create_nmodel_batched_array(mjm.body_subtreemass, dtype=float), - subtree_mass=create_nmodel_batched_array(subtree_mass, dtype=float), body_inertia=create_nmodel_batched_array(mjm.body_inertia, dtype=wp.vec3), body_invweight0=create_nmodel_batched_array(mjm.body_invweight0, dtype=wp.vec2), + body_gravcomp=create_nmodel_batched_array(mjm.body_gravcomp, dtype=float), body_contype=wp.array(mjm.body_contype, dtype=int), body_conaffinity=wp.array(mjm.body_conaffinity, dtype=int), - body_gravcomp=create_nmodel_batched_array(mjm.body_gravcomp, dtype=float), - body_fluid_ellipsoid=wp.array(body_fluid_ellipsoid, dtype=bool), + oct_aabb=wp.array2d(mjm.oct_aabb, dtype=wp.vec3), + oct_child=wp.array(mjm.oct_child, dtype=types.vec8i), + oct_coeff=wp.array(mjm.oct_coeff, dtype=types.vec8f), jnt_type=wp.array(mjm.jnt_type, dtype=int), jnt_qposadr=wp.array(mjm.jnt_qposadr, dtype=int), jnt_dofadr=wp.array(mjm.jnt_dofadr, dtype=int), jnt_bodyid=wp.array(mjm.jnt_bodyid, dtype=int), jnt_limited=wp.array(mjm.jnt_limited, dtype=int), jnt_actfrclimited=wp.array(mjm.jnt_actfrclimited, dtype=bool), + jnt_actgravcomp=wp.array(mjm.jnt_actgravcomp, dtype=int), jnt_solref=create_nmodel_batched_array(mjm.jnt_solref, dtype=wp.vec2), jnt_solimp=create_nmodel_batched_array(mjm.jnt_solimp, dtype=types.vec5), jnt_pos=create_nmodel_batched_array(mjm.jnt_pos, dtype=wp.vec3), @@ -589,18 +636,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: jnt_range=create_nmodel_batched_array(mjm.jnt_range, dtype=wp.vec2), jnt_actfrcrange=create_nmodel_batched_array(mjm.jnt_actfrcrange, dtype=wp.vec2), jnt_margin=create_nmodel_batched_array(mjm.jnt_margin, dtype=float), - # these jnt_limited adrs are used in constraint.py - jnt_limited_slide_hinge_adr=wp.array( - np.nonzero( - mjm.jnt_limited & ((mjm.jnt_type == mujoco.mjtJoint.mjJNT_SLIDE) | (mjm.jnt_type == mujoco.mjtJoint.mjJNT_HINGE)) - )[0], - dtype=int, - ), - jnt_limited_ball_adr=wp.array( - np.nonzero(mjm.jnt_limited & (mjm.jnt_type == mujoco.mjtJoint.mjJNT_BALL))[0], - dtype=int, - ), - jnt_actgravcomp=wp.array(mjm.jnt_actgravcomp, dtype=int), dof_bodyid=wp.array(mjm.dof_bodyid, dtype=int), dof_jntid=wp.array(mjm.dof_jntid, dtype=int), dof_parentid=wp.array(mjm.dof_parentid, dtype=int), @@ -611,29 +646,27 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: dof_frictionloss=create_nmodel_batched_array(mjm.dof_frictionloss, dtype=float), dof_solimp=create_nmodel_batched_array(mjm.dof_solimp, dtype=types.vec5), dof_solref=create_nmodel_batched_array(mjm.dof_solref, dtype=wp.vec2), - dof_tri_row=wp.array(dof_tri_row, dtype=int), - dof_tri_col=wp.array(dof_tri_col, dtype=int), geom_type=wp.array(mjm.geom_type, dtype=int), geom_contype=wp.array(mjm.geom_contype, dtype=int), geom_conaffinity=wp.array(mjm.geom_conaffinity, dtype=int), geom_condim=wp.array(mjm.geom_condim, dtype=int), geom_bodyid=wp.array(mjm.geom_bodyid, dtype=int), geom_dataid=wp.array(mjm.geom_dataid, dtype=int), - geom_group=wp.array(mjm.geom_group, dtype=int), geom_matid=create_nmodel_batched_array(mjm.geom_matid, dtype=int), + geom_group=wp.array(mjm.geom_group, dtype=int), geom_priority=wp.array(mjm.geom_priority, dtype=int), geom_solmix=create_nmodel_batched_array(mjm.geom_solmix, dtype=float), geom_solref=create_nmodel_batched_array(mjm.geom_solref, dtype=wp.vec2), geom_solimp=create_nmodel_batched_array(mjm.geom_solimp, dtype=types.vec5), geom_size=create_nmodel_batched_array(mjm.geom_size, dtype=wp.vec3), - geom_fluid=wp.array(geom_fluid_params, dtype=float), - geom_aabb=wp.array2d(mjm.geom_aabb, dtype=wp.vec3), + geom_aabb=create_nmodel_batched_array(mjm.geom_aabb, dtype=wp.vec3), geom_rbound=create_nmodel_batched_array(mjm.geom_rbound, dtype=float), geom_pos=create_nmodel_batched_array(mjm.geom_pos, dtype=wp.vec3), geom_quat=create_nmodel_batched_array(mjm.geom_quat, dtype=wp.quat), geom_friction=create_nmodel_batched_array(mjm.geom_friction, dtype=wp.vec3), geom_margin=create_nmodel_batched_array(mjm.geom_margin, dtype=float), geom_gap=create_nmodel_batched_array(mjm.geom_gap, dtype=float), + geom_fluid=wp.array(geom_fluid_params, dtype=float), geom_rgba=create_nmodel_batched_array(mjm.geom_rgba, dtype=wp.vec4), site_type=wp.array(mjm.site_type, dtype=int), site_bodyid=wp.array(mjm.site_bodyid, dtype=int), @@ -679,12 +712,12 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: flex_damping=wp.array(mjm.flex_damping, dtype=float), mesh_vertadr=wp.array(mjm.mesh_vertadr, dtype=int), mesh_vertnum=wp.array(mjm.mesh_vertnum, dtype=int), - mesh_vert=wp.array(mjm.mesh_vert, dtype=wp.vec3), - mesh_normaladr=wp.array(mjm.mesh_normaladr, dtype=int), - mesh_normal=wp.array(mjm.mesh_normal, dtype=wp.vec3), mesh_faceadr=wp.array(mjm.mesh_faceadr, dtype=int), - mesh_face=wp.array(mjm.mesh_face, dtype=wp.vec3i), + mesh_normaladr=wp.array(mjm.mesh_normaladr, dtype=int), mesh_graphadr=wp.array(mjm.mesh_graphadr, dtype=int), + mesh_vert=wp.array(mjm.mesh_vert, dtype=wp.vec3), + mesh_normal=wp.array(mjm.mesh_normal, dtype=wp.vec3), + mesh_face=wp.array(mjm.mesh_face, dtype=wp.vec3i), mesh_graph=wp.array(mjm.mesh_graph, dtype=int), mesh_quat=wp.array(mjm.mesh_quat, dtype=wp.quat), mesh_polynum=wp.array(mjm.mesh_polynum, dtype=int), @@ -696,16 +729,24 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: mesh_polymapadr=wp.array(mjm.mesh_polymapadr, dtype=int), mesh_polymapnum=wp.array(mjm.mesh_polymapnum, dtype=int), mesh_polymap=wp.array(mjm.mesh_polymap, dtype=int), - oct_aabb=wp.array2d(mjm.oct_aabb, dtype=wp.vec3), - oct_child=wp.array(mjm.oct_child, dtype=types.vec8i), - oct_coeff=wp.array(mjm.oct_coeff, dtype=types.vec8f), - nhfield=mjm.nhfield, - nhfielddata=mjm.nhfielddata, - hfield_adr=wp.array(mjm.hfield_adr, dtype=int), + hfield_size=wp.array(mjm.hfield_size, dtype=wp.vec4), hfield_nrow=wp.array(mjm.hfield_nrow, dtype=int), hfield_ncol=wp.array(mjm.hfield_ncol, dtype=int), - hfield_size=wp.array(mjm.hfield_size, dtype=wp.vec4), + hfield_adr=wp.array(mjm.hfield_adr, dtype=int), hfield_data=wp.array(mjm.hfield_data, dtype=float), + mat_texid=create_nmodel_batched_array(mjm.mat_texid, dtype=int), + mat_texrepeat=create_nmodel_batched_array(mjm.mat_texrepeat, dtype=wp.vec2), + mat_rgba=create_nmodel_batched_array(mjm.mat_rgba, dtype=wp.vec4), + pair_dim=wp.array(mjm.pair_dim, dtype=int), + pair_geom1=wp.array(mjm.pair_geom1, dtype=int), + pair_geom2=wp.array(mjm.pair_geom2, dtype=int), + pair_solref=create_nmodel_batched_array(mjm.pair_solref, dtype=wp.vec2), + pair_solreffriction=create_nmodel_batched_array(mjm.pair_solreffriction, dtype=wp.vec2), + pair_solimp=create_nmodel_batched_array(mjm.pair_solimp, dtype=types.vec5), + pair_margin=create_nmodel_batched_array(mjm.pair_margin, dtype=float), + pair_gap=create_nmodel_batched_array(mjm.pair_gap, dtype=float), + pair_friction=create_nmodel_batched_array(mjm.pair_friction, dtype=types.vec5), + exclude_signature=wp.array(mjm.exclude_signature, dtype=int), eq_type=wp.array(mjm.eq_type, dtype=int), eq_obj1id=wp.array(mjm.eq_obj1id, dtype=int), eq_obj2id=wp.array(mjm.eq_obj2id, dtype=int), @@ -714,13 +755,27 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: eq_solref=create_nmodel_batched_array(mjm.eq_solref, dtype=wp.vec2), eq_solimp=create_nmodel_batched_array(mjm.eq_solimp, dtype=types.vec5), eq_data=create_nmodel_batched_array(mjm.eq_data, dtype=types.vec11), - # pre-compute indices of equality constraints - eq_connect_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.CONNECT)[0], dtype=int), - eq_wld_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.WELD)[0], dtype=int), - eq_jnt_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.JOINT)[0], dtype=int), - eq_ten_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.TENDON)[0], dtype=int), - actuator_moment_tiles_nv=actuator_moment_tiles_nv, - actuator_moment_tiles_nu=actuator_moment_tiles_nu, + tendon_adr=wp.array(mjm.tendon_adr, dtype=int), + tendon_num=wp.array(mjm.tendon_num, dtype=int), + tendon_limited=wp.array(mjm.tendon_limited, dtype=int), + tendon_actfrclimited=wp.array(mjm.tendon_actfrclimited, dtype=bool), + tendon_solref_lim=create_nmodel_batched_array(mjm.tendon_solref_lim, dtype=wp.vec2f), + tendon_solimp_lim=create_nmodel_batched_array(mjm.tendon_solimp_lim, dtype=types.vec5), + tendon_solref_fri=create_nmodel_batched_array(mjm.tendon_solref_fri, dtype=wp.vec2f), + tendon_solimp_fri=create_nmodel_batched_array(mjm.tendon_solimp_fri, dtype=types.vec5), + tendon_range=create_nmodel_batched_array(mjm.tendon_range, dtype=wp.vec2f), + tendon_actfrcrange=create_nmodel_batched_array(mjm.tendon_actfrcrange, dtype=wp.vec2), + tendon_margin=create_nmodel_batched_array(mjm.tendon_margin, dtype=float), + tendon_stiffness=create_nmodel_batched_array(mjm.tendon_stiffness, dtype=float), + tendon_damping=create_nmodel_batched_array(mjm.tendon_damping, dtype=float), + tendon_armature=create_nmodel_batched_array(mjm.tendon_armature, dtype=float), + tendon_frictionloss=create_nmodel_batched_array(mjm.tendon_frictionloss, dtype=float), + tendon_lengthspring=create_nmodel_batched_array(mjm.tendon_lengthspring, dtype=wp.vec2), + tendon_length0=create_nmodel_batched_array(mjm.tendon_length0, dtype=float), + tendon_invweight0=create_nmodel_batched_array(mjm.tendon_invweight0, dtype=float), + wrap_type=wp.array(mjm.wrap_type, dtype=int), + wrap_objid=wp.array(mjm.wrap_objid, dtype=int), + wrap_prm=wp.array(mjm.wrap_prm, dtype=float), actuator_trntype=wp.array(mjm.actuator_trntype, dtype=int), actuator_dyntype=wp.array(mjm.actuator_dyntype, dtype=int), actuator_gaintype=wp.array(mjm.actuator_gaintype, dtype=int), @@ -742,53 +797,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: actuator_cranklength=wp.array(mjm.actuator_cranklength, dtype=float), actuator_acc0=wp.array(mjm.actuator_acc0, dtype=float), actuator_lengthrange=wp.array(mjm.actuator_lengthrange, dtype=wp.vec2), - exclude_signature=wp.array(mjm.exclude_signature, dtype=int), - nxn_geom_pair=wp.array(nxn_geom_pair, dtype=wp.vec2i), - nxn_geom_pair_filtered=wp.array(nxn_geom_pair_filtered, dtype=wp.vec2i), - nxn_pairid=wp.array(nxn_pairid, dtype=int), - nxn_pairid_filtered=wp.array(nxn_pairid_filtered, dtype=int), - pair_dim=wp.array(mjm.pair_dim, dtype=int), - pair_geom1=wp.array(mjm.pair_geom1, dtype=int), - pair_geom2=wp.array(mjm.pair_geom2, dtype=int), - pair_solref=create_nmodel_batched_array(mjm.pair_solref, dtype=wp.vec2), - pair_solreffriction=create_nmodel_batched_array(mjm.pair_solreffriction, dtype=wp.vec2), - pair_solimp=create_nmodel_batched_array(mjm.pair_solimp, dtype=types.vec5), - pair_margin=create_nmodel_batched_array(mjm.pair_margin, dtype=float), - pair_gap=create_nmodel_batched_array(mjm.pair_gap, dtype=float), - pair_friction=create_nmodel_batched_array(mjm.pair_friction, dtype=types.vec5), - condim_max=condim_max, # TODO(team): get max after filtering, - tendon_adr=wp.array(mjm.tendon_adr, dtype=int), - tendon_num=wp.array(mjm.tendon_num, dtype=int), - tendon_limited=wp.array(mjm.tendon_limited, dtype=int), - tendon_limited_adr=wp.array(np.nonzero(mjm.tendon_limited)[0], dtype=int), - tendon_actfrclimited=wp.array(mjm.tendon_actfrclimited, dtype=bool), - tendon_solref_lim=create_nmodel_batched_array(mjm.tendon_solref_lim, dtype=wp.vec2f), - tendon_solimp_lim=create_nmodel_batched_array(mjm.tendon_solimp_lim, dtype=types.vec5), - tendon_solref_fri=create_nmodel_batched_array(mjm.tendon_solref_fri, dtype=wp.vec2f), - tendon_solimp_fri=create_nmodel_batched_array(mjm.tendon_solimp_fri, dtype=types.vec5), - tendon_range=create_nmodel_batched_array(mjm.tendon_range, dtype=wp.vec2f), - tendon_actfrcrange=create_nmodel_batched_array(mjm.tendon_actfrcrange, dtype=wp.vec2), - tendon_margin=create_nmodel_batched_array(mjm.tendon_margin, dtype=float), - tendon_stiffness=create_nmodel_batched_array(mjm.tendon_stiffness, dtype=float), - tendon_damping=create_nmodel_batched_array(mjm.tendon_damping, dtype=float), - tendon_armature=create_nmodel_batched_array(mjm.tendon_armature, dtype=float), - tendon_frictionloss=create_nmodel_batched_array(mjm.tendon_frictionloss, dtype=float), - tendon_lengthspring=create_nmodel_batched_array(mjm.tendon_lengthspring, dtype=wp.vec2), - tendon_length0=create_nmodel_batched_array(mjm.tendon_length0, dtype=float), - tendon_invweight0=create_nmodel_batched_array(mjm.tendon_invweight0, dtype=float), - wrap_objid=wp.array(mjm.wrap_objid, dtype=int), - wrap_prm=wp.array(mjm.wrap_prm, dtype=float), - wrap_type=wp.array(mjm.wrap_type, dtype=int), - tendon_jnt_adr=wp.array(tendon_jnt_adr, dtype=int), - tendon_site_pair_adr=wp.array(tendon_site_pair_adr, dtype=int), - tendon_geom_adr=wp.array(tendon_geom_adr, dtype=int), - ten_wrapadr_site=wp.array(ten_wrapadr_site, dtype=int), - ten_wrapnum_site=wp.array(ten_wrapnum_site, dtype=int), - wrap_jnt_adr=wp.array(wrap_jnt_adr, dtype=int), - wrap_site_adr=wp.array(wrap_site_adr, dtype=int), - wrap_site_pair_adr=wp.array(wrap_site_pair_adr, dtype=int), - wrap_geom_adr=wp.array(wrap_geom_adr, dtype=int), - wrap_pulley_scale=wp.array(wrap_pulley_scale, dtype=float), sensor_type=wp.array(mjm.sensor_type, dtype=int), sensor_datatype=wp.array(mjm.sensor_datatype, dtype=int), sensor_objtype=wp.array(mjm.sensor_objtype, dtype=int), @@ -799,6 +807,63 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: sensor_dim=wp.array(mjm.sensor_dim, dtype=int), sensor_adr=wp.array(mjm.sensor_adr, dtype=int), sensor_cutoff=wp.array(mjm.sensor_cutoff, dtype=float), + plugin=wp.array(plugin_id, dtype=int), + plugin_attr=wp.array(plugin_attr, dtype=wp.vec3f), + M_rownnz=wp.array(mjm.M_rownnz, dtype=int), + M_rowadr=wp.array(mjm.M_rowadr, dtype=int), + M_colind=wp.array(mjm.M_colind, dtype=int), + mapM2M=wp.array(mjm.mapM2M, dtype=int), + # warp only fields: + nacttrnbody=np.sum(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY), + nsensorcollision=sum(nxn_pairid_collision >= 0), + nsensortaxel=sum(mjm.mesh_vertnum[mjm.sensor_objid[mjm.sensor_type == mujoco.mjtSensor.mjSENS_TACTILE]]), + nsensorcontact=np.sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT), + nrangefinder=sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER), + condim_max=condim_max, # TODO(team): get max after filtering, + nmaxpolygon=_compute_nmaxpolygon(mjm, tuple(geom_type_pair_count)), + nmaxmeshdeg=_compute_nmaxmeshdeg(mjm, tuple(geom_type_pair_count)), + has_sdf_geom=bool(np.any(mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF)), + block_dim=types.BlockDim(), + body_tree=body_tree, + mocap_bodyid=wp.array(mocap_bodyid, dtype=int), + body_fluid_ellipsoid=wp.array(body_fluid_ellipsoid, dtype=bool), + # these jnt_limited adrs are used in constraint.py + jnt_limited_slide_hinge_adr=wp.array( + np.nonzero( + mjm.jnt_limited & ((mjm.jnt_type == mujoco.mjtJoint.mjJNT_SLIDE) | (mjm.jnt_type == mujoco.mjtJoint.mjJNT_HINGE)) + )[0], + dtype=int, + ), + jnt_limited_ball_adr=wp.array( + np.nonzero(mjm.jnt_limited & (mjm.jnt_type == mujoco.mjtJoint.mjJNT_BALL))[0], + dtype=int, + ), + dof_tri_row=wp.array(dof_tri_row, dtype=int), + dof_tri_col=wp.array(dof_tri_col, dtype=int), + geom_pair_type_count=tuple(geom_type_pair_count), + geom_plugin_index=wp.array(geom_plugin_index, dtype=int), + nxn_geom_pair=wp.array(nxn_geom_pair, dtype=wp.vec2i), + nxn_geom_pair_filtered=wp.array(nxn_geom_pair_filtered, dtype=wp.vec2i), + nxn_pairid=wp.array(nxn_pairid, dtype=wp.vec2i), + nxn_pairid_filtered=wp.array(nxn_pairid_filtered, dtype=wp.vec2i), + eq_connect_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.CONNECT)[0], dtype=int), + eq_wld_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.WELD)[0], dtype=int), + eq_jnt_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.JOINT)[0], dtype=int), + eq_ten_adr=wp.array(np.nonzero(mjm.eq_type == types.EqType.TENDON)[0], dtype=int), + tendon_jnt_adr=wp.array(tendon_jnt_adr, dtype=int), + tendon_site_pair_adr=wp.array(tendon_site_pair_adr, dtype=int), + tendon_geom_adr=wp.array(tendon_geom_adr, dtype=int), + tendon_limited_adr=wp.array(np.nonzero(mjm.tendon_limited)[0], dtype=int), + ten_wrapadr_site=wp.array(ten_wrapadr_site, dtype=int), + ten_wrapnum_site=wp.array(ten_wrapnum_site, dtype=int), + wrap_jnt_adr=wp.array(wrap_jnt_adr, dtype=int), + wrap_site_adr=wp.array(wrap_site_adr, dtype=int), + wrap_site_pair_adr=wp.array(wrap_site_pair_adr, dtype=int), + wrap_geom_adr=wp.array(wrap_geom_adr, dtype=int), + wrap_pulley_scale=wp.array(wrap_pulley_scale, dtype=float), + actuator_moment_tiles_nv=actuator_moment_tiles_nv, + actuator_moment_tiles_nu=actuator_moment_tiles_nu, + actuator_trntype_body_adr=wp.array(np.nonzero(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)[0], dtype=int), sensor_pos_adr=wp.array( np.nonzero( (mjm.sensor_needstage == mujoco.mjtStage.mjSTAGE_POS) @@ -879,16 +944,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: sensor_rangefinder_bodyid=wp.array( mjm.site_bodyid[mjm.sensor_objid[mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER]], dtype=int ), - plugin=wp.array(plugin_id, dtype=int), - plugin_attr=wp.array(plugin_attr, dtype=wp.vec3f), - geom_plugin_index=wp.array(geom_plugin_index, dtype=int), - mat_texid=create_nmodel_batched_array(mjm.mat_texid, dtype=int), - mat_texrepeat=create_nmodel_batched_array(mjm.mat_texrepeat, dtype=wp.vec2), - mat_rgba=create_nmodel_batched_array(mjm.mat_rgba, dtype=wp.vec4), - actuator_trntype_body_adr=wp.array(np.nonzero(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)[0], dtype=int), - block_dim=types.BlockDim(), - geom_pair_type_count=tuple(geom_type_pair_count), - has_sdf_geom=bool(np.any(mjm.geom_type == mujoco.mjtGeom.mjGEOM_SDF)), + sensor_collision_start_adr=wp.array(sensor_collision_start_adr, dtype=int), taxel_vertadr=wp.array( [ j + mjm.mesh_vertadr[mjm.sensor_objid[i]] @@ -907,11 +963,36 @@ def put_model(mjm: mujoco.MjModel) -> types.Model: ], dtype=int, ), + qM_tiles=qM_tiles, + qLD_updates=qLD_updates, + qM_fullm_i=wp.array(qM_fullm_i, dtype=int), + qM_fullm_j=wp.array(qM_fullm_j, dtype=int), + qM_mulm_i=wp.array(qM_mulm_i, dtype=int), + qM_mulm_j=wp.array(qM_mulm_j, dtype=int), + qM_madr_ij=wp.array(qM_madr_ij, dtype=int), ) return m +def _get_padded_sizes(nv: int, njmax: int, nworld: int, is_sparse: bool, tile_size: int): + # if dense - we just pad to the next multiple of 4 for nv, to get the fast load path. + # we pad to the next multiple of tile_size for njmax to avoid out of bounds accesses. + # if sparse - we pad to the next multiple of tile_size for njmax, and nv. + + def round_up(x, multiple): + return ((x + multiple - 1) // multiple) * multiple + + njmax_padded = round_up(njmax, tile_size) + + if is_sparse: + nv_padded = round_up(nv, tile_size) + else: + nv_padded = round_up(nv, 4) + + return njmax_padded, nv_padded + + def make_data( mjm: mujoco.MjModel, nworld: int = 1, @@ -919,30 +1000,27 @@ def make_data( njmax: Optional[int] = None, naconmax: Optional[int] = None, ) -> types.Data: - """ - Creates a data object on device. + """Creates a data object on device. Args: - mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host). - nworld (int, optional): Number of worlds. Defaults to 1. - nworld (int, optional): The number of worlds. Defaults to 1. - nconmax (int, optional): Number of contacts to allocate per world. Contacts exist in large - heterogenous arrays: one world may have more than nconmax contacts. - njmax (int, optional): Number of constraints to allocate per world. Constraint arrays are - batched by world: no world may have more than njmax constraints. - naconmax (int, optional): Number of contacts to allocate for all worlds. Overrides nconmax. + mjm: The model containing kinematic and dynamic information (host). + nworld: 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. + 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. Returns: - Data: The data object containing the current state and output arrays (device). + The data object containing the current state and output arrays (device). """ - # TODO(team): move nconmax, njmax to Model? # TODO(team): improve heuristic for nconmax and njmax nconmax = nconmax or 20 njmax = njmax or nconmax * 6 - if nworld < 1 or nworld > MAX_WORLDS: - raise ValueError(f"nworld must be >= 1 and <= {MAX_WORLDS}") + if nworld < 1: + raise ValueError(f"nworld must be >= 1") if naconmax is None: if nconmax < 0: @@ -957,52 +1035,49 @@ def make_data( if mujoco.mj_isSparse(mjm): qM = wp.zeros((nworld, 1, mjm.nM), dtype=float) qLD = wp.zeros((nworld, 1, mjm.nC), dtype=float) - qM_integration = wp.zeros((nworld, 1, mjm.nM), dtype=float) - qLD_integration = wp.zeros((nworld, 1, mjm.nM), dtype=float) else: qM = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float) qLD = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float) - qM_integration = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float) - qLD_integration = wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float) condim = np.concatenate((mjm.geom_condim, mjm.pair_dim)) condim_max = np.max(condim) if len(condim) > 0 else 0 - max_npolygon = _max_npolygon(mjm) - max_meshdegree = _max_meshdegree(mjm) - nsensorcontact = np.sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT) - nrangefinder = sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER) + + if mujoco.mj_isSparse(mjm): + tile_size = types.TILE_SIZE_JTDAJ_SPARSE + else: + tile_size = types.TILE_SIZE_JTDAJ_DENSE + + njmax_padded, nv_padded = _get_padded_sizes(mjm.nv, njmax, nworld, mujoco.mj_isSparse(mjm), tile_size) + + # static geoms (attached to the world) have their poses calculated once during make_data instead + # of during each physics step. this speeds up scenes with many static geoms (e.g. terrains) + # TODO(team): remove this when we introduce dof islands + sleeping + mjd = mujoco.MjData(mjm) + mujoco.mj_kinematics(mjm, mjd) + geom_xpos = wp.array(np.tile(mjd.geom_xpos, (nworld, 1)), shape=(nworld, mjm.ngeom), dtype=wp.vec3) + geom_xmat = wp.array(np.tile(mjd.geom_xmat, (nworld, 1)), shape=(nworld, mjm.ngeom), dtype=wp.mat33) return types.Data( - nworld=nworld, - naconmax=naconmax, - njmax=njmax, solver_niter=wp.zeros(nworld, dtype=int), - nacon=wp.zeros(1, dtype=int), ne=wp.zeros(nworld, dtype=int), - ne_connect=wp.zeros(nworld, dtype=int), # warp only - ne_weld=wp.zeros(nworld, dtype=int), # warp only - ne_jnt=wp.zeros(nworld, dtype=int), # warp only - ne_ten=wp.zeros(nworld, dtype=int), # warp only nf=wp.zeros(nworld, dtype=int), nl=wp.zeros(nworld, dtype=int), nefc=wp.zeros(nworld, dtype=int), - nsolving=wp.zeros(1, dtype=int), # warp only time=wp.zeros(nworld, dtype=float), energy=wp.zeros(nworld, dtype=wp.vec2), qpos=wp.zeros((nworld, mjm.nq), dtype=float), qvel=wp.zeros((nworld, mjm.nv), dtype=float), act=wp.zeros((nworld, mjm.na), dtype=float), qacc_warmstart=wp.zeros((nworld, mjm.nv), dtype=float), - qacc_discrete=wp.zeros((nworld, mjm.nv), dtype=float), ctrl=wp.zeros((nworld, mjm.nu), dtype=float), qfrc_applied=wp.zeros((nworld, mjm.nv), dtype=float), xfrc_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), - fluid_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), eq_active=wp.array(np.tile(mjm.eq_active0, (nworld, 1)), dtype=bool), mocap_pos=wp.zeros((nworld, mjm.nmocap), dtype=wp.vec3), mocap_quat=wp.zeros((nworld, mjm.nmocap), dtype=wp.quat), qacc=wp.zeros((nworld, mjm.nv), dtype=float), act_dot=wp.zeros((nworld, mjm.na), dtype=float), + sensordata=wp.zeros((nworld, mjm.nsensordata), dtype=float), xpos=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3), xquat=wp.zeros((nworld, mjm.nbody), dtype=wp.quat), xmat=wp.zeros((nworld, mjm.nbody), dtype=wp.mat33), @@ -1010,9 +1085,8 @@ def make_data( ximat=wp.zeros((nworld, mjm.nbody), dtype=wp.mat33), xanchor=wp.zeros((nworld, mjm.njnt), dtype=wp.vec3), xaxis=wp.zeros((nworld, mjm.njnt), dtype=wp.vec3), - geom_skip=wp.zeros(mjm.ngeom, dtype=bool), # warp only - geom_xpos=wp.zeros((nworld, mjm.ngeom), dtype=wp.vec3), - geom_xmat=wp.zeros((nworld, mjm.ngeom), dtype=wp.mat33), + geom_xpos=geom_xpos, + geom_xmat=geom_xmat, site_xpos=wp.zeros((nworld, mjm.nsite), dtype=wp.vec3), site_xmat=wp.zeros((nworld, mjm.nsite), dtype=wp.mat33), cam_xpos=wp.zeros((nworld, mjm.ncam), dtype=wp.vec3), @@ -1024,13 +1098,19 @@ def make_data( cinert=wp.zeros((nworld, mjm.nbody), dtype=types.vec10), flexvert_xpos=wp.zeros((nworld, mjm.nflexvert), dtype=wp.vec3), flexedge_length=wp.zeros((nworld, mjm.nflexedge), dtype=wp.float32), - flexedge_velocity=wp.zeros((nworld, mjm.nflexedge), dtype=wp.float32), + ten_wrapadr=wp.zeros((nworld, mjm.ntendon), dtype=int), + ten_wrapnum=wp.zeros((nworld, mjm.ntendon), dtype=int), + ten_J=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float), + ten_length=wp.zeros((nworld, mjm.ntendon), dtype=float), + wrap_obj=wp.zeros((nworld, mjm.nwrap), dtype=wp.vec2i), + wrap_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector), actuator_length=wp.zeros((nworld, mjm.nu), dtype=float), actuator_moment=wp.zeros((nworld, mjm.nu, mjm.nv), dtype=float), crb=wp.zeros((nworld, mjm.nbody), dtype=types.vec10), qM=qM, qLD=qLD, qLDiagInv=wp.zeros((nworld, mjm.nv), dtype=float), + flexedge_velocity=wp.zeros((nworld, mjm.nflexedge), dtype=wp.float32), ten_velocity=wp.zeros((nworld, mjm.ntendon), dtype=float), actuator_velocity=wp.zeros((nworld, mjm.nu), dtype=float), cvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), @@ -1043,13 +1123,15 @@ def make_data( qfrc_passive=wp.zeros((nworld, mjm.nv), dtype=float), subtree_linvel=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3), subtree_angmom=wp.zeros((nworld, mjm.nbody), dtype=wp.vec3), - subtree_bodyvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), # warp only actuator_force=wp.zeros((nworld, mjm.nu), dtype=float), qfrc_actuator=wp.zeros((nworld, mjm.nv), dtype=float), qfrc_smooth=wp.zeros((nworld, mjm.nv), dtype=float), qacc_smooth=wp.zeros((nworld, mjm.nv), dtype=float), qfrc_constraint=wp.zeros((nworld, mjm.nv), dtype=float), qfrc_inverse=wp.zeros((nworld, mjm.nv), dtype=float), + cacc=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + cfrc_int=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), + cfrc_ext=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), contact=types.Contact( dist=wp.zeros((naconmax,), dtype=float), pos=wp.zeros((naconmax,), dtype=wp.vec3f), @@ -1063,14 +1145,16 @@ def make_data( geom=wp.zeros((naconmax,), dtype=wp.vec2i), efc_address=wp.zeros((naconmax, np.maximum(1, 2 * (condim_max - 1))), dtype=int), worldid=wp.zeros((naconmax,), dtype=int), + type=wp.zeros((naconmax,), dtype=int), + geomcollisionid=wp.empty((naconmax,), dtype=int), ), efc=types.Constraint( type=wp.zeros((nworld, njmax), dtype=int), id=wp.zeros((nworld, njmax), dtype=int), - J=wp.zeros((nworld, njmax, mjm.nv), dtype=float), + J=wp.zeros((nworld, njmax_padded, nv_padded), dtype=float), pos=wp.zeros((nworld, njmax), dtype=float), margin=wp.zeros((nworld, njmax), dtype=float), - D=wp.zeros((nworld, njmax), dtype=float), + D=wp.zeros((nworld, njmax_padded), dtype=float), vel=wp.zeros((nworld, njmax), dtype=float), aref=wp.zeros((nworld, njmax), dtype=float), frictionloss=wp.zeros((nworld, njmax), dtype=float), @@ -1087,105 +1171,34 @@ def make_data( gauss=wp.zeros((nworld,), dtype=float), cost=wp.zeros((nworld,), dtype=float), prev_cost=wp.zeros((nworld,), dtype=float), - state=wp.zeros((nworld, njmax), dtype=int), + state=wp.zeros((nworld, njmax_padded), dtype=int), mv=wp.zeros((nworld, mjm.nv), dtype=float), jv=wp.zeros((nworld, njmax), dtype=float), quad=wp.zeros((nworld, njmax), dtype=wp.vec3f), quad_gauss=wp.zeros((nworld,), dtype=wp.vec3f), - h=wp.zeros((nworld, mjm.nv, mjm.nv), dtype=float), + h=wp.zeros((nworld, nv_padded, nv_padded), dtype=float), alpha=wp.zeros((nworld,), dtype=float), prev_grad=wp.zeros((nworld, mjm.nv), dtype=float), prev_Mgrad=wp.zeros((nworld, mjm.nv), dtype=float), beta=wp.zeros((nworld,), dtype=float), done=wp.zeros((nworld,), dtype=bool), - # linesearch - cost_candidate=wp.zeros((nworld, mjm.opt.ls_iterations), dtype=float), - ), - # RK4 - qpos_t0=wp.zeros((nworld, mjm.nq), dtype=float), - qvel_t0=wp.zeros((nworld, mjm.nv), dtype=float), - act_t0=wp.zeros((nworld, mjm.na), dtype=float), - qvel_rk=wp.zeros((nworld, mjm.nv), dtype=float), - qacc_rk=wp.zeros((nworld, mjm.nv), dtype=float), - act_dot_rk=wp.zeros((nworld, mjm.na), dtype=float), - # euler + implicit integration - qfrc_integration=wp.zeros((nworld, mjm.nv), dtype=float), - qacc_integration=wp.zeros((nworld, mjm.nv), dtype=float), - act_vel_integration=wp.zeros((nworld, mjm.nu), dtype=float), - qM_integration=qM_integration, - qLD_integration=qLD_integration, - qLDiagInv_integration=wp.zeros((nworld, mjm.nv), dtype=float), - # sweep-and-prune broadphase - sap_projection_lower=wp.zeros((nworld, mjm.ngeom, 2), dtype=float), - sap_projection_upper=wp.zeros((nworld, mjm.ngeom), dtype=float), - sap_sort_index=wp.zeros((nworld, mjm.ngeom, 2), dtype=int), - sap_range=wp.zeros((nworld, mjm.ngeom), dtype=int), - sap_cumulative_sum=wp.zeros((nworld, mjm.ngeom), dtype=int), - sap_segment_index=wp.array( - np.array([i * mjm.ngeom if i < nworld + 1 else 0 for i in range(2 * nworld)]).reshape((nworld, 2)), dtype=int ), + # warp only fields: + nworld=nworld, + naconmax=naconmax, + njmax=njmax, + nacon=wp.zeros(1, dtype=int), + ne_connect=wp.zeros(nworld, dtype=int), + ne_weld=wp.zeros(nworld, dtype=int), + ne_jnt=wp.zeros(nworld, dtype=int), + ne_ten=wp.zeros(nworld, dtype=int), + nsolving=wp.zeros(1, dtype=int), + subtree_bodyvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), # collision driver collision_pair=wp.zeros((naconmax,), dtype=wp.vec2i), - collision_pairid=wp.zeros((naconmax,), dtype=int), + collision_pairid=wp.zeros((naconmax,), dtype=wp.vec2i), collision_worldid=wp.zeros((naconmax,), dtype=int), ncollision=wp.zeros((1,), dtype=int), - # narrowphase (EPA polytope) - epa_vert=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=wp.vec3), - epa_vert1=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=wp.vec3), - epa_vert2=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=wp.vec3), - epa_vert_index1=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=int), - epa_vert_index2=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=int), - epa_face=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=wp.vec3i), - epa_pr=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=wp.vec3), - epa_norm2=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=float), - epa_index=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=int), - epa_map=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=int), - epa_horizon=wp.zeros(shape=(naconmax, 2 * types.MJ_MAX_EPAHORIZON), dtype=int), - multiccd_polygon=wp.zeros(shape=(naconmax, 2 * max_npolygon), dtype=wp.vec3), - multiccd_clipped=wp.zeros(shape=(naconmax, 2 * max_npolygon), dtype=wp.vec3), - multiccd_pnormal=wp.zeros(shape=(naconmax, max_npolygon), dtype=wp.vec3), - multiccd_pdist=wp.zeros(shape=(naconmax, max_npolygon), dtype=float), - multiccd_idx1=wp.zeros(shape=(naconmax, max_meshdegree), dtype=int), - multiccd_idx2=wp.zeros(shape=(naconmax, max_meshdegree), dtype=int), - multiccd_n1=wp.zeros(shape=(naconmax, max_meshdegree), dtype=wp.vec3), - multiccd_n2=wp.zeros(shape=(naconmax, max_meshdegree), dtype=wp.vec3), - multiccd_endvert=wp.zeros(shape=(naconmax, max_meshdegree), dtype=wp.vec3), - multiccd_face1=wp.zeros(shape=(naconmax, max_npolygon), dtype=wp.vec3), - multiccd_face2=wp.zeros(shape=(naconmax, max_npolygon), dtype=wp.vec3), - # rne_postconstraint - cacc=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), - cfrc_int=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), - cfrc_ext=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), - # tendon - ten_length=wp.zeros((nworld, mjm.ntendon), dtype=float), - ten_J=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float), - ten_Jdot=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float), - ten_bias_coef=wp.zeros((nworld, mjm.ntendon), dtype=float), - ten_wrapadr=wp.zeros((nworld, mjm.ntendon), dtype=int), - ten_wrapnum=wp.zeros((nworld, mjm.ntendon), dtype=int), - ten_actfrc=wp.zeros((nworld, mjm.ntendon), dtype=float), - wrap_obj=wp.zeros((nworld, mjm.nwrap), dtype=wp.vec2i), - wrap_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector), - wrap_geom_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector), - # sensors - sensordata=wp.zeros((nworld, mjm.nsensordata), dtype=float), - sensor_rangefinder_pnt=wp.zeros((nworld, nrangefinder), dtype=wp.vec3), - sensor_rangefinder_vec=wp.zeros((nworld, nrangefinder), dtype=wp.vec3), - sensor_rangefinder_dist=wp.zeros((nworld, nrangefinder), dtype=float), - sensor_rangefinder_geomid=wp.zeros((nworld, nrangefinder), dtype=int), - sensor_contact_nmatch=wp.zeros((nworld, nsensorcontact), dtype=int), - sensor_contact_matchid=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=int), - sensor_contact_criteria=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=float), - sensor_contact_direction=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=float), - # ray - ray_bodyexclude=wp.zeros(1, dtype=int), - ray_dist=wp.zeros((nworld, 1), dtype=float), - ray_geomid=wp.zeros((nworld, 1), dtype=int), - # mul_m - energy_vel_mul_m_skip=wp.zeros((nworld,), dtype=bool), - inverse_mul_m_skip=wp.zeros((nworld,), dtype=bool), - # actuator - actuator_trntype_body_ncon=wp.zeros((nworld, np.sum(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)), dtype=int), ) @@ -1197,21 +1210,20 @@ def put_data( njmax: Optional[int] = None, naconmax: Optional[int] = None, ) -> types.Data: - """ - Moves data from host to a device. + """Moves data from host to a device. Args: - mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host). - mjd (mujoco.MjData): The data object containing current state and output arrays (host). - nworld (int, optional): The number of worlds. Defaults to 1. - nconmax (int, optional): Number of contacts to allocate per world. Contacts exist in large - heterogenous arrays: one world may have more than nconmax contacts. - njmax (int, optional): Number of constraints to allocate per world. Constraint arrays are - batched by world: no world may have more than njmax constraints. - naconmax (int, optional): Number of contacts to allocate for all worlds. Overrides nconmax. + mjm: The model containing kinematic and dynamic information (host). + mjd: The data object containing current state and output arrays (host). + 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. + 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. Returns: - Data: The data object containing the current state and output arrays (device). + The data object containing the current state and output arrays (device). """ # TODO(team): move nconmax and njmax to Model? # TODO(team): decide what to do about uninitialized warp-only fields created by put_data @@ -1221,8 +1233,8 @@ def put_data( nconmax = nconmax or max(5, 4 * mjd.ncon) njmax = njmax or max(5, 4 * mjd.nefc) - if nworld < 1 or nworld > MAX_WORLDS: - raise ValueError(f"nworld must be >= 1 and <= {MAX_WORLDS}") + if nworld < 1: + raise ValueError(f"nworld must be >= 1") if naconmax is None: if nconmax < 0: @@ -1241,15 +1253,14 @@ def put_data( if mjd.nefc > njmax: raise ValueError(f"njmax overflow (njmax must be >= {mjd.nefc})") - max_npolygon = _max_npolygon(mjm) - max_meshdegree = _max_meshdegree(mjm) + # ensure static geom positions are computed + # TODO: remove once MjData creation semantics are fixed + mujoco.mj_kinematics(mjm, mjd) # calculate some fields that cannot be easily computed inline: if mujoco.mj_isSparse(mjm): qM = np.expand_dims(mjd.qM, axis=0) qLD = np.expand_dims(mjd.qLD, axis=0) - qM_integration = np.zeros((1, mjm.nM), dtype=float) - qLD_integration = np.zeros((1, mjm.nM), dtype=float) efc_J = np.zeros((mjd.nefc, mjm.nv)) mujoco.mju_sparse2dense(efc_J, mjd.efc_J, mjd.efc_J_rownnz, mjd.efc_J_rowadr, mjd.efc_J_colind) ten_J = np.zeros((mjm.ntendon, mjm.nv)) @@ -1267,8 +1278,6 @@ def put_data( qLD = np.zeros((mjm.nv, mjm.nv)) else: qLD = np.linalg.cholesky(qM) - qM_integration = np.zeros((mjm.nv, mjm.nv), dtype=float) - qLD_integration = np.zeros((mjm.nv, mjm.nv), dtype=float) efc_J = mjd.efc_J.reshape((mjd.nefc, mjm.nv)) ten_J = mjd.ten_J.reshape((mjm.ntendon, mjm.nv)) @@ -1305,10 +1314,17 @@ def put_data( ne_jnt = int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_JOINT) & mjd.eq_active)) ne_ten = int(np.sum((mjm.eq_type == mujoco.mjtEq.mjEQ_TENDON) & mjd.eq_active)) + if mujoco.mj_isSparse(mjm): + tile_size = types.TILE_SIZE_JTDAJ_SPARSE + else: + tile_size = types.TILE_SIZE_JTDAJ_DENSE + + njmax_padded, nv_padded = _get_padded_sizes(mjm.nv, njmax, nworld, mujoco.mj_isSparse(mjm), tile_size) + efc_type_fill = np.zeros((nworld, njmax)) efc_id_fill = np.zeros((nworld, njmax)) - efc_J_fill = np.zeros((nworld, njmax, mjm.nv)) - efc_D_fill = np.zeros((nworld, njmax)) + efc_J_fill = np.zeros((nworld, njmax_padded, nv_padded)) + efc_D_fill = np.zeros((nworld, njmax_padded)) efc_vel_fill = np.zeros((nworld, njmax)) efc_pos_fill = np.zeros((nworld, njmax)) efc_aref_fill = np.zeros((nworld, njmax)) @@ -1319,7 +1335,7 @@ def put_data( nefc = mjd.nefc efc_type_fill[:, :nefc] = np.tile(mjd.efc_type, (nworld, 1)) efc_id_fill[:, :nefc] = np.tile(mjd.efc_id, (nworld, 1)) - efc_J_fill[:, :nefc, :] = np.tile(efc_J, (nworld, 1, 1)) + efc_J_fill[:, :nefc, : mjm.nv] = np.tile(efc_J, (nworld, 1, 1)) efc_D_fill[:, :nefc] = np.tile(mjd.efc_D, (nworld, 1)) efc_vel_fill[:, :nefc] = np.tile(mjd.efc_vel, (nworld, 1)) efc_pos_fill[:, :nefc] = np.tile(mjd.efc_pos, (nworld, 1)) @@ -1328,9 +1344,6 @@ def put_data( efc_force_fill[:, :nefc] = np.tile(mjd.efc_force, (nworld, 1)) efc_margin_fill[:, :nefc] = np.tile(mjd.efc_margin, (nworld, 1)) - nsensorcontact = np.sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT) - nrangefinder = sum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_RANGEFINDER) - # some helper functions to simplify the data field definitions below def arr(x, dtype=None): @@ -1357,36 +1370,26 @@ def put_data( return arr(np.pad(x, width), dtype) return types.Data( - nworld=nworld, - naconmax=naconmax, - njmax=njmax, solver_niter=tile(mjd.solver_niter[0]), - nacon=arr([mjd.ncon * nworld]), ne=wp.full(shape=(nworld), value=mjd.ne), - ne_connect=wp.full(shape=(nworld), value=ne_connect), - ne_weld=wp.full(shape=(nworld), value=ne_weld), - ne_jnt=wp.full(shape=(nworld), value=ne_jnt), - ne_ten=wp.full(shape=(nworld), value=ne_ten), nf=wp.full(shape=(nworld), value=mjd.nf), nl=wp.full(shape=(nworld), value=mjd.nl), nefc=wp.full(shape=(nworld), value=mjd.nefc), - nsolving=arr([nworld]), time=arr(mjd.time * np.ones(nworld)), energy=tile(mjd.energy, dtype=wp.vec2), qpos=tile(mjd.qpos), qvel=tile(mjd.qvel), act=tile(mjd.act), qacc_warmstart=tile(mjd.qacc_warmstart), - qacc_discrete=wp.zeros((nworld, mjm.nv), dtype=float), ctrl=tile(mjd.ctrl), qfrc_applied=tile(mjd.qfrc_applied), xfrc_applied=tile(mjd.xfrc_applied, dtype=wp.spatial_vector), - fluid_applied=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), eq_active=tile(mjd.eq_active.astype(bool)), mocap_pos=tile(mjd.mocap_pos, dtype=wp.vec3), mocap_quat=tile(mjd.mocap_quat, dtype=wp.quat), qacc=tile(mjd.qacc), act_dot=tile(mjd.act_dot), + sensordata=tile(mjd.sensordata), xpos=tile(mjd.xpos, dtype=wp.vec3), xquat=tile(mjd.xquat, dtype=wp.quat), xmat=tile(mjd.xmat, dtype=wp.mat33), @@ -1394,7 +1397,6 @@ def put_data( ximat=tile(mjd.ximat, dtype=wp.mat33), xanchor=tile(mjd.xanchor, dtype=wp.vec3), xaxis=tile(mjd.xaxis, dtype=wp.vec3), - geom_skip=wp.zeros(mjm.ngeom, dtype=bool), # warp only geom_xpos=tile(mjd.geom_xpos, dtype=wp.vec3), geom_xmat=tile(mjd.geom_xmat, dtype=wp.mat33), site_xpos=tile(mjd.site_xpos, dtype=wp.vec3), @@ -1408,13 +1410,19 @@ def put_data( cinert=tile(mjd.cinert, dtype=types.vec10), flexvert_xpos=tile(mjd.flexvert_xpos, dtype=wp.vec3), flexedge_length=tile(mjd.flexedge_length), - flexedge_velocity=tile(mjd.flexedge_velocity), + ten_wrapadr=tile(mjd.ten_wrapadr), + ten_wrapnum=tile(mjd.ten_wrapnum), + ten_J=tile(ten_J), + ten_length=tile(mjd.ten_length), + wrap_obj=tile(mjd.wrap_obj, dtype=wp.vec2i), + wrap_xpos=tile(mjd.wrap_xpos, dtype=wp.spatial_vector), actuator_length=tile(mjd.actuator_length), actuator_moment=tile(actuator_moment), crb=tile(mjd.crb, dtype=types.vec10), qM=tile(qM), qLD=tile(qLD), qLDiagInv=tile(mjd.qLDiagInv), + flexedge_velocity=tile(mjd.flexedge_velocity), ten_velocity=tile(mjd.ten_velocity), actuator_velocity=tile(mjd.actuator_velocity), cvel=tile(mjd.cvel, dtype=wp.spatial_vector), @@ -1427,13 +1435,15 @@ def put_data( qfrc_passive=tile(mjd.qfrc_passive), subtree_linvel=tile(mjd.subtree_linvel, dtype=wp.vec3), subtree_angmom=tile(mjd.subtree_angmom, dtype=wp.vec3), - subtree_bodyvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), actuator_force=tile(mjd.actuator_force), qfrc_actuator=tile(mjd.qfrc_actuator), qfrc_smooth=tile(mjd.qfrc_smooth), qacc_smooth=tile(mjd.qacc_smooth), qfrc_constraint=tile(mjd.qfrc_constraint), qfrc_inverse=tile(mjd.qfrc_inverse), + cacc=tile(mjd.cacc, dtype=wp.spatial_vector), + cfrc_int=tile(mjd.cfrc_int, dtype=wp.spatial_vector), + cfrc_ext=tile(mjd.cfrc_ext, dtype=wp.spatial_vector), contact=types.Contact( dist=padtile(mjd.contact.dist, naconmax), pos=padtile(mjd.contact.pos, naconmax, dtype=wp.vec3), @@ -1447,6 +1457,8 @@ def put_data( geom=padtile(mjd.contact.geom, naconmax, dtype=wp.vec2i), efc_address=arr(contact_efc_address), worldid=arr(contact_worldid), + type=wp.ones((naconmax,), dtype=int), # TODO(team): set values + geomcollisionid=wp.empty((naconmax,), dtype=int), # TODO(team): set values ), efc=types.Constraint( type=wp.array2d(efc_type_fill, dtype=int), @@ -1471,102 +1483,34 @@ def put_data( gauss=wp.empty(shape=(nworld,), dtype=float), cost=wp.empty(shape=(nworld,), dtype=float), prev_cost=wp.empty(shape=(nworld,), dtype=float), - state=wp.empty(shape=(nworld, njmax), dtype=int), + state=wp.zeros(shape=(nworld, njmax_padded), dtype=int), mv=wp.empty(shape=(nworld, mjm.nv), dtype=float), jv=wp.empty(shape=(nworld, njmax), dtype=float), quad=wp.empty(shape=(nworld, njmax), dtype=wp.vec3f), quad_gauss=wp.empty(shape=(nworld,), dtype=wp.vec3f), - h=wp.empty(shape=(nworld, mjm.nv, mjm.nv), dtype=float), + h=wp.zeros(shape=(nworld, nv_padded, nv_padded), dtype=float), alpha=wp.empty(shape=(nworld,), dtype=float), prev_grad=wp.empty(shape=(nworld, mjm.nv), dtype=float), prev_Mgrad=wp.empty(shape=(nworld, mjm.nv), dtype=float), beta=wp.empty(shape=(nworld,), dtype=float), done=wp.empty(shape=(nworld,), dtype=bool), - cost_candidate=wp.empty(shape=(nworld, mjm.opt.ls_iterations), dtype=float), ), - # TODO(team): skip allocation if integrator != RK4 - qpos_t0=wp.empty((nworld, mjm.nq), dtype=float), - qvel_t0=wp.empty((nworld, mjm.nv), dtype=float), - act_t0=wp.empty((nworld, mjm.na), dtype=float), - qvel_rk=wp.empty((nworld, mjm.nv), dtype=float), - qacc_rk=wp.empty((nworld, mjm.nv), dtype=float), - act_dot_rk=wp.empty((nworld, mjm.na), dtype=float), - # TODO(team): skip allocation if integrator != euler | implicit - qfrc_integration=wp.zeros((nworld, mjm.nv), dtype=float), - qacc_integration=wp.zeros((nworld, mjm.nv), dtype=float), - act_vel_integration=wp.zeros((nworld, mjm.nu), dtype=float), - qM_integration=tile(qM_integration), - qLD_integration=tile(qLD_integration), - qLDiagInv_integration=wp.zeros((nworld, mjm.nv), dtype=float), - # TODO(team): skip allocation if broadphase != sap - sap_projection_lower=wp.zeros((nworld, mjm.ngeom, 2), dtype=float), - sap_projection_upper=wp.zeros((nworld, mjm.ngeom), dtype=float), - sap_sort_index=wp.zeros((nworld, mjm.ngeom, 2), dtype=int), - sap_range=wp.zeros((nworld, mjm.ngeom), dtype=int), - sap_cumulative_sum=wp.zeros((nworld, mjm.ngeom), dtype=int), - sap_segment_index=arr(np.array([i * mjm.ngeom if i < nworld + 1 else 0 for i in range(2 * nworld)]).reshape((nworld, 2))), + # warp only fields: + nworld=nworld, + naconmax=naconmax, + njmax=njmax, + nacon=arr([mjd.ncon * nworld]), + ne_connect=wp.full(shape=(nworld), value=ne_connect), + ne_weld=wp.full(shape=(nworld), value=ne_weld), + ne_jnt=wp.full(shape=(nworld), value=ne_jnt), + ne_ten=wp.full(shape=(nworld), value=ne_ten), + nsolving=arr([nworld]), + subtree_bodyvel=wp.zeros((nworld, mjm.nbody), dtype=wp.spatial_vector), # collision driver collision_pair=wp.empty(naconmax, dtype=wp.vec2i), - collision_pairid=wp.empty(naconmax, dtype=int), + collision_pairid=wp.empty(naconmax, dtype=wp.vec2i), collision_worldid=wp.empty(naconmax, dtype=int), ncollision=wp.zeros(1, dtype=int), - # narrowphase (EPA polytope) - epa_vert=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=wp.vec3), - epa_vert1=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=wp.vec3), - epa_vert2=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=wp.vec3), - epa_vert_index1=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=int), - epa_vert_index2=wp.zeros(shape=(naconmax, 5 + mjm.opt.ccd_iterations), dtype=int), - epa_face=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=wp.vec3i), - epa_pr=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=wp.vec3), - epa_norm2=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=float), - epa_index=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=int), - epa_map=wp.zeros(shape=(naconmax, 6 + types.MJ_MAX_EPAFACES * mjm.opt.ccd_iterations), dtype=int), - epa_horizon=wp.zeros(shape=(naconmax, 2 * types.MJ_MAX_EPAHORIZON), dtype=int), - multiccd_polygon=wp.zeros(shape=(naconmax, 2 * max_npolygon), dtype=wp.vec3), - multiccd_clipped=wp.zeros(shape=(naconmax, 2 * max_npolygon), dtype=wp.vec3), - multiccd_pnormal=wp.zeros(shape=(naconmax, max_npolygon), dtype=wp.vec3), - multiccd_pdist=wp.zeros(shape=(naconmax, max_npolygon), dtype=float), - multiccd_idx1=wp.zeros(shape=(naconmax, max_meshdegree), dtype=int), - multiccd_idx2=wp.zeros(shape=(naconmax, max_meshdegree), dtype=int), - multiccd_n1=wp.zeros(shape=(naconmax, max_meshdegree), dtype=wp.vec3), - multiccd_n2=wp.zeros(shape=(naconmax, max_meshdegree), dtype=wp.vec3), - multiccd_endvert=wp.zeros(shape=(naconmax, max_meshdegree), dtype=wp.vec3), - multiccd_face1=wp.zeros(shape=(naconmax, max_npolygon), dtype=wp.vec3), - multiccd_face2=wp.zeros(shape=(naconmax, max_npolygon), dtype=wp.vec3), - # rne_postconstraint but also smooth - cacc=tile(mjd.cacc, dtype=wp.spatial_vector), - cfrc_int=tile(mjd.cfrc_int, dtype=wp.spatial_vector), - cfrc_ext=tile(mjd.cfrc_ext, dtype=wp.spatial_vector), - # tendon - ten_length=tile(mjd.ten_length), - ten_J=tile(ten_J), - ten_Jdot=wp.zeros((nworld, mjm.ntendon, mjm.nv), dtype=float), - ten_bias_coef=wp.zeros((nworld, mjm.ntendon), dtype=float), - ten_wrapadr=tile(mjd.ten_wrapadr), - ten_wrapnum=tile(mjd.ten_wrapnum), - ten_actfrc=wp.zeros((nworld, mjm.ntendon), dtype=float), - wrap_obj=tile(mjd.wrap_obj, dtype=wp.vec2i), - wrap_xpos=tile(mjd.wrap_xpos, dtype=wp.spatial_vector), - wrap_geom_xpos=wp.zeros((nworld, mjm.nwrap), dtype=wp.spatial_vector), - # sensors - sensordata=tile(mjd.sensordata), - sensor_rangefinder_pnt=wp.zeros((nworld, nrangefinder), dtype=wp.vec3), - sensor_rangefinder_vec=wp.zeros((nworld, nrangefinder), dtype=wp.vec3), - sensor_rangefinder_dist=wp.zeros((nworld, nrangefinder), dtype=float), - sensor_rangefinder_geomid=wp.zeros((nworld, nrangefinder), dtype=int), - sensor_contact_nmatch=wp.zeros((nworld, nsensorcontact), dtype=int), - sensor_contact_matchid=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=int), - sensor_contact_criteria=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=float), - sensor_contact_direction=wp.zeros((nworld, nsensorcontact, types.MJ_MAXCONPAIR), dtype=float), - # ray - ray_bodyexclude=wp.zeros(1, dtype=int), - ray_dist=wp.zeros((nworld, 1), dtype=float), - ray_geomid=wp.zeros((nworld, 1), dtype=int), - # mul_m - energy_vel_mul_m_skip=wp.zeros((nworld,), dtype=bool), - inverse_mul_m_skip=wp.zeros((nworld,), dtype=bool), - # actuator - actuator_trntype_body_ncon=wp.zeros((nworld, np.sum(mjm.actuator_trntype == mujoco.mjtTrn.mjTRN_BODY)), dtype=int), ) @@ -1578,86 +1522,119 @@ def get_data_into( """Gets data from a device into an existing mujoco.MjData. Args: - result (mujoco.MjData): The data object containing the current state and output arrays - (host). - mjm (mujoco.MjModel): The model containing kinematic and dynamic information (host). - d (Data): The data object containing the current state and output arrays (device). + result: The data object containing the current state and output arrays (host). + mjm: The model containing kinematic and dynamic information (host). + d: The data object containing the current state and output arrays (device). """ if d.nworld > 1: raise NotImplementedError("only nworld == 1 supported for now") - result.solver_niter[0] = d.solver_niter.numpy()[0] - - nacon = d.nacon.numpy()[0] - nefc = d.nefc.numpy()[0] + # nacon and nefc can overflow. in that case, only pull up to the max contacts and constraints + nacon = min(d.nacon.numpy()[0], d.naconmax) + nefc = min(d.nefc.numpy()[0], d.njmax) if nacon != result.ncon or nefc != result.nefc: - mujoco._functions._realloc_con_efc(result, ncon=nacon, nefc=nefc) + # TODO(team): if sparse, set nJ based on sparse efc_J + mujoco._functions._realloc_con_efc(result, ncon=nacon, nefc=nefc, nJ=nefc * mjm.nv) + ne = d.ne.numpy()[0] + nf = d.nf.numpy()[0] + nl = d.nl.numpy()[0] + + # efc indexing + # mujoco expects contigious efc ordering for contacts + # this ordering is not guarenteed with mujoco warp, we enforce order here + if nacon > 0: + efc_idx_efl = np.arange(ne + nf + nl) + + contact_dim = d.contact.dim.numpy() + contact_efc_address = d.contact.efc_address.numpy() + + efc_idx_c = [] + contact_efc_address_ordered = [ne + nf + nl] + for i in range(nacon): + dim = contact_dim[i] + if mjm.opt.cone == mujoco.mjtCone.mjCONE_PYRAMIDAL: + ndim = np.maximum(1, 2 * (dim - 1)) + else: + ndim = dim + efc_idx_c.append(contact_efc_address[i, :ndim]) + if i < nacon - 1: + contact_efc_address_ordered.append(contact_efc_address_ordered[-1] + ndim) + efc_idx = np.concatenate((efc_idx_efl, *efc_idx_c)) + contact_efc_address_ordered = np.array(contact_efc_address_ordered) + else: + efc_idx = np.array(np.arange(nefc)) + contact_efc_address_ordered = np.empty(0) + + efc_idx = efc_idx[:nefc] # dont emit indices for overflow constraints + + result.solver_niter[0] = d.solver_niter.numpy()[0] + result.ncon = nacon + result.ne = ne + result.nf = nf + result.nl = nl result.time = d.time.numpy()[0] - result.energy = d.energy.numpy()[0] - result.ne = d.ne.numpy()[0] + result.energy[:] = d.energy.numpy()[0] result.qpos[:] = d.qpos.numpy()[0] result.qvel[:] = d.qvel.numpy()[0] - result.qacc_warmstart = d.qacc_warmstart.numpy()[0] - result.qfrc_applied = d.qfrc_applied.numpy()[0] - result.mocap_pos = d.mocap_pos.numpy()[0] - result.mocap_quat = d.mocap_quat.numpy()[0] - result.qacc = d.qacc.numpy()[0] - result.xanchor = d.xanchor.numpy()[0] - result.xaxis = d.xaxis.numpy()[0] - result.xmat = d.xmat.numpy().reshape((-1, 9)) - result.xpos = d.xpos.numpy()[0] - result.xquat = d.xquat.numpy()[0] - result.xipos = d.xipos.numpy()[0] - result.ximat = d.ximat.numpy().reshape((-1, 9)) - result.subtree_com = d.subtree_com.numpy()[0] - result.geom_xpos = d.geom_xpos.numpy()[0] - result.geom_xmat = d.geom_xmat.numpy().reshape((-1, 9)) - result.site_xpos = d.site_xpos.numpy()[0] - result.site_xmat = d.site_xmat.numpy().reshape((-1, 9)) - result.cam_xpos = d.cam_xpos.numpy()[0] - result.cam_xmat = d.cam_xmat.numpy().reshape((-1, 9)) - result.light_xpos = d.light_xpos.numpy()[0] - result.light_xdir = d.light_xdir.numpy()[0] - result.cinert = d.cinert.numpy()[0] - result.flexvert_xpos = d.flexvert_xpos.numpy()[0] - result.flexedge_length = d.flexedge_length.numpy()[0] - result.flexedge_velocity = d.flexedge_velocity.numpy()[0] - result.cdof = d.cdof.numpy()[0] - result.crb = d.crb.numpy()[0] - result.qLDiagInv = d.qLDiagInv.numpy()[0] - result.ctrl = d.ctrl.numpy()[0] - result.ten_velocity = d.ten_velocity.numpy()[0] - result.actuator_velocity = d.actuator_velocity.numpy()[0] - result.actuator_force = d.actuator_force.numpy()[0] - result.actuator_length = d.actuator_length.numpy()[0] + result.act[:] = d.act.numpy()[0] + result.qacc_warmstart[:] = d.qacc_warmstart.numpy()[0] + result.ctrl[:] = d.ctrl.numpy()[0] + result.qfrc_applied[:] = d.qfrc_applied.numpy()[0] + result.xfrc_applied[:] = d.xfrc_applied.numpy()[0] + result.eq_active[:] = d.eq_active.numpy()[0] + result.mocap_pos[:] = d.mocap_pos.numpy()[0] + result.mocap_quat[:] = d.mocap_quat.numpy()[0] + result.qacc[:] = d.qacc.numpy()[0] + result.act_dot[:] = d.act_dot.numpy()[0] + result.xpos[:] = d.xpos.numpy()[0] + result.xquat[:] = d.xquat.numpy()[0] + result.xmat[:] = d.xmat.numpy().reshape((-1, 9)) + result.xipos[:] = d.xipos.numpy()[0] + result.ximat[:] = d.ximat.numpy().reshape((-1, 9)) + result.xanchor[:] = d.xanchor.numpy()[0] + result.xaxis[:] = d.xaxis.numpy()[0] + result.geom_xpos[:] = d.geom_xpos.numpy()[0] + result.geom_xmat[:] = d.geom_xmat.numpy().reshape((-1, 9)) + result.site_xpos[:] = d.site_xpos.numpy()[0] + result.site_xmat[:] = d.site_xmat.numpy().reshape((-1, 9)) + result.cam_xpos[:] = d.cam_xpos.numpy()[0] + result.cam_xmat[:] = d.cam_xmat.numpy().reshape((-1, 9)) + result.light_xpos[:] = d.light_xpos.numpy()[0] + result.light_xdir[:] = d.light_xdir.numpy()[0] + result.subtree_com[:] = d.subtree_com.numpy()[0] + result.cdof[:] = d.cdof.numpy()[0] + result.cinert[:] = d.cinert.numpy()[0] + result.flexvert_xpos[:] = d.flexvert_xpos.numpy()[0] + result.flexedge_length[:] = d.flexedge_length.numpy()[0] + result.flexedge_velocity[:] = d.flexedge_velocity.numpy()[0] + result.actuator_length[:] = d.actuator_length.numpy()[0] mujoco.mju_dense2sparse( - result.actuator_moment, - d.actuator_moment.numpy()[0], - result.moment_rownnz, - result.moment_rowadr, - result.moment_colind, + result.actuator_moment, d.actuator_moment.numpy()[0], result.moment_rownnz, result.moment_rowadr, result.moment_colind ) - result.cvel = d.cvel.numpy()[0] - result.cdof_dot = d.cdof_dot.numpy()[0] - result.qfrc_bias = d.qfrc_bias.numpy()[0] - result.qfrc_fluid = d.qfrc_fluid.numpy()[0] - result.qfrc_passive = d.qfrc_passive.numpy()[0] - result.subtree_linvel = d.subtree_linvel.numpy()[0] - result.subtree_angmom = d.subtree_angmom.numpy()[0] - result.qfrc_spring = d.qfrc_spring.numpy()[0] - result.qfrc_damper = d.qfrc_damper.numpy()[0] - result.qfrc_gravcomp = d.qfrc_gravcomp.numpy()[0] - result.qfrc_fluid = d.qfrc_fluid.numpy()[0] - result.qfrc_actuator = d.qfrc_actuator.numpy()[0] - result.qfrc_smooth = d.qfrc_smooth.numpy()[0] - result.qfrc_constraint = d.qfrc_constraint.numpy()[0] - result.qfrc_inverse = d.qfrc_inverse.numpy()[0] - result.qacc_smooth = d.qacc_smooth.numpy()[0] - result.act = d.act.numpy()[0] - result.act_dot = d.act_dot.numpy()[0] + result.crb[:] = d.crb.numpy()[0] + result.qLDiagInv[:] = d.qLDiagInv.numpy()[0] + result.ten_velocity[:] = d.ten_velocity.numpy()[0] + result.actuator_velocity[:] = d.actuator_velocity.numpy()[0] + result.cvel[:] = d.cvel.numpy()[0] + result.cdof_dot[:] = d.cdof_dot.numpy()[0] + result.qfrc_bias[:] = d.qfrc_bias.numpy()[0] + result.qfrc_spring[:] = d.qfrc_spring.numpy()[0] + result.qfrc_damper[:] = d.qfrc_damper.numpy()[0] + result.qfrc_gravcomp[:] = d.qfrc_gravcomp.numpy()[0] + result.qfrc_fluid[:] = d.qfrc_fluid.numpy()[0] + result.qfrc_passive[:] = d.qfrc_passive.numpy()[0] + result.subtree_linvel[:] = d.subtree_linvel.numpy()[0] + result.subtree_angmom[:] = d.subtree_angmom.numpy()[0] + result.actuator_force[:] = d.actuator_force.numpy()[0] + result.qfrc_actuator[:] = d.qfrc_actuator.numpy()[0] + result.qfrc_smooth[:] = d.qfrc_smooth.numpy()[0] + result.qacc_smooth[:] = d.qacc_smooth.numpy()[0] + result.qfrc_constraint[:] = d.qfrc_constraint.numpy()[0] + result.qfrc_inverse[:] = d.qfrc_inverse.numpy()[0] + # contact result.contact.dist[:] = d.contact.dist.numpy()[:nacon] result.contact.pos[:] = d.contact.pos.numpy()[:nacon] result.contact.frame[:] = d.contact.frame.numpy()[:nacon].reshape((-1, 9)) @@ -1667,16 +1644,15 @@ def get_data_into( result.contact.solreffriction[:] = d.contact.solreffriction.numpy()[:nacon] result.contact.solimp[:] = d.contact.solimp.numpy()[:nacon] result.contact.dim[:] = d.contact.dim.numpy()[:nacon] - result.contact.efc_address[:] = d.contact.efc_address.numpy()[:nacon, 0] + result.contact.geom[:] = d.contact.geom.numpy()[:nacon] + result.contact.efc_address[:] = contact_efc_address_ordered[:nacon] if mujoco.mj_isSparse(mjm): result.qM[:] = d.qM.numpy()[0, 0] result.qLD[:] = d.qLD.numpy()[0, 0] - # TODO(team): set efc_J after fix to _realloc_con_efc lands - # efc_J = d.efc_J.numpy()[0, :nefc] - # mujoco.mju_dense2sparse( - # result.efc_J, efc_J, result.efc_J_rownnz, result.efc_J_rowadr, result.efc_J_colind - # ) + if nefc > 0: + efc_J = d.efc.J.numpy()[0, efc_idx, : mjm.nv] + mujoco.mju_dense2sparse(result.efc_J, efc_J, result.efc_J_rownnz, result.efc_J_rowadr, result.efc_J_colind) else: qM = d.qM.numpy() adr = 0 @@ -1687,34 +1663,26 @@ def get_data_into( j = mjm.dof_parentid[j] adr += 1 mujoco.mj_factorM(mjm, result) - # TODO(team): set efc_J after fix to _realloc_con_efc lands - # if nefc > 0: - # result.efc_J[:nefc * mjm.nv] = d.efc_J.numpy()[:nefc].flatten() - result.xfrc_applied[:] = d.xfrc_applied.numpy()[0] - result.eq_active[:] = d.eq_active.numpy()[0] + if nefc > 0: + result.efc_J[: nefc * mjm.nv] = d.efc.J.numpy()[0, :nefc, : mjm.nv].flatten() - # TODO(team): set these efc_* fields after fix to _realloc_con_efc - # Safely copy only up to the minimum of the destination and source sizes - # n = min(result.efc_D.shape[0], d.efc.D.numpy()[:nefc].shape[0]) - # result.efc_D[:n] = d.efc.D.numpy()[:nefc][:n] - # n_pos = min(result.efc_pos.shape[0], d.efc.pos.numpy()[:nefc].shape[0]) - # result.efc_pos[:n_pos] = d.efc.pos.numpy()[:nefc][:n_pos] - - # n_aref = min(result.efc_aref.shape[0], d.efc.aref.numpy()[:nefc].shape[0]) - # result.efc_aref[:n_aref] = d.efc.aref.numpy()[:nefc][:n_aref] - - # n_force = min(result.efc_force.shape[0], d.efc.force.numpy()[:nefc].shape[0]) - # result.efc_force[:n_force] = d.efc.force.numpy()[:nefc][:n_force] - - # n_margin = min(result.efc_margin.shape[0], d.efc.margin.numpy()[:nefc].shape[0]) - # result.efc_margin[:n_margin] = d.efc.margin.numpy()[:nefc][:n_margin] + # efc + result.efc_type[:] = d.efc.type.numpy()[0, efc_idx] + result.efc_id[:] = d.efc.id.numpy()[0, efc_idx] + result.efc_pos[:] = d.efc.pos.numpy()[0, efc_idx] + result.efc_margin[:] = d.efc.margin.numpy()[0, efc_idx] + result.efc_D[:] = d.efc.D.numpy()[0, efc_idx] + result.efc_vel[:] = d.efc.vel.numpy()[0, efc_idx] + result.efc_aref[:] = d.efc.aref.numpy()[0, efc_idx] + result.efc_frictionloss[:] = d.efc.frictionloss.numpy()[0, efc_idx] + result.efc_state[:] = d.efc.state.numpy()[0, efc_idx] + result.efc_force[:] = d.efc.force.numpy()[0, efc_idx] + # rne_postconstraint result.cacc[:] = d.cacc.numpy()[0] result.cfrc_int[:] = d.cfrc_int.numpy()[0] result.cfrc_ext[:] = d.cfrc_ext.numpy()[0] - # TODO: other efc_ fields, anything else missing - # tendon result.ten_length[:] = d.ten_length.numpy()[0] result.ten_J[:] = d.ten_J.numpy()[0] @@ -1727,449 +1695,258 @@ def get_data_into( result.sensordata[:] = d.sensordata.numpy() -# TODO(thowell): shared @wp.func for _reset kernel? +def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): + """Clear data, set defaults; optionally by world. + Args: + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + reset: Per-world bitmask. Reset if True. + """ -@wp.kernel -def _reset_xfrc_applied_all(xfrc_applied_out: wp.array2d(dtype=wp.spatial_vector)): - worldid, bodyid, elemid = wp.tid() - xfrc_applied_out[worldid, bodyid][elemid] = 0.0 + @nested_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() + if wp.static(reset is not None): + if not reset_in[worldid]: + return -@wp.kernel -def _reset_xfrc_applied(reset_in: wp.array(dtype=bool), xfrc_applied_out: wp.array2d(dtype=wp.spatial_vector)): - worldid, bodyid, elemid = wp.tid() + xfrc_applied_out[worldid, bodyid][elemid] = 0.0 - if not reset_in[worldid]: - return + @nested_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() - xfrc_applied_out[worldid, bodyid][elemid] = 0.0 + if wp.static(reset is not None): + if not reset_in[worldid]: + return + qM_out[worldid, elemid1, elemid2] = 0.0 -@wp.kernel -def _reset_qM_all(qM_out: wp.array3d(dtype=float)): - worldid, elemid1, elemid2 = wp.tid() - qM_out[worldid, elemid1, elemid2] = 0.0 + @nested_kernel(module="unique", enable_backward=False) + def reset_nworld( + # Model: + nq: int, + nv: int, + nu: int, + na: int, + neq: int, + nsensordata: int, + qpos0: wp.array2d(dtype=float), + eq_active0: wp.array(dtype=bool), + # Data in: + nworld_in: int, + # In: + reset_in: wp.array(dtype=bool), + # Data out: + solver_niter_out: wp.array(dtype=int), + ne_out: wp.array(dtype=int), + nf_out: wp.array(dtype=int), + nl_out: wp.array(dtype=int), + nefc_out: wp.array(dtype=int), + time_out: wp.array(dtype=float), + energy_out: wp.array(dtype=wp.vec2), + qpos_out: wp.array2d(dtype=float), + qvel_out: wp.array2d(dtype=float), + act_out: wp.array2d(dtype=float), + qacc_warmstart_out: wp.array2d(dtype=float), + ctrl_out: wp.array2d(dtype=float), + qfrc_applied_out: wp.array2d(dtype=float), + eq_active_out: wp.array2d(dtype=bool), + qacc_out: wp.array2d(dtype=float), + act_dot_out: wp.array2d(dtype=float), + sensordata_out: wp.array2d(dtype=float), + 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), + nsolving_out: wp.array(dtype=int), + ): + worldid = wp.tid() + if wp.static(reset is not None): + if not reset_in[worldid]: + return -@wp.kernel -def _reset_qM(reset_in: wp.array(dtype=bool), qM_out: wp.array3d(dtype=float)): - worldid, elemid1, elemid2 = wp.tid() + solver_niter_out[worldid] = 0 + 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 + nf_out[worldid] = 0 + nl_out[worldid] = 0 + nefc_out[worldid] = 0 + if worldid == 0: + nsolving_out[0] = nworld_in + time_out[worldid] = 0.0 + energy_out[worldid] = wp.vec2(0.0, 0.0) + for i in range(nq): + qpos_out[worldid, i] = qpos0[worldid, i] + if i < nv: + qvel_out[worldid, i] = 0.0 + qacc_warmstart_out[worldid, i] = 0.0 + qfrc_applied_out[worldid, i] = 0.0 + qacc_out[worldid, i] = 0.0 + for i in range(nu): + ctrl_out[worldid, i] = 0.0 + if i < na: + act_out[worldid, i] = 0.0 + act_dot_out[worldid, i] = 0.0 + for i in range(neq): + eq_active_out[worldid, i] = eq_active0[i] + for i in range(nsensordata): + sensordata_out[worldid, i] = 0.0 - if not reset_in[worldid]: - return + @nested_kernel(module="unique", enable_backward=False) + def reset_mocap( + # Model: + body_mocapid: wp.array(dtype=int), + body_pos: wp.array2d(dtype=wp.vec3), + body_quat: wp.array2d(dtype=wp.quat), + # In: + reset_in: wp.array(dtype=bool), + # Data out: + mocap_pos_out: wp.array2d(dtype=wp.vec3), + mocap_quat_out: wp.array2d(dtype=wp.quat), + ): + worldid, bodyid = wp.tid() - qM_out[worldid, elemid1, elemid2] = 0.0 + if wp.static(reset is not None): + if not reset_in[worldid]: + return + mocapid = body_mocapid[bodyid] -@wp.kernel -def _reset_nworld_all( - # Model: - nq: int, - nv: int, - nu: int, - na: int, - neq: int, - nsensordata: int, - qpos0: wp.array2d(dtype=float), - eq_active0: wp.array(dtype=bool), - # Data in: - nworld_in: int, - # Data out: - solver_niter_out: wp.array(dtype=int), - nacon_out: wp.array(dtype=int), - ne_out: wp.array(dtype=int), - ne_connect_out: wp.array(dtype=int), - ne_weld_out: wp.array(dtype=int), - ne_jnt_out: wp.array(dtype=int), - ne_ten_out: wp.array(dtype=int), - nf_out: wp.array(dtype=int), - nl_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - nsolving_out: wp.array(dtype=int), - time_out: wp.array(dtype=float), - energy_out: wp.array(dtype=wp.vec2), - qpos_out: wp.array2d(dtype=float), - qvel_out: wp.array2d(dtype=float), - act_out: wp.array2d(dtype=float), - qacc_warmstart_out: wp.array2d(dtype=float), - ctrl_out: wp.array2d(dtype=float), - qfrc_applied_out: wp.array2d(dtype=float), - eq_active_out: wp.array2d(dtype=bool), - qacc_out: wp.array2d(dtype=float), - act_dot_out: wp.array2d(dtype=float), - sensordata_out: wp.array2d(dtype=float), -): - worldid = wp.tid() + if mocapid >= 0: + mocap_pos_out[worldid, mocapid] = body_pos[worldid, bodyid] + mocap_quat_out[worldid, mocapid] = body_quat[worldid, bodyid] - solver_niter_out[worldid] = 0 - 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 - nf_out[worldid] = 0 - nl_out[worldid] = 0 - nefc_out[worldid] = 0 - if worldid == 0: - nsolving_out[0] = nworld_in - time_out[worldid] = 0.0 - energy_out[worldid] = wp.vec2(0.0, 0.0) - for i in range(nq): - qpos_out[worldid, i] = qpos0[worldid, i] - if i < nv: - qvel_out[worldid, i] = 0.0 - qacc_warmstart_out[worldid, i] = 0.0 - qfrc_applied_out[worldid, i] = 0.0 - qacc_out[worldid, i] = 0.0 - for i in range(nu): - ctrl_out[worldid, i] = 0.0 - if i < na: - act_out[worldid, i] = 0.0 - act_dot_out[worldid, i] = 0.0 - for i in range(neq): - eq_active_out[worldid, i] = eq_active0[i] - for i in range(nsensordata): - sensordata_out[worldid, i] = 0.0 + @nested_kernel(module="unique", enable_backward=False) + def reset_contact( + # Data in: + nacon_in: wp.array(dtype=int), + # In: + reset_in: wp.array(dtype=bool), + nefcaddress: int, + # Data out: + contact_dist_out: wp.array(dtype=float), + contact_pos_out: wp.array(dtype=wp.vec3), + contact_frame_out: wp.array(dtype=wp.mat33), + contact_includemargin_out: wp.array(dtype=float), + contact_friction_out: wp.array(dtype=types.vec5), + contact_solref_out: wp.array(dtype=wp.vec2), + contact_solreffriction_out: wp.array(dtype=wp.vec2), + contact_solimp_out: wp.array(dtype=types.vec5), + contact_dim_out: wp.array(dtype=int), + contact_geom_out: wp.array(dtype=wp.vec2i), + contact_efc_address_out: wp.array2d(dtype=int), + contact_worldid_out: wp.array(dtype=int), + contact_type_out: wp.array(dtype=int), + contact_geomcollisionid_out: wp.array(dtype=int), + ): + conid = wp.tid() - -@wp.kernel -def _reset_nworld( - # Model: - nq: int, - nv: int, - nu: int, - na: int, - neq: int, - nsensordata: int, - qpos0: wp.array2d(dtype=float), - eq_active0: wp.array(dtype=bool), - # Data in: - nworld_in: int, - # In: - reset_in: wp.array(dtype=bool), - # Data out: - solver_niter_out: wp.array(dtype=int), - nacon_out: wp.array(dtype=int), - ne_out: wp.array(dtype=int), - ne_connect_out: wp.array(dtype=int), - ne_weld_out: wp.array(dtype=int), - ne_jnt_out: wp.array(dtype=int), - ne_ten_out: wp.array(dtype=int), - nf_out: wp.array(dtype=int), - nl_out: wp.array(dtype=int), - nefc_out: wp.array(dtype=int), - nsolving_out: wp.array(dtype=int), - time_out: wp.array(dtype=float), - energy_out: wp.array(dtype=wp.vec2), - qpos_out: wp.array2d(dtype=float), - qvel_out: wp.array2d(dtype=float), - act_out: wp.array2d(dtype=float), - qacc_warmstart_out: wp.array2d(dtype=float), - ctrl_out: wp.array2d(dtype=float), - qfrc_applied_out: wp.array2d(dtype=float), - eq_active_out: wp.array2d(dtype=bool), - qacc_out: wp.array2d(dtype=float), - act_dot_out: wp.array2d(dtype=float), - sensordata_out: wp.array2d(dtype=float), -): - worldid = wp.tid() - - if not reset_in[worldid]: - return - - solver_niter_out[worldid] = 0 - 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 - nf_out[worldid] = 0 - nl_out[worldid] = 0 - nefc_out[worldid] = 0 - if worldid == 0: - nsolving_out[0] = nworld_in - time_out[worldid] = 0.0 - energy_out[worldid] = wp.vec2(0.0, 0.0) - for i in range(nq): - qpos_out[worldid, i] = qpos0[worldid, i] - if i < nv: - qvel_out[worldid, i] = 0.0 - qacc_warmstart_out[worldid, i] = 0.0 - qfrc_applied_out[worldid, i] = 0.0 - qacc_out[worldid, i] = 0.0 - for i in range(nu): - ctrl_out[worldid, i] = 0.0 - if i < na: - act_out[worldid, i] = 0.0 - act_dot_out[worldid, i] = 0.0 - for i in range(neq): - eq_active_out[worldid, i] = eq_active0[i] - for i in range(nsensordata): - sensordata_out[worldid, i] = 0.0 - - -@wp.kernel -def _reset_mocap_all( - # Model: - body_mocapid: wp.array(dtype=int), - body_pos: wp.array2d(dtype=wp.vec3), - body_quat: wp.array2d(dtype=wp.quat), - # Data out: - mocap_pos_out: wp.array2d(dtype=wp.vec3), - mocap_quat_out: wp.array2d(dtype=wp.quat), -): - worldid, bodyid = wp.tid() - - mocapid = body_mocapid[bodyid] - - if mocapid >= 0: - mocap_pos_out[worldid, mocapid] = body_pos[worldid, bodyid] - mocap_quat_out[worldid, mocapid] = body_quat[worldid, bodyid] - - -@wp.kernel -def _reset_mocap( - # Model: - body_mocapid: wp.array(dtype=int), - body_pos: wp.array2d(dtype=wp.vec3), - body_quat: wp.array2d(dtype=wp.quat), - # In: - reset_in: wp.array(dtype=bool), - # Data out: - mocap_pos_out: wp.array2d(dtype=wp.vec3), - mocap_quat_out: wp.array2d(dtype=wp.quat), -): - worldid, bodyid = wp.tid() - - if not reset_in[worldid]: - return - - mocapid = body_mocapid[bodyid] - - if mocapid >= 0: - mocap_pos_out[worldid, mocapid] = body_pos[worldid, bodyid] - mocap_quat_out[worldid, mocapid] = body_quat[worldid, bodyid] - - -@wp.kernel -def _reset_contact_all( - # Data in: - nacon_in: wp.array(dtype=int), - # In: - nefcaddress: int, - # Data out: - contact_dist_out: wp.array(dtype=float), - contact_pos_out: wp.array(dtype=wp.vec3), - contact_frame_out: wp.array(dtype=wp.mat33), - contact_includemargin_out: wp.array(dtype=float), - contact_friction_out: wp.array(dtype=types.vec5), - contact_solref_out: wp.array(dtype=wp.vec2), - contact_solreffriction_out: wp.array(dtype=wp.vec2), - contact_solimp_out: wp.array(dtype=types.vec5), - contact_dim_out: wp.array(dtype=int), - contact_geom_out: wp.array(dtype=wp.vec2i), - contact_efc_address_out: wp.array2d(dtype=int), - contact_worldid_out: wp.array(dtype=int), -): - conid = wp.tid() - - if conid >= nacon_in[0]: - return - - contact_dist_out[conid] = 0.0 - contact_pos_out[conid] = wp.vec3(0.0, 0.0, 0.0) - contact_frame_out[conid] = wp.mat33(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0) - contact_includemargin_out[conid] = 0.0 - contact_friction_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0) - contact_solref_out[conid] = wp.vec2(0.0, 0.0) - contact_solreffriction_out[conid] = wp.vec2(0.0, 0.0) - contact_solimp_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0) - contact_dim_out[conid] = 0 - contact_geom_out[conid] = wp.vec2i(0, 0) - for i in range(nefcaddress): - contact_efc_address_out[conid, i] = 0 - contact_worldid_out[conid] = 0 - - -@wp.kernel -def _reset_contact( - # Data in: - nacon_in: wp.array(dtype=int), - # In: - reset_in: wp.array(dtype=bool), - nefcaddress: int, - # Data out: - contact_dist_out: wp.array(dtype=float), - contact_pos_out: wp.array(dtype=wp.vec3), - contact_frame_out: wp.array(dtype=wp.mat33), - contact_includemargin_out: wp.array(dtype=float), - contact_friction_out: wp.array(dtype=types.vec5), - contact_solref_out: wp.array(dtype=wp.vec2), - contact_solreffriction_out: wp.array(dtype=wp.vec2), - contact_solimp_out: wp.array(dtype=types.vec5), - contact_dim_out: wp.array(dtype=int), - contact_geom_out: wp.array(dtype=wp.vec2i), - contact_efc_address_out: wp.array2d(dtype=int), - contact_worldid_out: wp.array(dtype=int), -): - conid = wp.tid() - - if conid >= nacon_in[0]: - return - - worldid = contact_worldid_out[conid] - if worldid >= 0: - if not reset_in[worldid]: + if conid >= nacon_in[0]: return - contact_dist_out[conid] = 0.0 - contact_pos_out[conid] = wp.vec3(0.0) - contact_frame_out[conid] = wp.mat33(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0) - contact_includemargin_out[conid] = 0.0 - contact_friction_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0) - contact_solref_out[conid] = wp.vec2(0.0, 0.0) - contact_solreffriction_out[conid] = wp.vec2(0.0, 0.0) - contact_solimp_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0) - contact_dim_out[conid] = 0 - contact_geom_out[conid] = wp.vec2i(0, 0) - for i in range(nefcaddress): - contact_efc_address_out[conid, i] = 0 - contact_worldid_out[conid] = 0 + worldid = contact_worldid_out[conid] + if wp.static(reset is not None): + if worldid >= 0: + if not reset_in[worldid]: + return + contact_dist_out[conid] = 0.0 + contact_pos_out[conid] = wp.vec3(0.0) + contact_frame_out[conid] = wp.mat33(0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + contact_includemargin_out[conid] = 0.0 + contact_friction_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0) + contact_solref_out[conid] = wp.vec2(0.0, 0.0) + contact_solreffriction_out[conid] = wp.vec2(0.0, 0.0) + contact_solimp_out[conid] = types.vec5(0.0, 0.0, 0.0, 0.0, 0.0) + contact_dim_out[conid] = 0 + contact_geom_out[conid] = wp.vec2i(0, 0) + for i in range(nefcaddress): + contact_efc_address_out[conid, i] = 0 + contact_worldid_out[conid] = 0 + contact_type_out[conid] = 0 + contact_geomcollisionid_out[conid] = 0 -def reset_data(m: types.Model, d: types.Data, reset: Optional[wp.array] = None): - """Clear data, set defaults.""" - if m.opt.is_sparse: - qM_dim = (1, m.nM) - else: - qM_dim = (m.nv, m.nv) + reset_input = reset or wp.ones(d.nworld, dtype=bool) - if reset is not None: - wp.launch(_reset_xfrc_applied, dim=(d.nworld, m.nbody, 6), inputs=[reset], outputs=[d.xfrc_applied]) - wp.launch(_reset_qM, dim=(d.nworld, qM_dim[0], qM_dim[1]), inputs=[reset], outputs=[d.qM]) + wp.launch(reset_xfrc_applied, dim=(d.nworld, m.nbody, 6), inputs=[reset_input], outputs=[d.xfrc_applied]) + wp.launch( + reset_qM, + dim=(d.nworld, 1 if m.opt.is_sparse else m.nv, m.nM if m.opt.is_sparse else m.nv), + inputs=[reset_input], + outputs=[d.qM], + ) - # set mocap_pos/quat = body_pos/quat for mocap bodies - wp.launch( - _reset_mocap, - dim=(d.nworld, m.nbody), - inputs=[m.body_mocapid, m.body_pos, m.body_quat, reset], - outputs=[d.mocap_pos, d.mocap_quat], - ) + # set mocap_pos/quat = body_pos/quat for mocap bodies + wp.launch( + reset_mocap, + dim=(d.nworld, m.nbody), + inputs=[m.body_mocapid, m.body_pos, m.body_quat, reset_input], + outputs=[d.mocap_pos, d.mocap_quat], + ) - # clear contacts - wp.launch( - _reset_contact, - dim=d.naconmax, - inputs=[d.nacon, reset, d.contact.efc_address.shape[1]], - outputs=[ - d.contact.dist, - d.contact.pos, - d.contact.frame, - d.contact.includemargin, - d.contact.friction, - d.contact.solref, - d.contact.solreffriction, - d.contact.solimp, - d.contact.dim, - d.contact.geom, - d.contact.efc_address, - d.contact.worldid, - ], - ) + # clear contacts + wp.launch( + reset_contact, + dim=d.naconmax, + inputs=[d.nacon, reset_input, d.contact.efc_address.shape[1]], + outputs=[ + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.includemargin, + d.contact.friction, + d.contact.solref, + d.contact.solreffriction, + d.contact.solimp, + d.contact.dim, + d.contact.geom, + d.contact.efc_address, + d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + ], + ) - wp.launch( - _reset_nworld, - dim=d.nworld, - inputs=[m.nq, m.nv, m.nu, m.na, m.neq, m.nsensordata, m.qpos0, m.eq_active0, d.nworld, reset], - outputs=[ - d.solver_niter, - d.nacon, - d.ne, - d.ne_connect, - d.ne_weld, - d.ne_jnt, - d.ne_ten, - d.nf, - d.nl, - d.nefc, - d.nsolving, - d.time, - d.energy, - d.qpos, - d.qvel, - d.act, - d.qacc_warmstart, - d.ctrl, - d.qfrc_applied, - d.eq_active, - d.qacc, - d.act_dot, - d.sensordata, - ], - ) - else: - wp.launch(_reset_xfrc_applied_all, dim=(d.nworld, m.nbody, 6), outputs=[d.xfrc_applied]) - wp.launch(_reset_qM_all, dim=(d.nworld, qM_dim[0], qM_dim[1]), outputs=[d.qM]) - wp.launch( - _reset_mocap_all, - dim=(d.nworld, m.nbody), - inputs=[m.body_mocapid, m.body_pos, m.body_quat], - outputs=[d.mocap_pos, d.mocap_quat], - ) - wp.launch( - _reset_contact_all, - dim=d.naconmax, - inputs=[d.nacon, d.contact.efc_address.shape[1]], - outputs=[ - d.contact.dist, - d.contact.pos, - d.contact.frame, - d.contact.includemargin, - d.contact.friction, - d.contact.solref, - d.contact.solreffriction, - d.contact.solimp, - d.contact.dim, - d.contact.geom, - d.contact.efc_address, - d.contact.worldid, - ], - ) - wp.launch( - _reset_nworld_all, - dim=d.nworld, - inputs=[m.nq, m.nv, m.nu, m.na, m.neq, m.nsensordata, m.qpos0, m.eq_active0, d.nworld], - outputs=[ - d.solver_niter, - d.nacon, - d.ne, - d.ne_connect, - d.ne_weld, - d.ne_jnt, - d.ne_ten, - d.nf, - d.nl, - d.nefc, - d.nsolving, - d.time, - d.energy, - d.qpos, - d.qvel, - d.act, - d.qacc_warmstart, - d.ctrl, - d.qfrc_applied, - d.eq_active, - d.qacc, - d.act_dot, - d.sensordata, - ], - ) + wp.launch( + reset_nworld, + dim=d.nworld, + inputs=[m.nq, m.nv, m.nu, m.na, m.neq, m.nsensordata, m.qpos0, m.eq_active0, d.nworld, reset_input], + outputs=[ + d.solver_niter, + d.ne, + d.nf, + d.nl, + d.nefc, + d.time, + d.energy, + d.qpos, + d.qvel, + d.act, + d.qacc_warmstart, + d.ctrl, + d.qfrc_applied, + d.eq_active, + d.qacc, + d.act_dot, + d.sensordata, + d.nacon, + d.ne_connect, + d.ne_weld, + d.ne_jnt, + d.ne_ten, + d.nsolving, + ], + ) def override_model(model: Union[types.Model, mujoco.MjModel], overrides: Union[dict[str, Any], Sequence[str]]): @@ -2181,7 +1958,6 @@ def override_model(model: Union[types.Model, mujoco.MjModel], overrides: Union[d opt.cone = pyramidal opt.disableflags = contact | spring """ - enum_fields = { "opt.broadphase": types.BroadphaseType, "opt.broadphase_filter": types.BroadphaseFilter, diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py index 4254f2cb..d3b7d063 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py @@ -96,7 +96,6 @@ def inert_vec(i: types.vec10, v: wp.spatial_vector) -> wp.spatial_vector: @wp.func def motion_cross(u: wp.spatial_vector, v: wp.spatial_vector) -> wp.spatial_vector: """Cross product of two motions.""" - u0 = wp.vec3(u[0], u[1], u[2]) u1 = wp.vec3(u[3], u[4], u[5]) v0 = wp.vec3(v[0], v[1], v[2]) @@ -111,7 +110,6 @@ def motion_cross(u: wp.spatial_vector, v: wp.spatial_vector) -> wp.spatial_vecto @wp.func def motion_cross_force(v: wp.spatial_vector, f: wp.spatial_vector) -> wp.spatial_vector: """Cross product of a motion and a force.""" - v0 = wp.vec3(v[0], v[1], v[2]) v1 = wp.vec3(v[3], v[4], v[5]) f0 = wp.vec3(f[0], f[1], f[2]) @@ -249,7 +247,6 @@ def closest_segment_point_and_dist(a: wp.vec3, b: wp.vec3, pt: wp.vec3) -> Tuple @wp.func def closest_segment_to_segment_points(a0: wp.vec3, a1: wp.vec3, b0: wp.vec3, b1: wp.vec3) -> Tuple[wp.vec3, wp.vec3]: """Returns closest points between two line segments.""" - dir_a, len_a = normalize_with_norm(a1 - a0) dir_b, len_b = normalize_with_norm(b1 - b0) 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 8834a4d3..459c473b 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/passive.py @@ -86,8 +86,8 @@ def _spring_damper_dof_passive( ): worldid, jntid = wp.tid() dofid = jnt_dofadr[jntid] - stiffness = jnt_stiffness[worldid, jntid] - damping = dof_damping[worldid, dofid] + stiffness = jnt_stiffness[worldid % jnt_stiffness.shape[0], jntid] + damping = dof_damping[worldid % dof_damping.shape[0], dofid] has_stiffness = stiffness != 0.0 and not opt_disableflags & DisableBit.SPRING has_damping = damping != 0.0 and not opt_disableflags & DisableBit.DAMPER @@ -103,14 +103,15 @@ def _spring_damper_dof_passive( jnttype = jnt_type[jntid] qposid = jnt_qposadr[jntid] + qpos_spring_id = worldid % qpos_spring.shape[0] if jnttype == JointType.FREE: # spring if has_stiffness: dif = wp.vec3( - qpos_in[worldid, qposid + 0] - qpos_spring[worldid, qposid + 0], - qpos_in[worldid, qposid + 1] - qpos_spring[worldid, qposid + 1], - qpos_in[worldid, qposid + 2] - qpos_spring[worldid, qposid + 2], + qpos_in[worldid, qposid + 0] - qpos_spring[qpos_spring_id, qposid + 0], + qpos_in[worldid, qposid + 1] - qpos_spring[qpos_spring_id, qposid + 1], + qpos_in[worldid, qposid + 2] - qpos_spring[qpos_spring_id, qposid + 2], ) qfrc_spring_out[worldid, dofid + 0] = -stiffness * dif[0] qfrc_spring_out[worldid, dofid + 1] = -stiffness * dif[1] @@ -123,10 +124,10 @@ def _spring_damper_dof_passive( ) rot = wp.normalize(rot) ref = wp.quat( - qpos_spring[worldid, qposid + 3], - qpos_spring[worldid, qposid + 4], - qpos_spring[worldid, qposid + 5], - qpos_spring[worldid, qposid + 6], + qpos_spring[qpos_spring_id, qposid + 3], + qpos_spring[qpos_spring_id, qposid + 4], + qpos_spring[qpos_spring_id, qposid + 5], + qpos_spring[qpos_spring_id, qposid + 6], ) dif = math.quat_sub(rot, ref) qfrc_spring_out[worldid, dofid + 3] = -stiffness * dif[0] @@ -152,10 +153,10 @@ def _spring_damper_dof_passive( ) rot = wp.normalize(rot) ref = wp.quat( - qpos_spring[worldid, qposid + 0], - qpos_spring[worldid, qposid + 1], - qpos_spring[worldid, qposid + 2], - qpos_spring[worldid, qposid + 3], + qpos_spring[qpos_spring_id, qposid + 0], + qpos_spring[qpos_spring_id, qposid + 1], + qpos_spring[qpos_spring_id, qposid + 2], + qpos_spring[qpos_spring_id, qposid + 3], ) dif = math.quat_sub(rot, ref) qfrc_spring_out[worldid, dofid + 0] = -stiffness * dif[0] @@ -170,7 +171,7 @@ def _spring_damper_dof_passive( else: # mjJNT_SLIDE, mjJNT_HINGE # spring if has_stiffness: - fdif = qpos_in[worldid, qposid] - qpos_spring[worldid, qposid] + fdif = qpos_in[worldid, qposid] - qpos_spring[qpos_spring_id, qposid] qfrc_spring_out[worldid, dofid] = -stiffness * fdif # damper @@ -185,9 +186,9 @@ def _spring_damper_tendon_passive( tendon_damping: wp.array2d(dtype=float), tendon_lengthspring: wp.array2d(dtype=wp.vec2), # Data in: - ten_velocity_in: wp.array2d(dtype=float), - ten_length_in: wp.array2d(dtype=float), ten_J_in: wp.array3d(dtype=float), + ten_length_in: wp.array2d(dtype=float), + ten_velocity_in: wp.array2d(dtype=float), # In: dsbl_spring: bool, dsbl_damper: bool, @@ -197,8 +198,8 @@ def _spring_damper_tendon_passive( ): worldid, tenid, dofid = wp.tid() - stiffness = tendon_stiffness[worldid, tenid] - damping = tendon_damping[worldid, tenid] + stiffness = tendon_stiffness[worldid % tendon_stiffness.shape[0], tenid] + damping = tendon_damping[worldid % tendon_damping.shape[0], tenid] has_stiffness = stiffness != 0.0 and not dsbl_spring has_damping = damping != 0.0 and not dsbl_damper @@ -211,7 +212,7 @@ def _spring_damper_tendon_passive( if has_stiffness: # compute spring force along tendon length = ten_length_in[worldid, tenid] - lengthspring = tendon_lengthspring[worldid, tenid] + lengthspring = tendon_lengthspring[worldid % tendon_lengthspring.shape[0], tenid] lower = lengthspring[0] upper = lengthspring[1] @@ -251,12 +252,11 @@ def _gravity_force( ): worldid, bodyid, dofid = wp.tid() bodyid += 1 # skip world body - gravcomp = body_gravcomp[worldid, bodyid] - gravity = opt_gravity[worldid] + gravcomp = body_gravcomp[worldid % body_gravcomp.shape[0], bodyid] + gravity = opt_gravity[worldid % opt_gravity.shape[0]] if gravcomp: - force = -gravity * body_mass[worldid, bodyid] * 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) @@ -266,18 +266,18 @@ def _gravity_force( @wp.kernel def _fluid_force( # Model: - opt_wind: wp.array(dtype=wp.vec3), opt_density: wp.array(dtype=float), opt_viscosity: wp.array(dtype=float), + opt_wind: wp.array(dtype=wp.vec3), body_rootid: wp.array(dtype=int), body_geomnum: wp.array(dtype=int), body_geomadr: wp.array(dtype=int), body_mass: wp.array2d(dtype=float), body_inertia: wp.array2d(dtype=wp.vec3), - body_fluid_ellipsoid: wp.array(dtype=bool), geom_type: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), geom_fluid: wp.array2d(dtype=float), + body_fluid_ellipsoid: wp.array(dtype=bool), # Data in: xipos_in: wp.array2d(dtype=wp.vec3), ximat_in: wp.array2d(dtype=wp.mat33), @@ -285,11 +285,10 @@ def _fluid_force( geom_xmat_in: wp.array2d(dtype=wp.mat33), subtree_com_in: wp.array2d(dtype=wp.vec3), cvel_in: wp.array2d(dtype=wp.spatial_vector), - # Data out: + # Out: fluid_applied_out: wp.array2d(dtype=wp.spatial_vector), ): """Computes body-space fluid forces for both inertia-box and ellipsoid models.""" - worldid, bodyid = wp.tid() zero_force = wp.spatial_vector(wp.vec3(0.0), wp.vec3(0.0)) @@ -297,9 +296,9 @@ def _fluid_force( fluid_applied_out[worldid, bodyid] = zero_force return - wind = opt_wind[worldid] - density = opt_density[worldid] - viscosity = opt_viscosity[worldid] + wind = opt_wind[worldid % opt_wind.shape[0]] + density = opt_density[worldid % opt_density.shape[0]] + viscosity = opt_viscosity[worldid % opt_viscosity.shape[0]] # Body kinematics xipos = xipos_in[worldid, bodyid] @@ -324,7 +323,7 @@ def _fluid_force( if coef <= 0.0: continue - size = geom_size[worldid, geomid] + size = geom_size[worldid % geom_size.shape[0], geomid] semiaxes = _geom_semiaxes(size, geom_type[geomid]) geom_rot = geom_xmat_in[worldid, geomid] geom_rotT = wp.transpose(geom_rot) @@ -451,8 +450,8 @@ def _fluid_force( has_density = density > 0.0 if has_viscosity or has_density: - inertia = body_inertia[worldid, bodyid] - mass = body_mass[worldid, bodyid] + inertia = body_inertia[worldid % body_inertia.shape[0], bodyid] + mass = body_mass[worldid % body_mass.shape[0], bodyid] scl = 6.0 / mass box0 = wp.sqrt(wp.max(MJ_MINVAL, inertia[1] + inertia[2] - inertia[0]) * scl) box1 = wp.sqrt(wp.max(MJ_MINVAL, inertia[0] + inertia[2] - inertia[1]) * scl) @@ -487,22 +486,24 @@ def _fluid_force( def _fluid(m: Model, d: Data): + fluid_applied = wp.empty((d.nworld, m.nbody), dtype=wp.spatial_vector) + wp.launch( _fluid_force, dim=(d.nworld, m.nbody), inputs=[ - m.opt.wind, m.opt.density, m.opt.viscosity, + m.opt.wind, m.body_rootid, m.body_geomnum, m.body_geomadr, m.body_mass, m.body_inertia, - m.body_fluid_ellipsoid, m.geom_type, m.geom_size, m.geom_fluid, + m.body_fluid_ellipsoid, d.xipos, d.ximat, d.geom_xpos, @@ -510,12 +511,10 @@ def _fluid(m: Model, d: Data): d.subtree_com, d.cvel, ], - outputs=[ - d.fluid_applied, - ], + outputs=[fluid_applied], ) - support.apply_ft(m, d, d.fluid_applied, d.qfrc_fluid, False) + support.apply_ft(m, d, fluid_applied, d.qfrc_fluid, False) @wp.kernel @@ -574,7 +573,7 @@ def _flex_elasticity( qfrc_spring_out: wp.array2d(dtype=float), ): worldid, elemid = wp.tid() - timestep = opt_timestep[worldid] + timestep = opt_timestep[worldid % opt_timestep.shape[0]] f = 0 # TODO(quaglino): this should become a function of t dim = flex_dim[f] @@ -727,9 +726,9 @@ def passive(m: Model, d: Data): m.tendon_stiffness, m.tendon_damping, m.tendon_lengthspring, - d.ten_velocity, - d.ten_length, d.ten_J, + d.ten_length, + d.ten_velocity, dsbl_spring, dsbl_damper, ], 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 3962e90e..15959db7 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/ray.py @@ -13,7 +13,7 @@ # limitations under the License. # ============================================================================== -from typing import Tuple +from typing import Optional, Tuple import warp as wp @@ -40,7 +40,6 @@ def _ray_map(pos: wp.vec3, mat: wp.mat33, pnt: wp.vec3, vec: wp.vec3) -> Tuple[w Returns: 3D point and 3D direction in local geom frame """ - matT = wp.transpose(mat) lpnt = matT @ (pnt - pos) lvec = matT @ vec @@ -53,8 +52,8 @@ def _ray_eliminate( # Model: body_weldid: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), - geom_group: wp.array(dtype=int), geom_matid: wp.array(dtype=int), # kernel_analyzer: ignore + geom_group: wp.array(dtype=int), geom_rgba: wp.array(dtype=wp.vec4), # kernel_analyzer: ignore mat_rgba: wp.array(dtype=wp.vec4), # kernel_analyzer: ignore # In: @@ -184,7 +183,6 @@ def _ray_triangle(v0: wp.vec3, v1: wp.vec3, v2: wp.vec3, pnt: wp.vec3, vec: wp.v @wp.func def _ray_plane(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: """Returns the distance at which a ray intersects with a plane.""" - # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) @@ -222,7 +220,6 @@ def _ray_sphere(pos: wp.vec3, dist_sqr: float, pnt: wp.vec3, vec: wp.vec3) -> fl @wp.func def _ray_capsule(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3) -> float: """Returns the distance at which a ray intersects with a capsule.""" - # bounding sphere test ssz = size[0] + size[1] if _ray_sphere(pos, ssz * ssz, pnt, vec) < 0.0: @@ -279,7 +276,6 @@ 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) -> float: """Returns the distance at which a ray intersects with an ellipsoid.""" - # map to local frame lpnt, lvec = _ray_map(pos, mat, pnt, vec) @@ -395,10 +391,10 @@ def _ray_hfield( # Model: geom_type: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), - hfield_adr: wp.array(dtype=int), + hfield_size: wp.array(dtype=wp.vec4), hfield_nrow: wp.array(dtype=int), hfield_ncol: wp.array(dtype=int), - hfield_size: wp.array(dtype=wp.vec4), + hfield_adr: wp.array(dtype=int), hfield_data: wp.array(dtype=float), # In: pos: wp.vec3, @@ -540,8 +536,8 @@ def ray_mesh( # Model: nmeshface: int, mesh_vertadr: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), mesh_faceadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), mesh_face: wp.array(dtype=wp.vec3i), # In: data_id: int, @@ -606,7 +602,6 @@ def ray_mesh( @wp.func def ray_geom(pos: wp.vec3, mat: wp.mat33, size: wp.vec3, pnt: wp.vec3, vec: wp.vec3, geomtype: int) -> float: """Returns distance along ray to intersection with geom, or infinity if none.""" - # TODO(team): static loop unrolling to remove unnecessary branching if geomtype == GeomType.PLANE: return _ray_plane(pos, mat, size, pnt, vec) @@ -633,19 +628,19 @@ def _ray_geom_mesh( geom_type: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), - geom_group: 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), - hfield_adr: wp.array(dtype=int), + mesh_vertadr: wp.array(dtype=int), + mesh_faceadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_face: wp.array(dtype=wp.vec3i), + hfield_size: wp.array(dtype=wp.vec4), hfield_nrow: wp.array(dtype=int), hfield_ncol: wp.array(dtype=int), - hfield_size: wp.array(dtype=wp.vec4), + hfield_adr: wp.array(dtype=int), hfield_data: wp.array(dtype=float), - mesh_vertadr: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), - mesh_faceadr: wp.array(dtype=int), - mesh_face: wp.array(dtype=wp.vec3i), mat_rgba: wp.array2d(dtype=wp.vec4), # Data in: geom_xpos_in: wp.array2d(dtype=wp.vec3), @@ -662,10 +657,10 @@ def _ray_geom_mesh( if not _ray_eliminate( body_weldid, geom_bodyid, + geom_matid[worldid % geom_matid.shape[0]], geom_group, - geom_matid[worldid], - geom_rgba[worldid], - mat_rgba[worldid], + geom_rgba[worldid % geom_rgba.shape[0]], + mat_rgba[worldid % mat_rgba.shape[0]], geomid, geomgroup, flg_static, @@ -679,8 +674,8 @@ def _ray_geom_mesh( return ray_mesh( nmeshface, mesh_vertadr, - mesh_vert, mesh_faceadr, + mesh_vert, mesh_face, geom_dataid[geomid], pos, @@ -692,10 +687,10 @@ def _ray_geom_mesh( return _ray_hfield( geom_type, geom_dataid, - hfield_adr, + hfield_size, hfield_nrow, hfield_ncol, - hfield_size, + hfield_adr, hfield_data, pos, mat, @@ -704,7 +699,7 @@ def _ray_geom_mesh( geomid, ) else: - return ray_geom(pos, mat, geom_size[worldid, geomid], pnt, vec, type) + return ray_geom(pos, mat, geom_size[worldid % geom_size.shape[0], geomid], pnt, vec, type) else: return wp.inf @@ -718,19 +713,19 @@ def _ray( geom_type: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), - geom_group: 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), - hfield_adr: wp.array(dtype=int), + mesh_vertadr: wp.array(dtype=int), + mesh_faceadr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), + mesh_face: wp.array(dtype=wp.vec3i), + hfield_size: wp.array(dtype=wp.vec4), hfield_nrow: wp.array(dtype=int), hfield_ncol: wp.array(dtype=int), - hfield_size: wp.array(dtype=wp.vec4), + hfield_adr: wp.array(dtype=int), hfield_data: wp.array(dtype=float), - mesh_vertadr: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), - mesh_faceadr: wp.array(dtype=int), - mesh_face: wp.array(dtype=wp.vec3i), mat_rgba: wp.array2d(dtype=wp.vec4), # Data in: geom_xpos_in: wp.array2d(dtype=wp.vec3), @@ -761,19 +756,19 @@ def _ray( geom_type, geom_bodyid, geom_dataid, - geom_group, geom_matid, + geom_group, geom_size, geom_rgba, - hfield_adr, + mesh_vertadr, + mesh_faceadr, + mesh_vert, + mesh_face, + hfield_size, hfield_nrow, hfield_ncol, - hfield_size, + hfield_adr, hfield_data, - mesh_vertadr, - mesh_vert, - mesh_faceadr, - mesh_face, mat_rgba, geom_xpos_in, geom_xmat_in, @@ -810,41 +805,38 @@ def ray( d: Data, pnt: wp.array2d(dtype=wp.vec3), vec: wp.array2d(dtype=wp.vec3), - geomgroup: vec6 = None, + geomgroup: Optional[vec6] = None, flg_static: bool = True, bodyexclude: int = -1, -) -> tuple[wp.array2d(dtype=float), wp.array2d(dtype=int)]: +) -> Tuple[wp.array, wp.array]: """Returns the distance at which rays intersect with primitive geoms. Args: - m (Model): The model containing kinematic and dynamic information (device). - d (Data): The data object containing the current state and output arrays (device). - pnt (wp.array2d(dtype=wp.vec3)): Ray origin points. - vec (wp.array2d(dtype=wp.vec3)): Ray directions. - geomgroup (vec6, optional): Group inclusion/exclusion mask. - If all are wp.inf, ignore. - flg_static (bool, optional): If True, allows rays to intersect with static geoms. - Defaults to True. - bodyexclude (int, optional): Ignore geoms on specified body id (-1 to disable). - Defaults to -1. + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + pnt: Ray origin points. + vec: Ray directions. + geomgroup: Group inclusion/exclusion mask. If all are wp.inf, ignore. + flg_static: If True, allows rays to intersect with static geoms. + bodyexclude: Ignore geoms on specified body id (-1 to disable). Returns: - wp.array2d(dtype=float): Distances from ray origins to geom surfaces. - wp.array2d(dtype=int): IDs of intersected geoms (-1 if none). + Distances from ray origins to geom surfaces and IDs of intersected geoms (-1 if none). """ - + assert pnt.shape[0] == 1 assert pnt.shape[0] == vec.shape[0] - assert d.ray_dist.shape[1] == d.ray_geomid.shape[1] - assert pnt.shape[0] == d.ray_dist.shape[1] if geomgroup is None: geomgroup = vec6(-1, -1, -1, -1, -1, -1) - d.ray_bodyexclude.fill_(bodyexclude) + ray_bodyexclude = wp.empty(1, dtype=int) + ray_bodyexclude.fill_(bodyexclude) + ray_dist = wp.empty((d.nworld, 1), dtype=float) + ray_geomid = wp.empty((d.nworld, 1), dtype=int) - rays(m, d, pnt, vec, geomgroup, flg_static, d.ray_bodyexclude, d.ray_dist, d.ray_geomid) + rays(m, d, pnt, vec, geomgroup, flg_static, ray_bodyexclude, ray_dist, ray_geomid) - return d.ray_dist, d.ray_geomid + return ray_dist, ray_geomid def rays( @@ -868,19 +860,19 @@ def rays( m.geom_type, m.geom_bodyid, m.geom_dataid, - m.geom_group, m.geom_matid, + m.geom_group, m.geom_size, m.geom_rgba, - m.hfield_adr, + m.mesh_vertadr, + m.mesh_faceadr, + m.mesh_vert, + m.mesh_face, + m.hfield_size, m.hfield_nrow, m.hfield_ncol, - m.hfield_size, + m.hfield_adr, m.hfield_data, - m.mesh_vertadr, - m.mesh_vert, - m.mesh_faceadr, - m.mesh_face, m.mat_rgba, d.geom_xpos, d.geom_xmat, 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 4c24310f..70387cb3 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py @@ -21,17 +21,15 @@ from mujoco.mjx.third_party.mujoco_warp._src import math from mujoco.mjx.third_party.mujoco_warp._src import ray from mujoco.mjx.third_party.mujoco_warp._src import smooth from mujoco.mjx.third_party.mujoco_warp._src import support -from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import ccd -from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom 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_MINVAL from mujoco.mjx.third_party.mujoco_warp._src.types import ConeType from mujoco.mjx.third_party.mujoco_warp._src.types import ConstraintType +from mujoco.mjx.third_party.mujoco_warp._src.types import ContactType from mujoco.mjx.third_party.mujoco_warp._src.types import Data from mujoco.mjx.third_party.mujoco_warp._src.types import DataType from mujoco.mjx.third_party.mujoco_warp._src.types import DisableBit -from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType 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 @@ -119,7 +117,7 @@ def _magnetometer( worldid: int, objid: int, ) -> wp.vec3: - magnetic = opt_magnetic[worldid] + magnetic = opt_magnetic[worldid % opt_magnetic.shape[0]] return wp.transpose(site_xmat_in[worldid, objid]) @ magnetic @@ -198,9 +196,9 @@ def _sensor_rangefinder_init( # Data in: site_xpos_in: wp.array2d(dtype=wp.vec3), site_xmat_in: wp.array2d(dtype=wp.mat33), - # Data out: - sensor_rangefinder_pnt_out: wp.array2d(dtype=wp.vec3), - sensor_rangefinder_vec_out: wp.array2d(dtype=wp.vec3), + # Out: + pnt_out: wp.array2d(dtype=wp.vec3), + vec_out: wp.array2d(dtype=wp.vec3), ): worldid, rfid = wp.tid() sensorid = sensor_rangefinder_adr[rfid] @@ -208,8 +206,8 @@ def _sensor_rangefinder_init( site_xpos = site_xpos_in[worldid, objid] site_xmat = site_xmat_in[worldid, objid] - sensor_rangefinder_pnt_out[worldid, rfid] = site_xpos - sensor_rangefinder_vec_out[worldid, rfid] = wp.vec3(site_xmat[0, 2], site_xmat[1, 2], site_xmat[2, 2]) + pnt_out[worldid, rfid] = site_xpos + vec_out[worldid, rfid] = wp.vec3(site_xmat[0, 2], site_xmat[1, 2], site_xmat[2, 2]) @wp.func @@ -408,16 +406,20 @@ def _frame_quat( refid: int, reftype: int, ) -> wp.quat: + body_iquat_id = worldid % body_iquat.shape[0] + geom_quat_id = worldid % geom_quat.shape[0] + site_quat_id = worldid % site_quat.shape[0] + cam_quat_id = worldid % cam_quat.shape[0] if objtype == ObjType.BODY: - quat = math.mul_quat(xquat_in[worldid, objid], body_iquat[worldid, objid]) + quat = math.mul_quat(xquat_in[worldid, objid], body_iquat[body_iquat_id, objid]) elif objtype == ObjType.XBODY: quat = xquat_in[worldid, objid] elif objtype == ObjType.GEOM: - quat = math.mul_quat(xquat_in[worldid, geom_bodyid[objid]], geom_quat[worldid, objid]) + quat = math.mul_quat(xquat_in[worldid, geom_bodyid[objid]], geom_quat[geom_quat_id, objid]) elif objtype == ObjType.SITE: - quat = math.mul_quat(xquat_in[worldid, site_bodyid[objid]], site_quat[worldid, objid]) + quat = math.mul_quat(xquat_in[worldid, site_bodyid[objid]], site_quat[site_quat_id, objid]) elif objtype == ObjType.CAMERA: - quat = math.mul_quat(xquat_in[worldid, cam_bodyid[objid]], cam_quat[worldid, objid]) + quat = math.mul_quat(xquat_in[worldid, cam_bodyid[objid]], cam_quat[cam_quat_id, objid]) else: # UNKNOWN quat = wp.quat(1.0, 0.0, 0.0, 0.0) @@ -425,15 +427,15 @@ def _frame_quat( return quat if reftype == ObjType.BODY: - refquat = math.mul_quat(xquat_in[worldid, refid], body_iquat[worldid, refid]) + refquat = math.mul_quat(xquat_in[worldid, refid], body_iquat[body_iquat_id, refid]) elif reftype == ObjType.XBODY: refquat = xquat_in[worldid, refid] elif reftype == ObjType.GEOM: - refquat = math.mul_quat(xquat_in[worldid, geom_bodyid[refid]], geom_quat[worldid, refid]) + refquat = math.mul_quat(xquat_in[worldid, geom_bodyid[refid]], geom_quat[geom_quat_id, refid]) elif reftype == ObjType.SITE: - refquat = math.mul_quat(xquat_in[worldid, site_bodyid[refid]], site_quat[worldid, refid]) + refquat = math.mul_quat(xquat_in[worldid, site_bodyid[refid]], site_quat[site_quat_id, refid]) elif reftype == ObjType.CAMERA: - refquat = math.mul_quat(xquat_in[worldid, cam_bodyid[refid]], cam_quat[worldid, refid]) + refquat = math.mul_quat(xquat_in[worldid, cam_bodyid[refid]], cam_quat[cam_quat_id, refid]) else: # UNKNOWN refquat = wp.quat(1.0, 0.0, 0.0, 0.0) @@ -453,17 +455,14 @@ def _clock(time_in: wp.array(dtype=float), worldid: int) -> float: @wp.kernel def _sensor_pos( # Model: - opt_ccd_tolerance: wp.array(dtype=float), + ngeom: int, opt_magnetic: wp.array(dtype=wp.vec3), - opt_ccd_iterations: int, body_geomnum: wp.array(dtype=int), body_geomadr: wp.array(dtype=int), body_iquat: wp.array2d(dtype=wp.quat), jnt_qposadr: wp.array(dtype=int), geom_type: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), - geom_dataid: wp.array(dtype=int), - geom_size: wp.array2d(dtype=wp.vec3), geom_quat: wp.array2d(dtype=wp.quat), site_type: wp.array(dtype=int), site_bodyid: wp.array(dtype=int), @@ -475,20 +474,6 @@ def _sensor_pos( cam_resolution: wp.array(dtype=wp.vec2i), cam_sensorsize: wp.array(dtype=wp.vec2), cam_intrinsic: wp.array(dtype=wp.vec4), - mesh_vertadr: wp.array(dtype=int), - mesh_vertnum: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), - mesh_graphadr: wp.array(dtype=int), - mesh_graph: wp.array(dtype=int), - mesh_polynum: wp.array(dtype=int), - mesh_polyadr: wp.array(dtype=int), - mesh_polynormal: wp.array(dtype=wp.vec3), - mesh_polyvertadr: wp.array(dtype=int), - mesh_polyvertnum: wp.array(dtype=int), - mesh_polyvert: wp.array(dtype=int), - mesh_polymapadr: wp.array(dtype=int), - mesh_polymapnum: wp.array(dtype=int), - mesh_polymap: wp.array(dtype=int), sensor_type: wp.array(dtype=int), sensor_datatype: wp.array(dtype=int), sensor_objtype: wp.array(dtype=int), @@ -498,6 +483,7 @@ def _sensor_pos( sensor_adr: wp.array(dtype=int), sensor_cutoff: wp.array(dtype=float), sensor_pos_adr: wp.array(dtype=int), + sensor_collision_start_adr: wp.array(dtype=int), rangefinder_sensor_adr: wp.array(dtype=int), collision_sensor_adr: wp.array(dtype=int), # Data in: @@ -516,33 +502,19 @@ def _sensor_pos( cam_xpos_in: wp.array2d(dtype=wp.vec3), cam_xmat_in: wp.array2d(dtype=wp.mat33), subtree_com_in: wp.array2d(dtype=wp.vec3), - actuator_length_in: wp.array2d(dtype=float), - epa_vert_in: wp.array2d(dtype=wp.vec3), - epa_vert1_in: wp.array2d(dtype=wp.vec3), - epa_vert2_in: wp.array2d(dtype=wp.vec3), - epa_vert_index1_in: wp.array2d(dtype=int), - epa_vert_index2_in: wp.array2d(dtype=int), - epa_face_in: wp.array2d(dtype=wp.vec3i), - epa_pr_in: wp.array2d(dtype=wp.vec3), - epa_norm2_in: wp.array2d(dtype=float), - epa_index_in: wp.array2d(dtype=int), - epa_map_in: wp.array2d(dtype=int), - epa_horizon_in: wp.array2d(dtype=int), - multiccd_polygon_in: wp.array2d(dtype=wp.vec3), - multiccd_clipped_in: wp.array2d(dtype=wp.vec3), - multiccd_pnormal_in: wp.array2d(dtype=wp.vec3), - multiccd_pdist_in: wp.array2d(dtype=float), - multiccd_idx1_in: wp.array2d(dtype=int), - multiccd_idx2_in: wp.array2d(dtype=int), - multiccd_n1_in: wp.array2d(dtype=wp.vec3), - multiccd_n2_in: wp.array2d(dtype=wp.vec3), - multiccd_endvert_in: wp.array2d(dtype=wp.vec3), - multiccd_face1_in: wp.array2d(dtype=wp.vec3), - multiccd_face2_in: wp.array2d(dtype=wp.vec3), ten_length_in: wp.array2d(dtype=float), - sensor_rangefinder_dist_in: wp.array2d(dtype=float), + actuator_length_in: wp.array2d(dtype=float), + contact_dist_in: wp.array(dtype=float), + contact_pos_in: wp.array(dtype=wp.vec3), + contact_frame_in: wp.array(dtype=wp.mat33), + contact_geom_in: wp.array(dtype=wp.vec2i), + contact_worldid_in: wp.array(dtype=int), + contact_type_in: wp.array(dtype=int), + nacon_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=wp.vec2i), # In: - nsensor_collision: int, + rangefinder_dist_in: wp.array2d(dtype=float), + sensor_collision_in: wp.array4d(dtype=float), # Data out: sensordata_out: wp.array2d(dtype=float), ): @@ -562,7 +534,7 @@ def _sensor_pos( ) _write_vector(sensor_type, sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 2, vec2, out) elif sensortype == SensorType.RANGEFINDER: - val = sensor_rangefinder_dist_in[worldid, rangefinder_sensor_adr[sensorid]] + val = rangefinder_dist_in[worldid, rangefinder_sensor_adr[sensorid]] _write_scalar(sensor_type, sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) elif sensortype == SensorType.JOINTPOS: val = _joint_pos(jnt_qposadr, qpos_in, worldid, objid) @@ -635,25 +607,21 @@ def _sensor_pos( elif sensortype == SensorType.SUBTREECOM: vec3 = _subtree_com(subtree_com_in, worldid, objid) _write_vector(sensor_type, sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out) - elif ( - sensortype == int(SensorType.GEOMDIST.value) - or sensortype == int(SensorType.GEOMNORMAL.value) - or sensortype == int(SensorType.GEOMFROMTO.value) - ): + elif sensortype == SensorType.GEOMDIST or sensortype == SensorType.GEOMNORMAL or sensortype == SensorType.GEOMFROMTO: objtype = sensor_objtype[sensorid] + objid = sensor_objid[sensorid] reftype = sensor_reftype[sensorid] refid = sensor_refid[sensorid] - cutoff = sensor_cutoff[sensorid] - # initialize - dist = cutoff - fromto = vec6(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + dist = float(1.0e32) + pnts = vec6(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) + flip = bool(False) - # settings - tolerance = opt_ccd_tolerance[worldid] + collision_sensorid = collision_sensor_adr[sensorid] + collision_start_adr = sensor_collision_start_adr[collision_sensorid] - # get lists of geoms to collide + # check for flip direction if objtype == int(ObjType.BODY.value): n1 = body_geomnum[objid] id1 = body_geomadr[objid] @@ -667,114 +635,51 @@ def _sensor_pos( n2 = 1 id2 = refid - tid = worldid * nsensor_collision + collision_sensor_adr[sensorid] + for geom1 in range(n1): + for geom2 in range(n2): + collisionid = collision_start_adr + geom1 * n2 + geom2 + for i in range(8): + dist_new = sensor_collision_in[worldid, collisionid, i, 0] - # collide all pairs - for geom1id in range(id1, id1 + n1): - geomtype1 = geom_type[geom1id] - geom1_dataid = geom_dataid[geom1id] - pos1 = geom_xpos_in[worldid, geom1id] - geom1 = geom( - geomtype1, - geom1_dataid, - geom_size[worldid, geom1id], - mesh_vertadr, - mesh_vertnum, - mesh_vert, - mesh_graphadr, - mesh_graph, - mesh_polynum, - mesh_polyadr, - mesh_polynormal, - mesh_polyvertadr, - mesh_polyvertnum, - mesh_polyvert, - mesh_polymapadr, - mesh_polymapnum, - mesh_polymap, - pos1, - geom_xmat_in[worldid, geom1id], - ) - for geom2id in range(id2, id2 + n2): - geomtype2 = geom_type[geom2id] - geom2_dataid = geom_dataid[geom2id] - pos2 = geom_xpos_in[worldid, geom2id] - geom2 = geom( - geomtype2, - geom2_dataid, - geom_size[worldid, geom2id], - mesh_vertadr, - mesh_vertnum, - mesh_vert, - mesh_graphadr, - mesh_graph, - mesh_polynum, - mesh_polyadr, - mesh_polynormal, - mesh_polyvertadr, - mesh_polyvertnum, - mesh_polyvert, - mesh_polymapadr, - mesh_polymapnum, - mesh_polymap, - pos2, - geom_xmat_in[worldid, geom2id], - ) + if dist_new <= dist: + dist = dist_new - dist_new, _, witness1_new, witness2_new = ccd( - False, # no multiccd - tolerance, - cutoff, - opt_ccd_iterations, - geom1, - geom2, - geomtype1, - geomtype2, - pos1, - pos2, - epa_vert_in[tid], - epa_vert1_in[tid], - epa_vert2_in[tid], - epa_vert_index1_in[tid], - epa_vert_index2_in[tid], - epa_face_in[tid], - epa_pr_in[tid], - epa_norm2_in[tid], - epa_index_in[tid], - epa_map_in[tid], - epa_horizon_in[tid], - # TODO(team): since multiccd will always be off, empty arrays? - multiccd_polygon_in[tid], - multiccd_clipped_in[tid], - multiccd_pnormal_in[tid], - multiccd_pdist_in[tid], - multiccd_idx1_in[tid], - multiccd_idx2_in[tid], - multiccd_n1_in[tid], - multiccd_n2_in[tid], - multiccd_endvert_in[tid], - multiccd_face1_in[tid], - multiccd_face2_in[tid], - ) + if sensortype == SensorType.GEOMNORMAL or sensortype == SensorType.GEOMFROMTO: + pnts = vec6( + sensor_collision_in[worldid, collisionid, i, 1], + sensor_collision_in[worldid, collisionid, i, 2], + sensor_collision_in[worldid, collisionid, i, 3], + sensor_collision_in[worldid, collisionid, i, 4], + sensor_collision_in[worldid, collisionid, i, 5], + sensor_collision_in[worldid, collisionid, i, 6], + ) - if dist_new < dist: - dist = dist_new - fromto = vec6( - witness1_new[0][0], - witness1_new[0][1], - witness1_new[0][2], - witness2_new[0][0], - witness2_new[0][1], - witness2_new[0][2], - ) + geomid1 = id1 + geom1 + geomid2 = id2 + geom2 + if geom_type[geomid2] < geom_type[geomid1]: + flip = True + elif geom_type[geomid1] == geom_type[geomid2]: + if geomid2 < geomid1: + flip = True if sensortype == int(SensorType.GEOMDIST.value): _write_scalar(sensor_type, sensor_datatype, sensor_adr, sensor_cutoff, sensorid, dist, out) elif sensortype == int(SensorType.GEOMNORMAL.value): - normal = wp.vec3(fromto[3] - fromto[0], fromto[4] - fromto[1], fromto[5] - fromto[2]) - normal = wp.normalize(normal) + if dist <= sensor_cutoff[sensorid]: + normal = wp.normalize(wp.vec3(pnts[3] - pnts[0], pnts[4] - pnts[1], pnts[5] - pnts[2])) + if flip: + normal *= -1.0 + else: + normal = wp.vec3(0.0, 0.0, 0.0) _write_vector(sensor_type, sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, normal, out) elif sensortype == int(SensorType.GEOMFROMTO.value): + if dist <= sensor_cutoff[sensorid]: + if flip: + fromto = vec6(pnts[3], pnts[4], pnts[5], pnts[0], pnts[1], pnts[2]) + else: + fromto = pnts + else: + fromto = vec6(0.0, 0.0, 0.0, 0.0, 0.0, 0.0) _write_vector(sensor_type, sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 6, fromto, out) elif sensortype == SensorType.INSIDESITE: objtype = sensor_objtype[sensorid] @@ -804,42 +709,89 @@ def _sensor_pos( _write_scalar(sensor_type, sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, out) +@wp.kernel +def _sensor_collision( + # Model: + ngeom: int, + # Data in: + contact_dist_in: wp.array(dtype=float), + contact_pos_in: wp.array(dtype=wp.vec3), + contact_frame_in: wp.array(dtype=wp.mat33), + contact_geom_in: wp.array(dtype=wp.vec2i), + contact_worldid_in: wp.array(dtype=int), + contact_type_in: wp.array(dtype=int), + contact_geomcollisionid_in: wp.array(dtype=int), + nacon_in: wp.array(dtype=int), + collision_pairid_in: wp.array(dtype=wp.vec2i), + # Out: + sensor_collision_out: wp.array4d(dtype=float), +): + conid = wp.tid() + + if conid >= nacon_in[0]: + return + + if not contact_type_in[conid] & ContactType.SENSOR: + return + + geom = contact_geom_in[conid] + if geom[0] <= geom[1]: + pairid = math.upper_tri_index(ngeom, geom[0], geom[1]) + else: + pairid = math.upper_tri_index(ngeom, geom[1], geom[0]) + + worldid = contact_worldid_in[conid] + collisionid = collision_pairid_in[pairid][1] + geomcollisionid = contact_geomcollisionid_in[conid] + + dist = contact_dist_in[conid] + pos = contact_pos_in[conid] + frame = contact_frame_in[conid] + normal = wp.vec3(frame[0, 0], frame[0, 1], frame[0, 2]) + pnt1 = pos - 0.5 * dist * normal + pnt2 = pos + 0.5 * dist * normal + + sensor_collision_out[worldid, collisionid, geomcollisionid, 0] = dist + sensor_collision_out[worldid, collisionid, geomcollisionid, 1] = pnt1[0] + sensor_collision_out[worldid, collisionid, geomcollisionid, 2] = pnt1[1] + sensor_collision_out[worldid, collisionid, geomcollisionid, 3] = pnt1[2] + sensor_collision_out[worldid, collisionid, geomcollisionid, 4] = pnt2[0] + sensor_collision_out[worldid, collisionid, geomcollisionid, 5] = pnt2[1] + sensor_collision_out[worldid, collisionid, geomcollisionid, 6] = pnt2[2] + + @event_scope def sensor_pos(m: Model, d: Data): """Compute position-dependent sensor values.""" - if m.opt.disableflags & DisableBit.SENSOR: return # rangefinder + rangefinder_dist = wp.empty((d.nworld, m.nrangefinder), dtype=float) if m.sensor_rangefinder_adr.size > 0: + rangefinder_pnt = wp.empty((d.nworld, m.nrangefinder), dtype=wp.vec3) + rangefinder_vec = wp.empty((d.nworld, m.nrangefinder), dtype=wp.vec3) + rangefinder_geomid = wp.empty((d.nworld, m.nrangefinder), dtype=int) + # get position and direction wp.launch( _sensor_rangefinder_init, dim=(d.nworld, m.sensor_rangefinder_adr.size), - inputs=[ - m.sensor_objid, - m.sensor_rangefinder_adr, - d.site_xpos, - d.site_xmat, - ], - outputs=[ - d.sensor_rangefinder_pnt, - d.sensor_rangefinder_vec, - ], + inputs=[m.sensor_objid, m.sensor_rangefinder_adr, d.site_xpos, d.site_xmat], + outputs=[rangefinder_pnt, rangefinder_vec], ) # get distances ray.rays( m, d, - d.sensor_rangefinder_pnt, - d.sensor_rangefinder_vec, + rangefinder_pnt, + rangefinder_vec, vec6(wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf), True, m.sensor_rangefinder_bodyid, - d.sensor_rangefinder_dist, - d.sensor_rangefinder_geomid, + rangefinder_dist, + rangefinder_geomid, ) if m.sensor_e_potential: @@ -848,21 +800,40 @@ def sensor_pos(m: Model, d: Data): if m.sensor_e_kinetic: 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) + if m.nsensorcollision: + wp.launch( + _sensor_collision, + dim=d.naconmax, + inputs=[ + m.ngeom, + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.geom, + d.contact.worldid, + d.contact.type, + d.contact.geomcollisionid, + d.nacon, + d.collision_pairid, + ], + outputs=[sensor_collision], + ) + wp.launch( _sensor_pos, dim=(d.nworld, m.sensor_pos_adr.size), inputs=[ - m.opt.ccd_tolerance, + m.ngeom, m.opt.magnetic, - m.opt.ccd_iterations, m.body_geomnum, m.body_geomadr, m.body_iquat, m.jnt_qposadr, m.geom_type, m.geom_bodyid, - m.geom_dataid, - m.geom_size, m.geom_quat, m.site_type, m.site_bodyid, @@ -874,20 +845,6 @@ def sensor_pos(m: Model, d: Data): m.cam_resolution, m.cam_sensorsize, m.cam_intrinsic, - m.mesh_vertadr, - m.mesh_vertnum, - m.mesh_vert, - m.mesh_graphadr, - m.mesh_graph, - m.mesh_polynum, - m.mesh_polyadr, - m.mesh_polynormal, - m.mesh_polyvertadr, - m.mesh_polyvertnum, - m.mesh_polyvert, - m.mesh_polymapadr, - m.mesh_polymapnum, - m.mesh_polymap, m.sensor_type, m.sensor_datatype, m.sensor_objtype, @@ -897,6 +854,7 @@ def sensor_pos(m: Model, d: Data): m.sensor_adr, m.sensor_cutoff, m.sensor_pos_adr, + m.sensor_collision_start_adr, m.rangefinder_sensor_adr, m.collision_sensor_adr, d.time, @@ -914,32 +872,18 @@ def sensor_pos(m: Model, d: Data): d.cam_xpos, d.cam_xmat, d.subtree_com, - d.actuator_length, - d.epa_vert, - d.epa_vert1, - d.epa_vert2, - d.epa_vert_index1, - d.epa_vert_index2, - d.epa_face, - d.epa_pr, - d.epa_norm2, - d.epa_index, - d.epa_map, - d.epa_horizon, - d.multiccd_polygon, - d.multiccd_clipped, - d.multiccd_pnormal, - d.multiccd_pdist, - d.multiccd_idx1, - d.multiccd_idx2, - d.multiccd_n1, - d.multiccd_n2, - d.multiccd_endvert, - d.multiccd_face1, - d.multiccd_face2, d.ten_length, - d.sensor_rangefinder_dist, - m.collision_sensor_adr.size, + d.actuator_length, + d.contact.dist, + d.contact.pos, + d.contact.frame, + d.contact.geom, + d.contact.worldid, + d.contact.type, + d.nacon, + d.collision_pairid, + rangefinder_dist, + sensor_collision, ], outputs=[d.sensordata], ) @@ -1440,7 +1384,6 @@ def _sensor_vel( @event_scope def sensor_vel(m: Model, d: Data): """Compute velocity-dependent sensor values.""" - if m.opt.disableflags & DisableBit.SENSOR: return @@ -1775,8 +1718,6 @@ def _sensor_acc( sensor_acc_adr: wp.array(dtype=int), sensor_adr_to_contact_adr: wp.array(dtype=int), # Data in: - njmax_in: int, - nacon_in: wp.array(dtype=int), xpos_in: wp.array2d(dtype=wp.vec3), xipos_in: wp.array2d(dtype=wp.vec3), geom_xpos_in: wp.array2d(dtype=wp.vec3), @@ -1787,6 +1728,8 @@ def _sensor_acc( cvel_in: wp.array2d(dtype=wp.spatial_vector), actuator_force_in: wp.array2d(dtype=float), qfrc_actuator_in: wp.array2d(dtype=float), + cacc_in: wp.array2d(dtype=wp.spatial_vector), + cfrc_int_in: wp.array2d(dtype=wp.spatial_vector), contact_dist_in: wp.array(dtype=float), contact_pos_in: wp.array(dtype=wp.vec3), contact_frame_in: wp.array(dtype=wp.mat33), @@ -1794,8 +1737,9 @@ def _sensor_acc( contact_dim_in: wp.array(dtype=int), contact_efc_address_in: wp.array2d(dtype=int), efc_force_in: wp.array2d(dtype=float), - cacc_in: wp.array2d(dtype=wp.spatial_vector), - cfrc_int_in: wp.array2d(dtype=wp.spatial_vector), + njmax_in: int, + nacon_in: wp.array(dtype=int), + # In: sensor_contact_nmatch_in: wp.array2d(dtype=int), sensor_contact_matchid_in: wp.array3d(dtype=int), sensor_contact_direction_in: wp.array3d(dtype=float), @@ -1865,13 +1809,13 @@ def _sensor_acc( contact_forcetorque = support.contact_force_fn( opt_cone, - njmax_in, - nacon_in, contact_frame_in, contact_friction_in, contact_dim_in, contact_efc_address_in, efc_force_in, + njmax_in, + nacon_in, worldid, cid, False, @@ -1895,13 +1839,13 @@ def _sensor_acc( contact_forcetorque = support.contact_force_fn( opt_cone, - njmax_in, - nacon_in, contact_frame_in, contact_friction_in, contact_dim_in, contact_efc_address_in, efc_force_in, + njmax_in, + nacon_in, worldid, cid, False, @@ -1972,13 +1916,13 @@ def _sensor_acc( if force or torque: contact_forcetorque = support.contact_force_fn( opt_cone, - njmax_in, - nacon_in, contact_frame_in, contact_friction_in, contact_dim_in, contact_efc_address_in, efc_force_in, + njmax_in, + nacon_in, worldid, cid, False, @@ -2082,7 +2026,6 @@ def _sensor_touch( sensor_adr: wp.array(dtype=int), sensor_touch_adr: wp.array(dtype=int), # Data in: - nacon_in: wp.array(dtype=int), site_xpos_in: wp.array2d(dtype=wp.vec3), site_xmat_in: wp.array2d(dtype=wp.mat33), contact_pos_in: wp.array(dtype=wp.vec3), @@ -2092,6 +2035,7 @@ def _sensor_touch( contact_efc_address_in: wp.array2d(dtype=int), contact_worldid_in: wp.array(dtype=int), efc_force_in: wp.array2d(dtype=float), + nacon_in: wp.array(dtype=int), # Data out: sensordata_out: wp.array2d(dtype=float), ): @@ -2161,17 +2105,17 @@ def _sensor_tactile( # Model: body_rootid: wp.array(dtype=int), body_weldid: wp.array(dtype=int), + oct_child: wp.array(dtype=vec8i), + oct_aabb: wp.array2d(dtype=wp.vec3), + oct_coeff: wp.array(dtype=vec8f), geom_type: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), mesh_vertadr: wp.array(dtype=int), - mesh_vert: wp.array(dtype=wp.vec3), mesh_normaladr: wp.array(dtype=int), + mesh_vert: wp.array(dtype=wp.vec3), mesh_normal: wp.array(dtype=wp.vec3), mesh_quat: wp.array(dtype=wp.quat), - oct_aabb: wp.array2d(dtype=wp.vec3), - oct_child: wp.array(dtype=vec8i), - oct_coeff: wp.array(dtype=vec8f), sensor_objid: wp.array(dtype=int), sensor_refid: wp.array(dtype=int), sensor_dim: wp.array(dtype=int), @@ -2182,13 +2126,13 @@ def _sensor_tactile( taxel_vertadr: wp.array(dtype=int), taxel_sensorid: wp.array(dtype=int), # Data in: - nacon_in: wp.array(dtype=int), geom_xpos_in: wp.array2d(dtype=wp.vec3), geom_xmat_in: wp.array2d(dtype=wp.mat33), subtree_com_in: wp.array2d(dtype=wp.vec3), cvel_in: wp.array2d(dtype=wp.spatial_vector), contact_geom_in: wp.array(dtype=wp.vec2i), contact_worldid_in: wp.array(dtype=int), + nacon_in: wp.array(dtype=int), # Data out: sensordata_out: wp.array2d(dtype=float), ): @@ -2236,7 +2180,7 @@ def _sensor_tactile( contact_type = geom_type[geom] plugin_attributes, plugin_index, volume_data, mesh_data = get_sdf_params( - oct_aabb, oct_child, oct_coeff, plugin, plugin_attr, contact_type, geom_size[worldid, geom], plugin_id, mesh_id + oct_child, oct_aabb, oct_coeff, plugin, plugin_attr, contact_type, geom_size[worldid, geom], plugin_id, mesh_id ) depth = wp.min(sdf(contact_type, lpos, plugin_attributes, plugin_index, volume_data, mesh_data), 0.0) @@ -2307,8 +2251,6 @@ def _contact_match( sensor_intprm: wp.array2d(dtype=int), sensor_contact_adr: wp.array(dtype=int), # Data in: - njmax_in: int, - nacon_in: wp.array(dtype=int), site_xpos_in: wp.array2d(dtype=wp.vec3), site_xmat_in: wp.array2d(dtype=wp.mat33), contact_dist_in: wp.array(dtype=float), @@ -2319,8 +2261,11 @@ def _contact_match( contact_geom_in: wp.array(dtype=wp.vec2i), contact_efc_address_in: wp.array2d(dtype=int), contact_worldid_in: wp.array(dtype=int), + contact_type_in: wp.array(dtype=int), efc_force_in: wp.array2d(dtype=float), - # Data out: + njmax_in: int, + nacon_in: wp.array(dtype=int), + # Out: sensor_contact_nmatch_out: wp.array2d(dtype=int), sensor_contact_matchid_out: wp.array3d(dtype=int), sensor_contact_criteria_out: wp.array3d(dtype=float), @@ -2332,6 +2277,9 @@ def _contact_match( if contactid >= nacon_in[0]: return + if not contact_type_in[contactid] & ContactType.CONSTRAINT: + return + # sensor information objtype = sensor_objtype[sensorid] objid = sensor_objid[sensorid] @@ -2402,13 +2350,13 @@ def _contact_match( elif reduce == 2: # maxforce contact_force = support.contact_force_fn( opt_cone, - njmax_in, - nacon_in, contact_frame_in, contact_friction_in, contact_dim_in, contact_efc_address_in, efc_force_in, + njmax_in, + nacon_in, worldid, contactid, False, @@ -2477,7 +2425,6 @@ def sensor_acc(m: Model, d: Data): m.sensor_objid, m.sensor_adr, m.sensor_touch_adr, - d.nacon, d.site_xpos, d.site_xmat, d.contact.pos, @@ -2487,6 +2434,7 @@ def sensor_acc(m: Model, d: Data): d.contact.efc_address, d.contact.worldid, d.efc.force, + d.nacon, ], outputs=[ d.sensordata, @@ -2499,17 +2447,17 @@ def sensor_acc(m: Model, d: Data): inputs=[ m.body_rootid, m.body_weldid, + m.oct_child, + m.oct_aabb, + m.oct_coeff, m.geom_type, m.geom_bodyid, m.geom_size, m.mesh_vertadr, - m.mesh_vert, m.mesh_normaladr, + m.mesh_vert, m.mesh_normal, m.mesh_quat, - m.oct_aabb, - m.oct_child, - m.oct_coeff, m.sensor_objid, m.sensor_refid, m.sensor_dim, @@ -2519,24 +2467,28 @@ def sensor_acc(m: Model, d: Data): m.geom_plugin_index, m.taxel_vertadr, m.taxel_sensorid, - d.nacon, d.geom_xpos, d.geom_xmat, d.subtree_com, d.cvel, d.contact.geom, d.contact.worldid, + d.nacon, ], outputs=[ d.sensordata, ], ) - if m.sensor_contact_adr.size: - # match criteria - d.sensor_contact_nmatch.zero_() - d.sensor_contact_matchid.fill_(-1) - d.sensor_contact_criteria.fill_(1.0e32) + sensor_contact_nmatch = wp.empty((d.nworld, m.nsensorcontact), dtype=int) + sensor_contact_matchid = wp.empty((d.nworld, m.nsensorcontact, m.opt.contact_sensor_maxmatch), dtype=int) + sensor_contact_direction = wp.empty((d.nworld, m.nsensorcontact, m.opt.contact_sensor_maxmatch), dtype=float) + if m.nsensorcontact: + sensor_contact_criteria = wp.empty((d.nworld, m.nsensorcontact, m.opt.contact_sensor_maxmatch), dtype=float) + # TODO(team): fill_ operations in one kernel? + sensor_contact_nmatch.fill_(0) + sensor_contact_matchid.fill_(-1) + sensor_contact_criteria.fill_(1.0e32) wp.launch( _contact_match, @@ -2554,8 +2506,6 @@ def sensor_acc(m: Model, d: Data): m.sensor_refid, m.sensor_intprm, m.sensor_contact_adr, - d.njmax, - d.nacon, d.site_xpos, d.site_xmat, d.contact.dist, @@ -2566,30 +2516,20 @@ def sensor_acc(m: Model, d: Data): d.contact.geom, d.contact.efc_address, d.contact.worldid, + d.contact.type, d.efc.force, + d.njmax, + d.nacon, ], - outputs=[ - d.sensor_contact_nmatch, - d.sensor_contact_matchid, - d.sensor_contact_criteria, - d.sensor_contact_direction, - ], + outputs=[sensor_contact_nmatch, sensor_contact_matchid, sensor_contact_criteria, sensor_contact_direction], ) # sorting wp.launch_tiled( _contact_sort(m.opt.contact_sensor_maxmatch), dim=(d.nworld, m.sensor_contact_adr.size), - inputs=[ - m.sensor_intprm, - m.sensor_contact_adr, - d.sensor_contact_nmatch, - d.sensor_contact_matchid, - d.sensor_contact_criteria, - ], - outputs=[ - d.sensor_contact_matchid, - ], + inputs=[m.sensor_intprm, m.sensor_contact_adr, sensor_contact_nmatch, sensor_contact_matchid, sensor_contact_criteria], + outputs=[sensor_contact_matchid], block_dim=m.block_dim.contact_sort, ) @@ -2616,8 +2556,6 @@ def sensor_acc(m: Model, d: Data): m.sensor_cutoff, m.sensor_acc_adr, m.sensor_adr_to_contact_adr, - d.njmax, - d.nacon, d.xpos, d.xipos, d.geom_xpos, @@ -2628,6 +2566,8 @@ def sensor_acc(m: Model, d: Data): d.cvel, d.actuator_force, d.qfrc_actuator, + d.cacc, + d.cfrc_int, d.contact.dist, d.contact.pos, d.contact.frame, @@ -2635,11 +2575,11 @@ def sensor_acc(m: Model, d: Data): d.contact.dim, d.contact.efc_address, d.efc.force, - d.cacc, - d.cfrc_int, - d.sensor_contact_nmatch, - d.sensor_contact_matchid, - d.sensor_contact_direction, + d.njmax, + d.nacon, + sensor_contact_nmatch, + sensor_contact_matchid, + sensor_contact_direction, ], outputs=[d.sensordata], ) @@ -2717,11 +2657,11 @@ def _energy_pos_gravity( energy_out: wp.array(dtype=wp.vec2), ): worldid, bodyid = wp.tid() - gravity = opt_gravity[worldid] + gravity = opt_gravity[worldid % opt_gravity.shape[0]] bodyid += 1 # skip world body energy = wp.vec2( - body_mass[worldid, bodyid] * wp.dot(gravity, xipos_in[worldid, bodyid]), + body_mass[worldid % body_mass.shape[0], bodyid] * wp.dot(gravity, xipos_in[worldid, bodyid]), 0.0, ) @@ -2741,19 +2681,21 @@ def _energy_pos_passive_joint( energy_out: wp.array(dtype=wp.vec2), ): worldid, jntid = wp.tid() - stiffness = jnt_stiffness[worldid, jntid] + jnt_stiffness_id = worldid % jnt_stiffness.shape[0] + stiffness = jnt_stiffness[jnt_stiffness_id, jntid] if stiffness == 0.0: return padr = jnt_qposadr[jntid] jnttype = jnt_type[jntid] + qpos_spring_id = worldid % qpos_spring.shape[0] if jnttype == JointType.FREE: dif0 = wp.vec3( - qpos_in[worldid, padr + 0] - qpos_spring[worldid, padr + 0], - qpos_in[worldid, padr + 1] - qpos_spring[worldid, padr + 1], - qpos_in[worldid, padr + 2] - qpos_spring[worldid, padr + 2], + qpos_in[worldid, padr + 0] - qpos_spring[qpos_spring_id, padr + 0], + qpos_in[worldid, padr + 1] - qpos_spring[qpos_spring_id, padr + 1], + qpos_in[worldid, padr + 2] - qpos_spring[qpos_spring_id, padr + 2], ) # convert quaternion difference into angular "velocity" @@ -2766,10 +2708,10 @@ def _energy_pos_passive_joint( quat1 = wp.normalize(quat1) quat_spring = wp.quat( - qpos_spring[worldid, padr + 3], - qpos_spring[worldid, padr + 4], - qpos_spring[worldid, padr + 5], - qpos_spring[worldid, padr + 6], + qpos_spring[qpos_spring_id, padr + 3], + qpos_spring[qpos_spring_id, padr + 4], + qpos_spring[qpos_spring_id, padr + 5], + qpos_spring[qpos_spring_id, padr + 6], ) dif1 = math.quat_sub(quat1, quat_spring) @@ -2791,10 +2733,10 @@ def _energy_pos_passive_joint( quat = wp.normalize(quat) quat_spring = wp.quat( - qpos_spring[worldid, padr + 0], - qpos_spring[worldid, padr + 1], - qpos_spring[worldid, padr + 2], - qpos_spring[worldid, padr + 3], + qpos_spring[qpos_spring_id, padr + 0], + qpos_spring[qpos_spring_id, padr + 1], + qpos_spring[qpos_spring_id, padr + 2], + qpos_spring[qpos_spring_id, padr + 3], ) dif = math.quat_sub(quat, quat_spring) @@ -2804,7 +2746,7 @@ def _energy_pos_passive_joint( ) wp.atomic_add(energy_out, worldid, energy) elif jnttype == JointType.SLIDE or jnttype == JointType.HINGE: - dif_ = qpos_in[worldid, padr] - qpos_spring[worldid, padr] + dif_ = qpos_in[worldid, padr] - qpos_spring[qpos_spring_id, padr] energy = wp.vec2( 0.5 * stiffness * dif_ * dif_, 0.0, @@ -2824,7 +2766,8 @@ def _energy_pos_passive_tendon( ): worldid, tenid = wp.tid() - stiffness = tendon_stiffness[worldid, tenid] + tendon_stiffness_id = worldid % tendon_stiffness.shape[0] + stiffness = tendon_stiffness[tendon_stiffness_id, tenid] if stiffness == 0.0: return @@ -2832,7 +2775,8 @@ def _energy_pos_passive_tendon( length = ten_length_in[worldid, tenid] # compute spring displacement - lengthspring = tendon_lengthspring[worldid, tenid] + tendon_lengthspring_id = worldid % tendon_lengthspring.shape[0] + lengthspring = tendon_lengthspring[tendon_lengthspring_id, tenid] lower = lengthspring[0] upper = lengthspring[1] @@ -2917,12 +2861,10 @@ def _energy_vel_kinetic(nv: int): def energy_vel(m: Model, d: Data): """Velocity-dependent energy (kinetic).""" - # kinetic energy: 0.5 * qvel.T @ M @ qvel # M @ qvel - skip = wp.zeros(d.nworld, dtype=bool) - support.mul_m(m, d, d.efc.mv, d.qvel, skip) + support.mul_m(m, d, d.efc.mv, d.qvel) wp.launch_tiled( _energy_vel_kinetic(m.nv), 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 864684f5..729f27a1 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py @@ -13,6 +13,7 @@ # limitations under the License. # ============================================================================== + import warp as wp from mujoco.mjx.third_party.mujoco_warp._src import math @@ -92,12 +93,15 @@ def _kinematics_level( jntadr = body_jntadr[bodyid] jntnum = body_jntnum[bodyid] 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] if jntnum == 0: # no joints - apply fixed translation and rotation relative to parent pid = body_parentid[bodyid] - xpos = (xmat_in[worldid, pid] * body_pos[worldid, bodyid]) + xpos_in[worldid, pid] - xquat = math.mul_quat(xquat_in[worldid, pid], body_quat[worldid, bodyid]) + xpos = (xmat_in[worldid, pid] * body_pos[body_pos_id, bodyid]) + xpos_in[worldid, pid] + xquat = math.mul_quat(xquat_in[worldid, pid], body_quat[body_quat_id, bodyid]) elif jntnum == 1 and jnt_type[jntadr] == JointType.FREE: # free joint qadr = jnt_qposadr[jntadr] @@ -105,19 +109,21 @@ def _kinematics_level( 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[worldid, jntadr] + 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] - xpos = (xmat_in[worldid, pid] * body_pos[worldid, bodyid]) + xpos_in[worldid, pid] - xquat = math.mul_quat(xquat_in[worldid, pid], body_quat[worldid, bodyid]) + xpos = (xmat_in[worldid, pid] * body_pos[body_pos_id, bodyid]) + xpos_in[worldid, pid] + xquat = math.mul_quat(xquat_in[worldid, pid], body_quat[body_quat_id, bodyid]) for _ in range(jntnum): qadr = jnt_qposadr[jntadr] jnt_type_ = jnt_type[jntadr] - jnt_axis_ = jnt_axis[worldid, jntadr] - xanchor = math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat) + xpos + jnt_axis_ = jnt_axis[jnt_axis_id, jntadr] + xanchor = math.rot_vec_quat(jnt_pos[jnt_pos_id, jntadr], xquat) + xpos xaxis = math.rot_vec_quat(jnt_axis_, xquat) if jnt_type_ == JointType.BALL: @@ -130,15 +136,15 @@ def _kinematics_level( qloc = wp.normalize(qloc) xquat = math.mul_quat(xquat, qloc) # correct for off-center rotation - xpos = xanchor - math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat) + xpos = xanchor - math.rot_vec_quat(jnt_pos[jnt_pos_id, jntadr], xquat) elif jnt_type_ == JointType.SLIDE: - xpos += xaxis * (qpos[qadr] - qpos0[worldid, qadr]) + xpos += xaxis * (qpos[qadr] - qpos0[qpos0_id, qadr]) elif jnt_type_ == JointType.HINGE: - qpos0_ = qpos0[worldid, qadr] + qpos0_ = qpos0[qpos0_id, qadr] qloc_ = math.axis_angle_to_quat(jnt_axis_, qpos[qadr] - qpos0_) xquat = math.mul_quat(xquat, qloc_) # correct for off-center rotation - xpos = xanchor - math.rot_vec_quat(jnt_pos[worldid, jntadr], xquat) + xpos = xanchor - math.rot_vec_quat(jnt_pos[jnt_pos_id, jntadr], xquat) xanchor_out[worldid, jntadr] = xanchor xaxis_out[worldid, jntadr] = xaxis @@ -147,37 +153,38 @@ def _kinematics_level( xpos_out[worldid, bodyid] = xpos xquat_out[worldid, bodyid] = wp.normalize(xquat) xmat_out[worldid, bodyid] = math.quat_to_mat(xquat) - xipos_out[worldid, bodyid] = xpos + math.rot_vec_quat(body_ipos[worldid, bodyid], xquat) - ximat_out[worldid, bodyid] = math.quat_to_mat(math.mul_quat(xquat, body_iquat[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 _geom_local_to_global( # Model: + body_rootid: wp.array(dtype=int), + body_weldid: wp.array(dtype=int), + body_mocapid: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), geom_pos: wp.array2d(dtype=wp.vec3), geom_quat: wp.array2d(dtype=wp.quat), # Data in: xpos_in: wp.array2d(dtype=wp.vec3), xquat_in: wp.array2d(dtype=wp.quat), - geom_skip_in: wp.array(dtype=bool), # Data out: - geom_skip_out: wp.array(dtype=bool), geom_xpos_out: wp.array2d(dtype=wp.vec3), geom_xmat_out: wp.array2d(dtype=wp.mat33), ): worldid, geomid = wp.tid() bodyid = geom_bodyid[geomid] - if not geom_skip_in[geomid]: - # Calculate only if necessary - xpos = xpos_in[worldid, bodyid] - xquat = xquat_in[worldid, bodyid] - geom_xpos_out[worldid, geomid] = xpos + math.rot_vec_quat(geom_pos[worldid, geomid], xquat) - geom_xmat_out[worldid, geomid] = math.quat_to_mat(math.mul_quat(xquat, geom_quat[worldid, geomid])) - if bodyid == 0: - # static geom pose are calculated only once - geom_skip_out[geomid] = True + if body_weldid[bodyid] == 0 and body_mocapid[body_rootid[bodyid]] == -1: + # geoms attached to the world are static (unless they are descended from mcocap bodies) + # for such static geoms, geom_xpos and geom_xquat are computed only once during make_data + return + + xpos = xpos_in[worldid, bodyid] + xquat = xquat_in[worldid, bodyid] + geom_xpos_out[worldid, geomid] = xpos + math.rot_vec_quat(geom_pos[worldid % geom_pos.shape[0], geomid], xquat) + geom_xmat_out[worldid, geomid] = math.quat_to_mat(math.mul_quat(xquat, geom_quat[worldid % geom_quat.shape[0], geomid])) @wp.kernel @@ -197,8 +204,8 @@ def _site_local_to_global( bodyid = site_bodyid[siteid] xpos = xpos_in[worldid, bodyid] xquat = xquat_in[worldid, bodyid] - site_xpos_out[worldid, siteid] = xpos + math.rot_vec_quat(site_pos[worldid, siteid], xquat) - site_xmat_out[worldid, siteid] = math.quat_to_mat(math.mul_quat(xquat, site_quat[worldid, siteid])) + site_xpos_out[worldid, siteid] = xpos + math.rot_vec_quat(site_pos[worldid % site_pos.shape[0], siteid], xquat) + site_xmat_out[worldid, siteid] = math.quat_to_mat(math.mul_quat(xquat, site_quat[worldid % site_quat.shape[0], siteid])) @wp.kernel @@ -268,14 +275,13 @@ def _mocap( xpos_out[worldid, bodyid] = xpos xquat_out[worldid, bodyid] = mocap_quat xmat_out[worldid, bodyid] = math.quat_to_mat(mocap_quat) - xipos_out[worldid, bodyid] = xpos + math.rot_vec_quat(body_ipos[worldid, bodyid], mocap_quat) - ximat_out[worldid, bodyid] = math.quat_to_mat(math.mul_quat(mocap_quat, body_iquat[worldid, bodyid])) + xipos_out[worldid, bodyid] = xpos + math.rot_vec_quat(body_ipos[worldid % body_ipos.shape[0], bodyid], mocap_quat) + ximat_out[worldid, bodyid] = math.quat_to_mat(math.mul_quat(mocap_quat, body_iquat[worldid % body_iquat.shape[0], bodyid])) @event_scope def kinematics(m: Model, d: Data): - """ - Computes forward kinematics for all bodies, sites, geoms, and flexible elements. + """Computes forward kinematics for all bodies, sites, geoms, and flexible elements. This function updates the global positions and orientations of all bodies, as well as the derived positions and orientations of geoms, sites, and flexible elements, based on the @@ -320,8 +326,8 @@ def kinematics(m: Model, d: Data): wp.launch( _geom_local_to_global, dim=(d.nworld, m.ngeom), - inputs=[m.geom_bodyid, m.geom_pos, m.geom_quat, d.xpos, d.xquat, d.geom_skip], - outputs=[d.geom_skip, d.geom_xpos, d.geom_xmat], + inputs=[m.body_rootid, m.body_weldid, m.body_mocapid, m.geom_bodyid, m.geom_pos, m.geom_quat, d.xpos, d.xquat], + outputs=[d.geom_xpos, d.geom_xmat], ) wp.launch( @@ -350,7 +356,7 @@ def _subtree_com_init( subtree_com_out: wp.array2d(dtype=wp.vec3), ): worldid, bodyid = wp.tid() - subtree_com_out[worldid, bodyid] = xipos_in[worldid, bodyid] * body_mass[worldid, bodyid] + subtree_com_out[worldid, bodyid] = xipos_in[worldid, bodyid] * body_mass[worldid % body_mass.shape[0], bodyid] @wp.kernel @@ -367,13 +373,14 @@ def _subtree_com_acc( worldid, nodeid = wp.tid() bodyid = body_tree_[nodeid] pid = body_parentid[bodyid] - wp.atomic_add(subtree_com_out, worldid, pid, subtree_com_in[worldid, bodyid]) + if bodyid != 0: + wp.atomic_add(subtree_com_out, worldid, pid, subtree_com_in[worldid, bodyid]) @wp.kernel def _subtree_div( # Model: - subtree_mass: wp.array2d(dtype=float), + body_subtreemass: wp.array2d(dtype=float), # Data in: subtree_com_in: wp.array2d(dtype=wp.vec3), # Data out: @@ -381,7 +388,7 @@ def _subtree_div( ): worldid, bodyid = wp.tid() com = subtree_com_in[worldid, bodyid] - mass = subtree_mass[worldid, bodyid] + mass = body_subtreemass[worldid % body_subtreemass.shape[0], bodyid] if mass != 0.0: subtree_com_out[worldid, bodyid] = com / mass @@ -401,8 +408,8 @@ def _cinert( ): worldid, bodyid = wp.tid() mat = ximat_in[worldid, bodyid] - inert = body_inertia[worldid, bodyid] - mass = body_mass[worldid, bodyid] + inert = body_inertia[worldid % body_inertia.shape[0], bodyid] + mass = body_mass[worldid % body_mass.shape[0], bodyid] dif = xipos_in[worldid, bodyid] - subtree_com_in[worldid, body_rootid[bodyid]] # express inertia in com-based frame (mju_inertCom) @@ -479,12 +486,11 @@ def _cdof( @event_scope def com_pos(m: Model, d: Data): - """ - Computes subtree center of mass positions. Transforms inertia and motion to global frame - centered at subtree CoM. + """Computes subtree center of mass positions. - Accumulates the mass-weighted positions up the kinematic tree, divides by total mass, and - computes composite inertias and motion degrees of freedom in the subtree CoM frame. + Transforms inertia and motion to global frame centered at subtree CoM. Accumulates the + mass-weighted positions up the kinematic tree, divides by total mass, and computes composite + inertias and motion degrees of freedom in the subtree CoM frame. """ wp.launch(_subtree_com_init, dim=(d.nworld, m.nbody), inputs=[m.body_mass, d.xipos], outputs=[d.subtree_com]) @@ -497,7 +503,7 @@ def com_pos(m: Model, d: Data): outputs=[d.subtree_com], ) - wp.launch(_subtree_div, dim=(d.nworld, m.nbody), inputs=[m.subtree_mass, d.subtree_com], outputs=[d.subtree_com]) + wp.launch(_subtree_div, dim=(d.nworld, m.nbody), inputs=[m.body_subtreemass, d.subtree_com], outputs=[d.subtree_com]) wp.launch( _cinert, dim=(d.nworld, m.nbody), @@ -532,26 +538,30 @@ def _cam_local_to_global( cam_xmat_out: wp.array2d(dtype=wp.mat33), ): worldid, camid = wp.tid() + cam_pos_id = worldid % cam_pos.shape[0] + cam_quat_id = worldid % cam_quat.shape[0] is_target_cam = (cam_mode[camid] == CamLightType.TARGETBODY) or (cam_mode[camid] == CamLightType.TARGETBODYCOM) invalid_target = is_target_cam and (cam_targetbodyid[camid] < 0) if invalid_target: bodyid = cam_bodyid[camid] xpos = xpos_in[worldid, bodyid] xquat = xquat_in[worldid, bodyid] - cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat) - cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[worldid, camid])) + cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[cam_pos_id, camid], xquat) + cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[cam_quat_id, camid])) elif cam_mode[camid] == CamLightType.TRACK: - cam_xmat_out[worldid, camid] = cam_mat0[worldid, camid] + cam_xmat_out[worldid, camid] = cam_mat0[worldid % cam_mat0.shape[0], camid] body_xpos = xpos_in[worldid, cam_bodyid[camid]] - cam_xpos_out[worldid, camid] = body_xpos + cam_pos0[worldid, camid] + cam_xpos_out[worldid, camid] = body_xpos + cam_pos0[worldid % cam_pos0.shape[0], camid] elif cam_mode[camid] == CamLightType.TRACKCOM: - cam_xmat_out[worldid, camid] = cam_mat0[worldid, camid] - cam_xpos_out[worldid, camid] = subtree_com_in[worldid, cam_bodyid[camid]] + cam_poscom0[worldid, camid] + cam_xmat_out[worldid, camid] = cam_mat0[worldid % cam_mat0.shape[0], camid] + cam_xpos_out[worldid, camid] = ( + subtree_com_in[worldid, cam_bodyid[camid]] + cam_poscom0[worldid % cam_poscom0.shape[0], camid] + ) elif cam_mode[camid] == CamLightType.TARGETBODY or cam_mode[camid] == CamLightType.TARGETBODYCOM: bodyid = cam_bodyid[camid] xpos = xpos_in[worldid, bodyid] xquat = xquat_in[worldid, bodyid] - cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat) + cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[cam_pos_id, camid], xquat) pos = xpos_in[worldid, cam_targetbodyid[camid]] if cam_mode[camid] == CamLightType.TARGETBODYCOM: pos = subtree_com_in[worldid, cam_targetbodyid[camid]] @@ -571,8 +581,8 @@ def _cam_local_to_global( bodyid = cam_bodyid[camid] xpos = xpos_in[worldid, bodyid] xquat = xquat_in[worldid, bodyid] - cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat) - cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[worldid, camid])) + cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[cam_pos_id, camid], xquat) + cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[cam_quat_id, camid])) @wp.kernel @@ -595,27 +605,31 @@ def _light_local_to_global( light_xdir_out: wp.array2d(dtype=wp.vec3), ): worldid, lightid = wp.tid() + light_pos_id = worldid % light_pos.shape[0] + light_dir_id = worldid % light_dir.shape[0] is_target_light = (light_mode[lightid] == CamLightType.TARGETBODY) or (light_mode[lightid] == CamLightType.TARGETBODYCOM) invalid_target = is_target_light and (light_targetbodyid[lightid] < 0) if invalid_target: bodyid = light_bodyid[lightid] xpos = xpos_in[worldid, bodyid] xquat = xquat_in[worldid, bodyid] - light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat) - light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[worldid, lightid], xquat) + light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[light_pos_id, lightid], xquat) + light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[light_dir_id, lightid], xquat) return elif light_mode[lightid] == CamLightType.TRACK: - light_xdir_out[worldid, lightid] = light_dir0[worldid, lightid] + light_xdir_out[worldid, lightid] = light_dir0[worldid % light_dir0.shape[0], lightid] body_xpos = xpos_in[worldid, light_bodyid[lightid]] - light_xpos_out[worldid, lightid] = body_xpos + light_pos0[worldid, lightid] + light_xpos_out[worldid, lightid] = body_xpos + light_pos0[worldid % light_pos0.shape[0], lightid] elif light_mode[lightid] == CamLightType.TRACKCOM: - light_xdir_out[worldid, lightid] = light_dir0[worldid, lightid] - light_xpos_out[worldid, lightid] = subtree_com_in[worldid, light_bodyid[lightid]] + light_poscom0[worldid, lightid] + light_xdir_out[worldid, lightid] = light_dir0[worldid % light_dir0.shape[0], lightid] + light_xpos_out[worldid, lightid] = ( + subtree_com_in[worldid, light_bodyid[lightid]] + light_poscom0[worldid % light_poscom0.shape[0], lightid] + ) elif light_mode[lightid] == CamLightType.TARGETBODY or light_mode[lightid] == CamLightType.TARGETBODYCOM: bodyid = light_bodyid[lightid] xpos = xpos_in[worldid, bodyid] xquat = xquat_in[worldid, bodyid] - light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat) + light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[light_pos_id, lightid], xquat) pos = xpos_in[worldid, light_targetbodyid[lightid]] if light_mode[lightid] == CamLightType.TARGETBODYCOM: pos = subtree_com_in[worldid, light_targetbodyid[lightid]] @@ -624,16 +638,15 @@ def _light_local_to_global( bodyid = light_bodyid[lightid] xpos = xpos_in[worldid, bodyid] xquat = xquat_in[worldid, bodyid] - light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat) - light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[worldid, lightid], xquat) + light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[light_pos_id, lightid], xquat) + light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[light_dir_id, lightid], xquat) light_xdir_out[worldid, lightid] = wp.normalize(light_xdir_out[worldid, lightid]) @event_scope def camlight(m: Model, d: Data): - """ - Computes camera and light positions and orientations. + """Computes camera and light positions and orientations. Updates the global positions and orientations for all cameras and lights in the model, including special handling for tracking and target modes. @@ -709,7 +722,7 @@ def _qM_sparse( qM_out: wp.array3d(dtype=float), ): worldid, dofid = wp.tid() - madr_ij = dof_Madr[dofid] + madr_ij = dof_Madr[dofid] # dof_Madr is not batched bodyid = dof_bodyid[dofid] # init M(i,i) with armature inertia @@ -739,8 +752,8 @@ def _qM_dense( ): worldid, dofid = wp.tid() bodyid = dof_bodyid[dofid] - # init M(i,i) with armature inertia - M = dof_armature[worldid, dofid] + # init M(i,i) with armature inertia. + M = dof_armature[worldid % dof_armature.shape[0], dofid] # precompute buf = crb_body_i * cdof_i buf = math.inert_vec(crb_in[worldid, bodyid], cdof_in[worldid, dofid]) @@ -760,8 +773,7 @@ def _qM_dense( @event_scope def crb(m: Model, d: Data): - """ - Computes composite rigid body inertias for each body and the joint-space inertia matrix. + """Computes composite rigid body inertias for each body and the joint-space inertia matrix. Accumulates composite rigid body inertias up the kinematic tree and computes the joint-space inertia matrix in either sparse or dense format, depending on model options. @@ -800,7 +812,7 @@ def _tendon_armature( ): worldid, tenid, dofid = wp.tid() - if opt_is_sparse: + if opt_is_sparse: # opt_is_sparse is not batched madr_ij = dof_Madr[dofid] armature = tendon_armature[worldid, tenid] @@ -962,7 +974,7 @@ def _cacc_world( cacc_out: wp.array2d(dtype=wp.spatial_vector), ): worldid = wp.tid() - cacc_out[worldid, 0] = wp.spatial_vector(wp.vec3(0.0), -gravity[worldid]) + cacc_out[worldid, 0] = wp.spatial_vector(wp.vec3(0.0), -gravity[worldid % gravity.shape[0]]) def _rne_cacc_world(m: Model, d: Data): @@ -1086,17 +1098,15 @@ def _qfrc_bias( @event_scope def rne(m: Model, d: Data, flg_acc: bool = False): - """ - Computes inverse dynamics using the recursive Newton-Euler algorithm. + """Computes inverse dynamics using the recursive Newton-Euler algorithm. - Computes the bias forces (qfrc_bias) and internal forces (cfrc_int) for the current state, + Computes the bias forces (`qfrc_bias`) and internal forces (`cfrc_int`) for the current state, including the effects of gravity and optionally joint accelerations. Args: - m (Model): The model containing kinematic and dynamic information. - d (Data): The data object containing the current state and output arrays. - flg_acc (bool, optional): If True, includes joint accelerations in the computation. - Defaults to False. + m: The model containing kinematic and dynamic information. + d: The data object containing the current state and output arrays. + flg_acc: If True, includes joint accelerations in the computation. """ _rne_cacc_world(m, d) _rne_cacc_forward(m, d, flg_acc=flg_acc) @@ -1137,13 +1147,13 @@ def _cfrc_ext_equality( eq_objtype: wp.array(dtype=int), eq_data: wp.array2d(dtype=vec11), # Data in: - ne_connect_in: wp.array(dtype=int), - ne_weld_in: wp.array(dtype=int), xpos_in: wp.array2d(dtype=wp.vec3), xmat_in: wp.array2d(dtype=wp.mat33), subtree_com_in: wp.array2d(dtype=wp.vec3), efc_id_in: wp.array2d(dtype=int), efc_force_in: wp.array2d(dtype=float), + ne_connect_in: wp.array(dtype=int), + ne_weld_in: wp.array(dtype=int), # Data out: cfrc_ext_out: wp.array2d(dtype=wp.spatial_vector), ): @@ -1242,8 +1252,6 @@ def _cfrc_ext_contact( body_rootid: wp.array(dtype=int), geom_bodyid: wp.array(dtype=int), # Data in: - njmax_in: int, - nacon_in: wp.array(dtype=int), subtree_com_in: wp.array2d(dtype=wp.vec3), contact_pos_in: wp.array(dtype=wp.vec3), contact_frame_in: wp.array(dtype=wp.mat33), @@ -1253,6 +1261,8 @@ def _cfrc_ext_contact( contact_efc_address_in: wp.array2d(dtype=int), contact_worldid_in: wp.array(dtype=int), efc_force_in: wp.array2d(dtype=float), + njmax_in: int, + nacon_in: wp.array(dtype=int), # Data out: cfrc_ext_out: wp.array2d(dtype=wp.spatial_vector), ): @@ -1273,13 +1283,13 @@ def _cfrc_ext_contact( # contact force in world frame force = support.contact_force_fn( opt_cone, - njmax_in, - nacon_in, contact_frame_in, contact_friction_in, contact_dim_in, contact_efc_address_in, efc_force_in, + njmax_in, + nacon_in, worldid, contactid, to_world_frame=True, @@ -1299,10 +1309,9 @@ def _cfrc_ext_contact( @event_scope def rne_postconstraint(m: Model, d: Data): - """ - Computes the recursive Newton-Euler algorithm after constraints are applied. + """Computes the recursive Newton-Euler algorithm after constraints are applied. - Computes cacc, cfrc_ext, and cfrc_int, including the effects of applied forces, equality + Computes `cacc`, `cfrc_ext`, and `cfrc_int`, including the effects of applied forces, equality constraints, and contacts. """ # cfrc_ext = perturb @@ -1324,13 +1333,13 @@ def rne_postconstraint(m: Model, d: Data): m.eq_obj2id, m.eq_objtype, m.eq_data, - d.ne_connect, - d.ne_weld, d.xpos, d.xmat, d.subtree_com, d.efc.id, d.efc.force, + d.ne_connect, + d.ne_weld, ], outputs=[d.cfrc_ext], ) @@ -1343,8 +1352,6 @@ def rne_postconstraint(m: Model, d: Data): m.opt.cone, m.body_rootid, m.geom_bodyid, - d.njmax, - d.nacon, d.subtree_com, d.contact.pos, d.contact.frame, @@ -1354,6 +1361,8 @@ def rne_postconstraint(m: Model, d: Data): d.contact.efc_address, d.contact.worldid, d.efc.force, + d.njmax, + d.nacon, ], outputs=[d.cfrc_ext], ) @@ -1383,16 +1392,16 @@ def _tendon_dot( tendon_adr: wp.array(dtype=int), tendon_num: wp.array(dtype=int), tendon_armature: wp.array2d(dtype=float), + wrap_type: wp.array(dtype=int), wrap_objid: wp.array(dtype=int), wrap_prm: wp.array(dtype=float), - wrap_type: wp.array(dtype=int), # Data in: site_xpos_in: wp.array2d(dtype=wp.vec3), subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), cvel_in: wp.array2d(dtype=wp.spatial_vector), cdof_dot_in: wp.array2d(dtype=wp.spatial_vector), - # Data out: + # Out: ten_Jdot_out: wp.array3d(dtype=float), ): worldid, tenid = wp.tid() @@ -1548,8 +1557,9 @@ def _tendon_bias_coef( tendon_armature: wp.array2d(dtype=float), # Data in: qvel_in: wp.array2d(dtype=float), + # In: ten_Jdot_in: wp.array3d(dtype=float), - # Data out: + # Out: ten_bias_coef_out: wp.array2d(dtype=float), ): worldid, tenid, dofid = wp.tid() @@ -1571,6 +1581,7 @@ def _tendon_bias_qfrc( tendon_armature: wp.array2d(dtype=float), # Data in: ten_J_in: wp.array3d(dtype=float), + # In: ten_bias_coef_in: wp.array2d(dtype=float), # Out: qfrc_out: wp.array2d(dtype=float), @@ -1590,8 +1601,15 @@ def _tendon_bias_qfrc( @event_scope def tendon_bias(m: Model, d: Data, qfrc: wp.array2d(dtype=float)): - """Add bias force due to tendon armature.""" - d.ten_Jdot.zero_() + """Add bias force due to tendon armature. + + Args: + m: The model containing kinematic and dynamic information. + d: The data object containing the current state and output arrays. + qfrc: Force. + """ + # time derivative of tendon Jacobian + ten_Jdot = wp.zeros((d.nworld, m.ntendon, m.nv), dtype=float) wp.launch( _tendon_dot, dim=(d.nworld, m.ntendon), @@ -1607,45 +1625,32 @@ def tendon_bias(m: Model, d: Data, qfrc: wp.array2d(dtype=float)): m.tendon_adr, m.tendon_num, m.tendon_armature, + m.wrap_type, m.wrap_objid, m.wrap_prm, - m.wrap_type, d.site_xpos, d.subtree_com, d.cdof, d.cvel, d.cdof_dot, ], - outputs=[ - d.ten_Jdot, - ], + outputs=[ten_Jdot], ) - d.ten_bias_coef.zero_() + # tendon bias force coefficients + ten_bias_coef = wp.zeros((d.nworld, m.ntendon), dtype=float) wp.launch( _tendon_bias_coef, dim=(d.nworld, m.ntendon, m.nv), - inputs=[ - m.tendon_armature, - d.qvel, - d.ten_Jdot, - ], - outputs=[ - d.ten_bias_coef, - ], + inputs=[m.tendon_armature, d.qvel, ten_Jdot], + outputs=[ten_bias_coef], ) wp.launch( _tendon_bias_qfrc, dim=(d.nworld, m.ntendon, m.nv), - inputs=[ - m.tendon_armature, - d.ten_J, - d.ten_bias_coef, - ], - outputs=[ - qfrc, - ], + inputs=[m.tendon_armature, d.ten_J, ten_bias_coef], + outputs=[qfrc], ) @@ -1726,8 +1731,7 @@ def _comvel_level( @event_scope def com_vel(m: Model, d: Data): - """ - Computes the spatial velocities (cvel) and the derivative cdof_dot for all bodies. + """Computes the spatial velocities (cvel) and the derivative `cdof_dot` for all bodies. Propagates velocities down the kinematic tree, updating the spatial velocity and derivative for each body. @@ -1759,14 +1763,14 @@ def _transmission( dof_parentid: wp.array(dtype=int), site_bodyid: wp.array(dtype=int), site_quat: wp.array2d(dtype=wp.quat), + tendon_adr: wp.array(dtype=int), + tendon_num: wp.array(dtype=int), + wrap_type: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), actuator_trntype: wp.array(dtype=int), actuator_trnid: wp.array(dtype=wp.vec2i), actuator_gear: wp.array2d(dtype=wp.spatial_vector), actuator_cranklength: wp.array(dtype=float), - tendon_adr: wp.array(dtype=int), - tendon_num: wp.array(dtype=int), - wrap_objid: wp.array(dtype=int), - wrap_type: wp.array(dtype=int), # Data in: qpos_in: wp.array2d(dtype=float), xquat_in: wp.array2d(dtype=wp.quat), @@ -1774,15 +1778,16 @@ def _transmission( site_xmat_in: wp.array2d(dtype=wp.mat33), subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), - ten_length_in: wp.array2d(dtype=float), ten_J_in: wp.array3d(dtype=float), + ten_length_in: wp.array2d(dtype=float), # Data out: actuator_length_out: wp.array2d(dtype=float), actuator_moment_out: wp.array3d(dtype=float), ): worldid, actid = wp.tid() trntype = actuator_trntype[actid] - gear = actuator_gear[worldid, actid] + actuator_gear_id = worldid % actuator_gear.shape[0] + gear = actuator_gear[actuator_gear_id, actid] if trntype == TrnType.JOINT or trntype == TrnType.JOINTINPARENT: qpos = qpos_in[worldid] jntid = actuator_trnid[actid][0] @@ -1931,6 +1936,7 @@ def _transmission( refid = trnid[1] gear = actuator_gear[worldid, actid] + site_quat_id = worldid % site_quat.shape[0] gear_translation = wp.spatial_top(gear) gear_rotational = wp.spatial_bottom(gear) @@ -2003,8 +2009,8 @@ def _transmission( if rotational_transmission: # get site and refsite quats from parent bodies (avoid converting matrix to quat) - quat = math.mul_quat(site_quat[worldid, siteid], xquat_in[worldid, bodyid]) - refquat = math.mul_quat(site_quat[worldid, refid], xquat_in[worldid, bodyrefid]) + quat = math.mul_quat(site_quat[site_quat_id, siteid], xquat_in[worldid, bodyid]) + refquat = math.mul_quat(site_quat[site_quat_id, refid], xquat_in[worldid, bodyrefid]) # convert difference to expmap (axis-angle) vec = math.quat_sub(quat, refquat) @@ -2077,7 +2083,6 @@ def _transmission_body_moment( actuator_trnid: wp.array(dtype=wp.vec2i), actuator_trntype_body_adr: wp.array(dtype=int), # Data in: - nacon_in: wp.array(dtype=int), subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), contact_dist_in: wp.array(dtype=float), @@ -2089,8 +2094,10 @@ def _transmission_body_moment( contact_efc_address_in: wp.array2d(dtype=int), contact_worldid_in: wp.array(dtype=int), efc_J_in: wp.array3d(dtype=float), + nacon_in: wp.array(dtype=int), # Data out: actuator_moment_out: wp.array3d(dtype=float), + # Out: actuator_trntype_body_ncon_out: wp.array2d(dtype=int), ): trnbodyid, conid, dofid = wp.tid() @@ -2181,7 +2188,7 @@ def _transmission_body_moment( def _transmission_body_moment_scale( # Model: actuator_trntype_body_adr: wp.array(dtype=int), - # Data in: + # In: actuator_trntype_body_ncon_in: wp.array2d(dtype=int), # Data out: actuator_moment_out: wp.array3d(dtype=float), @@ -2197,8 +2204,7 @@ def _transmission_body_moment_scale( @event_scope def transmission(m: Model, d: Data): - """ - Computes actuator/transmission lengths and moments. + """Computes actuator/transmission lengths and moments. Updates the actuator length and moments for all actuators in the model, including joint and tendon transmissions. @@ -2220,34 +2226,32 @@ def transmission(m: Model, d: Data): m.dof_parentid, m.site_bodyid, m.site_quat, + m.tendon_adr, + m.tendon_num, + m.wrap_type, + m.wrap_objid, m.actuator_trntype, m.actuator_trnid, m.actuator_gear, m.actuator_cranklength, - m.tendon_adr, - m.tendon_num, - m.wrap_objid, - m.wrap_type, d.qpos, d.xquat, d.site_xpos, d.site_xmat, d.subtree_com, d.cdof, - d.ten_length, d.ten_J, + d.ten_length, ], outputs=[d.actuator_length, d.actuator_moment], ) - if m.actuator_trntype_body_adr.size > 0: - # reset number of active contacts - d.actuator_trntype_body_ncon.zero_() - + if m.nacttrnbody: # compute moments + ncon = wp.zeros((d.nworld, m.nacttrnbody), dtype=int) wp.launch( _transmission_body_moment, - dim=(m.actuator_trntype_body_adr.size, d.naconmax, m.nv), + dim=(m.nacttrnbody, d.naconmax, m.nv), inputs=[ m.opt.cone, m.body_parentid, @@ -2256,7 +2260,6 @@ def transmission(m: Model, d: Data): m.geom_bodyid, m.actuator_trnid, m.actuator_trntype_body_adr, - d.nacon, d.subtree_com, d.cdof, d.contact.dist, @@ -2268,21 +2271,16 @@ def transmission(m: Model, d: Data): d.contact.efc_address, d.contact.worldid, d.efc.J, + d.nacon, ], - outputs=[ - d.actuator_moment, - d.actuator_trntype_body_ncon, - ], + outputs=[d.actuator_moment, ncon], ) # scale moments wp.launch( _transmission_body_moment_scale, - dim=(d.nworld, m.actuator_trntype_body_adr.size, m.nv), - inputs=[ - m.actuator_trntype_body_adr, - d.actuator_trntype_body_ncon, - ], + dim=(d.nworld, m.nacttrnbody, m.nv), + inputs=[m.actuator_trntype_body_adr, ncon], outputs=[d.actuator_moment], ) @@ -2334,8 +2332,7 @@ def _solve_LD_sparse( x: wp.array2d(dtype=float), y: wp.array2d(dtype=float), ): - """Computes sparse backsubstitution: x = inv(L'*D*L)*y""" - + """Computes sparse backsubstitution: x = inv(L'*D*L)*y.""" wp.copy(x, y) for qLD_updates in reversed(m.qLD_updates): wp.launch(_solve_LD_sparse_x_acc_up, dim=(d.nworld, qLD_updates.size), inputs=[L, qLD_updates], outputs=[x]) @@ -2372,7 +2369,7 @@ def _tile_cholesky_solve(tile: TileSet): def _solve_LD_dense(m: Model, d: Data, L: wp.array3d(dtype=float), x: wp.array2d(dtype=float), y: wp.array2d(dtype=float)): - """Computes dense backsubstitution: x = inv(L'*L)*y""" + """Computes dense backsubstitution: x = inv(L'*L)*y.""" for tile in m.qM_tiles: wp.launch_tiled( _tile_cholesky_solve(tile), @@ -2391,19 +2388,19 @@ def solve_LD( x: wp.array2d(dtype=float), y: wp.array2d(dtype=float), ): - """ - Computes backsubstitution to solve a linear system of the form x = inv(L'*D*L) * y, - where L and D are the factors from the Cholesky factorization of the inertia matrix. + """Computes backsubstitution to solve a linear system of the form x = inv(L'*D*L) * y. + + L and D are the factors from the Cholesky factorization of the inertia matrix. This function dispatches to either a sparse or dense solver depending on Model options. Args: - m (Model): The model containing factorization and sparsity information. - d (Data): The data object containing workspace and factorization results. - L (array3d): Lower-triangular factor from the factorization (sparse or dense). - D (array2d): Diagonal factor from the factorization (only used for sparse). - x (array2d): Output array for the solution. - y (array2d): Input right-hand side array. + m: The model containing factorization and sparsity information. + d: The data object containing workspace and factorization results. + L: Lower-triangular factor from the factorization (sparse or dense). + D: Diagonal factor from the factorization (only used for sparse). + x: Output array for the solution. + y: Input right-hand side array. """ if m.opt.is_sparse: _solve_LD_sparse(m, d, L, D, x, y) @@ -2413,14 +2410,13 @@ def solve_LD( @event_scope def solve_m(m: Model, d: Data, x: wp.array2d(dtype=float), y: wp.array2d(dtype=float)): - """ - Computes backsubstitution: x = qLD * y. + """Computes backsubstitution: x = qLD * y. Args: - m (Model): The model containing inertia and factorization information. - d (Data): The data object containing factorization results. - x (array2d): Output array for the solution. - y (array2d): Input right-hand side array. + m: The model containing inertia and factorization information. + d: The data object containing factorization results. + x: Output array for the solution. + y: Input right-hand side array. """ solve_LD(m, d, d.qLD, d.qLDiagInv, x, y) @@ -2473,21 +2469,21 @@ def _factor_solve_i_dense( def factor_solve_i(m, d, M, L, D, x, y): - """ - Factorizes and solves the linear system: x = inv(L'*D*L) * y or x = inv(L'*L) * y, - where M is an inertia-like matrix and L, D are its Cholesky-like factors. + """Factorizes and solves the linear system: x = inv(L'*D*L) * y or x = inv(L'*L) * y. + + M is an inertia-like matrix and L, D are its Cholesky-like factors. This function first factorizes the matrix M (sparse or dense depending on model options), then solves the system for x given right-hand side y. Args: - m (Model): The model containing factorization and sparsity information. - d (Data): The data object containing workspace and factorization results. - M (array3d): The inertia-like matrix to factorize. - L (array3d): Output lower-triangular factor from the factorization (sparse or dense). - D (array2d): Output diagonal factor from the factorization (only used for sparse). - x (array2d): Output array for the solution. - y (array2d): Input right-hand side array. + m: The model containing factorization and sparsity information. + d: The data object containing workspace and factorization results. + M: The inertia-like matrix to factorize. + L: Output lower-triangular factor from the factorization (sparse or dense). + D: Output diagonal factor from the factorization (only used for sparse). + x: Output array for the solution. + y: Input right-hand side array. """ if m.opt.is_sparse: _factor_i_sparse(m, d, M, L, D) @@ -2513,6 +2509,8 @@ def _subtree_vel_forward( subtree_bodyvel_out: wp.array2d(dtype=wp.spatial_vector), ): worldid, bodyid = wp.tid() + body_mass_id = worldid % body_mass.shape[0] + body_inertia_id = worldid % body_inertia.shape[0] cvel = cvel_in[worldid, bodyid] ang = wp.spatial_top(cvel) @@ -2524,11 +2522,11 @@ def _subtree_vel_forward( # update linear velocity lin -= wp.cross(xipos - subtree_com_root, ang) - subtree_linvel_out[worldid, bodyid] = body_mass[worldid, bodyid] * lin + subtree_linvel_out[worldid, bodyid] = body_mass[body_mass_id, bodyid] * lin dv = wp.transpose(ximat) @ ang - dv[0] *= body_inertia[worldid, bodyid][0] - dv[1] *= body_inertia[worldid, bodyid][1] - dv[2] *= body_inertia[worldid, bodyid][2] + dv[0] *= body_inertia[body_inertia_id, bodyid][0] + dv[1] *= body_inertia[body_inertia_id, bodyid][1] + dv[2] *= body_inertia[body_inertia_id, bodyid][2] subtree_angmom_out[worldid, bodyid] = ximat @ dv subtree_bodyvel_out[worldid, bodyid] = wp.spatial_vector(ang, lin) @@ -2550,7 +2548,7 @@ def _linear_momentum( if bodyid: pid = body_parentid[bodyid] wp.atomic_add(subtree_linvel_out[worldid], pid, subtree_linvel_in[worldid, bodyid]) - subtree_linvel_out[worldid, bodyid] /= wp.max(MJ_MINVAL, body_subtreemass[worldid, bodyid]) + subtree_linvel_out[worldid, bodyid] /= wp.max(MJ_MINVAL, body_subtreemass[worldid % body_subtreemass.shape[0], bodyid]) @wp.kernel @@ -2582,9 +2580,9 @@ def _angular_momentum( com_parent = subtree_com_in[worldid, pid] vel = subtree_bodyvel_in[worldid, bodyid] linvel = subtree_linvel_in[worldid, bodyid] - linvel_parent = subtree_linvel_in[worldid, pid] - mass = body_mass[worldid, bodyid] - subtreemass = body_subtreemass[worldid, bodyid] + linvel_parent = subtree_linvel_in[worldid, pid] # Data field + mass = body_mass[worldid % body_mass.shape[0], bodyid] + subtreemass = body_subtreemass[worldid % body_subtreemass.shape[0], bodyid] # momentum wrt body i dx = xipos - com @@ -2607,13 +2605,11 @@ def _angular_momentum( def subtree_vel(m: Model, d: Data): - """ - Computes subtree linear velocity and angular momentum. + """Computes subtree linear velocity and angular momentum. Computes the linear momentum and angular momentum for each subtree, accumulating contributions up the kinematic tree. """ - # bodywise quantities wp.launch( _subtree_vel_forward, @@ -2661,8 +2657,8 @@ def _joint_tendon( # Data in: qpos_in: wp.array2d(dtype=float), # Data out: - ten_length_out: wp.array2d(dtype=float), ten_J_out: wp.array3d(dtype=float), + ten_length_out: wp.array2d(dtype=float), ): worldid, wrapid = wp.tid() @@ -2698,8 +2694,8 @@ def _spatial_site_tendon( subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), # Data out: - ten_length_out: wp.array2d(dtype=float), ten_J_out: wp.array3d(dtype=float), + ten_length_out: wp.array2d(dtype=float), ): worldid, elementid = wp.tid() @@ -2745,9 +2741,9 @@ def _spatial_geom_tendon( geom_bodyid: wp.array(dtype=int), geom_size: wp.array2d(dtype=wp.vec3), site_bodyid: wp.array(dtype=int), + wrap_type: wp.array(dtype=int), wrap_objid: wp.array(dtype=int), wrap_prm: wp.array(dtype=float), - wrap_type: wp.array(dtype=int), tendon_geom_adr: wp.array(dtype=int), wrap_geom_adr: wp.array(dtype=int), wrap_pulley_scale: wp.array(dtype=float), @@ -2758,8 +2754,9 @@ def _spatial_geom_tendon( subtree_com_in: wp.array2d(dtype=wp.vec3), cdof_in: wp.array2d(dtype=wp.spatial_vector), # Data out: - ten_length_out: wp.array2d(dtype=float), ten_J_out: wp.array3d(dtype=float), + ten_length_out: wp.array2d(dtype=float), + # Out: wrap_geom_xpos_out: wp.array2d(dtype=wp.spatial_vector), ): worldid, elementid = wp.tid() @@ -2887,10 +2884,11 @@ def _spatial_tendon_wrap( ntendon: int, tendon_adr: wp.array(dtype=int), tendon_num: wp.array(dtype=int), - wrap_objid: wp.array(dtype=int), wrap_type: wp.array(dtype=int), + wrap_objid: wp.array(dtype=int), # Data in: site_xpos_in: wp.array2d(dtype=wp.vec3), + # In: wrap_geom_xpos_in: wp.array2d(dtype=wp.spatial_vector), # Data out: ten_wrapadr_out: wp.array2d(dtype=int), @@ -3040,8 +3038,7 @@ def _spatial_tendon_wrap( def tendon(m: Model, d: Data): - """ - Computes tendon lengths and moments. + """Computes tendon lengths and moments. Updates the tendon length and moment arrays for all tendons in the model, including joint, site, and geom tendons. @@ -3052,12 +3049,15 @@ def tendon(m: Model, d: Data): d.ten_length.zero_() d.ten_J.zero_() + # Cartesian 3D points fro geom wrap points + wrap_geom_xpos = wp.empty((d.nworld, m.nwrap), dtype=wp.spatial_vector) + # process joint tendons wp.launch( _joint_tendon, dim=(d.nworld, m.wrap_jnt_adr.size), inputs=[m.jnt_qposadr, m.jnt_dofadr, m.wrap_objid, m.wrap_prm, m.tendon_jnt_adr, m.wrap_jnt_adr, d.qpos], - outputs=[d.ten_length, d.ten_J], + outputs=[d.ten_J, d.ten_length], ) spatial_site = m.wrap_site_pair_adr.size > 0 @@ -3085,7 +3085,7 @@ def tendon(m: Model, d: Data): d.subtree_com, d.cdof, ], - outputs=[d.ten_length, d.ten_J], + outputs=[d.ten_J, d.ten_length], ) # process spatial geom tendons @@ -3100,9 +3100,9 @@ def tendon(m: Model, d: Data): m.geom_bodyid, m.geom_size, m.site_bodyid, + m.wrap_type, m.wrap_objid, m.wrap_prm, - m.wrap_type, m.tendon_geom_adr, m.wrap_geom_adr, m.wrap_pulley_scale, @@ -3112,13 +3112,13 @@ def tendon(m: Model, d: Data): d.subtree_com, d.cdof, ], - outputs=[d.ten_length, d.ten_J, d.wrap_geom_xpos], + outputs=[d.ten_J, d.ten_length, wrap_geom_xpos], ) if spatial_site or spatial_geom: wp.launch( _spatial_tendon_wrap, dim=(d.nworld,), - inputs=[m.ntendon, m.tendon_adr, m.tendon_num, m.wrap_objid, m.wrap_type, d.site_xpos, d.wrap_geom_xpos], + inputs=[m.ntendon, m.tendon_adr, m.tendon_num, m.wrap_type, m.wrap_objid, d.site_xpos, wrap_geom_xpos], outputs=[d.ten_wrapadr, d.ten_wrapnum, d.wrap_obj, d.wrap_xpos], ) 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 6f9d1364..360a5d8d 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py @@ -297,7 +297,6 @@ def linesearch_iterative( opt_ls_iterations: int, stat_meaninertia: float, # Data in: - njmax_in: int, ne_in: wp.array(dtype=int), nf_in: wp.array(dtype=int), nefc_in: wp.array(dtype=int), @@ -313,6 +312,7 @@ def linesearch_iterative( 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), ): @@ -321,7 +321,7 @@ def linesearch_iterative( if efc_done_in[worldid]: return - impratio = opt_impratio[worldid] + impratio = opt_impratio[worldid % opt_impratio.shape[0]] efc_type = efc_type_in[worldid] efc_id = efc_id_in[worldid] efc_D = efc_D_in[worldid] @@ -330,8 +330,8 @@ def linesearch_iterative( efc_jv = efc_jv_in[worldid] efc_quad = efc_quad_in[worldid] efc_quad_gauss = efc_quad_gauss_in[worldid] - tolerance = opt_tolerance[worldid] - ls_tolerance = opt_ls_tolerance[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]) @@ -466,7 +466,6 @@ def _linesearch_iterative(m: types.Model, d: types.Data): m.opt.ls_tolerance, m.opt.ls_iterations, m.stat.meaninertia, - d.njmax, d.ne, d.nf, d.nefc, @@ -482,6 +481,7 @@ def _linesearch_iterative(m: types.Model, d: types.Data): d.efc.quad, d.efc.quad_gauss, d.efc.done, + d.njmax, ], outputs=[d.efc.alpha], ) @@ -496,12 +496,10 @@ def _log_scale(min_value: float, max_value: float, num_values: int, i: int) -> f @wp.kernel def linesearch_parallel_fused( # Model: - nlsp: int, opt_impratio: wp.array(dtype=float), + opt_ls_iterations: int, opt_ls_parallel_min_step: float, # Data in: - njmax_in: int, - nacon_in: wp.array(dtype=int), ne_in: wp.array(dtype=int), nf_in: wp.array(dtype=int), nefc_in: wp.array(dtype=int), @@ -516,15 +514,17 @@ def linesearch_parallel_fused( efc_quad_in: wp.array2d(dtype=wp.vec3), efc_quad_gauss_in: wp.array(dtype=wp.vec3), efc_done_in: wp.array(dtype=bool), - # Data out: - efc_cost_candidate_out: wp.array2d(dtype=float), + njmax_in: int, + nacon_in: wp.array(dtype=int), + # Out: + cost_out: wp.array2d(dtype=float), ): worldid, alphaid = wp.tid() if efc_done_in[worldid]: return - alpha = _log_scale(opt_ls_parallel_min_step, 1.0, nlsp, alphaid) + alpha = _log_scale(opt_ls_parallel_min_step, 1.0, opt_ls_iterations, alphaid) out = _eval_cost(efc_quad_gauss_in[worldid], alpha) @@ -569,7 +569,7 @@ def linesearch_parallel_fused( continue friction = contact_friction_in[conid] - mu = friction[0] / wp.sqrt(opt_impratio[worldid]) + mu = friction[0] / wp.sqrt(opt_impratio[worldid % opt_impratio.shape[0]]) # unpack quad efcid1 = contact_efc_address_in[conid, 1] @@ -613,17 +613,18 @@ def linesearch_parallel_fused( if x < 0.0: out += _eval_cost(efc_quad_in[worldid, efcid], alpha) - efc_cost_candidate_out[worldid, alphaid] = out + cost_out[worldid, alphaid] = out @wp.kernel def linesearch_parallel_best_alpha( # Model: - nlsp: int, + opt_ls_iterations: int, opt_ls_parallel_min_step: float, # Data in: efc_done_in: wp.array(dtype=bool), - efc_cost_candidate_in: wp.array2d(dtype=float), + # In: + cost_in: wp.array2d(dtype=float), # Data out: efc_alpha_out: wp.array(dtype=float), ): @@ -632,29 +633,25 @@ def linesearch_parallel_best_alpha( if efc_done_in[worldid]: return - # TODO(team): investigate alternatives to wp.argmin - # TODO(thowell): how did this use to work? bestid = int(0) best_cost = float(wp.inf) - for i in range(nlsp): - cost = efc_cost_candidate_in[worldid, i] + 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, nlsp, bestid) + efc_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): +def _linesearch_parallel(m: types.Model, d: types.Data, cost: wp.array2d(dtype=float)): wp.launch( linesearch_parallel_fused, - dim=(d.nworld, m.nlsp), + dim=(d.nworld, m.opt.ls_iterations), inputs=[ - m.nlsp, m.opt.impratio, + m.opt.ls_iterations, m.opt.ls_parallel_min_step, - d.njmax, - d.nacon, d.ne, d.nf, d.nefc, @@ -669,14 +666,16 @@ def _linesearch_parallel(m: types.Model, d: types.Data): d.efc.quad, d.efc.quad_gauss, d.efc.done, + d.njmax, + d.nacon, ], - outputs=[d.efc.cost_candidate], + outputs=[cost], ) wp.launch( linesearch_parallel_best_alpha, dim=(d.nworld), - inputs=[m.nlsp, m.opt.ls_parallel_min_step, d.efc.done, d.efc.cost_candidate], + inputs=[m.opt.ls_iterations, m.opt.ls_parallel_min_step, d.efc.done, cost], outputs=[d.efc.alpha], ) @@ -771,7 +770,6 @@ def linesearch_prepare_quad( # Model: opt_impratio: wp.array(dtype=float), # Data in: - nacon_in: wp.array(dtype=int), nefc_in: wp.array(dtype=int), contact_friction_in: wp.array(dtype=types.vec5), contact_dim_in: wp.array(dtype=int), @@ -782,6 +780,7 @@ def linesearch_prepare_quad( 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), ): @@ -904,9 +903,9 @@ def linesearch_jaref( @event_scope -def _linesearch(m: types.Model, d: types.Data): +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, d.efc.done) + support.mul_m(m, d, d.efc.mv, d.efc.search, skip=d.efc.done) # jv = efc_J @ search # TODO(team): is there a better way of doing batched matmuls with dynamic array sizes? @@ -951,7 +950,6 @@ def _linesearch(m: types.Model, d: types.Data): dim=(d.nworld, d.njmax), inputs=[ m.opt.impratio, - d.nacon, d.nefc, d.contact.friction, d.contact.dim, @@ -962,12 +960,13 @@ def _linesearch(m: types.Model, d: types.Data): d.efc.Jaref, d.efc.jv, d.efc.done, + d.nacon, ], outputs=[d.efc.quad], ) if m.opt.ls_parallel: - _linesearch_parallel(m, d) + _linesearch_parallel(m, d, cost) else: _linesearch_iterative(m, d) @@ -1064,7 +1063,6 @@ def update_constraint_efc( # Model: opt_impratio: wp.array(dtype=float), # Data in: - nacon_in: wp.array(dtype=int), ne_in: wp.array(dtype=int), nf_in: wp.array(dtype=int), nefc_in: wp.array(dtype=int), @@ -1077,6 +1075,7 @@ def update_constraint_efc( 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), # Data out: efc_force_out: wp.array2d(dtype=float), efc_cost_out: wp.array(dtype=float), @@ -1187,11 +1186,11 @@ def update_constraint_efc( @wp.kernel def update_constraint_init_qfrc_constraint( # Data in: - njmax_in: int, 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, # Data out: qfrc_constraint_out: wp.array2d(dtype=float), ): @@ -1251,7 +1250,6 @@ def update_constraint_gauss_cost(nv: int, dofs_per_thread: int): def _update_constraint(m: types.Model, d: types.Data): """Update constraint arrays after each solve iteration.""" - wp.launch( update_constraint_init_cost, dim=(d.nworld), @@ -1264,7 +1262,6 @@ def _update_constraint(m: types.Model, d: types.Data): dim=(d.nworld, d.njmax), inputs=[ m.opt.impratio, - d.nacon, d.ne, d.nf, d.nefc, @@ -1277,6 +1274,7 @@ def _update_constraint(m: types.Model, d: types.Data): d.efc.frictionloss, d.efc.Jaref, d.efc.done, + d.nacon, ], outputs=[d.efc.force, d.efc.cost, d.efc.state], ) @@ -1285,7 +1283,7 @@ def _update_constraint(m: types.Model, d: types.Data): wp.launch( update_constraint_init_qfrc_constraint, dim=(d.nworld, m.nv), - inputs=[d.njmax, d.nefc, d.efc.J, d.efc.force, d.efc.done], + inputs=[d.nefc, d.efc.J, d.efc.force, d.efc.done, d.njmax], outputs=[d.qfrc_constraint], ) @@ -1364,100 +1362,150 @@ def update_gradient_set_h_qM_lower_sparse( efc_h_out[worldid, i, j] += qM_in[worldid, 0, elementid] -@wp.kernel -def update_gradient_JTDAJ_sparse( - # Data in: - njmax_in: int, - 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), - # Data out: - efc_h_out: wp.array3d(dtype=float), -): - worldid, elementid = wp.tid() - - if efc_done_in[worldid]: - return - - nefc = nefc_in[worldid] - - dofi = (int(sqrt(float(1 + 8 * elementid))) - 1) // 2 - dofj = elementid - (dofi * (dofi + 1)) // 2 - - # To optimize the loop, data for the next iteration is prefetched - # This allows to parallelize memory load and computation to hide memory latency - efc_state = efc_state_in[worldid, 0] - efc_D = efc_D_in[worldid, 0] - # TODO(team): sparse efc_J - efc_Ji = efc_J_in[worldid, 0, dofi] - efc_Jj = efc_J_in[worldid, 0, dofj] - sum_h = float(0.0) - for efcid in range(min(njmax_in, nefc) - 1): - if efc_state == types.ConstraintState.QUADRATIC and efc_D != 0.0: - sum_h += efc_Ji * efc_Jj * efc_D - - jj = efcid + 1 - efc_D = efc_D_in[worldid, jj] - efc_Ji = efc_J_in[worldid, jj, dofi] - efc_Jj = efc_J_in[worldid, jj, dofj] - efc_state = efc_state_in[worldid, jj] - - # Adding the contribution from the last constraint row - if efc_state == types.ConstraintState.QUADRATIC and efc_D != 0.0: - sum_h += efc_Ji * efc_Jj * efc_D - - efc_h_out[worldid, dofi, dofj] = sum_h +@wp.func +def state_check(D: float, state: int) -> float: + if state == types.ConstraintState.QUADRATIC.value: + return D + else: + return 0.0 -@wp.kernel -def update_gradient_JTDAJ_dense( - # Data in: - njmax_in: int, - nefc_in: wp.array(dtype=int), - qM_in: wp.array3d(dtype=float), - 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), - # Data out: - efc_h_out: wp.array3d(dtype=float), -): - worldid, elementid = wp.tid() +@wp.func +def active_check(tid: int, threshold: int) -> float: + if tid >= threshold: + return 0.0 + else: + return 1.0 - if efc_done_in[worldid]: - return - nefc = nefc_in[worldid] +@cache_kernel +def update_gradient_JTDAJ_sparse_tiled(tile_size: int, njmax: int): + TILE_SIZE = tile_size - dofi = (int(sqrt(float(1 + 8 * elementid))) - 1) // 2 - dofj = elementid - (dofi * (dofi + 1)) // 2 + @nested_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), + # Data out: + efc_h_out: wp.array3d(dtype=float), + ): + worldid, elementid = wp.tid() - # To optimize the loop, data for the next iteration is prefetched - # This allows to parallelize memory load and computation to hide memory latency - efc_state = efc_state_in[worldid, 0] - efc_D = efc_D_in[worldid, 0] - # TODO(team): sparse efc_J - efc_Ji = efc_J_in[worldid, 0, dofi] - efc_Jj = efc_J_in[worldid, 0, dofj] - sum_h = float(0.0) - for efcid in range(min(njmax_in, nefc) - 1): - if efc_state == types.ConstraintState.QUADRATIC and efc_D != 0.0: - sum_h += efc_Ji * efc_Jj * efc_D + if efc_done_in[worldid]: + return - jj = efcid + 1 - efc_D = efc_D_in[worldid, jj] - efc_Ji = efc_J_in[worldid, jj, dofi] - efc_Jj = efc_J_in[worldid, jj, dofj] - efc_state = efc_state_in[worldid, jj] + nefc = nefc_in[worldid] - # Adding the contribution from the last constraint row - if efc_state == types.ConstraintState.QUADRATIC and efc_D != 0.0: - sum_h += efc_Ji * efc_Jj * efc_D + # get lower diagonal index + i = (int(sqrt(float(1 + 8 * elementid))) - 1) // 2 + j = elementid - (i * (i + 1)) // 2 - qM = qM_in[worldid, dofi, dofj] - efc_h_out[worldid, dofi, dofj] = qM + sum_h + offset_i = i * TILE_SIZE + offset_j = j * TILE_SIZE + + sum_val = wp.tile_zeros(shape=(TILE_SIZE, TILE_SIZE), dtype=wp.float32) + + # Each tile processes looping over all constraints, producing 1 output tile + for k in range(0, njmax, TILE_SIZE): + if k >= nefc: + break + + # 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, TILE_SIZE), offset=(k, offset_i), bounds_check=False) + + if offset_i != offset_j: + J_kj = wp.tile_load(efc_J_in[worldid], shape=(TILE_SIZE, TILE_SIZE), offset=(k, offset_j), bounds_check=False) + else: + wp.tile_assign(J_kj, J_ki, (0, 0)) + + D_k = wp.tile_load(efc_D_in[worldid], shape=TILE_SIZE, offset=k, bounds_check=False) + state = wp.tile_load(efc_state_in[worldid], shape=TILE_SIZE, offset=k, bounds_check=False) + + D_k = wp.tile_map(state_check, D_k, state) + + # force unused elements to be zero + tid_tile = wp.tile_arange(TILE_SIZE, dtype=int) + threshold_tile = wp.tile_ones(shape=TILE_SIZE, dtype=int) * (nefc - k) + + 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=(TILE_SIZE, TILE_SIZE))) + + sum_val += wp.tile_matmul(J_ki, J_kj) + + # AD: setting bounds_check to True explicitly here because for some reason it was + # slower to disable it. + wp.tile_store(efc_h_out[worldid], sum_val, offset=(offset_i, offset_j), bounds_check=True) + + return kernel + + +@cache_kernel +def update_gradient_JTDAJ_dense_tiled(nv: int, tile_size: int, njmax: int): + if njmax < tile_size: + tile_size = njmax + + TILE_SIZE_K = tile_size + + @nested_kernel(module="unique", enable_backward=False) + def kernel( + # Data in: + nefc_in: wp.array(dtype=int), + qM_in: wp.array3d(dtype=float), + 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), + # Data out: + efc_h_out: wp.array3d(dtype=float), + ): + worldid = wp.tid() + + if efc_done_in[worldid]: + return + + nefc = nefc_in[worldid] + + sum_val = wp.tile_load(qM_in[worldid], shape=(nv, nv), bounds_check=False) + + # Each tile processes one output tile by looping over all constraints + for k in range(0, njmax, TILE_SIZE_K): + if k >= nefc: + break + + # 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), offset=(k, 0), bounds_check=False) + J_kj = J_ki + + # state check + D_k = wp.tile_load(efc_D_in[worldid], shape=TILE_SIZE_K, offset=k, bounds_check=False) + state = wp.tile_load(efc_state_in[worldid], shape=TILE_SIZE_K, offset=k, bounds_check=False) + + D_k = wp.tile_map(state_check, D_k, state) + + # force unused elements to be zero + tid_tile = wp.tile_arange(TILE_SIZE_K, dtype=int) + threshold_tile = wp.tile_ones(shape=TILE_SIZE_K, dtype=int) * (nefc - k) + + 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, TILE_SIZE_K))) + + sum_val += wp.tile_matmul(J_ki, J_kj) + + wp.tile_store(efc_h_out[worldid], sum_val, bounds_check=False) + + return kernel # TODO(thowell): combine with JTDAJ ? @@ -1468,8 +1516,6 @@ def update_gradient_JTCJ( dof_tri_row: wp.array(dtype=int), dof_tri_col: wp.array(dtype=int), # Data in: - naconmax_in: int, - nacon_in: wp.array(dtype=int), contact_dist_in: wp.array(dtype=float), contact_includemargin_in: wp.array(dtype=float), contact_friction_in: wp.array(dtype=types.vec5), @@ -1481,6 +1527,8 @@ def update_gradient_JTCJ( 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: nblocks_perblock: int, dim_block: int, @@ -1670,13 +1718,13 @@ def _update_gradient(m: types.Model, d: types.Data): smooth.solve_m(m, d, d.efc.Mgrad, d.efc.grad) elif m.opt.solver == types.SolverType.NEWTON: # h = qM + (efc_J.T * efc_D * active) @ efc_J - lower_triangle_dim = int(m.nv * (m.nv + 1) / 2) if m.opt.is_sparse: - wp.launch( - update_gradient_JTDAJ_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( + update_gradient_JTDAJ_sparse_tiled(types.TILE_SIZE_JTDAJ_SPARSE, d.njmax), dim=(d.nworld, lower_triangle_dim), inputs=[ - d.njmax, d.nefc, d.efc.J, d.efc.D, @@ -1684,7 +1732,9 @@ def _update_gradient(m: types.Model, d: types.Data): d.efc.done, ], outputs=[d.efc.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), @@ -1692,11 +1742,10 @@ def _update_gradient(m: types.Model, d: types.Data): outputs=[d.efc.h], ) else: - wp.launch( - update_gradient_JTDAJ_dense, - dim=(d.nworld, lower_triangle_dim), + wp.launch_tiled( + update_gradient_JTDAJ_dense_tiled(m.nv, types.TILE_SIZE_JTDAJ_DENSE, d.njmax), + dim=d.nworld, inputs=[ - d.njmax, d.nefc, d.qM, d.efc.J, @@ -1705,6 +1754,7 @@ def _update_gradient(m: types.Model, d: types.Data): d.efc.done, ], outputs=[d.efc.h], + block_dim=m.block_dim.update_gradient_JTDAJ_dense, ) if m.opt.cone == types.ConeType.ELLIPTIC: @@ -1738,8 +1788,6 @@ def _update_gradient(m: types.Model, d: types.Data): m.opt.impratio, m.dof_tri_row, m.dof_tri_col, - d.naconmax, - d.nacon, d.contact.dist, d.contact.includemargin, d.contact.friction, @@ -1751,6 +1799,8 @@ def _update_gradient(m: types.Model, d: types.Data): d.efc.Jaref, d.efc.state, d.efc.done, + d.naconmax, + d.nacon, nblocks_perblock, dim_block, ], @@ -1888,8 +1938,8 @@ def solve_done( efc_done_in: wp.array(dtype=bool), # Data out: solver_niter_out: wp.array(dtype=int), - nsolving_out: wp.array(dtype=int), efc_done_out: wp.array(dtype=bool), + nsolving_out: wp.array(dtype=int), ): worldid = wp.tid() @@ -1897,7 +1947,7 @@ def solve_done( return solver_niter_out[worldid] += 1 - tolerance = opt_tolerance[worldid] + tolerance = opt_tolerance[worldid % opt_tolerance.shape[0]] improvement = _rescale(nv, stat_meaninertia, efc_prev_cost_in[worldid] - efc_cost_in[worldid]) gradient = _rescale(nv, stat_meaninertia, wp.math.sqrt(efc_grad_dot_in[worldid])) @@ -1913,8 +1963,9 @@ def solve_done( def _solver_iteration( m: types.Model, d: types.Data, + step_size_cost: wp.array2d(dtype=float), ): - _linesearch(m, d) + _linesearch(m, d, step_size_cost) if m.opt.solver == types.SolverType.CG: wp.launch( @@ -1958,7 +2009,7 @@ def _solver_iteration( d.efc.prev_cost, d.efc.done, ], - outputs=[d.solver_niter, d.nsolving, d.efc.done], + outputs=[d.solver_niter, d.efc.done, d.nsolving], ) @@ -1979,7 +2030,7 @@ def create_context(m: types.Model, d: types.Data, grad: bool = True): ) # Ma = qM @ qacc - support.mul_m(m, d, d.efc.Ma, d.qacc, d.efc.done) + support.mul_m(m, d, d.efc.Ma, d.qacc, skip=d.efc.done) _update_constraint(m, d) @@ -1989,7 +2040,7 @@ def create_context(m: types.Model, d: types.Data, grad: bool = True): @event_scope def solve(m: types.Model, d: types.Data): - if d.njmax == 0: + if d.njmax == 0 or m.nv == 0: wp.copy(d.qacc, d.qacc_smooth) d.solver_niter.fill_(0) else: @@ -2014,6 +2065,8 @@ def _solve(m: types.Model, d: types.Data): outputs=[d.efc.search, d.efc.search_dot], ) + step_size_cost = wp.empty((d.nworld, m.opt.ls_iterations if m.opt.ls_parallel else 0), dtype=float) + 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. @@ -2028,10 +2081,11 @@ def _solve(m: types.Model, d: types.Data): while_body=_solver_iteration, m=m, d=d, + step_size_cost=step_size_cost, ) 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) + _solver_iteration(m, d, step_size_cost) 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 d21f0df0..a2e3ca8a 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/support.py @@ -13,19 +13,18 @@ # limitations under the License. # ============================================================================== -from typing import Tuple +from typing import Optional, Tuple import warp as wp -from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom from mujoco.mjx.third_party.mujoco_warp._src.math import motion_cross 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 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.types import vec6 from mujoco.mjx.third_party.mujoco_warp._src.warp_util import cache_kernel from mujoco.mjx.third_party.mujoco_warp._src.warp_util import event_scope from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_kernel @@ -33,63 +32,73 @@ from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_k wp.set_module_options({"enable_backward": False}) -@wp.kernel -def mul_m_sparse_diag( - # Model: - dof_Madr: wp.array(dtype=int), - # Data in: - qM_in: wp.array3d(dtype=float), - # In: - vec: wp.array2d(dtype=float), - skip: wp.array(dtype=bool), - # Out: - res: wp.array2d(dtype=float), -): - """Diagonal update for sparse matmul.""" - worldid, dofid = wp.tid() +@cache_kernel +def mul_m_sparse_diag(check_skip: bool): + @nested_kernel(module="unique", enable_backward=False) + def _mul_m_sparse_diag( + # Model: + dof_Madr: wp.array(dtype=int), + # Data in: + qM_in: wp.array3d(dtype=float), + # In: + vec: wp.array2d(dtype=float), + skip: wp.array(dtype=bool), + # Out: + res: wp.array2d(dtype=float), + ): + """Diagonal update for sparse matmul.""" + worldid, dofid = wp.tid() - if skip[worldid]: - return + if wp.static(check_skip): + if skip[worldid]: + return - res[worldid, dofid] = qM_in[worldid, 0, dof_Madr[dofid]] * vec[worldid, dofid] + res[worldid, dofid] = qM_in[worldid, 0, dof_Madr[dofid]] * vec[worldid, dofid] - -@wp.kernel -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), - # Data in: - qM_in: wp.array3d(dtype=float), - # In: - vec: wp.array2d(dtype=float), - skip: wp.array(dtype=bool), - # Out: - res: wp.array2d(dtype=float), -): - """Off-diagonal update for sparse matmul.""" - worldid, elementid = wp.tid() - - 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_diag @cache_kernel -def mul_m_dense(tile: TileSet): - """Returns a matmul kernel for some tile size""" +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), + # Data in: + qM_in: wp.array3d(dtype=float), + # In: + vec: wp.array2d(dtype=float), + skip: wp.array(dtype=bool), + # Out: + res: wp.array2d(dtype=float), + ): + """Off-diagonal update for sparse matmul.""" + worldid, elementid = 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 kernel( + def _mul_m_dense( # Data In: qM_in: wp.array3d(dtype=float), # In: @@ -102,8 +111,9 @@ def mul_m_dense(tile: TileSet): worldid, nodeid = wp.tid() TILE_SIZE = wp.static(tile.size) - if skip[worldid]: - return + 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)) @@ -111,7 +121,7 @@ def mul_m_dense(tile: TileSet): res_tile = wp.tile_matmul(qM_tile, vec_tile) wp.tile_store(res[worldid], res_tile, offset=(dofid, 0)) - return kernel + return _mul_m_dense @event_scope @@ -120,33 +130,35 @@ def mul_m( d: Data, res: wp.array2d(dtype=float), vec: wp.array2d(dtype=float), - skip: wp.array(dtype=bool), - M: wp.array3d(dtype=float) = None, + skip: Optional[wp.array] = None, + M: Optional[wp.array] = None, ): - """Multiply vectors by inertia matrix. + """Multiply vectors by inertia matrix; optionally skip per world. Args: - m (Model): The model containing kinematic and dynamic information (device). - d (Data): The data object containing the current state and output arrays (device). - res (wp.array2d(dtype=float)): Result: qM @ vec. - vec (wp.array2d(dtype=float)): Input vector to multiply by qM. - skip (wp.array(dtype=flooat)): Skip output. - M (wp.array3d(dtype=float), optional): Input matrix: M @ vec. + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + res: Result: qM @ vec. + vec: Input vector to multiply by qM. + skip: Per-world bitmask to skip computing output. + M: Input matrix: M @ vec. """ + check_skip = skip is not None + skip = skip or wp.empty(0, dtype=bool) if M is None: M = d.qM if m.opt.is_sparse: wp.launch( - mul_m_sparse_diag, + mul_m_sparse_diag(check_skip), dim=(d.nworld, m.nv), inputs=[m.dof_Madr, M, vec, skip], outputs=[res], ) wp.launch( - mul_m_sparse_ij, + 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], outputs=[res], @@ -155,7 +167,7 @@ def mul_m( else: for tile in m.qM_tiles: wp.launch_tiled( - mul_m_dense(tile), + mul_m_dense(tile, check_skip), dim=(d.nworld, tile.adr.size), inputs=[ M, @@ -225,13 +237,12 @@ def apply_ft(m: Model, d: Data, ft: wp.array2d(dtype=wp.spatial_vector), qfrc: w @event_scope def xfrc_accumulate(m: Model, d: Data, qfrc: wp.array2d(dtype=float)): - """ - Map applied forces at each body via Jacobians to dof space and accumulate. + """Map applied forces at each body via Jacobians to dof space and accumulate. Args: - m (Model): The model containing kinematic and dynamic information (device). - d (Data): The data object containing the current state and output arrays (device). - qfrc (wp.array2d(dtype=float)): Total applied force mapped to dof space. + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + qfrc: Total applied force mapped to dof space. """ apply_ft(m, d, d.xfrc_applied, qfrc, True) @@ -295,13 +306,13 @@ def contact_force_fn( # Model: opt_cone: int, # Data in: - njmax_in: int, - nacon_in: wp.array(dtype=int), contact_frame_in: wp.array(dtype=wp.mat33), contact_friction_in: wp.array(dtype=vec5), contact_dim_in: wp.array(dtype=int), contact_efc_address_in: wp.array2d(dtype=int), efc_force_in: wp.array2d(dtype=float), + njmax_in: int, + nacon_in: wp.array(dtype=int), # In: worldid: int, contact_id: int, @@ -340,14 +351,14 @@ def contact_force_kernel( # Model: opt_cone: int, # Data in: - njmax_in: int, - nacon_in: wp.array(dtype=int), contact_frame_in: wp.array(dtype=wp.mat33), contact_friction_in: wp.array(dtype=vec5), contact_dim_in: wp.array(dtype=int), contact_efc_address_in: wp.array2d(dtype=int), contact_worldid_in: wp.array(dtype=int), efc_force_in: wp.array2d(dtype=float), + njmax_in: int, + nacon_in: wp.array(dtype=int), # In: contact_ids: wp.array(dtype=int), to_world_frame: bool, @@ -365,13 +376,13 @@ def contact_force_kernel( out[tid] = contact_force_fn( opt_cone, - njmax_in, - nacon_in, contact_frame_in, contact_friction_in, contact_dim_in, contact_efc_address_in, efc_force_in, + njmax_in, + nacon_in, worldid, contactid, to_world_frame, @@ -379,35 +390,30 @@ def contact_force_kernel( def contact_force( - m: Model, - d: Data, - contact_ids: wp.array(dtype=int), - to_world_frame: bool, - force: wp.array(dtype=wp.spatial_vector), + m: Model, d: Data, contact_ids: wp.array(dtype=int), to_world_frame: bool, force: wp.array(dtype=wp.spatial_vector) ): - """ - Compute forces for contacts in Data. + """Compute forces for contacts in Data. Args: - m (Model): The model containing kinematic and dynamic information (device). - d (Data): The data object containing the current state and output arrays (device). - contact_ids (wp.array(dtype=int)): IDs for each contact. - to_world_frame (bool): If True, map force from contact to world frame. - force (wp.array(dtype=wp.spatial_vector)): Contact forces. + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output arrays (device). + contact_ids: IDs for each contact. + to_world_frame: If True, map force from contact to world frame. + force: Contact forces. """ wp.launch( contact_force_kernel, - dim=(contact_ids.size,), + dim=contact_ids.size, inputs=[ m.opt.cone, - d.njmax, - d.nacon, d.contact.frame, d.contact.friction, d.contact.dim, d.contact.efc_address, d.contact.worldid, d.efc.force, + d.njmax, + d.nacon, contact_ids, to_world_frame, ], @@ -532,3 +538,288 @@ def jac_dot( jacr = cdof_dot_ang return jacp, jacr + + +def get_state(m: Model, d: Data, state: wp.array2d(dtype=float), sig: int, active: Optional[wp.array] = None): + """Copy concatenated state components specified by sig from Data into state. + + The bits of the integer sig correspond to element fields of State. + + Args: + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output information (device). + state: Concatenation of state components. + sig: Bitflag specifying state components. + active: Per-world bitmask for getting state. + """ + if sig >= (1 << State.NSTATE): + raise ValueError(f"invalid state signature {sig} >= 2^mjNSTATE") + + @nested_kernel(module="unique", enable_backward=False) + def _get_state( + # Model: + nq: int, + nv: int, + nu: int, + na: int, + nbody: int, + neq: int, + nmocap: int, + # Data in: + time_in: wp.array(dtype=float), + qpos_in: wp.array2d(dtype=float), + qvel_in: wp.array2d(dtype=float), + act_in: wp.array2d(dtype=float), + qacc_warmstart_in: wp.array2d(dtype=float), + ctrl_in: wp.array2d(dtype=float), + qfrc_applied_in: wp.array2d(dtype=float), + xfrc_applied_in: wp.array2d(dtype=wp.spatial_vector), + eq_active_in: wp.array2d(dtype=bool), + mocap_pos_in: wp.array2d(dtype=wp.vec3), + mocap_quat_in: wp.array2d(dtype=wp.quat), + # In: + sig_in: int, + active_in: wp.array(dtype=bool), + # Out: + state_out: wp.array2d(dtype=float), + ): + worldid = wp.tid() + + if wp.static(active is not None): + if not active_in[worldid]: + return + + adr = int(0) + for i in range(State.NSTATE.value): + element = 1 << i + if element & sig_in: + if element == State.TIME: + state_out[worldid, adr] = time_in[worldid] + adr += 1 + elif element == State.QPOS: + for j in range(nq): + state_out[worldid, adr + j] = qpos_in[worldid, j] + adr += nq + elif element == State.QVEL: + for j in range(nv): + state_out[worldid, adr + j] = qvel_in[worldid, j] + adr += nv + elif element == State.ACT: + for j in range(na): + state_out[worldid, adr + j] = act_in[worldid, j] + adr += na + elif element == State.WARMSTART: + for j in range(nv): + state_out[worldid, adr + j] = qacc_warmstart_in[worldid, j] + adr += nv + elif element == State.CTRL: + for j in range(nu): + state_out[worldid, adr + j] = ctrl_in[worldid, j] + adr += nu + elif element == State.QFRC_APPLIED: + for j in range(nv): + state_out[worldid, adr + j] = qfrc_applied_in[worldid, j] + adr += nv + elif element == State.XFRC_APPLIED: + for j in range(nbody): + xfrc = xfrc_applied_in[worldid, j] + state_out[worldid, adr + 0] = xfrc[0] + state_out[worldid, adr + 1] = xfrc[1] + state_out[worldid, adr + 2] = xfrc[2] + state_out[worldid, adr + 3] = xfrc[3] + state_out[worldid, adr + 4] = xfrc[4] + state_out[worldid, adr + 5] = xfrc[5] + adr += 6 + elif element == State.EQ_ACTIVE: + for j in range(neq): + state_out[worldid, adr + j] = float(eq_active_in[worldid, j]) + adr += j + elif element == State.MOCAP_POS: + for j in range(nmocap): + pos = mocap_pos_in[worldid, j] + state_out[worldid, adr + 0] = pos[0] + state_out[worldid, adr + 1] = pos[1] + state_out[worldid, adr + 2] = pos[2] + adr += 3 + elif element == State.MOCAP_QUAT: + for j in range(nmocap): + quat = mocap_quat_in[worldid, j] + state_out[worldid, adr + 0] = quat[0] + state_out[worldid, adr + 1] = quat[1] + state_out[worldid, adr + 2] = quat[2] + state_out[worldid, adr + 3] = quat[3] + adr += 4 + + wp.launch( + _get_state, + dim=d.nworld, + inputs=[ + m.nq, + m.nv, + m.nu, + m.na, + m.nbody, + m.neq, + m.nmocap, + d.time, + d.qpos, + d.qvel, + d.act, + d.qacc_warmstart, + d.ctrl, + d.qfrc_applied, + d.xfrc_applied, + d.eq_active, + d.mocap_pos, + d.mocap_quat, + int(sig), + active or wp.ones(d.nworld, dtype=bool), + ], + outputs=[state], + ) + + +def set_state(m: Model, d: Data, state: wp.array2d(dtype=float), sig: int, active: Optional[wp.array] = None): + """Copy concatenated state components specified by sig from state into Data. + + The bits of the integer sig correspond to element fields of State. + + Args: + m: The model containing kinematic and dynamic information (device). + d: The data object containing the current state and output information (device). + state: Concatenation of state components. + sig: Bitflag specifying state components. + active: Per-world bitmask for setting state. + """ + if sig >= (1 << State.NSTATE): + raise ValueError(f"invalid state signature {sig} >= 2^mjNSTATE") + + @nested_kernel(module="unique", enable_backward=False) + def _set_state( + # Model: + nq: int, + nv: int, + nu: int, + na: int, + nbody: int, + neq: int, + nmocap: int, + # In: + sig_in: int, + active_in: wp.array(dtype=bool), + state_in: wp.array2d(dtype=float), + # Data out: + time_out: wp.array(dtype=float), + qpos_out: wp.array2d(dtype=float), + qvel_out: wp.array2d(dtype=float), + act_out: wp.array2d(dtype=float), + qacc_warmstart_out: wp.array2d(dtype=float), + ctrl_out: wp.array2d(dtype=float), + qfrc_applied_out: wp.array2d(dtype=float), + xfrc_applied_out: wp.array2d(dtype=wp.spatial_vector), + eq_active_out: wp.array2d(dtype=bool), + mocap_pos_out: wp.array2d(dtype=wp.vec3), + mocap_quat_out: wp.array2d(dtype=wp.quat), + ): + worldid = wp.tid() + + if wp.static(active is not None): + if not active_in[worldid]: + return + + adr = int(0) + for i in range(State.NSTATE.value): + element = 1 << i + if element & sig_in: + if element == State.TIME: + time_out[worldid] = state_in[worldid, adr] + adr += 1 + elif element == State.QPOS: + for j in range(nq): + qpos_out[worldid, j] = state_in[worldid, adr + j] + adr += nq + elif element == State.QVEL: + for j in range(nv): + qvel_out[worldid, j] = state_in[worldid, adr + j] + adr += nv + elif element == State.ACT: + for j in range(na): + act_out[worldid, j] = state_in[worldid, adr + j] + adr += na + elif element == State.WARMSTART: + for j in range(nv): + qacc_warmstart_out[worldid, j] = state_in[worldid, adr + j] + adr += nv + elif element == State.CTRL: + for j in range(nu): + ctrl_out[worldid, j] = state_in[worldid, adr + j] + adr += nu + elif element == State.QFRC_APPLIED: + for j in range(nv): + qfrc_applied_out[worldid, j] = state_in[worldid, adr + j] + adr += nv + elif element == State.XFRC_APPLIED: + for j in range(nbody): + xfrc = wp.spatial_vector( + state_in[worldid, adr + 0], + state_in[worldid, adr + 1], + state_in[worldid, adr + 2], + state_in[worldid, adr + 3], + state_in[worldid, adr + 4], + state_in[worldid, adr + 5], + ) + xfrc_applied_out[worldid, j] = xfrc + adr += 6 + elif element == State.EQ_ACTIVE: + for j in range(neq): + eq_active_out[worldid, j] = bool(state_in[worldid, adr + j]) + adr += j + elif element == State.MOCAP_POS: + for j in range(nmocap): + pos = wp.vec3( + state_in[worldid, adr + 1], + state_in[worldid, adr + 0], + state_in[worldid, adr + 2], + ) + mocap_pos_out[worldid, j] = pos + adr += 3 + elif element == State.MOCAP_QUAT: + for j in range(nmocap): + quat = wp.quat( + state_in[worldid, adr + 0], + state_in[worldid, adr + 1], + state_in[worldid, adr + 2], + state_in[worldid, adr + 3], + ) + mocap_quat_out[worldid, j] = quat + adr += 4 + + wp.launch( + _set_state, + dim=d.nworld, + inputs=[ + m.nq, + m.nv, + m.nu, + m.na, + m.nbody, + m.neq, + m.nmocap, + int(sig), + active or wp.ones(d.nworld, dtype=bool), + state, + ], + outputs=[ + d.time, + d.qpos, + d.qvel, + d.act, + d.qacc_warmstart, + d.ctrl, + d.qfrc_applied, + d.xfrc_applied, + d.eq_active, + d.mocap_pos, + d.mocap_quat, + ], + ) 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 3e356b95..01b93b58 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py @@ -29,38 +29,39 @@ MJ_MAX_EPAHORIZON = 12 # maximum average number of trianglarfaces EPA can insert at each iteration MJ_MAX_EPAFACES = 5 +TILE_SIZE_JTDAJ_SPARSE = 16 +TILE_SIZE_JTDAJ_DENSE = 16 + # TODO(team): add check that all wp.launch_tiled 'block_dim' settings are configurable @dataclasses.dataclass class BlockDim: - """ - Block dimension 'block_dim' settings for wp.launch_tiled. + """Block dimension 'block_dim' settings for wp.launch_tiled. TODO(team): experimental and may be removed """ # collision_driver segmented_sort: int = 128 - # derivative - qderiv_actuator_passive_actuation: int = 64 - qderiv_actuator_passive_no_actuation: int = 256 # forward - euler_dense: int = 256 + euler_dense: int = 32 actuator_velocity: int = 32 - tendon_velocity: int = 256 + tendon_velocity: int = 32 # ray ray: int = 64 # sensor contact_sort: int = 64 - energy_vel_kinetic: int = 256 + energy_vel_kinetic: int = 32 # smooth - cholesky_factorize: int = 256 - cholesky_solve: int = 256 - cholesky_factorize_solve: int = 256 + cholesky_factorize: int = 32 + cholesky_solve: int = 32 + cholesky_factorize_solve: int = 32 # solver update_gradient_cholesky: int = 64 + update_gradient_JTDAJ_sparse: int = 64 + update_gradient_JTDAJ_dense: int = 96 # support - mul_m_dense: int = 256 + mul_m_dense: int = 32 class BroadphaseType(enum.IntEnum): @@ -81,10 +82,10 @@ class BroadphaseFilter(enum.IntFlag): """Bitmask specifying which collision functions to run during broadphase. Attributes: - PLANE: collision between bounding sphere and plane. - SPHERE: collision between bounding spheres. - AABB: collision between axis-aligned bounding boxes. - OBB: collision between oriented bounding boxes. + PLANE: collision between bounding sphere and plane + SPHERE: collision between bounding spheres + AABB: collision between axis-aligned bounding boxes + OBB: collision between oriented bounding boxes """ PLANE = 1 @@ -526,6 +527,49 @@ class WrapType(enum.IntEnum): CYLINDER = mujoco.mjtWrap.mjWRAP_CYLINDER +class State(enum.IntEnum): + """State component elements as integer bitflags. + + Includes several convenient combinations of these flags. + + Attributes: + TIME: time + QPOS: position + QVEL: velocity + ACT: actuator activation + WARMSTART: acceleration used for warmstart + CTRL: control + QFRC_APPLIED: applied generalized force + XFRC_APPLIED: applied Cartesian force/torque + EQ_ACTIVE: enable/disable constraints + MOCAP_POS: positions of mocap bodies + MOCAP_QUAT: orientations of mocap bodies + NSTATE: number of state elements + PHYSICS: QPOS | QVEL | ACT + FULLPHYSICS: TIME | PHYSICS | PLUGIN + USER: CTRL | QFRC_APPLIED | XFRC_APPLIED | EQ_ACTIVE | MOCAP_POS | MOCAP_QUAT | USERDATA + INTEGRATION: FULLPHYSICS | USER | WARMSTART + """ + + TIME = mujoco.mjtState.mjSTATE_TIME + QPOS = mujoco.mjtState.mjSTATE_QPOS + QVEL = mujoco.mjtState.mjSTATE_QVEL + ACT = mujoco.mjtState.mjSTATE_ACT + WARMSTART = mujoco.mjtState.mjSTATE_WARMSTART + CTRL = mujoco.mjtState.mjSTATE_CTRL + QFRC_APPLIED = mujoco.mjtState.mjSTATE_QFRC_APPLIED + XFRC_APPLIED = mujoco.mjtState.mjSTATE_XFRC_APPLIED + EQ_ACTIVE = mujoco.mjtState.mjSTATE_EQ_ACTIVE + MOCAP_POS = mujoco.mjtState.mjSTATE_MOCAP_POS + MOCAP_QUAT = mujoco.mjtState.mjSTATE_MOCAP_QUAT + NSTATE = mujoco.mjtState.mjNSTATE + PHYSICS = mujoco.mjtState.mjSTATE_PHYSICS + FULLPHYSICS = mujoco.mjtState.mjSTATE_FULLPHYSICS + USER = mujoco.mjtState.mjSTATE_USER + INTEGRATION = mujoco.mjtState.mjSTATE_INTEGRATION + # unsupported: USERDATA, PLUGIN + + class vec5f(wp.types.vector(length=5, dtype=float)): pass @@ -566,28 +610,30 @@ class Option: tolerance: main solver tolerance ls_tolerance: CG/Newton linesearch tolerance ccd_tolerance: convex collision detection tolerance + density: density of medium + viscosity: viscosity of medium gravity: gravitational acceleration + wind: wind (for lift, drag, and viscosity) magnetic: global magnetic flux integrator: integration mode (IntegratorType) cone: type of friction cone (ConeType) solver: solver algorithm (SolverType) iterations: number of main solver iterations ls_iterations: maximum number of CG/Newton linesearch iterations + ccd_iterations: number of iterations in convex collision detection disableflags: bit flags for disabling standard features enableflags: bit flags for enabling optional features - is_sparse: whether to use sparse representations - ccd_iterations: number of iterations in convex collision detection - ls_parallel: evaluate engine solver step sizes in parallel - ls_parallel_min_step: minimum step size for solver linesearch - wind: wind (for lift, drag, and viscosity) - has_fluid: True if wind, density, or viscosity are non-zero at put_model time - density: density of medium - viscosity: viscosity of medium - broadphase: broadphase type (BroadphaseType) - broadphase_filter: broadphase filter bitflag (BroadphaseFilter) - graph_conditional: flag to use cuda graph conditional, should be False when JAX is used sdf_initpoints: number of starting points for gradient descent sdf_iterations: max number of iterations for gradient descent + + warp only fields: + 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 run_collision_detection: if False, skips collision detection and allows user-populated contacts during the physics step (as opposed to DisableBit.CONTACT which explicitly zeros out the contacts at each step) @@ -601,31 +647,32 @@ class Option: tolerance: wp.array(dtype=float) ls_tolerance: wp.array(dtype=float) ccd_tolerance: wp.array(dtype=float) + density: wp.array(dtype=float) + viscosity: wp.array(dtype=float) gravity: wp.array(dtype=wp.vec3) + wind: wp.array(dtype=wp.vec3) magnetic: wp.array(dtype=wp.vec3) integrator: int cone: int solver: int iterations: int ls_iterations: int + ccd_iterations: int disableflags: int enableflags: int - is_sparse: bool - ccd_iterations: int - ls_parallel: bool # warp only - ls_parallel_min_step: float # warp only - wind: wp.array(dtype=wp.vec3) - has_fluid: bool - density: wp.array(dtype=float) - viscosity: wp.array(dtype=float) - broadphase: int # warp only - broadphase_filter: int # warp only - graph_conditional: bool # warp only sdf_initpoints: int sdf_iterations: int - run_collision_detection: bool # warp only + # warp only fields: + is_sparse: bool + ls_parallel: bool + ls_parallel_min_step: float + has_fluid: bool + broadphase: int + broadphase_filter: int + graph_conditional: bool + run_collision_detection: bool legacy_gjk: bool - contact_sensor_maxmatch: int # warp only + contact_sensor_maxmatch: int @dataclasses.dataclass @@ -661,21 +708,20 @@ class Constraint: Mgrad: M / grad (nworld, nv) search: linesearch vector (nworld, nv) search_dot: dot(search, search) (nworld,) - gauss: gauss Cost (nworld,) + gauss: Gauss Cost (nworld,) cost: constraint + Gauss cost (nworld,) prev_cost: cost from previous iter (nworld,) state: constraint state (nworld, njmax) 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) - h: cone hessian (nworld, nv, nv) + quad_gauss: quadratic cost Gauss coefficients (nworld, 3) + h: Hessian (nworld, nv, nv) alpha: line search step size (nworld,) prev_grad: previous grad (nworld, nv) prev_Mgrad: previous Mgrad (nworld, nv) - beta: polak-ribiere beta (nworld,) + beta: Polak-Ribiere beta (nworld,) done: solver done (nworld,) - cost_candidate: costs associated with step sizes (nworld, nlsp) """ type: wp.array2d(dtype=int) @@ -711,8 +757,6 @@ class Constraint: prev_Mgrad: wp.array2d(dtype=float) beta: wp.array(dtype=float) done: wp.array(dtype=bool) - # linesearch - cost_candidate: wp.array2d(dtype=float) @dataclasses.dataclass @@ -730,9 +774,6 @@ class TileSet: size: int -# TODO(team): make Model/Data fields sort order match mujoco - - @dataclasses.dataclass class Model: """Model definition and parameters. @@ -744,47 +785,39 @@ class Model: na: number of activation states nbody: number of bodies njnt: number of joints + nM: number of non-zeros in sparse inertia matrix + nC: number of non-zeros in sparse body-dof matrix ngeom: number of geoms nsite: number of sites ncam: number of cameras nlight: number of lights - nmat: number of materials - nexclude: number of excluded geom pairs - neq: number of equality constraints - nmocap: number of mocap bodies - ngravcomp: number of bodies with nonzero gravcomp - nM: number of non-zeros in sparse inertia matrix - nC: number of non-zeros in sparse reduced dof-dof matrix - ntendon: number of tendons - nwrap: number of wrap objects in all tendon paths - nsensor: number of sensors - nsensordata: number of elements in sensor data vector - nsensortaxel: number of taxels in all tactile sensors + nflex: number of flexes + nflexvert: number of vertices in all flexes + nflexedge: number of edges in all flexes + nflexelem: number of elements in all flexes + nflexelemdata: number of element vertex ids in all flexes nmeshvert: number of vertices for all meshes nmeshface: number of faces for all meshes nmeshgraph: number of ints in mesh auxiliary data nmeshpoly: number of polygons in all meshes nmeshpolyvert: number of vertices in all polygons nmeshpolymap: number of polygons in vertex map - nlsp: number of step sizes for parallel linsearch - npair: number of predefined geom pairs nhfield: number of heightfields nhfielddata: size of elevation data + nmat: number of materials + npair: number of predefined geom pairs + nexclude: number of excluded geom pairs + neq: number of equality constraints + ntendon: number of tendons + nwrap: number of wrap objects in all tendon paths + nsensor: number of sensors + nmocap: number of mocap bodies + ngravcomp: number of bodies with nonzero gravcomp + nsensordata: number of elements in sensor data vector opt: physics options stat: model statistics qpos0: qpos values at default pose (nworld, nq) qpos_spring: reference pose for springs (nworld, nq) - qM_fullm_i: sparse mass matrix addressing - qM_fullm_j: sparse mass matrix addressing - qM_mulm_i: sparse mass matrix addressing - qM_mulm_j: sparse mass matrix addressing - qM_madr_ij: sparse mass matrix addressing - M_rownnz: number of non-zeros in each row of qM (nv,) - M_rowadr: index of each row in qM (nv,) - M_colind: column indices of non-zeros in qM (nM,) - mapM2M: index mapping from M (legacy) to M (CSR) (nC) - qM_tiles: tiling configuration - body_tree: list of body ids by tree level body_parentid: id of body's parent (nbody,) body_rootid: id of root above body (nbody,) body_weldid: id of body that this body is welded to (nbody,) @@ -801,19 +834,21 @@ class Model: body_iquat: local orientation of inertia ellipsoid (nworld, nbody, 4) body_mass: mass (nworld, nbody,) body_subtreemass: mass of subtree starting at this body (nworld, nbody,) - subtree_mass: mass of subtree (nworld, nbody,) body_inertia: diagonal inertia in ipos/iquat frame (nworld, nbody, 3) body_invweight0: mean inv inert in qpos0 (trn, rot) (nworld, nbody, 2) + body_gravcomp: antigravity force, units of body weight (nworld, nbody) body_contype: OR over all geom contypes (nbody,) body_conaffinity: OR over all geom conaffinities (nbody,) - body_gravcomp: antigravity force, units of body weight (nworld, nbody) - body_fluid_ellipsoid: does body use ellipsoid fluid (nbody,) + oct_child: octree children (noct, 8) + oct_aabb: octree axis-aligned bounding boxes (noct, 6) + oct_coeff: octree interpolation coefficients (noct, 8) jnt_type: type of joint (JointType) (njnt,) jnt_qposadr: start addr in 'qpos' for joint's data (njnt,) jnt_dofadr: start addr in 'qvel' for joint's data (njnt,) jnt_bodyid: id of joint's body (njnt,) jnt_limited: does joint have limits (njnt,) jnt_actfrclimited: does joint have actuator force limits (njnt,) + jnt_actgravcomp: is gravcomp force applied via actuators (njnt,) jnt_solref: constraint solver reference: limit (nworld, njnt, mjNREF) jnt_solimp: constraint solver impedance: limit (nworld, njnt, mjNIMP) jnt_pos: local anchor position (nworld, njnt, 3) @@ -822,50 +857,41 @@ class Model: jnt_range: joint limits (nworld, njnt, 2) jnt_actfrcrange: range of total actuator force (nworld, njnt, 2) jnt_margin: min distance for limit detection (nworld, njnt) - jnt_limited_slide_hinge_adr: limited/slide/hinge jntadr - jnt_limited_ball_adr: limited/ball jntadr - jnt_actgravcomp: is gravcomp force applied via actuators (njnt,) 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_Madr: dof address in M-diagonal (nv,) + dof_solref: constraint solver reference: frictionloss (nworld, nv, NREF) + dof_solimp: constraint solver impedance: frictionloss (nworld, nv, NIMP) + dof_frictionloss: dof friction loss (nworld, nv) dof_armature: dof armature inertia/mass (nworld, nv) dof_damping: damping coefficient (nworld, nv) dof_invweight0: diag. inverse inertia in qpos0 (nworld, nv) - dof_frictionloss: dof friction loss (nworld, nv) - dof_solimp: constraint solver impedance: frictionloss (nworld, nv, NIMP) - dof_solref: constraint solver reference: frictionloss (nworld, nv, NREF) - dof_tri_row: np.tril_indices (mjm.nv)[0] - dof_tri_col: np.tril_indices (mjm.nv)[1] geom_type: geometric type (GeomType) (ngeom,) geom_contype: geom contact type (ngeom,) geom_conaffinity: geom contact affinity (ngeom,) geom_condim: contact dimensionality (1, 3, 4, 6) (ngeom,) geom_bodyid: id of geom's body (ngeom,) geom_dataid: id of geom's mesh/hfield; -1: none (ngeom,) - geom_group: geom group inclusion/exclusion mask (ngeom,) geom_matid: material id for rendering (nworld, ngeom,) + geom_group: geom group inclusion/exclusion mask (ngeom,) geom_priority: geom contact priority (ngeom,) geom_solmix: mixing coef for solref/imp in geom pair (nworld, ngeom,) geom_solref: constraint solver reference: contact (nworld, ngeom, mjNREF) geom_solimp: constraint solver impedance: contact (nworld, ngeom, mjNIMP) geom_size: geom-specific size parameters (ngeom, 3) - geom_fluid: fluid interaction parameters (ngeom, mjNFLUID) - geom_aabb: bounding box, (center, size) (ngeom, 6) + geom_aabb: bounding box, (center, size) (nworld, ngeom, 2, 3) geom_rbound: radius of bounding sphere (nworld, ngeom,) geom_pos: local position offset rel. to body (nworld, ngeom, 3) geom_quat: local orientation offset rel. to body (nworld, ngeom, 4) geom_friction: friction for (slide, spin, roll) (nworld, ngeom, 3) geom_margin: detect contact if dist bool: p4: 2D point from segment 2 Returns: - intersection status of line segments + Intersection status of line segments. """ # compute determinant, check det = (p4[1] - p3[1]) * (p2[0] - p1[0]) - (p4[0] - p3[0]) * (p2[1] - p1[1]) @@ -77,13 +77,13 @@ def length_circle(p0: wp.vec2, p1: wp.vec2, ind: int, radius: float) -> float: """Curve length along circle. Args: - p0: 2D point - p1: 2D point - ind: input for flip - radius: circle radius + p0: 2D point. + p1: 2D point. + ind: input for flip. + radius: circle radius. Returns: - curve length + Curve length. """ # compute angle between 0 and pi p0n, _ = math.normalize_with_norm(p0) @@ -104,12 +104,12 @@ def wrap_circle(end: wp.vec4, side: wp.vec2, radius: float) -> Tuple[float, wp.v """2D circle wrap. Args: - end: two 2D points - side: optional 2D side point, no side point: wp.vec2(wp.inf) - radius: circle radius + end: Two 2D points. + side: Optional 2D side point, no side point: wp.vec2(wp.inf). + radius: Circle radius. Returns: - length of circular wrap or -1.0 if no wrap, pair of 2D wrap points + Length of circular wrap or -1.0 if no wrap, pair of 2D wrap points. """ valid_side = wp.norm_l2(side) < wp.inf @@ -210,16 +210,15 @@ def wrap_inside( """2D inside wrap. Args: - end: two 2D points - radius: circle radius - maxiter: maximum number of solver iterations - zinit: initialization for solver - tolerance: solver convergence tolerance + end: Two 2D points. + radius: Circle radius. + maxiter: Maximum number of solver iterations. + zinit: Initialization for solver. + tolerance: Solver convergence tolerance. Returns: - 0.0 if wrap else -1.0, pair of 2D wrap points + 0.0 if wrap else -1.0, pair of 2D wrap points. """ - end0 = wp.vec2(end[0], end[1]) end1 = wp.vec2(end[2], end[3]) @@ -330,16 +329,16 @@ def wrap( """Wrap tendons around spheres and cylinders. Args: - x0: 3D endpoint - x1: 3D endpoint - pos: position of geom - mat: orientation of geom - radius: geom radius - type: wrap type (mjtWrap) - side: 3D position for sidesite, no side point: wp.vec3(wp.inf) + x0: 3D endpoint. + x1: 3D endpoint. + pos: Position of geom. + mat: Orientation of geom. + radius: Geom radius. + geomtype: Wrap type (mjtWrap). + side: 3D position for sidesite, no side point: wp.vec3(wp.inf). Returns: - length of circular wrap else -1.0 if no wrap, pair of 3D wrap points + 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: @@ -453,7 +452,6 @@ def wrap( @wp.func def muscle_gain_length(length: float, lmin: float, lmax: float) -> float: """Normalized muscle length-gain curve.""" - if (lmin > length) or (length > lmax): return 0.0 @@ -478,7 +476,6 @@ def muscle_gain_length(length: float, lmin: float, lmax: float) -> float: @wp.func def muscle_gain(len: float, vel: float, lengthrange: wp.vec2, acc0: float, prm: vec10) -> float: """Muscle active force, prm = (range[2], force, scale, lmin, lmax, vmax, fpmax, fvmax).""" - # unpack parameters range_ = wp.vec2(prm[0], prm[1]) force = prm[2] @@ -521,8 +518,8 @@ def muscle_gain(len: float, vel: float, lengthrange: wp.vec2, acc0: float, prm: def muscle_bias(len: float, lengthrange: wp.vec2, acc0: float, prm: vec10) -> float: """Calculates muscle passive force. - prm = (range[2], force, scale, lmin, lmax, vmax, fpmax, fvmax).""" - + prm = (range[2], force, scale, lmin, lmax, vmax, fpmax, fvmax). + """ # unpack parameters range_ = wp.vec2(prm[0], prm[1]) force = prm[2] @@ -555,7 +552,6 @@ def muscle_bias(len: float, lengthrange: wp.vec2, acc0: float, prm: vec10) -> fl @wp.func def _sigmoid(x: float) -> float: """Sigmoid function over 0 <= x <= 1 using quintic polynomial.""" - if x <= 0.0: return 0.0 @@ -570,7 +566,6 @@ def _sigmoid(x: float) -> float: @wp.func def muscle_dynamics_timescale(dctrl: float, tau_act: float, tau_deact: float, smooth_width: float) -> float: """Muscle time constant with optional smoothing.""" - # hard switching if smooth_width < MJ_MINVAL: if dctrl > 0.0: @@ -585,7 +580,6 @@ def muscle_dynamics_timescale(dctrl: float, tau_act: float, tau_deact: float, sm @wp.func def muscle_dynamics(control: float, activation: float, prm: vec10) -> float: """Muscle activation dynamics, prm = (tau_act, tau_deact, smooth_width).""" - # clamp control ctrlclamp = wp.clamp(control, 0.0, 1.0) 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 7b671433..10612283 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 @@ -19,7 +19,6 @@ from typing import Callable, Optional import warp as wp from warp.context import Module -from warp.context import assert_conditional_graph_support from warp.context import get_module _STACK = None @@ -130,8 +129,8 @@ def kernel( enable_backward: Optional[bool] = None, module: Optional[Module] = None, ): - """ - Decorator to register a Warp kernel from a Python function. + """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. @@ -220,9 +219,9 @@ def cache_kernel(func): return wrapper -def conditional_graph_supported(): - try: - assert_conditional_graph_support() - except Exception: - return False - return True +def check_toolkit_driver(): + if wp.context.runtime is None: + wp.context.init() + if wp.get_device().is_cuda: + if wp.context.runtime.toolkit_version < (12, 4) or wp.context.runtime.driver_version < (12, 4): + 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 b72c389b..f982afff 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/pyproject.toml @@ -28,9 +28,9 @@ requires-python = ">=3.9" dependencies = [ "absl-py", "etils[epath]", - "mujoco>=3.3.6.dev802089588", + "mujoco>=3.3.7", "numpy", - "warp-lang>=1.9.0.dev20250825", + "warp-lang>=1.9.1", ] [[tool.uv.index]] @@ -49,6 +49,7 @@ mujoco = {index = "mujoco"} [project.optional-dependencies] dev = [ + "asv", "pre-commit", "pytest", "pytest-xdist", @@ -81,8 +82,15 @@ extend-exclude = ["*.ipynb"] [tool.ruff.lint] select = [ + "D", # pydocstyle conventions and style "I", # isort "W", # pycodestyle + "F401", # unused imports +] + +ignore = [ + "D100", # missing docstring public module + "D103", # missing docstring public function ] [tool.ruff.lint.isort] @@ -93,6 +101,13 @@ single-line-exclusions = ["typing"] max-doc-length = 100 max-line-length = 128 +[tool.ruff.lint.pydocstyle] +convention = "google" + +[tool.ruff.lint.per-file-ignores] +"__init__.py" = ["D"] +"contrib/*" = ["D"] + [tool.ruff.format] docstring-code-format = true docstring-code-line-length = 100 diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py index 75b61720..7d1b9afd 100644 --- a/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py +++ b/mjx/mujoco/mjx/third_party/mujoco_warp/viewer.py @@ -45,6 +45,8 @@ from mujoco.mjx.third_party.mujoco_warp._src.io import override_model class EngineOptions(enum.IntEnum): + """Engine option.""" + WARP = 0 C = 1 @@ -142,15 +144,16 @@ def _main(argv: Sequence[str]) -> None: 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, ls_parallel = m.opt.iterations, m.opt.ls_iterations, m.opt.ls_parallel + 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_iterations: {ls_iterations} ls_parallel: {ls_parallel}\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.nconmax} njmax: {d.njmax}\n") + 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}...") diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py index 0675b0e5..d4c8a1a1 100644 --- a/mjx/mujoco/mjx/warp/collision_driver.py +++ b/mjx/mujoco/mjx/warp/collision_driver.py @@ -42,12 +42,13 @@ _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 nworld: int, block_dim: mjwp_types.BlockDim, - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_condim: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), geom_friction: wp.array2d(dtype=wp.vec3), @@ -85,10 +86,12 @@ def _collision_shim( mesh_vertadr: wp.array(dtype=int), mesh_vertnum: wp.array(dtype=int), ngeom: int, + nmaxmeshdeg: int, + nmaxpolygon: int, nmeshface: int, nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i), - nxn_pairid: wp.array(dtype=int), - nxn_pairid_filtered: wp.array(dtype=int), + nxn_pairid: wp.array(dtype=wp.vec2i), + nxn_pairid_filtered: wp.array(dtype=wp.vec2i), oct_aabb: wp.array2d(dtype=wp.vec3), oct_child: wp.array(dtype=mjwp_types.vec8i), oct_coeff: wp.array(dtype=mjwp_types.vec8f), @@ -106,57 +109,30 @@ def _collision_shim( opt__ccd_iterations: int, opt__ccd_tolerance: wp.array(dtype=float), opt__disableflags: int, - opt__graph_conditional: bool, opt__legacy_gjk: bool, opt__sdf_initpoints: int, opt__sdf_iterations: int, # Data naconmax: int, collision_pair: wp.array(dtype=wp.vec2i), - collision_pairid: wp.array(dtype=int), + collision_pairid: wp.array(dtype=wp.vec2i), collision_worldid: wp.array(dtype=int), - epa_face: wp.array2d(dtype=wp.vec3i), - epa_horizon: wp.array2d(dtype=int), - epa_index: wp.array2d(dtype=int), - epa_map: wp.array2d(dtype=int), - epa_norm2: wp.array2d(dtype=float), - epa_pr: wp.array2d(dtype=wp.vec3), - epa_vert: wp.array2d(dtype=wp.vec3), - epa_vert1: wp.array2d(dtype=wp.vec3), - epa_vert2: wp.array2d(dtype=wp.vec3), - epa_vert_index1: wp.array2d(dtype=int), - epa_vert_index2: wp.array2d(dtype=int), geom_xmat: wp.array2d(dtype=wp.mat33), geom_xpos: wp.array2d(dtype=wp.vec3), - multiccd_clipped: wp.array2d(dtype=wp.vec3), - multiccd_endvert: wp.array2d(dtype=wp.vec3), - multiccd_face1: wp.array2d(dtype=wp.vec3), - multiccd_face2: wp.array2d(dtype=wp.vec3), - multiccd_idx1: wp.array2d(dtype=int), - multiccd_idx2: wp.array2d(dtype=int), - multiccd_n1: wp.array2d(dtype=wp.vec3), - multiccd_n2: wp.array2d(dtype=wp.vec3), - multiccd_pdist: wp.array2d(dtype=float), - multiccd_pnormal: wp.array2d(dtype=wp.vec3), - multiccd_polygon: wp.array2d(dtype=wp.vec3), nacon: wp.array(dtype=int), ncollision: wp.array(dtype=int), - sap_cumulative_sum: wp.array2d(dtype=int), - sap_projection_lower: wp.array3d(dtype=float), - sap_projection_upper: wp.array2d(dtype=float), - sap_range: wp.array2d(dtype=int), - sap_segment_index: wp.array2d(dtype=int), - sap_sort_index: wp.array3d(dtype=int), contact__dim: wp.array(dtype=int), contact__dist: wp.array(dtype=float), contact__frame: wp.array(dtype=wp.mat33), contact__friction: wp.array(dtype=mjwp_types.vec5), contact__geom: wp.array(dtype=wp.vec2i), + contact__geomcollisionid: wp.array(dtype=int), contact__includemargin: wp.array(dtype=float), contact__pos: wp.array(dtype=wp.vec3), contact__solimp: wp.array(dtype=mjwp_types.vec5), contact__solref: wp.array(dtype=wp.vec2), contact__solreffriction: wp.array(dtype=wp.vec2), + contact__type: wp.array(dtype=int), contact__worldid: wp.array(dtype=int), ): _m.stat = _s @@ -202,6 +178,8 @@ def _collision_shim( _m.mesh_vertadr = mesh_vertadr _m.mesh_vertnum = mesh_vertnum _m.ngeom = ngeom + _m.nmaxmeshdeg = nmaxmeshdeg + _m.nmaxpolygon = nmaxpolygon _m.nmeshface = nmeshface _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered _m.nxn_pairid = nxn_pairid @@ -214,7 +192,6 @@ def _collision_shim( _m.opt.ccd_iterations = opt__ccd_iterations _m.opt.ccd_tolerance = opt__ccd_tolerance _m.opt.disableflags = opt__disableflags - _m.opt.graph_conditional = opt__graph_conditional _m.opt.legacy_gjk = opt__legacy_gjk _m.opt.sdf_initpoints = opt__sdf_initpoints _m.opt.sdf_iterations = opt__sdf_iterations @@ -235,45 +212,19 @@ def _collision_shim( _d.contact.frame = contact__frame _d.contact.friction = contact__friction _d.contact.geom = contact__geom + _d.contact.geomcollisionid = contact__geomcollisionid _d.contact.includemargin = contact__includemargin _d.contact.pos = contact__pos _d.contact.solimp = contact__solimp _d.contact.solref = contact__solref _d.contact.solreffriction = contact__solreffriction + _d.contact.type = contact__type _d.contact.worldid = contact__worldid - _d.epa_face = epa_face - _d.epa_horizon = epa_horizon - _d.epa_index = epa_index - _d.epa_map = epa_map - _d.epa_norm2 = epa_norm2 - _d.epa_pr = epa_pr - _d.epa_vert = epa_vert - _d.epa_vert1 = epa_vert1 - _d.epa_vert2 = epa_vert2 - _d.epa_vert_index1 = epa_vert_index1 - _d.epa_vert_index2 = epa_vert_index2 _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos - _d.multiccd_clipped = multiccd_clipped - _d.multiccd_endvert = multiccd_endvert - _d.multiccd_face1 = multiccd_face1 - _d.multiccd_face2 = multiccd_face2 - _d.multiccd_idx1 = multiccd_idx1 - _d.multiccd_idx2 = multiccd_idx2 - _d.multiccd_n1 = multiccd_n1 - _d.multiccd_n2 = multiccd_n2 - _d.multiccd_pdist = multiccd_pdist - _d.multiccd_pnormal = multiccd_pnormal - _d.multiccd_polygon = multiccd_polygon _d.nacon = nacon _d.naconmax = naconmax _d.ncollision = ncollision - _d.sap_cumulative_sum = sap_cumulative_sum - _d.sap_projection_lower = sap_projection_lower - _d.sap_projection_upper = sap_projection_upper - _d.sap_range = sap_range - _d.sap_segment_index = sap_segment_index - _d.sap_sort_index = sap_sort_index _d.nworld = nworld mjwarp.collision(_m, _d) @@ -283,101 +234,49 @@ def _collision_jax_impl(m: types.Model, d: types.Data): 'collision_pair': d._impl.collision_pair.shape, 'collision_pairid': d._impl.collision_pairid.shape, 'collision_worldid': d._impl.collision_worldid.shape, - 'epa_face': d._impl.epa_face.shape, - 'epa_horizon': d._impl.epa_horizon.shape, - 'epa_index': d._impl.epa_index.shape, - 'epa_map': d._impl.epa_map.shape, - 'epa_norm2': d._impl.epa_norm2.shape, - 'epa_pr': d._impl.epa_pr.shape, - 'epa_vert': d._impl.epa_vert.shape, - 'epa_vert1': d._impl.epa_vert1.shape, - 'epa_vert2': d._impl.epa_vert2.shape, - 'epa_vert_index1': d._impl.epa_vert_index1.shape, - 'epa_vert_index2': d._impl.epa_vert_index2.shape, 'geom_xmat': d.geom_xmat.shape, 'geom_xpos': d.geom_xpos.shape, - 'multiccd_clipped': d._impl.multiccd_clipped.shape, - 'multiccd_endvert': d._impl.multiccd_endvert.shape, - 'multiccd_face1': d._impl.multiccd_face1.shape, - 'multiccd_face2': d._impl.multiccd_face2.shape, - 'multiccd_idx1': d._impl.multiccd_idx1.shape, - 'multiccd_idx2': d._impl.multiccd_idx2.shape, - 'multiccd_n1': d._impl.multiccd_n1.shape, - 'multiccd_n2': d._impl.multiccd_n2.shape, - 'multiccd_pdist': d._impl.multiccd_pdist.shape, - 'multiccd_pnormal': d._impl.multiccd_pnormal.shape, - 'multiccd_polygon': d._impl.multiccd_polygon.shape, 'nacon': d._impl.nacon.shape, 'ncollision': d._impl.ncollision.shape, - 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape, - 'sap_projection_lower': d._impl.sap_projection_lower.shape, - 'sap_projection_upper': d._impl.sap_projection_upper.shape, - 'sap_range': d._impl.sap_range.shape, - 'sap_segment_index': d._impl.sap_segment_index.shape, - 'sap_sort_index': d._impl.sap_sort_index.shape, 'contact__dim': d._impl.contact__dim.shape, 'contact__dist': d._impl.contact__dist.shape, 'contact__frame': d._impl.contact__frame.shape, 'contact__friction': d._impl.contact__friction.shape, 'contact__geom': d._impl.contact__geom.shape, + 'contact__geomcollisionid': d._impl.contact__geomcollisionid.shape, 'contact__includemargin': d._impl.contact__includemargin.shape, 'contact__pos': d._impl.contact__pos.shape, 'contact__solimp': d._impl.contact__solimp.shape, 'contact__solref': d._impl.contact__solref.shape, 'contact__solreffriction': d._impl.contact__solreffriction.shape, + 'contact__type': d._impl.contact__type.shape, 'contact__worldid': d._impl.contact__worldid.shape, } jf = ffi.jax_callable_variadic_tuple( _collision_shim, - num_outputs=46, + num_outputs=20, output_dims=output_dims, vmap_method=None, in_out_argnames={ 'collision_pair', 'collision_pairid', 'collision_worldid', - 'epa_face', - 'epa_horizon', - 'epa_index', - 'epa_map', - 'epa_norm2', - 'epa_pr', - 'epa_vert', - 'epa_vert1', - 'epa_vert2', - 'epa_vert_index1', - 'epa_vert_index2', 'geom_xmat', 'geom_xpos', - 'multiccd_clipped', - 'multiccd_endvert', - 'multiccd_face1', - 'multiccd_face2', - 'multiccd_idx1', - 'multiccd_idx2', - 'multiccd_n1', - 'multiccd_n2', - 'multiccd_pdist', - 'multiccd_pnormal', - 'multiccd_polygon', 'nacon', 'ncollision', - 'sap_cumulative_sum', - 'sap_projection_lower', - 'sap_projection_upper', - 'sap_range', - 'sap_segment_index', - 'sap_sort_index', 'contact__dim', 'contact__dist', 'contact__frame', 'contact__friction', 'contact__geom', + 'contact__geomcollisionid', 'contact__includemargin', 'contact__pos', 'contact__solimp', 'contact__solref', 'contact__solreffriction', + 'contact__type', 'contact__worldid', }, ) @@ -422,6 +321,8 @@ def _collision_jax_impl(m: types.Model, d: types.Data): m.mesh_vertadr, m.mesh_vertnum, m.ngeom, + m._impl.nmaxmeshdeg, + m._impl.nmaxpolygon, m.nmeshface, m._impl.nxn_geom_pair_filtered, m._impl.nxn_pairid, @@ -443,7 +344,6 @@ def _collision_jax_impl(m: types.Model, d: types.Data): m.opt._impl.ccd_iterations, m.opt._impl.ccd_tolerance, m.opt.disableflags, - m.opt._impl.graph_conditional, m.opt._impl.legacy_gjk, m.opt._impl.sdf_initpoints, m.opt._impl.sdf_iterations, @@ -451,97 +351,45 @@ def _collision_jax_impl(m: types.Model, d: types.Data): d._impl.collision_pair, d._impl.collision_pairid, d._impl.collision_worldid, - d._impl.epa_face, - d._impl.epa_horizon, - d._impl.epa_index, - d._impl.epa_map, - d._impl.epa_norm2, - d._impl.epa_pr, - d._impl.epa_vert, - d._impl.epa_vert1, - d._impl.epa_vert2, - d._impl.epa_vert_index1, - d._impl.epa_vert_index2, d.geom_xmat, d.geom_xpos, - d._impl.multiccd_clipped, - d._impl.multiccd_endvert, - d._impl.multiccd_face1, - d._impl.multiccd_face2, - d._impl.multiccd_idx1, - d._impl.multiccd_idx2, - d._impl.multiccd_n1, - d._impl.multiccd_n2, - d._impl.multiccd_pdist, - d._impl.multiccd_pnormal, - d._impl.multiccd_polygon, d._impl.nacon, d._impl.ncollision, - d._impl.sap_cumulative_sum, - d._impl.sap_projection_lower, - d._impl.sap_projection_upper, - d._impl.sap_range, - d._impl.sap_segment_index, - d._impl.sap_sort_index, d._impl.contact__dim, d._impl.contact__dist, d._impl.contact__frame, d._impl.contact__friction, d._impl.contact__geom, + d._impl.contact__geomcollisionid, d._impl.contact__includemargin, d._impl.contact__pos, d._impl.contact__solimp, d._impl.contact__solref, d._impl.contact__solreffriction, + d._impl.contact__type, d._impl.contact__worldid, ) d = d.tree_replace({ '_impl.collision_pair': out[0], '_impl.collision_pairid': out[1], '_impl.collision_worldid': out[2], - '_impl.epa_face': out[3], - '_impl.epa_horizon': out[4], - '_impl.epa_index': out[5], - '_impl.epa_map': out[6], - '_impl.epa_norm2': out[7], - '_impl.epa_pr': out[8], - '_impl.epa_vert': out[9], - '_impl.epa_vert1': out[10], - '_impl.epa_vert2': out[11], - '_impl.epa_vert_index1': out[12], - '_impl.epa_vert_index2': out[13], - 'geom_xmat': out[14], - 'geom_xpos': out[15], - '_impl.multiccd_clipped': out[16], - '_impl.multiccd_endvert': out[17], - '_impl.multiccd_face1': out[18], - '_impl.multiccd_face2': out[19], - '_impl.multiccd_idx1': out[20], - '_impl.multiccd_idx2': out[21], - '_impl.multiccd_n1': out[22], - '_impl.multiccd_n2': out[23], - '_impl.multiccd_pdist': out[24], - '_impl.multiccd_pnormal': out[25], - '_impl.multiccd_polygon': out[26], - '_impl.nacon': out[27], - '_impl.ncollision': out[28], - '_impl.sap_cumulative_sum': out[29], - '_impl.sap_projection_lower': out[30], - '_impl.sap_projection_upper': out[31], - '_impl.sap_range': out[32], - '_impl.sap_segment_index': out[33], - '_impl.sap_sort_index': out[34], - '_impl.contact__dim': out[35], - '_impl.contact__dist': out[36], - '_impl.contact__frame': out[37], - '_impl.contact__friction': out[38], - '_impl.contact__geom': out[39], - '_impl.contact__includemargin': out[40], - '_impl.contact__pos': out[41], - '_impl.contact__solimp': out[42], - '_impl.contact__solref': out[43], - '_impl.contact__solreffriction': out[44], - '_impl.contact__worldid': out[45], + 'geom_xmat': out[3], + 'geom_xpos': out[4], + '_impl.nacon': out[5], + '_impl.ncollision': out[6], + '_impl.contact__dim': out[7], + '_impl.contact__dist': out[8], + '_impl.contact__frame': out[9], + '_impl.contact__friction': out[10], + '_impl.contact__geom': out[11], + '_impl.contact__geomcollisionid': out[12], + '_impl.contact__includemargin': out[13], + '_impl.contact__pos': out[14], + '_impl.contact__solimp': out[15], + '_impl.contact__solref': out[16], + '_impl.contact__solreffriction': out[17], + '_impl.contact__type': out[18], + '_impl.contact__worldid': out[19], }) return d diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py index ae66d6e7..156363ba 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 @@ -84,6 +85,7 @@ def _forward_shim( body_jntadr: wp.array(dtype=int), body_jntnum: wp.array(dtype=int), body_mass: wp.array2d(dtype=float), + body_mocapid: wp.array(dtype=int), body_parentid: wp.array(dtype=int), body_pos: wp.array2d(dtype=wp.vec3), body_quat: wp.array2d(dtype=wp.quat), @@ -140,7 +142,7 @@ def _forward_shim( flex_vertadr: wp.array(dtype=int), flex_vertbodyid: wp.array(dtype=int), flexedge_length0: wp.array(dtype=float), - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_bodyid: wp.array(dtype=int), geom_condim: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), @@ -216,6 +218,7 @@ def _forward_shim( mocap_bodyid: wp.array(dtype=int), nC: int, na: int, + nacttrnbody: int, nbody: int, ncam: int, neq: int, @@ -226,17 +229,22 @@ def _forward_shim( ngravcomp: int, njnt: int, nlight: int, - nlsp: int, + nmaxmeshdeg: int, + nmaxpolygon: int, nmeshface: int, nmocap: int, + nrangefinder: int, + nsensorcollision: int, + nsensorcontact: int, nsensortaxel: int, nsite: int, ntendon: int, nu: int, nv: int, + nwrap: int, nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i), - nxn_pairid: wp.array(dtype=int), - nxn_pairid_filtered: wp.array(dtype=int), + nxn_pairid: wp.array(dtype=wp.vec2i), + nxn_pairid_filtered: wp.array(dtype=wp.vec2i), oct_aabb: wp.array2d(dtype=wp.vec3), oct_child: wp.array(dtype=mjwp_types.vec8i), oct_coeff: wp.array(dtype=mjwp_types.vec8f), @@ -262,6 +270,7 @@ def _forward_shim( sensor_acc_adr: wp.array(dtype=int), sensor_adr: wp.array(dtype=int), sensor_adr_to_contact_adr: wp.array(dtype=int), + sensor_collision_start_adr: wp.array(dtype=int), sensor_contact_adr: wp.array(dtype=int), sensor_cutoff: wp.array(dtype=float), sensor_datatype: wp.array(dtype=int), @@ -290,7 +299,6 @@ def _forward_shim( site_quat: wp.array2d(dtype=wp.quat), site_size: wp.array(dtype=wp.vec3), site_type: wp.array(dtype=int), - subtree_mass: wp.array2d(dtype=float), taxel_sensorid: wp.array(dtype=int), taxel_vertadr: wp.array(dtype=int), tendon_actfrclimited: wp.array(dtype=bool), @@ -359,7 +367,6 @@ def _forward_shim( actuator_force: wp.array2d(dtype=float), actuator_length: wp.array2d(dtype=float), actuator_moment: wp.array3d(dtype=float), - actuator_trntype_body_ncon: wp.array2d(dtype=int), actuator_velocity: wp.array2d(dtype=float), cacc: wp.array2d(dtype=wp.spatial_vector), cam_xmat: wp.array2d(dtype=wp.mat33), @@ -370,46 +377,22 @@ def _forward_shim( 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=int), + 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), energy: wp.array(dtype=wp.vec2), - epa_face: wp.array2d(dtype=wp.vec3i), - epa_horizon: wp.array2d(dtype=int), - epa_index: wp.array2d(dtype=int), - epa_map: wp.array2d(dtype=int), - epa_norm2: wp.array2d(dtype=float), - epa_pr: wp.array2d(dtype=wp.vec3), - epa_vert: wp.array2d(dtype=wp.vec3), - epa_vert1: wp.array2d(dtype=wp.vec3), - epa_vert2: wp.array2d(dtype=wp.vec3), - epa_vert_index1: wp.array2d(dtype=int), - epa_vert_index2: wp.array2d(dtype=int), eq_active: wp.array2d(dtype=bool), flexedge_length: wp.array2d(dtype=float), flexedge_velocity: wp.array2d(dtype=float), flexvert_xpos: wp.array2d(dtype=wp.vec3), - fluid_applied: wp.array2d(dtype=wp.spatial_vector), - geom_skip: wp.array(dtype=bool), geom_xmat: wp.array2d(dtype=wp.mat33), geom_xpos: wp.array2d(dtype=wp.vec3), light_xdir: wp.array2d(dtype=wp.vec3), light_xpos: wp.array2d(dtype=wp.vec3), mocap_pos: wp.array2d(dtype=wp.vec3), mocap_quat: wp.array2d(dtype=wp.quat), - multiccd_clipped: wp.array2d(dtype=wp.vec3), - multiccd_endvert: wp.array2d(dtype=wp.vec3), - multiccd_face1: wp.array2d(dtype=wp.vec3), - multiccd_face2: wp.array2d(dtype=wp.vec3), - multiccd_idx1: wp.array2d(dtype=int), - multiccd_idx2: wp.array2d(dtype=int), - multiccd_n1: wp.array2d(dtype=wp.vec3), - multiccd_n2: wp.array2d(dtype=wp.vec3), - multiccd_pdist: wp.array2d(dtype=float), - multiccd_pnormal: wp.array2d(dtype=wp.vec3), - multiccd_polygon: wp.array2d(dtype=wp.vec3), nacon: wp.array(dtype=int), ncollision: wp.array(dtype=int), ne: wp.array(dtype=int), @@ -439,20 +422,6 @@ def _forward_shim( qfrc_spring: wp.array2d(dtype=float), qpos: wp.array2d(dtype=float), qvel: wp.array2d(dtype=float), - sap_cumulative_sum: wp.array2d(dtype=int), - sap_projection_lower: wp.array3d(dtype=float), - sap_projection_upper: wp.array2d(dtype=float), - sap_range: wp.array2d(dtype=int), - sap_segment_index: wp.array2d(dtype=int), - sap_sort_index: wp.array3d(dtype=int), - sensor_contact_criteria: wp.array3d(dtype=float), - sensor_contact_direction: wp.array3d(dtype=float), - sensor_contact_matchid: wp.array3d(dtype=int), - sensor_contact_nmatch: wp.array2d(dtype=int), - sensor_rangefinder_dist: wp.array2d(dtype=float), - sensor_rangefinder_geomid: wp.array2d(dtype=int), - sensor_rangefinder_pnt: wp.array2d(dtype=wp.vec3), - sensor_rangefinder_vec: wp.array2d(dtype=wp.vec3), sensordata: wp.array2d(dtype=float), site_xmat: wp.array2d(dtype=wp.mat33), site_xpos: wp.array2d(dtype=wp.vec3), @@ -462,15 +431,11 @@ def _forward_shim( subtree_com: wp.array2d(dtype=wp.vec3), subtree_linvel: wp.array2d(dtype=wp.vec3), ten_J: wp.array3d(dtype=float), - ten_Jdot: wp.array3d(dtype=float), - ten_actfrc: wp.array2d(dtype=float), - ten_bias_coef: wp.array2d(dtype=float), ten_length: wp.array2d(dtype=float), ten_velocity: wp.array2d(dtype=float), ten_wrapadr: wp.array2d(dtype=int), ten_wrapnum: wp.array2d(dtype=int), time: wp.array(dtype=float), - wrap_geom_xpos: wp.array2d(dtype=wp.spatial_vector), wrap_obj: wp.array2d(dtype=wp.vec2i), wrap_xpos: wp.array2d(dtype=wp.spatial_vector), xanchor: wp.array2d(dtype=wp.vec3), @@ -487,11 +452,13 @@ def _forward_shim( contact__frame: wp.array(dtype=wp.mat33), contact__friction: wp.array(dtype=mjwp_types.vec5), contact__geom: wp.array(dtype=wp.vec2i), + contact__geomcollisionid: wp.array(dtype=int), contact__includemargin: wp.array(dtype=float), contact__pos: wp.array(dtype=wp.vec3), contact__solimp: wp.array(dtype=mjwp_types.vec5), contact__solref: wp.array(dtype=wp.vec2), contact__solreffriction: wp.array(dtype=wp.vec2), + contact__type: wp.array(dtype=int), contact__worldid: wp.array(dtype=int), efc__D: wp.array2d(dtype=float), efc__J: wp.array3d(dtype=float), @@ -504,7 +471,6 @@ def _forward_shim( efc__cholesky_L_tmp: wp.array3d(dtype=float), efc__cholesky_y_tmp: wp.array2d(dtype=float), efc__cost: wp.array(dtype=float), - efc__cost_candidate: wp.array2d(dtype=float), efc__done: wp.array(dtype=bool), efc__force: wp.array2d(dtype=float), efc__frictionloss: wp.array2d(dtype=float), @@ -570,6 +536,7 @@ def _forward_shim( _m.body_jntadr = body_jntadr _m.body_jntnum = body_jntnum _m.body_mass = body_mass + _m.body_mocapid = body_mocapid _m.body_parentid = body_parentid _m.body_pos = body_pos _m.body_quat = body_quat @@ -702,6 +669,7 @@ def _forward_shim( _m.mocap_bodyid = mocap_bodyid _m.nC = nC _m.na = na + _m.nacttrnbody = nacttrnbody _m.nbody = nbody _m.ncam = ncam _m.neq = neq @@ -712,14 +680,19 @@ def _forward_shim( _m.ngravcomp = ngravcomp _m.njnt = njnt _m.nlight = nlight - _m.nlsp = nlsp + _m.nmaxmeshdeg = nmaxmeshdeg + _m.nmaxpolygon = nmaxpolygon _m.nmeshface = nmeshface _m.nmocap = nmocap + _m.nrangefinder = nrangefinder + _m.nsensorcollision = nsensorcollision + _m.nsensorcontact = nsensorcontact _m.nsensortaxel = nsensortaxel _m.nsite = nsite _m.ntendon = ntendon _m.nu = nu _m.nv = nv + _m.nwrap = nwrap _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered _m.nxn_pairid = nxn_pairid _m.nxn_pairid_filtered = nxn_pairid_filtered @@ -777,6 +750,7 @@ def _forward_shim( _m.sensor_acc_adr = sensor_acc_adr _m.sensor_adr = sensor_adr _m.sensor_adr_to_contact_adr = sensor_adr_to_contact_adr + _m.sensor_collision_start_adr = sensor_collision_start_adr _m.sensor_contact_adr = sensor_contact_adr _m.sensor_cutoff = sensor_cutoff _m.sensor_datatype = sensor_datatype @@ -806,7 +780,6 @@ def _forward_shim( _m.site_size = site_size _m.site_type = site_type _m.stat.meaninertia = stat__meaninertia - _m.subtree_mass = subtree_mass _m.taxel_sensorid = taxel_sensorid _m.taxel_vertadr = taxel_vertadr _m.tendon_actfrclimited = tendon_actfrclimited @@ -842,7 +815,6 @@ def _forward_shim( _d.actuator_force = actuator_force _d.actuator_length = actuator_length _d.actuator_moment = actuator_moment - _d.actuator_trntype_body_ncon = actuator_trntype_body_ncon _d.actuator_velocity = actuator_velocity _d.cacc = cacc _d.cam_xmat = cam_xmat @@ -861,11 +833,13 @@ def _forward_shim( _d.contact.frame = contact__frame _d.contact.friction = contact__friction _d.contact.geom = contact__geom + _d.contact.geomcollisionid = contact__geomcollisionid _d.contact.includemargin = contact__includemargin _d.contact.pos = contact__pos _d.contact.solimp = contact__solimp _d.contact.solref = contact__solref _d.contact.solreffriction = contact__solreffriction + _d.contact.type = contact__type _d.contact.worldid = contact__worldid _d.crb = crb _d.ctrl = ctrl @@ -881,7 +855,6 @@ def _forward_shim( _d.efc.cholesky_L_tmp = efc__cholesky_L_tmp _d.efc.cholesky_y_tmp = efc__cholesky_y_tmp _d.efc.cost = efc__cost - _d.efc.cost_candidate = efc__cost_candidate _d.efc.done = efc__done _d.efc.force = efc__force _d.efc.frictionloss = efc__frictionloss @@ -905,40 +878,16 @@ def _forward_shim( _d.efc.type = efc__type _d.efc.vel = efc__vel _d.energy = energy - _d.epa_face = epa_face - _d.epa_horizon = epa_horizon - _d.epa_index = epa_index - _d.epa_map = epa_map - _d.epa_norm2 = epa_norm2 - _d.epa_pr = epa_pr - _d.epa_vert = epa_vert - _d.epa_vert1 = epa_vert1 - _d.epa_vert2 = epa_vert2 - _d.epa_vert_index1 = epa_vert_index1 - _d.epa_vert_index2 = epa_vert_index2 _d.eq_active = eq_active _d.flexedge_length = flexedge_length _d.flexedge_velocity = flexedge_velocity _d.flexvert_xpos = flexvert_xpos - _d.fluid_applied = fluid_applied - _d.geom_skip = geom_skip _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos _d.light_xdir = light_xdir _d.light_xpos = light_xpos _d.mocap_pos = mocap_pos _d.mocap_quat = mocap_quat - _d.multiccd_clipped = multiccd_clipped - _d.multiccd_endvert = multiccd_endvert - _d.multiccd_face1 = multiccd_face1 - _d.multiccd_face2 = multiccd_face2 - _d.multiccd_idx1 = multiccd_idx1 - _d.multiccd_idx2 = multiccd_idx2 - _d.multiccd_n1 = multiccd_n1 - _d.multiccd_n2 = multiccd_n2 - _d.multiccd_pdist = multiccd_pdist - _d.multiccd_pnormal = multiccd_pnormal - _d.multiccd_polygon = multiccd_polygon _d.nacon = nacon _d.naconmax = naconmax _d.ncollision = ncollision @@ -970,20 +919,6 @@ def _forward_shim( _d.qfrc_spring = qfrc_spring _d.qpos = qpos _d.qvel = qvel - _d.sap_cumulative_sum = sap_cumulative_sum - _d.sap_projection_lower = sap_projection_lower - _d.sap_projection_upper = sap_projection_upper - _d.sap_range = sap_range - _d.sap_segment_index = sap_segment_index - _d.sap_sort_index = sap_sort_index - _d.sensor_contact_criteria = sensor_contact_criteria - _d.sensor_contact_direction = sensor_contact_direction - _d.sensor_contact_matchid = sensor_contact_matchid - _d.sensor_contact_nmatch = sensor_contact_nmatch - _d.sensor_rangefinder_dist = sensor_rangefinder_dist - _d.sensor_rangefinder_geomid = sensor_rangefinder_geomid - _d.sensor_rangefinder_pnt = sensor_rangefinder_pnt - _d.sensor_rangefinder_vec = sensor_rangefinder_vec _d.sensordata = sensordata _d.site_xmat = site_xmat _d.site_xpos = site_xpos @@ -993,15 +928,11 @@ def _forward_shim( _d.subtree_com = subtree_com _d.subtree_linvel = subtree_linvel _d.ten_J = ten_J - _d.ten_Jdot = ten_Jdot - _d.ten_actfrc = ten_actfrc - _d.ten_bias_coef = ten_bias_coef _d.ten_length = ten_length _d.ten_velocity = ten_velocity _d.ten_wrapadr = ten_wrapadr _d.ten_wrapnum = ten_wrapnum _d.time = time - _d.wrap_geom_xpos = wrap_geom_xpos _d.wrap_obj = wrap_obj _d.wrap_xpos = wrap_xpos _d.xanchor = xanchor @@ -1023,7 +954,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'actuator_force': d.actuator_force.shape, 'actuator_length': d._impl.actuator_length.shape, 'actuator_moment': d._impl.actuator_moment.shape, - 'actuator_trntype_body_ncon': d._impl.actuator_trntype_body_ncon.shape, 'actuator_velocity': d._impl.actuator_velocity.shape, 'cacc': d._impl.cacc.shape, 'cam_xmat': d.cam_xmat.shape, @@ -1040,40 +970,16 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'ctrl': d.ctrl.shape, 'cvel': d.cvel.shape, 'energy': d._impl.energy.shape, - 'epa_face': d._impl.epa_face.shape, - 'epa_horizon': d._impl.epa_horizon.shape, - 'epa_index': d._impl.epa_index.shape, - 'epa_map': d._impl.epa_map.shape, - 'epa_norm2': d._impl.epa_norm2.shape, - 'epa_pr': d._impl.epa_pr.shape, - 'epa_vert': d._impl.epa_vert.shape, - 'epa_vert1': d._impl.epa_vert1.shape, - 'epa_vert2': d._impl.epa_vert2.shape, - 'epa_vert_index1': d._impl.epa_vert_index1.shape, - 'epa_vert_index2': d._impl.epa_vert_index2.shape, 'eq_active': d.eq_active.shape, 'flexedge_length': d._impl.flexedge_length.shape, 'flexedge_velocity': d._impl.flexedge_velocity.shape, 'flexvert_xpos': d._impl.flexvert_xpos.shape, - 'fluid_applied': d._impl.fluid_applied.shape, - 'geom_skip': d._impl.geom_skip.shape, 'geom_xmat': d.geom_xmat.shape, 'geom_xpos': d.geom_xpos.shape, 'light_xdir': d._impl.light_xdir.shape, 'light_xpos': d._impl.light_xpos.shape, 'mocap_pos': d.mocap_pos.shape, 'mocap_quat': d.mocap_quat.shape, - 'multiccd_clipped': d._impl.multiccd_clipped.shape, - 'multiccd_endvert': d._impl.multiccd_endvert.shape, - 'multiccd_face1': d._impl.multiccd_face1.shape, - 'multiccd_face2': d._impl.multiccd_face2.shape, - 'multiccd_idx1': d._impl.multiccd_idx1.shape, - 'multiccd_idx2': d._impl.multiccd_idx2.shape, - 'multiccd_n1': d._impl.multiccd_n1.shape, - 'multiccd_n2': d._impl.multiccd_n2.shape, - 'multiccd_pdist': d._impl.multiccd_pdist.shape, - 'multiccd_pnormal': d._impl.multiccd_pnormal.shape, - 'multiccd_polygon': d._impl.multiccd_polygon.shape, 'nacon': d._impl.nacon.shape, 'ncollision': d._impl.ncollision.shape, 'ne': d._impl.ne.shape, @@ -1103,20 +1009,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'qfrc_spring': d._impl.qfrc_spring.shape, 'qpos': d.qpos.shape, 'qvel': d.qvel.shape, - 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape, - 'sap_projection_lower': d._impl.sap_projection_lower.shape, - 'sap_projection_upper': d._impl.sap_projection_upper.shape, - 'sap_range': d._impl.sap_range.shape, - 'sap_segment_index': d._impl.sap_segment_index.shape, - 'sap_sort_index': d._impl.sap_sort_index.shape, - 'sensor_contact_criteria': d._impl.sensor_contact_criteria.shape, - 'sensor_contact_direction': d._impl.sensor_contact_direction.shape, - 'sensor_contact_matchid': d._impl.sensor_contact_matchid.shape, - 'sensor_contact_nmatch': d._impl.sensor_contact_nmatch.shape, - 'sensor_rangefinder_dist': d._impl.sensor_rangefinder_dist.shape, - 'sensor_rangefinder_geomid': d._impl.sensor_rangefinder_geomid.shape, - 'sensor_rangefinder_pnt': d._impl.sensor_rangefinder_pnt.shape, - 'sensor_rangefinder_vec': d._impl.sensor_rangefinder_vec.shape, 'sensordata': d.sensordata.shape, 'site_xmat': d.site_xmat.shape, 'site_xpos': d.site_xpos.shape, @@ -1126,15 +1018,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'subtree_com': d.subtree_com.shape, 'subtree_linvel': d._impl.subtree_linvel.shape, 'ten_J': d._impl.ten_J.shape, - 'ten_Jdot': d._impl.ten_Jdot.shape, - 'ten_actfrc': d._impl.ten_actfrc.shape, - 'ten_bias_coef': d._impl.ten_bias_coef.shape, 'ten_length': d.ten_length.shape, 'ten_velocity': d._impl.ten_velocity.shape, 'ten_wrapadr': d._impl.ten_wrapadr.shape, 'ten_wrapnum': d._impl.ten_wrapnum.shape, 'time': d.time.shape, - 'wrap_geom_xpos': d._impl.wrap_geom_xpos.shape, 'wrap_obj': d._impl.wrap_obj.shape, 'wrap_xpos': d._impl.wrap_xpos.shape, 'xanchor': d.xanchor.shape, @@ -1151,11 +1039,13 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__frame': d._impl.contact__frame.shape, 'contact__friction': d._impl.contact__friction.shape, 'contact__geom': d._impl.contact__geom.shape, + 'contact__geomcollisionid': d._impl.contact__geomcollisionid.shape, 'contact__includemargin': d._impl.contact__includemargin.shape, 'contact__pos': d._impl.contact__pos.shape, 'contact__solimp': d._impl.contact__solimp.shape, 'contact__solref': d._impl.contact__solref.shape, 'contact__solreffriction': d._impl.contact__solreffriction.shape, + 'contact__type': d._impl.contact__type.shape, 'contact__worldid': d._impl.contact__worldid.shape, 'efc__D': d._impl.efc__D.shape, 'efc__J': d._impl.efc__J.shape, @@ -1168,7 +1058,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'efc__cholesky_L_tmp': d._impl.efc__cholesky_L_tmp.shape, 'efc__cholesky_y_tmp': d._impl.efc__cholesky_y_tmp.shape, 'efc__cost': d._impl.efc__cost.shape, - 'efc__cost_candidate': d._impl.efc__cost_candidate.shape, 'efc__done': d._impl.efc__done.shape, 'efc__force': d._impl.efc__force.shape, 'efc__frictionloss': d._impl.efc__frictionloss.shape, @@ -1194,7 +1083,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _forward_shim, - num_outputs=173, + num_outputs=131, output_dims=output_dims, vmap_method=None, in_out_argnames={ @@ -1203,7 +1092,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'actuator_force', 'actuator_length', 'actuator_moment', - 'actuator_trntype_body_ncon', 'actuator_velocity', 'cacc', 'cam_xmat', @@ -1220,40 +1108,16 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'ctrl', 'cvel', 'energy', - 'epa_face', - 'epa_horizon', - 'epa_index', - 'epa_map', - 'epa_norm2', - 'epa_pr', - 'epa_vert', - 'epa_vert1', - 'epa_vert2', - 'epa_vert_index1', - 'epa_vert_index2', 'eq_active', 'flexedge_length', 'flexedge_velocity', 'flexvert_xpos', - 'fluid_applied', - 'geom_skip', 'geom_xmat', 'geom_xpos', 'light_xdir', 'light_xpos', 'mocap_pos', 'mocap_quat', - 'multiccd_clipped', - 'multiccd_endvert', - 'multiccd_face1', - 'multiccd_face2', - 'multiccd_idx1', - 'multiccd_idx2', - 'multiccd_n1', - 'multiccd_n2', - 'multiccd_pdist', - 'multiccd_pnormal', - 'multiccd_polygon', 'nacon', 'ncollision', 'ne', @@ -1283,20 +1147,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'qfrc_spring', 'qpos', 'qvel', - 'sap_cumulative_sum', - 'sap_projection_lower', - 'sap_projection_upper', - 'sap_range', - 'sap_segment_index', - 'sap_sort_index', - 'sensor_contact_criteria', - 'sensor_contact_direction', - 'sensor_contact_matchid', - 'sensor_contact_nmatch', - 'sensor_rangefinder_dist', - 'sensor_rangefinder_geomid', - 'sensor_rangefinder_pnt', - 'sensor_rangefinder_vec', 'sensordata', 'site_xmat', 'site_xpos', @@ -1306,15 +1156,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'subtree_com', 'subtree_linvel', 'ten_J', - 'ten_Jdot', - 'ten_actfrc', - 'ten_bias_coef', 'ten_length', 'ten_velocity', 'ten_wrapadr', 'ten_wrapnum', 'time', - 'wrap_geom_xpos', 'wrap_obj', 'wrap_xpos', 'xanchor', @@ -1331,11 +1177,13 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'contact__frame', 'contact__friction', 'contact__geom', + 'contact__geomcollisionid', 'contact__includemargin', 'contact__pos', 'contact__solimp', 'contact__solref', 'contact__solreffriction', + 'contact__type', 'contact__worldid', 'efc__D', 'efc__J', @@ -1348,7 +1196,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'efc__cholesky_L_tmp', 'efc__cholesky_y_tmp', 'efc__cost', - 'efc__cost_candidate', 'efc__done', 'efc__force', 'efc__frictionloss', @@ -1413,6 +1260,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.body_jntadr, m.body_jntnum, m.body_mass, + m.body_mocapid, m.body_parentid, m.body_pos, m.body_quat, @@ -1545,6 +1393,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.mocap_bodyid, m.nC, m.na, + m._impl.nacttrnbody, m.nbody, m.ncam, m.neq, @@ -1555,14 +1404,19 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.ngravcomp, m.njnt, m.nlight, - m._impl.nlsp, + m._impl.nmaxmeshdeg, + m._impl.nmaxpolygon, m.nmeshface, m.nmocap, + m._impl.nrangefinder, + m._impl.nsensorcollision, + m._impl.nsensorcontact, m._impl.nsensortaxel, m.nsite, m.ntendon, m.nu, m.nv, + m.nwrap, m._impl.nxn_geom_pair_filtered, m._impl.nxn_pairid, m._impl.nxn_pairid_filtered, @@ -1591,6 +1445,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m._impl.sensor_acc_adr, m.sensor_adr, m._impl.sensor_adr_to_contact_adr, + m._impl.sensor_collision_start_adr, m._impl.sensor_contact_adr, m.sensor_cutoff, m.sensor_datatype, @@ -1619,7 +1474,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): m.site_quat, m.site_size, m.site_type, - m._impl.subtree_mass, m._impl.taxel_sensorid, m._impl.taxel_vertadr, m.tendon_actfrclimited, @@ -1687,7 +1541,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d.actuator_force, d._impl.actuator_length, d._impl.actuator_moment, - d._impl.actuator_trntype_body_ncon, d._impl.actuator_velocity, d._impl.cacc, d.cam_xmat, @@ -1704,40 +1557,16 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d.ctrl, d.cvel, d._impl.energy, - d._impl.epa_face, - d._impl.epa_horizon, - d._impl.epa_index, - d._impl.epa_map, - d._impl.epa_norm2, - d._impl.epa_pr, - d._impl.epa_vert, - d._impl.epa_vert1, - d._impl.epa_vert2, - d._impl.epa_vert_index1, - d._impl.epa_vert_index2, d.eq_active, d._impl.flexedge_length, d._impl.flexedge_velocity, d._impl.flexvert_xpos, - d._impl.fluid_applied, - d._impl.geom_skip, d.geom_xmat, d.geom_xpos, d._impl.light_xdir, d._impl.light_xpos, d.mocap_pos, d.mocap_quat, - d._impl.multiccd_clipped, - d._impl.multiccd_endvert, - d._impl.multiccd_face1, - d._impl.multiccd_face2, - d._impl.multiccd_idx1, - d._impl.multiccd_idx2, - d._impl.multiccd_n1, - d._impl.multiccd_n2, - d._impl.multiccd_pdist, - d._impl.multiccd_pnormal, - d._impl.multiccd_polygon, d._impl.nacon, d._impl.ncollision, d._impl.ne, @@ -1767,20 +1596,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.qfrc_spring, d.qpos, d.qvel, - d._impl.sap_cumulative_sum, - d._impl.sap_projection_lower, - d._impl.sap_projection_upper, - d._impl.sap_range, - d._impl.sap_segment_index, - d._impl.sap_sort_index, - d._impl.sensor_contact_criteria, - d._impl.sensor_contact_direction, - d._impl.sensor_contact_matchid, - d._impl.sensor_contact_nmatch, - d._impl.sensor_rangefinder_dist, - d._impl.sensor_rangefinder_geomid, - d._impl.sensor_rangefinder_pnt, - d._impl.sensor_rangefinder_vec, d.sensordata, d.site_xmat, d.site_xpos, @@ -1790,15 +1605,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d.subtree_com, d._impl.subtree_linvel, d._impl.ten_J, - d._impl.ten_Jdot, - d._impl.ten_actfrc, - d._impl.ten_bias_coef, d.ten_length, d._impl.ten_velocity, d._impl.ten_wrapadr, d._impl.ten_wrapnum, d.time, - d._impl.wrap_geom_xpos, d._impl.wrap_obj, d._impl.wrap_xpos, d.xanchor, @@ -1815,11 +1626,13 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.contact__frame, d._impl.contact__friction, d._impl.contact__geom, + d._impl.contact__geomcollisionid, d._impl.contact__includemargin, d._impl.contact__pos, d._impl.contact__solimp, d._impl.contact__solref, d._impl.contact__solreffriction, + d._impl.contact__type, d._impl.contact__worldid, d._impl.efc__D, d._impl.efc__J, @@ -1832,7 +1645,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data): d._impl.efc__cholesky_L_tmp, d._impl.efc__cholesky_y_tmp, d._impl.efc__cost, - d._impl.efc__cost_candidate, d._impl.efc__done, d._impl.efc__force, d._impl.efc__frictionloss, @@ -1862,174 +1674,132 @@ def _forward_jax_impl(m: types.Model, d: types.Data): 'actuator_force': out[2], '_impl.actuator_length': out[3], '_impl.actuator_moment': out[4], - '_impl.actuator_trntype_body_ncon': out[5], - '_impl.actuator_velocity': out[6], - '_impl.cacc': out[7], - 'cam_xmat': out[8], - 'cam_xpos': out[9], - '_impl.cdof': out[10], - '_impl.cdof_dot': out[11], - '_impl.cfrc_ext': out[12], - '_impl.cfrc_int': out[13], - '_impl.cinert': out[14], - '_impl.collision_pair': out[15], - '_impl.collision_pairid': out[16], - '_impl.collision_worldid': out[17], - '_impl.crb': out[18], - 'ctrl': out[19], - 'cvel': out[20], - '_impl.energy': out[21], - '_impl.epa_face': out[22], - '_impl.epa_horizon': out[23], - '_impl.epa_index': out[24], - '_impl.epa_map': out[25], - '_impl.epa_norm2': out[26], - '_impl.epa_pr': out[27], - '_impl.epa_vert': out[28], - '_impl.epa_vert1': out[29], - '_impl.epa_vert2': out[30], - '_impl.epa_vert_index1': out[31], - '_impl.epa_vert_index2': out[32], - 'eq_active': out[33], - '_impl.flexedge_length': out[34], - '_impl.flexedge_velocity': out[35], - '_impl.flexvert_xpos': out[36], - '_impl.fluid_applied': out[37], - '_impl.geom_skip': out[38], - 'geom_xmat': out[39], - 'geom_xpos': out[40], - '_impl.light_xdir': out[41], - '_impl.light_xpos': out[42], - 'mocap_pos': out[43], - 'mocap_quat': out[44], - '_impl.multiccd_clipped': out[45], - '_impl.multiccd_endvert': out[46], - '_impl.multiccd_face1': out[47], - '_impl.multiccd_face2': out[48], - '_impl.multiccd_idx1': out[49], - '_impl.multiccd_idx2': out[50], - '_impl.multiccd_n1': out[51], - '_impl.multiccd_n2': out[52], - '_impl.multiccd_pdist': out[53], - '_impl.multiccd_pnormal': out[54], - '_impl.multiccd_polygon': out[55], - '_impl.nacon': out[56], - '_impl.ncollision': out[57], - '_impl.ne': out[58], - '_impl.ne_connect': out[59], - '_impl.ne_jnt': out[60], - '_impl.ne_ten': out[61], - '_impl.ne_weld': out[62], - '_impl.nefc': out[63], - '_impl.nf': out[64], - '_impl.nl': out[65], - '_impl.nsolving': out[66], - '_impl.qLD': out[67], - '_impl.qLDiagInv': out[68], - '_impl.qM': out[69], - 'qacc': out[70], - 'qacc_smooth': out[71], - 'qacc_warmstart': out[72], - 'qfrc_actuator': out[73], - 'qfrc_applied': out[74], - 'qfrc_bias': out[75], - 'qfrc_constraint': out[76], - '_impl.qfrc_damper': out[77], - 'qfrc_fluid': out[78], - 'qfrc_gravcomp': out[79], - 'qfrc_passive': out[80], - 'qfrc_smooth': out[81], - '_impl.qfrc_spring': out[82], - 'qpos': out[83], - 'qvel': out[84], - '_impl.sap_cumulative_sum': out[85], - '_impl.sap_projection_lower': out[86], - '_impl.sap_projection_upper': out[87], - '_impl.sap_range': out[88], - '_impl.sap_segment_index': out[89], - '_impl.sap_sort_index': out[90], - '_impl.sensor_contact_criteria': out[91], - '_impl.sensor_contact_direction': out[92], - '_impl.sensor_contact_matchid': out[93], - '_impl.sensor_contact_nmatch': out[94], - '_impl.sensor_rangefinder_dist': out[95], - '_impl.sensor_rangefinder_geomid': out[96], - '_impl.sensor_rangefinder_pnt': out[97], - '_impl.sensor_rangefinder_vec': out[98], - 'sensordata': out[99], - 'site_xmat': out[100], - 'site_xpos': out[101], - '_impl.solver_niter': out[102], - '_impl.subtree_angmom': out[103], - '_impl.subtree_bodyvel': out[104], - 'subtree_com': out[105], - '_impl.subtree_linvel': out[106], - '_impl.ten_J': out[107], - '_impl.ten_Jdot': out[108], - '_impl.ten_actfrc': out[109], - '_impl.ten_bias_coef': out[110], - 'ten_length': out[111], - '_impl.ten_velocity': out[112], - '_impl.ten_wrapadr': out[113], - '_impl.ten_wrapnum': out[114], - 'time': out[115], - '_impl.wrap_geom_xpos': out[116], - '_impl.wrap_obj': out[117], - '_impl.wrap_xpos': out[118], - 'xanchor': out[119], - 'xaxis': out[120], - 'xfrc_applied': out[121], - 'ximat': out[122], - 'xipos': out[123], - 'xmat': out[124], - 'xpos': out[125], - 'xquat': out[126], - '_impl.contact__dim': out[127], - '_impl.contact__dist': out[128], - '_impl.contact__efc_address': out[129], - '_impl.contact__frame': out[130], - '_impl.contact__friction': out[131], - '_impl.contact__geom': out[132], - '_impl.contact__includemargin': out[133], - '_impl.contact__pos': out[134], - '_impl.contact__solimp': out[135], - '_impl.contact__solref': out[136], - '_impl.contact__solreffriction': out[137], - '_impl.contact__worldid': out[138], - '_impl.efc__D': out[139], - '_impl.efc__J': out[140], - '_impl.efc__Jaref': out[141], - '_impl.efc__Ma': out[142], - '_impl.efc__Mgrad': out[143], - '_impl.efc__alpha': out[144], - '_impl.efc__aref': out[145], - '_impl.efc__beta': out[146], - '_impl.efc__cholesky_L_tmp': out[147], - '_impl.efc__cholesky_y_tmp': out[148], - '_impl.efc__cost': out[149], - '_impl.efc__cost_candidate': out[150], - '_impl.efc__done': out[151], - '_impl.efc__force': out[152], - '_impl.efc__frictionloss': out[153], - '_impl.efc__gauss': out[154], - '_impl.efc__grad': out[155], - '_impl.efc__grad_dot': out[156], - '_impl.efc__h': out[157], - '_impl.efc__id': out[158], - '_impl.efc__jv': out[159], - '_impl.efc__margin': out[160], - '_impl.efc__mv': out[161], - '_impl.efc__pos': out[162], - '_impl.efc__prev_Mgrad': out[163], - '_impl.efc__prev_cost': out[164], - '_impl.efc__prev_grad': out[165], - '_impl.efc__quad': out[166], - '_impl.efc__quad_gauss': out[167], - '_impl.efc__search': out[168], - '_impl.efc__search_dot': out[169], - '_impl.efc__state': out[170], - '_impl.efc__type': out[171], - '_impl.efc__vel': out[172], + '_impl.actuator_velocity': out[5], + '_impl.cacc': out[6], + 'cam_xmat': out[7], + 'cam_xpos': out[8], + '_impl.cdof': out[9], + '_impl.cdof_dot': out[10], + '_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], + 'ctrl': out[18], + 'cvel': out[19], + '_impl.energy': out[20], + 'eq_active': out[21], + '_impl.flexedge_length': out[22], + '_impl.flexedge_velocity': out[23], + '_impl.flexvert_xpos': out[24], + 'geom_xmat': out[25], + 'geom_xpos': out[26], + '_impl.light_xdir': out[27], + '_impl.light_xpos': out[28], + 'mocap_pos': out[29], + 'mocap_quat': out[30], + '_impl.nacon': out[31], + '_impl.ncollision': out[32], + '_impl.ne': out[33], + '_impl.ne_connect': out[34], + '_impl.ne_jnt': out[35], + '_impl.ne_ten': out[36], + '_impl.ne_weld': out[37], + '_impl.nefc': out[38], + '_impl.nf': out[39], + '_impl.nl': out[40], + '_impl.nsolving': out[41], + '_impl.qLD': out[42], + '_impl.qLDiagInv': out[43], + '_impl.qM': out[44], + 'qacc': out[45], + 'qacc_smooth': out[46], + 'qacc_warmstart': out[47], + 'qfrc_actuator': out[48], + 'qfrc_applied': out[49], + 'qfrc_bias': out[50], + 'qfrc_constraint': out[51], + '_impl.qfrc_damper': out[52], + 'qfrc_fluid': out[53], + 'qfrc_gravcomp': out[54], + 'qfrc_passive': out[55], + 'qfrc_smooth': out[56], + '_impl.qfrc_spring': out[57], + 'qpos': out[58], + 'qvel': out[59], + 'sensordata': out[60], + 'site_xmat': out[61], + 'site_xpos': out[62], + '_impl.solver_niter': out[63], + '_impl.subtree_angmom': out[64], + '_impl.subtree_bodyvel': out[65], + 'subtree_com': out[66], + '_impl.subtree_linvel': out[67], + '_impl.ten_J': out[68], + 'ten_length': out[69], + '_impl.ten_velocity': out[70], + '_impl.ten_wrapadr': out[71], + '_impl.ten_wrapnum': out[72], + 'time': out[73], + '_impl.wrap_obj': out[74], + '_impl.wrap_xpos': out[75], + 'xanchor': out[76], + 'xaxis': out[77], + 'xfrc_applied': out[78], + 'ximat': out[79], + 'xipos': out[80], + 'xmat': out[81], + 'xpos': out[82], + 'xquat': out[83], + '_impl.contact__dim': out[84], + '_impl.contact__dist': out[85], + '_impl.contact__efc_address': out[86], + '_impl.contact__frame': out[87], + '_impl.contact__friction': out[88], + '_impl.contact__geom': out[89], + '_impl.contact__geomcollisionid': out[90], + '_impl.contact__includemargin': out[91], + '_impl.contact__pos': out[92], + '_impl.contact__solimp': out[93], + '_impl.contact__solref': out[94], + '_impl.contact__solreffriction': out[95], + '_impl.contact__type': out[96], + '_impl.contact__worldid': out[97], + '_impl.efc__D': out[98], + '_impl.efc__J': out[99], + '_impl.efc__Jaref': out[100], + '_impl.efc__Ma': out[101], + '_impl.efc__Mgrad': out[102], + '_impl.efc__alpha': out[103], + '_impl.efc__aref': out[104], + '_impl.efc__beta': out[105], + '_impl.efc__cholesky_L_tmp': out[106], + '_impl.efc__cholesky_y_tmp': out[107], + '_impl.efc__cost': out[108], + '_impl.efc__done': out[109], + '_impl.efc__force': out[110], + '_impl.efc__frictionloss': out[111], + '_impl.efc__gauss': out[112], + '_impl.efc__grad': out[113], + '_impl.efc__grad_dot': out[114], + '_impl.efc__h': out[115], + '_impl.efc__id': out[116], + '_impl.efc__jv': out[117], + '_impl.efc__margin': out[118], + '_impl.efc__mv': out[119], + '_impl.efc__pos': out[120], + '_impl.efc__prev_Mgrad': out[121], + '_impl.efc__prev_cost': out[122], + '_impl.efc__prev_grad': out[123], + '_impl.efc__quad': out[124], + '_impl.efc__quad_gauss': out[125], + '_impl.efc__search': out[126], + '_impl.efc__search_dot': out[127], + '_impl.efc__state': out[128], + '_impl.efc__type': out[129], + '_impl.efc__vel': out[130], }) return d @@ -2064,6 +1834,7 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _step_shim( # Model @@ -2106,6 +1877,7 @@ def _step_shim( body_jntadr: wp.array(dtype=int), body_jntnum: wp.array(dtype=int), body_mass: wp.array2d(dtype=float), + body_mocapid: wp.array(dtype=int), body_parentid: wp.array(dtype=int), body_pos: wp.array2d(dtype=wp.vec3), body_quat: wp.array2d(dtype=wp.quat), @@ -2162,7 +1934,7 @@ def _step_shim( flex_vertadr: wp.array(dtype=int), flex_vertbodyid: wp.array(dtype=int), flexedge_length0: wp.array(dtype=float), - geom_aabb: wp.array2d(dtype=wp.vec3), + geom_aabb: wp.array3d(dtype=wp.vec3), geom_bodyid: wp.array(dtype=int), geom_condim: wp.array(dtype=int), geom_dataid: wp.array(dtype=int), @@ -2237,7 +2009,9 @@ def _step_shim( mesh_vertnum: wp.array(dtype=int), mocap_bodyid: wp.array(dtype=int), nC: int, + nM: int, na: int, + nacttrnbody: int, nbody: int, ncam: int, neq: int, @@ -2248,17 +2022,22 @@ def _step_shim( ngravcomp: int, njnt: int, nlight: int, - nlsp: int, + nmaxmeshdeg: int, + nmaxpolygon: int, nmeshface: int, nmocap: int, + nrangefinder: int, + nsensorcollision: int, + nsensorcontact: int, nsensortaxel: int, nsite: int, ntendon: int, nu: int, nv: int, + nwrap: int, nxn_geom_pair_filtered: wp.array(dtype=wp.vec2i), - nxn_pairid: wp.array(dtype=int), - nxn_pairid_filtered: wp.array(dtype=int), + nxn_pairid: wp.array(dtype=wp.vec2i), + nxn_pairid_filtered: wp.array(dtype=wp.vec2i), oct_aabb: wp.array2d(dtype=wp.vec3), oct_child: wp.array(dtype=mjwp_types.vec8i), oct_coeff: wp.array(dtype=mjwp_types.vec8f), @@ -2284,6 +2063,7 @@ def _step_shim( sensor_acc_adr: wp.array(dtype=int), sensor_adr: wp.array(dtype=int), sensor_adr_to_contact_adr: wp.array(dtype=int), + sensor_collision_start_adr: wp.array(dtype=int), sensor_contact_adr: wp.array(dtype=int), sensor_cutoff: wp.array(dtype=float), sensor_datatype: wp.array(dtype=int), @@ -2312,7 +2092,6 @@ def _step_shim( site_quat: wp.array2d(dtype=wp.quat), site_size: wp.array(dtype=wp.vec3), site_type: wp.array(dtype=int), - subtree_mass: wp.array2d(dtype=float), taxel_sensorid: wp.array(dtype=int), taxel_vertadr: wp.array(dtype=int), tendon_actfrclimited: wp.array(dtype=bool), @@ -2379,12 +2158,9 @@ def _step_shim( njmax: int, act: wp.array2d(dtype=float), act_dot: wp.array2d(dtype=float), - act_dot_rk: wp.array2d(dtype=float), - act_t0: wp.array2d(dtype=float), actuator_force: wp.array2d(dtype=float), actuator_length: wp.array2d(dtype=float), actuator_moment: wp.array3d(dtype=float), - actuator_trntype_body_ncon: wp.array2d(dtype=int), actuator_velocity: wp.array2d(dtype=float), cacc: wp.array2d(dtype=wp.spatial_vector), cam_xmat: wp.array2d(dtype=wp.mat33), @@ -2395,47 +2171,22 @@ def _step_shim( 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=int), + 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), energy: wp.array(dtype=wp.vec2), - epa_face: wp.array2d(dtype=wp.vec3i), - epa_horizon: wp.array2d(dtype=int), - epa_index: wp.array2d(dtype=int), - epa_map: wp.array2d(dtype=int), - epa_norm2: wp.array2d(dtype=float), - epa_pr: wp.array2d(dtype=wp.vec3), - epa_vert: wp.array2d(dtype=wp.vec3), - epa_vert1: wp.array2d(dtype=wp.vec3), - epa_vert2: wp.array2d(dtype=wp.vec3), - epa_vert_index1: wp.array2d(dtype=int), - epa_vert_index2: wp.array2d(dtype=int), eq_active: wp.array2d(dtype=bool), flexedge_length: wp.array2d(dtype=float), flexedge_velocity: wp.array2d(dtype=float), flexvert_xpos: wp.array2d(dtype=wp.vec3), - fluid_applied: wp.array2d(dtype=wp.spatial_vector), - geom_skip: wp.array(dtype=bool), geom_xmat: wp.array2d(dtype=wp.mat33), geom_xpos: wp.array2d(dtype=wp.vec3), - inverse_mul_m_skip: wp.array(dtype=bool), light_xdir: wp.array2d(dtype=wp.vec3), light_xpos: wp.array2d(dtype=wp.vec3), mocap_pos: wp.array2d(dtype=wp.vec3), mocap_quat: wp.array2d(dtype=wp.quat), - multiccd_clipped: wp.array2d(dtype=wp.vec3), - multiccd_endvert: wp.array2d(dtype=wp.vec3), - multiccd_face1: wp.array2d(dtype=wp.vec3), - multiccd_face2: wp.array2d(dtype=wp.vec3), - multiccd_idx1: wp.array2d(dtype=int), - multiccd_idx2: wp.array2d(dtype=int), - multiccd_n1: wp.array2d(dtype=wp.vec3), - multiccd_n2: wp.array2d(dtype=wp.vec3), - multiccd_pdist: wp.array2d(dtype=float), - multiccd_pnormal: wp.array2d(dtype=wp.vec3), - multiccd_polygon: wp.array2d(dtype=wp.vec3), nacon: wp.array(dtype=int), ncollision: wp.array(dtype=int), ne: wp.array(dtype=int), @@ -2448,14 +2199,9 @@ def _step_shim( nl: wp.array(dtype=int), nsolving: wp.array(dtype=int), qLD: wp.array3d(dtype=float), - qLD_integration: wp.array3d(dtype=float), qLDiagInv: wp.array2d(dtype=float), - qLDiagInv_integration: wp.array2d(dtype=float), qM: wp.array3d(dtype=float), - qM_integration: wp.array3d(dtype=float), qacc: wp.array2d(dtype=float), - qacc_integration: wp.array2d(dtype=float), - qacc_rk: wp.array2d(dtype=float), qacc_smooth: wp.array2d(dtype=float), qacc_warmstart: wp.array2d(dtype=float), qfrc_actuator: wp.array2d(dtype=float), @@ -2465,29 +2211,11 @@ def _step_shim( qfrc_damper: wp.array2d(dtype=float), qfrc_fluid: wp.array2d(dtype=float), qfrc_gravcomp: wp.array2d(dtype=float), - qfrc_integration: wp.array2d(dtype=float), qfrc_passive: wp.array2d(dtype=float), qfrc_smooth: wp.array2d(dtype=float), qfrc_spring: wp.array2d(dtype=float), qpos: wp.array2d(dtype=float), - qpos_t0: wp.array2d(dtype=float), qvel: wp.array2d(dtype=float), - qvel_rk: wp.array2d(dtype=float), - qvel_t0: wp.array2d(dtype=float), - sap_cumulative_sum: wp.array2d(dtype=int), - sap_projection_lower: wp.array3d(dtype=float), - sap_projection_upper: wp.array2d(dtype=float), - sap_range: wp.array2d(dtype=int), - sap_segment_index: wp.array2d(dtype=int), - sap_sort_index: wp.array3d(dtype=int), - sensor_contact_criteria: wp.array3d(dtype=float), - sensor_contact_direction: wp.array3d(dtype=float), - sensor_contact_matchid: wp.array3d(dtype=int), - sensor_contact_nmatch: wp.array2d(dtype=int), - sensor_rangefinder_dist: wp.array2d(dtype=float), - sensor_rangefinder_geomid: wp.array2d(dtype=int), - sensor_rangefinder_pnt: wp.array2d(dtype=wp.vec3), - sensor_rangefinder_vec: wp.array2d(dtype=wp.vec3), sensordata: wp.array2d(dtype=float), site_xmat: wp.array2d(dtype=wp.mat33), site_xpos: wp.array2d(dtype=wp.vec3), @@ -2497,15 +2225,11 @@ def _step_shim( subtree_com: wp.array2d(dtype=wp.vec3), subtree_linvel: wp.array2d(dtype=wp.vec3), ten_J: wp.array3d(dtype=float), - ten_Jdot: wp.array3d(dtype=float), - ten_actfrc: wp.array2d(dtype=float), - ten_bias_coef: wp.array2d(dtype=float), ten_length: wp.array2d(dtype=float), ten_velocity: wp.array2d(dtype=float), ten_wrapadr: wp.array2d(dtype=int), ten_wrapnum: wp.array2d(dtype=int), time: wp.array(dtype=float), - wrap_geom_xpos: wp.array2d(dtype=wp.spatial_vector), wrap_obj: wp.array2d(dtype=wp.vec2i), wrap_xpos: wp.array2d(dtype=wp.spatial_vector), xanchor: wp.array2d(dtype=wp.vec3), @@ -2522,11 +2246,13 @@ def _step_shim( contact__frame: wp.array(dtype=wp.mat33), contact__friction: wp.array(dtype=mjwp_types.vec5), contact__geom: wp.array(dtype=wp.vec2i), + contact__geomcollisionid: wp.array(dtype=int), contact__includemargin: wp.array(dtype=float), contact__pos: wp.array(dtype=wp.vec3), contact__solimp: wp.array(dtype=mjwp_types.vec5), contact__solref: wp.array(dtype=wp.vec2), contact__solreffriction: wp.array(dtype=wp.vec2), + contact__type: wp.array(dtype=int), contact__worldid: wp.array(dtype=int), efc__D: wp.array2d(dtype=float), efc__J: wp.array3d(dtype=float), @@ -2539,7 +2265,6 @@ def _step_shim( efc__cholesky_L_tmp: wp.array3d(dtype=float), efc__cholesky_y_tmp: wp.array2d(dtype=float), efc__cost: wp.array(dtype=float), - efc__cost_candidate: wp.array2d(dtype=float), efc__done: wp.array(dtype=bool), efc__force: wp.array2d(dtype=float), efc__frictionloss: wp.array2d(dtype=float), @@ -2605,6 +2330,7 @@ def _step_shim( _m.body_jntadr = body_jntadr _m.body_jntnum = body_jntnum _m.body_mass = body_mass + _m.body_mocapid = body_mocapid _m.body_parentid = body_parentid _m.body_pos = body_pos _m.body_quat = body_quat @@ -2736,7 +2462,9 @@ def _step_shim( _m.mesh_vertnum = mesh_vertnum _m.mocap_bodyid = mocap_bodyid _m.nC = nC + _m.nM = nM _m.na = na + _m.nacttrnbody = nacttrnbody _m.nbody = nbody _m.ncam = ncam _m.neq = neq @@ -2747,14 +2475,19 @@ def _step_shim( _m.ngravcomp = ngravcomp _m.njnt = njnt _m.nlight = nlight - _m.nlsp = nlsp + _m.nmaxmeshdeg = nmaxmeshdeg + _m.nmaxpolygon = nmaxpolygon _m.nmeshface = nmeshface _m.nmocap = nmocap + _m.nrangefinder = nrangefinder + _m.nsensorcollision = nsensorcollision + _m.nsensorcontact = nsensorcontact _m.nsensortaxel = nsensortaxel _m.nsite = nsite _m.ntendon = ntendon _m.nu = nu _m.nv = nv + _m.nwrap = nwrap _m.nxn_geom_pair_filtered = nxn_geom_pair_filtered _m.nxn_pairid = nxn_pairid _m.nxn_pairid_filtered = nxn_pairid_filtered @@ -2813,6 +2546,7 @@ def _step_shim( _m.sensor_acc_adr = sensor_acc_adr _m.sensor_adr = sensor_adr _m.sensor_adr_to_contact_adr = sensor_adr_to_contact_adr + _m.sensor_collision_start_adr = sensor_collision_start_adr _m.sensor_contact_adr = sensor_contact_adr _m.sensor_cutoff = sensor_cutoff _m.sensor_datatype = sensor_datatype @@ -2842,7 +2576,6 @@ def _step_shim( _m.site_size = site_size _m.site_type = site_type _m.stat.meaninertia = stat__meaninertia - _m.subtree_mass = subtree_mass _m.taxel_sensorid = taxel_sensorid _m.taxel_vertadr = taxel_vertadr _m.tendon_actfrclimited = tendon_actfrclimited @@ -2875,12 +2608,9 @@ def _step_shim( _m.wrap_type = wrap_type _d.act = act _d.act_dot = act_dot - _d.act_dot_rk = act_dot_rk - _d.act_t0 = act_t0 _d.actuator_force = actuator_force _d.actuator_length = actuator_length _d.actuator_moment = actuator_moment - _d.actuator_trntype_body_ncon = actuator_trntype_body_ncon _d.actuator_velocity = actuator_velocity _d.cacc = cacc _d.cam_xmat = cam_xmat @@ -2899,11 +2629,13 @@ def _step_shim( _d.contact.frame = contact__frame _d.contact.friction = contact__friction _d.contact.geom = contact__geom + _d.contact.geomcollisionid = contact__geomcollisionid _d.contact.includemargin = contact__includemargin _d.contact.pos = contact__pos _d.contact.solimp = contact__solimp _d.contact.solref = contact__solref _d.contact.solreffriction = contact__solreffriction + _d.contact.type = contact__type _d.contact.worldid = contact__worldid _d.crb = crb _d.ctrl = ctrl @@ -2919,7 +2651,6 @@ def _step_shim( _d.efc.cholesky_L_tmp = efc__cholesky_L_tmp _d.efc.cholesky_y_tmp = efc__cholesky_y_tmp _d.efc.cost = efc__cost - _d.efc.cost_candidate = efc__cost_candidate _d.efc.done = efc__done _d.efc.force = efc__force _d.efc.frictionloss = efc__frictionloss @@ -2943,41 +2674,16 @@ def _step_shim( _d.efc.type = efc__type _d.efc.vel = efc__vel _d.energy = energy - _d.epa_face = epa_face - _d.epa_horizon = epa_horizon - _d.epa_index = epa_index - _d.epa_map = epa_map - _d.epa_norm2 = epa_norm2 - _d.epa_pr = epa_pr - _d.epa_vert = epa_vert - _d.epa_vert1 = epa_vert1 - _d.epa_vert2 = epa_vert2 - _d.epa_vert_index1 = epa_vert_index1 - _d.epa_vert_index2 = epa_vert_index2 _d.eq_active = eq_active _d.flexedge_length = flexedge_length _d.flexedge_velocity = flexedge_velocity _d.flexvert_xpos = flexvert_xpos - _d.fluid_applied = fluid_applied - _d.geom_skip = geom_skip _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos - _d.inverse_mul_m_skip = inverse_mul_m_skip _d.light_xdir = light_xdir _d.light_xpos = light_xpos _d.mocap_pos = mocap_pos _d.mocap_quat = mocap_quat - _d.multiccd_clipped = multiccd_clipped - _d.multiccd_endvert = multiccd_endvert - _d.multiccd_face1 = multiccd_face1 - _d.multiccd_face2 = multiccd_face2 - _d.multiccd_idx1 = multiccd_idx1 - _d.multiccd_idx2 = multiccd_idx2 - _d.multiccd_n1 = multiccd_n1 - _d.multiccd_n2 = multiccd_n2 - _d.multiccd_pdist = multiccd_pdist - _d.multiccd_pnormal = multiccd_pnormal - _d.multiccd_polygon = multiccd_polygon _d.nacon = nacon _d.naconmax = naconmax _d.ncollision = ncollision @@ -2992,14 +2698,9 @@ def _step_shim( _d.nl = nl _d.nsolving = nsolving _d.qLD = qLD - _d.qLD_integration = qLD_integration _d.qLDiagInv = qLDiagInv - _d.qLDiagInv_integration = qLDiagInv_integration _d.qM = qM - _d.qM_integration = qM_integration _d.qacc = qacc - _d.qacc_integration = qacc_integration - _d.qacc_rk = qacc_rk _d.qacc_smooth = qacc_smooth _d.qacc_warmstart = qacc_warmstart _d.qfrc_actuator = qfrc_actuator @@ -3009,29 +2710,11 @@ def _step_shim( _d.qfrc_damper = qfrc_damper _d.qfrc_fluid = qfrc_fluid _d.qfrc_gravcomp = qfrc_gravcomp - _d.qfrc_integration = qfrc_integration _d.qfrc_passive = qfrc_passive _d.qfrc_smooth = qfrc_smooth _d.qfrc_spring = qfrc_spring _d.qpos = qpos - _d.qpos_t0 = qpos_t0 _d.qvel = qvel - _d.qvel_rk = qvel_rk - _d.qvel_t0 = qvel_t0 - _d.sap_cumulative_sum = sap_cumulative_sum - _d.sap_projection_lower = sap_projection_lower - _d.sap_projection_upper = sap_projection_upper - _d.sap_range = sap_range - _d.sap_segment_index = sap_segment_index - _d.sap_sort_index = sap_sort_index - _d.sensor_contact_criteria = sensor_contact_criteria - _d.sensor_contact_direction = sensor_contact_direction - _d.sensor_contact_matchid = sensor_contact_matchid - _d.sensor_contact_nmatch = sensor_contact_nmatch - _d.sensor_rangefinder_dist = sensor_rangefinder_dist - _d.sensor_rangefinder_geomid = sensor_rangefinder_geomid - _d.sensor_rangefinder_pnt = sensor_rangefinder_pnt - _d.sensor_rangefinder_vec = sensor_rangefinder_vec _d.sensordata = sensordata _d.site_xmat = site_xmat _d.site_xpos = site_xpos @@ -3041,15 +2724,11 @@ def _step_shim( _d.subtree_com = subtree_com _d.subtree_linvel = subtree_linvel _d.ten_J = ten_J - _d.ten_Jdot = ten_Jdot - _d.ten_actfrc = ten_actfrc - _d.ten_bias_coef = ten_bias_coef _d.ten_length = ten_length _d.ten_velocity = ten_velocity _d.ten_wrapadr = ten_wrapadr _d.ten_wrapnum = ten_wrapnum _d.time = time - _d.wrap_geom_xpos = wrap_geom_xpos _d.wrap_obj = wrap_obj _d.wrap_xpos = wrap_xpos _d.xanchor = xanchor @@ -3068,12 +2747,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): output_dims = { 'act': d.act.shape, 'act_dot': d.act_dot.shape, - 'act_dot_rk': d._impl.act_dot_rk.shape, - 'act_t0': d._impl.act_t0.shape, 'actuator_force': d.actuator_force.shape, 'actuator_length': d._impl.actuator_length.shape, 'actuator_moment': d._impl.actuator_moment.shape, - 'actuator_trntype_body_ncon': d._impl.actuator_trntype_body_ncon.shape, 'actuator_velocity': d._impl.actuator_velocity.shape, 'cacc': d._impl.cacc.shape, 'cam_xmat': d.cam_xmat.shape, @@ -3090,41 +2766,16 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'ctrl': d.ctrl.shape, 'cvel': d.cvel.shape, 'energy': d._impl.energy.shape, - 'epa_face': d._impl.epa_face.shape, - 'epa_horizon': d._impl.epa_horizon.shape, - 'epa_index': d._impl.epa_index.shape, - 'epa_map': d._impl.epa_map.shape, - 'epa_norm2': d._impl.epa_norm2.shape, - 'epa_pr': d._impl.epa_pr.shape, - 'epa_vert': d._impl.epa_vert.shape, - 'epa_vert1': d._impl.epa_vert1.shape, - 'epa_vert2': d._impl.epa_vert2.shape, - 'epa_vert_index1': d._impl.epa_vert_index1.shape, - 'epa_vert_index2': d._impl.epa_vert_index2.shape, 'eq_active': d.eq_active.shape, 'flexedge_length': d._impl.flexedge_length.shape, 'flexedge_velocity': d._impl.flexedge_velocity.shape, 'flexvert_xpos': d._impl.flexvert_xpos.shape, - 'fluid_applied': d._impl.fluid_applied.shape, - 'geom_skip': d._impl.geom_skip.shape, 'geom_xmat': d.geom_xmat.shape, 'geom_xpos': d.geom_xpos.shape, - 'inverse_mul_m_skip': d._impl.inverse_mul_m_skip.shape, 'light_xdir': d._impl.light_xdir.shape, 'light_xpos': d._impl.light_xpos.shape, 'mocap_pos': d.mocap_pos.shape, 'mocap_quat': d.mocap_quat.shape, - 'multiccd_clipped': d._impl.multiccd_clipped.shape, - 'multiccd_endvert': d._impl.multiccd_endvert.shape, - 'multiccd_face1': d._impl.multiccd_face1.shape, - 'multiccd_face2': d._impl.multiccd_face2.shape, - 'multiccd_idx1': d._impl.multiccd_idx1.shape, - 'multiccd_idx2': d._impl.multiccd_idx2.shape, - 'multiccd_n1': d._impl.multiccd_n1.shape, - 'multiccd_n2': d._impl.multiccd_n2.shape, - 'multiccd_pdist': d._impl.multiccd_pdist.shape, - 'multiccd_pnormal': d._impl.multiccd_pnormal.shape, - 'multiccd_polygon': d._impl.multiccd_polygon.shape, 'nacon': d._impl.nacon.shape, 'ncollision': d._impl.ncollision.shape, 'ne': d._impl.ne.shape, @@ -3137,14 +2788,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'nl': d._impl.nl.shape, 'nsolving': d._impl.nsolving.shape, 'qLD': d._impl.qLD.shape, - 'qLD_integration': d._impl.qLD_integration.shape, 'qLDiagInv': d._impl.qLDiagInv.shape, - 'qLDiagInv_integration': d._impl.qLDiagInv_integration.shape, 'qM': d._impl.qM.shape, - 'qM_integration': d._impl.qM_integration.shape, 'qacc': d.qacc.shape, - 'qacc_integration': d._impl.qacc_integration.shape, - 'qacc_rk': d._impl.qacc_rk.shape, 'qacc_smooth': d.qacc_smooth.shape, 'qacc_warmstart': d.qacc_warmstart.shape, 'qfrc_actuator': d.qfrc_actuator.shape, @@ -3154,29 +2800,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'qfrc_damper': d._impl.qfrc_damper.shape, 'qfrc_fluid': d.qfrc_fluid.shape, 'qfrc_gravcomp': d.qfrc_gravcomp.shape, - 'qfrc_integration': d._impl.qfrc_integration.shape, 'qfrc_passive': d.qfrc_passive.shape, 'qfrc_smooth': d.qfrc_smooth.shape, 'qfrc_spring': d._impl.qfrc_spring.shape, 'qpos': d.qpos.shape, - 'qpos_t0': d._impl.qpos_t0.shape, 'qvel': d.qvel.shape, - 'qvel_rk': d._impl.qvel_rk.shape, - 'qvel_t0': d._impl.qvel_t0.shape, - 'sap_cumulative_sum': d._impl.sap_cumulative_sum.shape, - 'sap_projection_lower': d._impl.sap_projection_lower.shape, - 'sap_projection_upper': d._impl.sap_projection_upper.shape, - 'sap_range': d._impl.sap_range.shape, - 'sap_segment_index': d._impl.sap_segment_index.shape, - 'sap_sort_index': d._impl.sap_sort_index.shape, - 'sensor_contact_criteria': d._impl.sensor_contact_criteria.shape, - 'sensor_contact_direction': d._impl.sensor_contact_direction.shape, - 'sensor_contact_matchid': d._impl.sensor_contact_matchid.shape, - 'sensor_contact_nmatch': d._impl.sensor_contact_nmatch.shape, - 'sensor_rangefinder_dist': d._impl.sensor_rangefinder_dist.shape, - 'sensor_rangefinder_geomid': d._impl.sensor_rangefinder_geomid.shape, - 'sensor_rangefinder_pnt': d._impl.sensor_rangefinder_pnt.shape, - 'sensor_rangefinder_vec': d._impl.sensor_rangefinder_vec.shape, 'sensordata': d.sensordata.shape, 'site_xmat': d.site_xmat.shape, 'site_xpos': d.site_xpos.shape, @@ -3186,15 +2814,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'subtree_com': d.subtree_com.shape, 'subtree_linvel': d._impl.subtree_linvel.shape, 'ten_J': d._impl.ten_J.shape, - 'ten_Jdot': d._impl.ten_Jdot.shape, - 'ten_actfrc': d._impl.ten_actfrc.shape, - 'ten_bias_coef': d._impl.ten_bias_coef.shape, 'ten_length': d.ten_length.shape, 'ten_velocity': d._impl.ten_velocity.shape, 'ten_wrapadr': d._impl.ten_wrapadr.shape, 'ten_wrapnum': d._impl.ten_wrapnum.shape, 'time': d.time.shape, - 'wrap_geom_xpos': d._impl.wrap_geom_xpos.shape, 'wrap_obj': d._impl.wrap_obj.shape, 'wrap_xpos': d._impl.wrap_xpos.shape, 'xanchor': d.xanchor.shape, @@ -3211,11 +2835,13 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__frame': d._impl.contact__frame.shape, 'contact__friction': d._impl.contact__friction.shape, 'contact__geom': d._impl.contact__geom.shape, + 'contact__geomcollisionid': d._impl.contact__geomcollisionid.shape, 'contact__includemargin': d._impl.contact__includemargin.shape, 'contact__pos': d._impl.contact__pos.shape, 'contact__solimp': d._impl.contact__solimp.shape, 'contact__solref': d._impl.contact__solref.shape, 'contact__solreffriction': d._impl.contact__solreffriction.shape, + 'contact__type': d._impl.contact__type.shape, 'contact__worldid': d._impl.contact__worldid.shape, 'efc__D': d._impl.efc__D.shape, 'efc__J': d._impl.efc__J.shape, @@ -3228,7 +2854,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'efc__cholesky_L_tmp': d._impl.efc__cholesky_L_tmp.shape, 'efc__cholesky_y_tmp': d._impl.efc__cholesky_y_tmp.shape, 'efc__cost': d._impl.efc__cost.shape, - 'efc__cost_candidate': d._impl.efc__cost_candidate.shape, 'efc__done': d._impl.efc__done.shape, 'efc__force': d._impl.efc__force.shape, 'efc__frictionloss': d._impl.efc__frictionloss.shape, @@ -3254,18 +2879,15 @@ def _step_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _step_shim, - num_outputs=185, + num_outputs=131, output_dims=output_dims, vmap_method=None, in_out_argnames={ 'act', 'act_dot', - 'act_dot_rk', - 'act_t0', 'actuator_force', 'actuator_length', 'actuator_moment', - 'actuator_trntype_body_ncon', 'actuator_velocity', 'cacc', 'cam_xmat', @@ -3282,41 +2904,16 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'ctrl', 'cvel', 'energy', - 'epa_face', - 'epa_horizon', - 'epa_index', - 'epa_map', - 'epa_norm2', - 'epa_pr', - 'epa_vert', - 'epa_vert1', - 'epa_vert2', - 'epa_vert_index1', - 'epa_vert_index2', 'eq_active', 'flexedge_length', 'flexedge_velocity', 'flexvert_xpos', - 'fluid_applied', - 'geom_skip', 'geom_xmat', 'geom_xpos', - 'inverse_mul_m_skip', 'light_xdir', 'light_xpos', 'mocap_pos', 'mocap_quat', - 'multiccd_clipped', - 'multiccd_endvert', - 'multiccd_face1', - 'multiccd_face2', - 'multiccd_idx1', - 'multiccd_idx2', - 'multiccd_n1', - 'multiccd_n2', - 'multiccd_pdist', - 'multiccd_pnormal', - 'multiccd_polygon', 'nacon', 'ncollision', 'ne', @@ -3329,14 +2926,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'nl', 'nsolving', 'qLD', - 'qLD_integration', 'qLDiagInv', - 'qLDiagInv_integration', 'qM', - 'qM_integration', 'qacc', - 'qacc_integration', - 'qacc_rk', 'qacc_smooth', 'qacc_warmstart', 'qfrc_actuator', @@ -3346,29 +2938,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'qfrc_damper', 'qfrc_fluid', 'qfrc_gravcomp', - 'qfrc_integration', 'qfrc_passive', 'qfrc_smooth', 'qfrc_spring', 'qpos', - 'qpos_t0', 'qvel', - 'qvel_rk', - 'qvel_t0', - 'sap_cumulative_sum', - 'sap_projection_lower', - 'sap_projection_upper', - 'sap_range', - 'sap_segment_index', - 'sap_sort_index', - 'sensor_contact_criteria', - 'sensor_contact_direction', - 'sensor_contact_matchid', - 'sensor_contact_nmatch', - 'sensor_rangefinder_dist', - 'sensor_rangefinder_geomid', - 'sensor_rangefinder_pnt', - 'sensor_rangefinder_vec', 'sensordata', 'site_xmat', 'site_xpos', @@ -3378,15 +2952,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'subtree_com', 'subtree_linvel', 'ten_J', - 'ten_Jdot', - 'ten_actfrc', - 'ten_bias_coef', 'ten_length', 'ten_velocity', 'ten_wrapadr', 'ten_wrapnum', 'time', - 'wrap_geom_xpos', 'wrap_obj', 'wrap_xpos', 'xanchor', @@ -3403,11 +2973,13 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'contact__frame', 'contact__friction', 'contact__geom', + 'contact__geomcollisionid', 'contact__includemargin', 'contact__pos', 'contact__solimp', 'contact__solref', 'contact__solreffriction', + 'contact__type', 'contact__worldid', 'efc__D', 'efc__J', @@ -3420,7 +2992,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): 'efc__cholesky_L_tmp', 'efc__cholesky_y_tmp', 'efc__cost', - 'efc__cost_candidate', 'efc__done', 'efc__force', 'efc__frictionloss', @@ -3485,6 +3056,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.body_jntadr, m.body_jntnum, m.body_mass, + m.body_mocapid, m.body_parentid, m.body_pos, m.body_quat, @@ -3616,7 +3188,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.mesh_vertnum, m._impl.mocap_bodyid, m.nC, + m.nM, m.na, + m._impl.nacttrnbody, m.nbody, m.ncam, m.neq, @@ -3627,14 +3201,19 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.ngravcomp, m.njnt, m.nlight, - m._impl.nlsp, + m._impl.nmaxmeshdeg, + m._impl.nmaxpolygon, m.nmeshface, m.nmocap, + m._impl.nrangefinder, + m._impl.nsensorcollision, + m._impl.nsensorcontact, m._impl.nsensortaxel, m.nsite, m.ntendon, m.nu, m.nv, + m.nwrap, m._impl.nxn_geom_pair_filtered, m._impl.nxn_pairid, m._impl.nxn_pairid_filtered, @@ -3663,6 +3242,7 @@ def _step_jax_impl(m: types.Model, d: types.Data): m._impl.sensor_acc_adr, m.sensor_adr, m._impl.sensor_adr_to_contact_adr, + m._impl.sensor_collision_start_adr, m._impl.sensor_contact_adr, m.sensor_cutoff, m.sensor_datatype, @@ -3691,7 +3271,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): m.site_quat, m.site_size, m.site_type, - m._impl.subtree_mass, m._impl.taxel_sensorid, m._impl.taxel_vertadr, m.tendon_actfrclimited, @@ -3757,12 +3336,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.njmax, d.act, d.act_dot, - d._impl.act_dot_rk, - d._impl.act_t0, d.actuator_force, d._impl.actuator_length, d._impl.actuator_moment, - d._impl.actuator_trntype_body_ncon, d._impl.actuator_velocity, d._impl.cacc, d.cam_xmat, @@ -3779,41 +3355,16 @@ def _step_jax_impl(m: types.Model, d: types.Data): d.ctrl, d.cvel, d._impl.energy, - d._impl.epa_face, - d._impl.epa_horizon, - d._impl.epa_index, - d._impl.epa_map, - d._impl.epa_norm2, - d._impl.epa_pr, - d._impl.epa_vert, - d._impl.epa_vert1, - d._impl.epa_vert2, - d._impl.epa_vert_index1, - d._impl.epa_vert_index2, d.eq_active, d._impl.flexedge_length, d._impl.flexedge_velocity, d._impl.flexvert_xpos, - d._impl.fluid_applied, - d._impl.geom_skip, d.geom_xmat, d.geom_xpos, - d._impl.inverse_mul_m_skip, d._impl.light_xdir, d._impl.light_xpos, d.mocap_pos, d.mocap_quat, - d._impl.multiccd_clipped, - d._impl.multiccd_endvert, - d._impl.multiccd_face1, - d._impl.multiccd_face2, - d._impl.multiccd_idx1, - d._impl.multiccd_idx2, - d._impl.multiccd_n1, - d._impl.multiccd_n2, - d._impl.multiccd_pdist, - d._impl.multiccd_pnormal, - d._impl.multiccd_polygon, d._impl.nacon, d._impl.ncollision, d._impl.ne, @@ -3826,14 +3377,9 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.nl, d._impl.nsolving, d._impl.qLD, - d._impl.qLD_integration, d._impl.qLDiagInv, - d._impl.qLDiagInv_integration, d._impl.qM, - d._impl.qM_integration, d.qacc, - d._impl.qacc_integration, - d._impl.qacc_rk, d.qacc_smooth, d.qacc_warmstart, d.qfrc_actuator, @@ -3843,29 +3389,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.qfrc_damper, d.qfrc_fluid, d.qfrc_gravcomp, - d._impl.qfrc_integration, d.qfrc_passive, d.qfrc_smooth, d._impl.qfrc_spring, d.qpos, - d._impl.qpos_t0, d.qvel, - d._impl.qvel_rk, - d._impl.qvel_t0, - d._impl.sap_cumulative_sum, - d._impl.sap_projection_lower, - d._impl.sap_projection_upper, - d._impl.sap_range, - d._impl.sap_segment_index, - d._impl.sap_sort_index, - d._impl.sensor_contact_criteria, - d._impl.sensor_contact_direction, - d._impl.sensor_contact_matchid, - d._impl.sensor_contact_nmatch, - d._impl.sensor_rangefinder_dist, - d._impl.sensor_rangefinder_geomid, - d._impl.sensor_rangefinder_pnt, - d._impl.sensor_rangefinder_vec, d.sensordata, d.site_xmat, d.site_xpos, @@ -3875,15 +3403,11 @@ def _step_jax_impl(m: types.Model, d: types.Data): d.subtree_com, d._impl.subtree_linvel, d._impl.ten_J, - d._impl.ten_Jdot, - d._impl.ten_actfrc, - d._impl.ten_bias_coef, d.ten_length, d._impl.ten_velocity, d._impl.ten_wrapadr, d._impl.ten_wrapnum, d.time, - d._impl.wrap_geom_xpos, d._impl.wrap_obj, d._impl.wrap_xpos, d.xanchor, @@ -3900,11 +3424,13 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.contact__frame, d._impl.contact__friction, d._impl.contact__geom, + d._impl.contact__geomcollisionid, d._impl.contact__includemargin, d._impl.contact__pos, d._impl.contact__solimp, d._impl.contact__solref, d._impl.contact__solreffriction, + d._impl.contact__type, d._impl.contact__worldid, d._impl.efc__D, d._impl.efc__J, @@ -3917,7 +3443,6 @@ def _step_jax_impl(m: types.Model, d: types.Data): d._impl.efc__cholesky_L_tmp, d._impl.efc__cholesky_y_tmp, d._impl.efc__cost, - d._impl.efc__cost_candidate, d._impl.efc__done, d._impl.efc__force, d._impl.efc__frictionloss, @@ -3944,189 +3469,135 @@ def _step_jax_impl(m: types.Model, d: types.Data): d = d.tree_replace({ 'act': out[0], 'act_dot': out[1], - '_impl.act_dot_rk': out[2], - '_impl.act_t0': out[3], - 'actuator_force': out[4], - '_impl.actuator_length': out[5], - '_impl.actuator_moment': out[6], - '_impl.actuator_trntype_body_ncon': out[7], - '_impl.actuator_velocity': out[8], - '_impl.cacc': out[9], - 'cam_xmat': out[10], - 'cam_xpos': out[11], - '_impl.cdof': out[12], - '_impl.cdof_dot': out[13], - '_impl.cfrc_ext': out[14], - '_impl.cfrc_int': out[15], - '_impl.cinert': out[16], - '_impl.collision_pair': out[17], - '_impl.collision_pairid': out[18], - '_impl.collision_worldid': out[19], - '_impl.crb': out[20], - 'ctrl': out[21], - 'cvel': out[22], - '_impl.energy': out[23], - '_impl.epa_face': out[24], - '_impl.epa_horizon': out[25], - '_impl.epa_index': out[26], - '_impl.epa_map': out[27], - '_impl.epa_norm2': out[28], - '_impl.epa_pr': out[29], - '_impl.epa_vert': out[30], - '_impl.epa_vert1': out[31], - '_impl.epa_vert2': out[32], - '_impl.epa_vert_index1': out[33], - '_impl.epa_vert_index2': out[34], - 'eq_active': out[35], - '_impl.flexedge_length': out[36], - '_impl.flexedge_velocity': out[37], - '_impl.flexvert_xpos': out[38], - '_impl.fluid_applied': out[39], - '_impl.geom_skip': out[40], - 'geom_xmat': out[41], - 'geom_xpos': out[42], - '_impl.inverse_mul_m_skip': out[43], - '_impl.light_xdir': out[44], - '_impl.light_xpos': out[45], - 'mocap_pos': out[46], - 'mocap_quat': out[47], - '_impl.multiccd_clipped': out[48], - '_impl.multiccd_endvert': out[49], - '_impl.multiccd_face1': out[50], - '_impl.multiccd_face2': out[51], - '_impl.multiccd_idx1': out[52], - '_impl.multiccd_idx2': out[53], - '_impl.multiccd_n1': out[54], - '_impl.multiccd_n2': out[55], - '_impl.multiccd_pdist': out[56], - '_impl.multiccd_pnormal': out[57], - '_impl.multiccd_polygon': out[58], - '_impl.nacon': out[59], - '_impl.ncollision': out[60], - '_impl.ne': out[61], - '_impl.ne_connect': out[62], - '_impl.ne_jnt': out[63], - '_impl.ne_ten': out[64], - '_impl.ne_weld': out[65], - '_impl.nefc': out[66], - '_impl.nf': out[67], - '_impl.nl': out[68], - '_impl.nsolving': out[69], - '_impl.qLD': out[70], - '_impl.qLD_integration': out[71], - '_impl.qLDiagInv': out[72], - '_impl.qLDiagInv_integration': out[73], - '_impl.qM': out[74], - '_impl.qM_integration': out[75], - 'qacc': out[76], - '_impl.qacc_integration': out[77], - '_impl.qacc_rk': out[78], - 'qacc_smooth': out[79], - 'qacc_warmstart': out[80], - 'qfrc_actuator': out[81], - 'qfrc_applied': out[82], - 'qfrc_bias': out[83], - 'qfrc_constraint': out[84], - '_impl.qfrc_damper': out[85], - 'qfrc_fluid': out[86], - 'qfrc_gravcomp': out[87], - '_impl.qfrc_integration': out[88], - 'qfrc_passive': out[89], - 'qfrc_smooth': out[90], - '_impl.qfrc_spring': out[91], - 'qpos': out[92], - '_impl.qpos_t0': out[93], - 'qvel': out[94], - '_impl.qvel_rk': out[95], - '_impl.qvel_t0': out[96], - '_impl.sap_cumulative_sum': out[97], - '_impl.sap_projection_lower': out[98], - '_impl.sap_projection_upper': out[99], - '_impl.sap_range': out[100], - '_impl.sap_segment_index': out[101], - '_impl.sap_sort_index': out[102], - '_impl.sensor_contact_criteria': out[103], - '_impl.sensor_contact_direction': out[104], - '_impl.sensor_contact_matchid': out[105], - '_impl.sensor_contact_nmatch': out[106], - '_impl.sensor_rangefinder_dist': out[107], - '_impl.sensor_rangefinder_geomid': out[108], - '_impl.sensor_rangefinder_pnt': out[109], - '_impl.sensor_rangefinder_vec': out[110], - 'sensordata': out[111], - 'site_xmat': out[112], - 'site_xpos': out[113], - '_impl.solver_niter': out[114], - '_impl.subtree_angmom': out[115], - '_impl.subtree_bodyvel': out[116], - 'subtree_com': out[117], - '_impl.subtree_linvel': out[118], - '_impl.ten_J': out[119], - '_impl.ten_Jdot': out[120], - '_impl.ten_actfrc': out[121], - '_impl.ten_bias_coef': out[122], - 'ten_length': out[123], - '_impl.ten_velocity': out[124], - '_impl.ten_wrapadr': out[125], - '_impl.ten_wrapnum': out[126], - 'time': out[127], - '_impl.wrap_geom_xpos': out[128], - '_impl.wrap_obj': out[129], - '_impl.wrap_xpos': out[130], - 'xanchor': out[131], - 'xaxis': out[132], - 'xfrc_applied': out[133], - 'ximat': out[134], - 'xipos': out[135], - 'xmat': out[136], - 'xpos': out[137], - 'xquat': out[138], - '_impl.contact__dim': out[139], - '_impl.contact__dist': out[140], - '_impl.contact__efc_address': out[141], - '_impl.contact__frame': out[142], - '_impl.contact__friction': out[143], - '_impl.contact__geom': out[144], - '_impl.contact__includemargin': out[145], - '_impl.contact__pos': out[146], - '_impl.contact__solimp': out[147], - '_impl.contact__solref': out[148], - '_impl.contact__solreffriction': out[149], - '_impl.contact__worldid': out[150], - '_impl.efc__D': out[151], - '_impl.efc__J': out[152], - '_impl.efc__Jaref': out[153], - '_impl.efc__Ma': out[154], - '_impl.efc__Mgrad': out[155], - '_impl.efc__alpha': out[156], - '_impl.efc__aref': out[157], - '_impl.efc__beta': out[158], - '_impl.efc__cholesky_L_tmp': out[159], - '_impl.efc__cholesky_y_tmp': out[160], - '_impl.efc__cost': out[161], - '_impl.efc__cost_candidate': out[162], - '_impl.efc__done': out[163], - '_impl.efc__force': out[164], - '_impl.efc__frictionloss': out[165], - '_impl.efc__gauss': out[166], - '_impl.efc__grad': out[167], - '_impl.efc__grad_dot': out[168], - '_impl.efc__h': out[169], - '_impl.efc__id': out[170], - '_impl.efc__jv': out[171], - '_impl.efc__margin': out[172], - '_impl.efc__mv': out[173], - '_impl.efc__pos': out[174], - '_impl.efc__prev_Mgrad': out[175], - '_impl.efc__prev_cost': out[176], - '_impl.efc__prev_grad': out[177], - '_impl.efc__quad': out[178], - '_impl.efc__quad_gauss': out[179], - '_impl.efc__search': out[180], - '_impl.efc__search_dot': out[181], - '_impl.efc__state': out[182], - '_impl.efc__type': out[183], - '_impl.efc__vel': out[184], + 'actuator_force': out[2], + '_impl.actuator_length': out[3], + '_impl.actuator_moment': out[4], + '_impl.actuator_velocity': out[5], + '_impl.cacc': out[6], + 'cam_xmat': out[7], + 'cam_xpos': out[8], + '_impl.cdof': out[9], + '_impl.cdof_dot': out[10], + '_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], + 'ctrl': out[18], + 'cvel': out[19], + '_impl.energy': out[20], + 'eq_active': out[21], + '_impl.flexedge_length': out[22], + '_impl.flexedge_velocity': out[23], + '_impl.flexvert_xpos': out[24], + 'geom_xmat': out[25], + 'geom_xpos': out[26], + '_impl.light_xdir': out[27], + '_impl.light_xpos': out[28], + 'mocap_pos': out[29], + 'mocap_quat': out[30], + '_impl.nacon': out[31], + '_impl.ncollision': out[32], + '_impl.ne': out[33], + '_impl.ne_connect': out[34], + '_impl.ne_jnt': out[35], + '_impl.ne_ten': out[36], + '_impl.ne_weld': out[37], + '_impl.nefc': out[38], + '_impl.nf': out[39], + '_impl.nl': out[40], + '_impl.nsolving': out[41], + '_impl.qLD': out[42], + '_impl.qLDiagInv': out[43], + '_impl.qM': out[44], + 'qacc': out[45], + 'qacc_smooth': out[46], + 'qacc_warmstart': out[47], + 'qfrc_actuator': out[48], + 'qfrc_applied': out[49], + 'qfrc_bias': out[50], + 'qfrc_constraint': out[51], + '_impl.qfrc_damper': out[52], + 'qfrc_fluid': out[53], + 'qfrc_gravcomp': out[54], + 'qfrc_passive': out[55], + 'qfrc_smooth': out[56], + '_impl.qfrc_spring': out[57], + 'qpos': out[58], + 'qvel': out[59], + 'sensordata': out[60], + 'site_xmat': out[61], + 'site_xpos': out[62], + '_impl.solver_niter': out[63], + '_impl.subtree_angmom': out[64], + '_impl.subtree_bodyvel': out[65], + 'subtree_com': out[66], + '_impl.subtree_linvel': out[67], + '_impl.ten_J': out[68], + 'ten_length': out[69], + '_impl.ten_velocity': out[70], + '_impl.ten_wrapadr': out[71], + '_impl.ten_wrapnum': out[72], + 'time': out[73], + '_impl.wrap_obj': out[74], + '_impl.wrap_xpos': out[75], + 'xanchor': out[76], + 'xaxis': out[77], + 'xfrc_applied': out[78], + 'ximat': out[79], + 'xipos': out[80], + 'xmat': out[81], + 'xpos': out[82], + 'xquat': out[83], + '_impl.contact__dim': out[84], + '_impl.contact__dist': out[85], + '_impl.contact__efc_address': out[86], + '_impl.contact__frame': out[87], + '_impl.contact__friction': out[88], + '_impl.contact__geom': out[89], + '_impl.contact__geomcollisionid': out[90], + '_impl.contact__includemargin': out[91], + '_impl.contact__pos': out[92], + '_impl.contact__solimp': out[93], + '_impl.contact__solref': out[94], + '_impl.contact__solreffriction': out[95], + '_impl.contact__type': out[96], + '_impl.contact__worldid': out[97], + '_impl.efc__D': out[98], + '_impl.efc__J': out[99], + '_impl.efc__Jaref': out[100], + '_impl.efc__Ma': out[101], + '_impl.efc__Mgrad': out[102], + '_impl.efc__alpha': out[103], + '_impl.efc__aref': out[104], + '_impl.efc__beta': out[105], + '_impl.efc__cholesky_L_tmp': out[106], + '_impl.efc__cholesky_y_tmp': out[107], + '_impl.efc__cost': out[108], + '_impl.efc__done': out[109], + '_impl.efc__force': out[110], + '_impl.efc__frictionloss': out[111], + '_impl.efc__gauss': out[112], + '_impl.efc__grad': out[113], + '_impl.efc__grad_dot': out[114], + '_impl.efc__h': out[115], + '_impl.efc__id': out[116], + '_impl.efc__jv': out[117], + '_impl.efc__margin': out[118], + '_impl.efc__mv': out[119], + '_impl.efc__pos': out[120], + '_impl.efc__prev_Mgrad': out[121], + '_impl.efc__prev_cost': out[122], + '_impl.efc__prev_grad': out[123], + '_impl.efc__quad': out[124], + '_impl.efc__quad_gauss': out[125], + '_impl.efc__search': out[126], + '_impl.efc__search_dot': out[127], + '_impl.efc__state': out[128], + '_impl.efc__type': out[129], + '_impl.efc__vel': out[130], }) return d diff --git a/mjx/mujoco/mjx/warp/smooth.py b/mjx/mujoco/mjx/warp/smooth.py index ce4bd16e..9a00ffb7 100644 --- a/mjx/mujoco/mjx/warp/smooth.py +++ b/mjx/mujoco/mjx/warp/smooth.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 _kinematics_shim( # Model @@ -51,10 +52,13 @@ def _kinematics_shim( body_iquat: wp.array2d(dtype=wp.quat), body_jntadr: wp.array(dtype=int), body_jntnum: wp.array(dtype=int), + body_mocapid: wp.array(dtype=int), body_parentid: wp.array(dtype=int), 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), flex_edge: wp.array(dtype=wp.vec2i), flex_vertadr: wp.array(dtype=int), flex_vertbodyid: wp.array(dtype=int), @@ -79,7 +83,6 @@ def _kinematics_shim( flexedge_length: wp.array2d(dtype=float), flexedge_velocity: wp.array2d(dtype=float), flexvert_xpos: wp.array2d(dtype=wp.vec3), - geom_skip: wp.array(dtype=bool), geom_xmat: wp.array2d(dtype=wp.mat33), geom_xpos: wp.array2d(dtype=wp.vec3), mocap_pos: wp.array2d(dtype=wp.vec3), @@ -105,10 +108,13 @@ def _kinematics_shim( _m.body_iquat = body_iquat _m.body_jntadr = body_jntadr _m.body_jntnum = body_jntnum + _m.body_mocapid = body_mocapid _m.body_parentid = body_parentid _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.flex_edge = flex_edge _m.flex_vertadr = flex_vertadr _m.flex_vertbodyid = flex_vertbodyid @@ -132,7 +138,6 @@ def _kinematics_shim( _d.flexedge_length = flexedge_length _d.flexedge_velocity = flexedge_velocity _d.flexvert_xpos = flexvert_xpos - _d.geom_skip = geom_skip _d.geom_xmat = geom_xmat _d.geom_xpos = geom_xpos _d.mocap_pos = mocap_pos @@ -157,7 +162,6 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): 'flexedge_length': d._impl.flexedge_length.shape, 'flexedge_velocity': d._impl.flexedge_velocity.shape, 'flexvert_xpos': d._impl.flexvert_xpos.shape, - 'geom_skip': d._impl.geom_skip.shape, 'geom_xmat': d.geom_xmat.shape, 'geom_xpos': d.geom_xpos.shape, 'mocap_pos': d.mocap_pos.shape, @@ -176,14 +180,13 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): } jf = ffi.jax_callable_variadic_tuple( _kinematics_shim, - num_outputs=19, + num_outputs=18, output_dims=output_dims, vmap_method=None, in_out_argnames={ 'flexedge_length', 'flexedge_velocity', 'flexvert_xpos', - 'geom_skip', 'geom_xmat', 'geom_xpos', 'mocap_pos', @@ -208,10 +211,13 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): m.body_iquat, m.body_jntadr, m.body_jntnum, + m.body_mocapid, m.body_parentid, m.body_pos, m.body_quat, + m.body_rootid, m._impl.body_tree, + m.body_weldid, m._impl.flex_edge, m._impl.flex_vertadr, m._impl.flex_vertbodyid, @@ -235,7 +241,6 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): d._impl.flexedge_length, d._impl.flexedge_velocity, d._impl.flexvert_xpos, - d._impl.geom_skip, d.geom_xmat, d.geom_xpos, d.mocap_pos, @@ -256,22 +261,21 @@ def _kinematics_jax_impl(m: types.Model, d: types.Data): '_impl.flexedge_length': out[0], '_impl.flexedge_velocity': out[1], '_impl.flexvert_xpos': out[2], - '_impl.geom_skip': out[3], - 'geom_xmat': out[4], - 'geom_xpos': out[5], - 'mocap_pos': out[6], - 'mocap_quat': out[7], - 'qpos': out[8], - 'qvel': out[9], - 'site_xmat': out[10], - 'site_xpos': out[11], - 'xanchor': out[12], - 'xaxis': out[13], - 'ximat': out[14], - 'xipos': out[15], - 'xmat': out[16], - 'xpos': out[17], - 'xquat': out[18], + 'geom_xmat': out[3], + 'geom_xpos': out[4], + 'mocap_pos': out[5], + 'mocap_quat': out[6], + 'qpos': out[7], + 'qvel': out[8], + 'site_xmat': out[9], + 'site_xpos': out[10], + 'xanchor': out[11], + 'xaxis': out[12], + 'ximat': out[13], + 'xipos': out[14], + 'xmat': out[15], + 'xpos': out[16], + 'xquat': out[17], }) return d @@ -306,6 +310,7 @@ _e = mjwarp.Constraint( **{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init} ) + @ffi.format_args_for_warp def _tendon_shim( # Model @@ -319,6 +324,7 @@ def _tendon_shim( jnt_qposadr: wp.array(dtype=int), ntendon: int, nv: int, + nwrap: int, site_bodyid: wp.array(dtype=int), tendon_adr: wp.array(dtype=int), tendon_geom_adr: wp.array(dtype=int), @@ -343,7 +349,6 @@ def _tendon_shim( ten_length: wp.array2d(dtype=float), ten_wrapadr: wp.array2d(dtype=int), ten_wrapnum: wp.array2d(dtype=int), - wrap_geom_xpos: wp.array2d(dtype=wp.spatial_vector), wrap_obj: wp.array2d(dtype=wp.vec2i), wrap_xpos: wp.array2d(dtype=wp.spatial_vector), ): @@ -360,6 +365,7 @@ def _tendon_shim( _m.jnt_qposadr = jnt_qposadr _m.ntendon = ntendon _m.nv = nv + _m.nwrap = nwrap _m.site_bodyid = site_bodyid _m.tendon_adr = tendon_adr _m.tendon_geom_adr = tendon_geom_adr @@ -383,7 +389,6 @@ def _tendon_shim( _d.ten_length = ten_length _d.ten_wrapadr = ten_wrapadr _d.ten_wrapnum = ten_wrapnum - _d.wrap_geom_xpos = wrap_geom_xpos _d.wrap_obj = wrap_obj _d.wrap_xpos = wrap_xpos _d.nworld = nworld @@ -402,13 +407,12 @@ def _tendon_jax_impl(m: types.Model, d: types.Data): 'ten_length': d.ten_length.shape, 'ten_wrapadr': d._impl.ten_wrapadr.shape, 'ten_wrapnum': d._impl.ten_wrapnum.shape, - 'wrap_geom_xpos': d._impl.wrap_geom_xpos.shape, 'wrap_obj': d._impl.wrap_obj.shape, 'wrap_xpos': d._impl.wrap_xpos.shape, } jf = ffi.jax_callable_variadic_tuple( _tendon_shim, - num_outputs=13, + num_outputs=12, output_dims=output_dims, vmap_method=None, in_out_argnames={ @@ -422,7 +426,6 @@ def _tendon_jax_impl(m: types.Model, d: types.Data): 'ten_length', 'ten_wrapadr', 'ten_wrapnum', - 'wrap_geom_xpos', 'wrap_obj', 'wrap_xpos', }, @@ -438,6 +441,7 @@ def _tendon_jax_impl(m: types.Model, d: types.Data): m.jnt_qposadr, m.ntendon, m.nv, + m.nwrap, m.site_bodyid, m.tendon_adr, m._impl.tendon_geom_adr, @@ -461,7 +465,6 @@ def _tendon_jax_impl(m: types.Model, d: types.Data): d.ten_length, d._impl.ten_wrapadr, d._impl.ten_wrapnum, - d._impl.wrap_geom_xpos, d._impl.wrap_obj, d._impl.wrap_xpos, ) @@ -476,9 +479,8 @@ def _tendon_jax_impl(m: types.Model, d: types.Data): 'ten_length': out[7], '_impl.ten_wrapadr': out[8], '_impl.ten_wrapnum': out[9], - '_impl.wrap_geom_xpos': out[10], - '_impl.wrap_obj': out[11], - '_impl.wrap_xpos': out[12], + '_impl.wrap_obj': out[10], + '_impl.wrap_xpos': out[11], }) return d diff --git a/mjx/mujoco/mjx/warp/test_util.py b/mjx/mujoco/mjx/warp/test_util.py index 3afb840b..1752f3c9 100644 --- a/mjx/mujoco/mjx/warp/test_util.py +++ b/mjx/mujoco/mjx/warp/test_util.py @@ -41,7 +41,7 @@ def assert_attr_eq(a, b, attr): def make_data( - m: mujoco.MjModel, worldid: int, nconmax: int = 1_000, njmax: int = 100 + m: mujoco.MjModel, worldid: int, nconmax: int = 1_000, njmax: int = 200 ): """Make data for a given worldid using keyframes when available.""" dx = mjx.make_data(m, impl='warp', nconmax=nconmax, njmax=njmax) @@ -149,19 +149,20 @@ def _mjx_efc(dx, worldid: int): keys = np.arange(nefc) if not keys.size: empty = np.array([]) - return 0, empty, empty, np.zeros((0, dx.qvel.shape[0])), empty, empty - efc_pos = select(dx._impl.efc__pos[:nefc]) - efc_type = select(dx._impl.efc__type[:nefc]) - efc_d = select(dx._impl.efc__D[:nefc]) + return 0, empty, empty, np.zeros((0, dx.qvel.shape[-1])), empty, empty + efc_pos = select(dx._impl.efc__pos)[:nefc] + efc_type = select(dx._impl.efc__type)[:nefc] + efc_d = select(dx._impl.efc__D)[:nefc] keys_sorted = np.lexsort((-efc_pos, efc_type, efc_d)) keys = keys[keys_sorted] nefc = len(keys) type_ = efc_type[keys] pos = efc_pos[keys] - j = select(dx._impl.efc__J[:nefc])[keys] - aref = select(dx._impl.efc__aref[:nefc])[keys] - d_ = select(dx._impl.efc__D[:nefc])[keys] + # MuJoCo Warp may pad efc_J for tiled ops. + j = select(dx._impl.efc__J)[:nefc][keys][:, :dx.qvel.shape[-1]] + aref = select(dx._impl.efc__aref)[:nefc][keys] + d_ = select(dx._impl.efc__D)[:nefc][keys] return nefc, type_, pos, j, aref, d_ diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py index b2d65558..4ba743ef 100644 --- a/mjx/mujoco/mjx/warp/types.py +++ b/mjx/mujoco/mjx/warp/types.py @@ -63,11 +63,11 @@ class BlockDim: energy_vel_kinetic: int euler_dense: int mul_m_dense: int - qderiv_actuator_passive_actuation: int - qderiv_actuator_passive_no_actuation: int ray: int segmented_sort: int tendon_velocity: int + update_gradient_JTDAJ_dense: int + update_gradient_JTDAJ_sparse: int update_gradient_cholesky: int def tree_flatten(self): @@ -154,12 +154,17 @@ class ModelWarp(PyTreeNode): mesh_polyvertadr: np.ndarray mesh_polyvertnum: np.ndarray mocap_bodyid: np.ndarray + nacttrnbody: int nflex: int nflexedge: int nflexelem: int nflexelemdata: int nflexvert: int - nlsp: int + nmaxmeshdeg: int + nmaxpolygon: int + nrangefinder: int + nsensorcollision: int + nsensorcontact: int nsensortaxel: int nxn_geom_pair: np.ndarray nxn_geom_pair_filtered: np.ndarray @@ -180,6 +185,7 @@ class ModelWarp(PyTreeNode): rangefinder_sensor_adr: np.ndarray sensor_acc_adr: np.ndarray sensor_adr_to_contact_adr: np.ndarray + sensor_collision_start_adr: np.ndarray sensor_contact_adr: np.ndarray sensor_e_kinetic: bool sensor_e_potential: bool @@ -194,7 +200,6 @@ class ModelWarp(PyTreeNode): sensor_tendonactfrc_adr: np.ndarray sensor_touch_adr: np.ndarray sensor_vel_adr: np.ndarray - subtree_mass: jax.Array taxel_sensorid: np.ndarray taxel_vertadr: np.ndarray ten_wrapadr_site: np.ndarray @@ -211,12 +216,8 @@ class ModelWarp(PyTreeNode): class DataWarp(PyTreeNode): """Derived fields from Data.""" - act_dot_rk: jax.Array - act_t0: jax.Array - act_vel_integration: jax.Array actuator_length: jax.Array actuator_moment: jax.Array - actuator_trntype_body_ncon: jax.Array actuator_velocity: jax.Array cacc: jax.Array cdof: jax.Array @@ -233,11 +234,13 @@ class DataWarp(PyTreeNode): contact__frame: jax.Array contact__friction: jax.Array contact__geom: jax.Array + contact__geomcollisionid: jax.Array contact__includemargin: jax.Array contact__pos: jax.Array contact__solimp: jax.Array contact__solref: jax.Array contact__solreffriction: jax.Array + contact__type: jax.Array contact__worldid: jax.Array crb: jax.Array efc__D: jax.Array @@ -251,7 +254,6 @@ class DataWarp(PyTreeNode): efc__cholesky_L_tmp: jax.Array efc__cholesky_y_tmp: jax.Array efc__cost: jax.Array - efc__cost_candidate: jax.Array efc__done: jax.Array efc__force: jax.Array efc__frictionloss: jax.Array @@ -275,37 +277,11 @@ class DataWarp(PyTreeNode): efc__type: jax.Array efc__vel: jax.Array energy: jax.Array - energy_vel_mul_m_skip: jax.Array - epa_face: jax.Array - epa_horizon: jax.Array - epa_index: jax.Array - epa_map: jax.Array - epa_norm2: jax.Array - epa_pr: jax.Array - epa_vert: jax.Array - epa_vert1: jax.Array - epa_vert2: jax.Array - epa_vert_index1: jax.Array - epa_vert_index2: jax.Array flexedge_length: jax.Array flexedge_velocity: jax.Array flexvert_xpos: jax.Array - fluid_applied: jax.Array - geom_skip: jax.Array - inverse_mul_m_skip: jax.Array light_xdir: jax.Array light_xpos: jax.Array - multiccd_clipped: jax.Array - multiccd_endvert: jax.Array - multiccd_face1: jax.Array - multiccd_face2: jax.Array - multiccd_idx1: jax.Array - multiccd_idx2: jax.Array - multiccd_n1: jax.Array - multiccd_n2: jax.Array - multiccd_pdist: jax.Array - multiccd_pnormal: jax.Array - multiccd_polygon: jax.Array nacon: jax.Array naconmax: int ncollision: jax.Array @@ -321,49 +297,18 @@ class DataWarp(PyTreeNode): nsolving: jax.Array nworld: int qLD: jax.Array - qLD_integration: jax.Array qLDiagInv: jax.Array - qLDiagInv_integration: jax.Array qM: jax.Array - qM_integration: jax.Array - qacc_discrete: jax.Array - qacc_integration: jax.Array - qacc_rk: jax.Array qfrc_damper: jax.Array - qfrc_integration: jax.Array qfrc_spring: jax.Array - qpos_t0: jax.Array - qvel_rk: jax.Array - qvel_t0: jax.Array - ray_bodyexclude: jax.Array - ray_dist: jax.Array - ray_geomid: jax.Array - sap_cumulative_sum: jax.Array - sap_projection_lower: jax.Array - sap_projection_upper: jax.Array - sap_range: jax.Array - sap_segment_index: jax.Array - sap_sort_index: jax.Array - sensor_contact_criteria: jax.Array - sensor_contact_direction: jax.Array - sensor_contact_matchid: jax.Array - sensor_contact_nmatch: jax.Array - sensor_rangefinder_dist: jax.Array - sensor_rangefinder_geomid: jax.Array - sensor_rangefinder_pnt: jax.Array - sensor_rangefinder_vec: jax.Array solver_niter: jax.Array subtree_angmom: jax.Array subtree_bodyvel: jax.Array subtree_linvel: jax.Array ten_J: jax.Array - ten_Jdot: jax.Array - ten_actfrc: jax.Array - ten_bias_coef: jax.Array ten_velocity: jax.Array ten_wrapadr: jax.Array ten_wrapnum: jax.Array - wrap_geom_xpos: jax.Array wrap_obj: jax.Array wrap_xpos: jax.Array shape = property(lambda self: self.cacc.shape) @@ -377,42 +322,20 @@ DATA_NON_VMAP = { 'contact__frame', 'contact__friction', 'contact__geom', + 'contact__geomcollisionid', 'contact__includemargin', 'contact__pos', 'contact__solimp', 'contact__solref', 'contact__solreffriction', + 'contact__type', 'contact__worldid', - 'epa_face', - 'epa_horizon', - 'epa_index', - 'epa_map', - 'epa_norm2', - 'epa_pr', - 'epa_vert', - 'epa_vert1', - 'epa_vert2', - 'epa_vert_index1', - 'epa_vert_index2', - 'geom_skip', - 'multiccd_clipped', - 'multiccd_endvert', - 'multiccd_face1', - 'multiccd_face2', - 'multiccd_idx1', - 'multiccd_idx2', - 'multiccd_n1', - 'multiccd_n2', - 'multiccd_pdist', - 'multiccd_pnormal', - 'multiccd_polygon', 'nacon', 'naconmax', 'ncollision', 'njmax', 'nsolving', 'nworld', - 'ray_bodyexclude', } def _to_elt(cont, _, d, axis): @@ -443,13 +366,9 @@ _NDIM = { 'Data': { 'act': 2, 'act_dot': 2, - 'act_dot_rk': 2, - 'act_t0': 2, - 'act_vel_integration': 2, 'actuator_force': 2, 'actuator_length': 2, 'actuator_moment': 3, - 'actuator_trntype_body_ncon': 2, 'actuator_velocity': 2, 'cacc': 3, 'cam_xmat': 4, @@ -460,7 +379,7 @@ _NDIM = { 'cfrc_int': 3, 'cinert': 3, 'collision_pair': 2, - 'collision_pairid': 1, + 'collision_pairid': 2, 'collision_worldid': 1, 'contact__dim': 1, 'contact__dist': 1, @@ -468,11 +387,13 @@ _NDIM = { 'contact__frame': 3, 'contact__friction': 2, 'contact__geom': 2, + 'contact__geomcollisionid': 1, 'contact__includemargin': 1, 'contact__pos': 2, 'contact__solimp': 2, 'contact__solref': 2, 'contact__solreffriction': 2, + 'contact__type': 1, 'contact__worldid': 1, 'crb': 3, 'ctrl': 2, @@ -488,7 +409,6 @@ _NDIM = { 'efc__cholesky_L_tmp': 3, 'efc__cholesky_y_tmp': 2, 'efc__cost': 1, - 'efc__cost_candidate': 2, 'efc__done': 1, 'efc__force': 2, 'efc__frictionloss': 2, @@ -512,42 +432,16 @@ _NDIM = { 'efc__type': 2, 'efc__vel': 2, 'energy': 2, - 'energy_vel_mul_m_skip': 1, - 'epa_face': 3, - 'epa_horizon': 2, - 'epa_index': 2, - 'epa_map': 2, - 'epa_norm2': 2, - 'epa_pr': 3, - 'epa_vert': 3, - 'epa_vert1': 3, - 'epa_vert2': 3, - 'epa_vert_index1': 2, - 'epa_vert_index2': 2, 'eq_active': 2, 'flexedge_length': 2, 'flexedge_velocity': 2, 'flexvert_xpos': 3, - 'fluid_applied': 3, - 'geom_skip': 1, 'geom_xmat': 4, 'geom_xpos': 3, - 'inverse_mul_m_skip': 1, 'light_xdir': 3, 'light_xpos': 3, 'mocap_pos': 3, 'mocap_quat': 3, - 'multiccd_clipped': 3, - 'multiccd_endvert': 3, - 'multiccd_face1': 3, - 'multiccd_face2': 3, - 'multiccd_idx1': 2, - 'multiccd_idx2': 2, - 'multiccd_n1': 3, - 'multiccd_n2': 3, - 'multiccd_pdist': 2, - 'multiccd_pnormal': 3, - 'multiccd_polygon': 3, 'nacon': 1, 'naconmax': 0, 'ncollision': 1, @@ -563,15 +457,9 @@ _NDIM = { 'nsolving': 1, 'nworld': 0, 'qLD': 3, - 'qLD_integration': 3, 'qLDiagInv': 2, - 'qLDiagInv_integration': 2, 'qM': 3, - 'qM_integration': 3, 'qacc': 2, - 'qacc_discrete': 2, - 'qacc_integration': 2, - 'qacc_rk': 2, 'qacc_smooth': 2, 'qacc_warmstart': 2, 'qfrc_actuator': 2, @@ -581,33 +469,12 @@ _NDIM = { 'qfrc_damper': 2, 'qfrc_fluid': 2, 'qfrc_gravcomp': 2, - 'qfrc_integration': 2, 'qfrc_inverse': 2, 'qfrc_passive': 2, 'qfrc_smooth': 2, 'qfrc_spring': 2, 'qpos': 2, - 'qpos_t0': 2, 'qvel': 2, - 'qvel_rk': 2, - 'qvel_t0': 2, - 'ray_bodyexclude': 1, - 'ray_dist': 2, - 'ray_geomid': 2, - 'sap_cumulative_sum': 2, - 'sap_projection_lower': 3, - 'sap_projection_upper': 2, - 'sap_range': 2, - 'sap_segment_index': 2, - 'sap_sort_index': 3, - 'sensor_contact_criteria': 3, - 'sensor_contact_direction': 3, - 'sensor_contact_matchid': 3, - 'sensor_contact_nmatch': 2, - 'sensor_rangefinder_dist': 2, - 'sensor_rangefinder_geomid': 2, - 'sensor_rangefinder_pnt': 3, - 'sensor_rangefinder_vec': 3, 'sensordata': 2, 'site_xmat': 4, 'site_xpos': 3, @@ -617,15 +484,11 @@ _NDIM = { 'subtree_com': 3, 'subtree_linvel': 3, 'ten_J': 3, - 'ten_Jdot': 3, - 'ten_actfrc': 2, - 'ten_bias_coef': 2, 'ten_length': 2, 'ten_velocity': 2, 'ten_wrapadr': 2, 'ten_wrapnum': 2, 'time': 1, - 'wrap_geom_xpos': 3, 'wrap_obj': 3, 'wrap_xpos': 3, 'xanchor': 3, @@ -673,11 +536,11 @@ _NDIM = { 'block_dim__energy_vel_kinetic': 0, 'block_dim__euler_dense': 0, 'block_dim__mul_m_dense': 0, - 'block_dim__qderiv_actuator_passive_actuation': 0, - 'block_dim__qderiv_actuator_passive_no_actuation': 0, 'block_dim__ray': 0, 'block_dim__segmented_sort': 0, 'block_dim__tendon_velocity': 0, + 'block_dim__update_gradient_JTDAJ_dense': 0, + 'block_dim__update_gradient_JTDAJ_sparse': 0, 'block_dim__update_gradient_cholesky': 0, 'body_conaffinity': 1, 'body_contype': 1, @@ -755,7 +618,7 @@ _NDIM = { 'flex_vertbodyid': 1, 'flex_vertnum': 1, 'flexedge_length0': 1, - 'geom_aabb': 3, + 'geom_aabb': 4, 'geom_bodyid': 1, 'geom_conaffinity': 1, 'geom_condim': 1, @@ -840,6 +703,7 @@ _NDIM = { 'nC': 0, 'nM': 0, 'na': 0, + 'nacttrnbody': 0, 'nbody': 0, 'ncam': 0, 'neq': 0, @@ -855,8 +719,9 @@ _NDIM = { 'nhfielddata': 0, 'njnt': 0, 'nlight': 0, - 'nlsp': 0, 'nmat': 0, + 'nmaxmeshdeg': 0, + 'nmaxpolygon': 0, 'nmeshface': 0, 'nmeshgraph': 0, 'nmeshpoly': 0, @@ -866,7 +731,10 @@ _NDIM = { 'nmocap': 0, 'npair': 0, 'nq': 0, + 'nrangefinder': 0, 'nsensor': 0, + 'nsensorcollision': 0, + 'nsensorcontact': 0, 'nsensordata': 0, 'nsensortaxel': 0, 'nsite': 0, @@ -876,8 +744,8 @@ _NDIM = { 'nwrap': 0, 'nxn_geom_pair': 2, 'nxn_geom_pair_filtered': 2, - 'nxn_pairid': 1, - 'nxn_pairid_filtered': 1, + 'nxn_pairid': 2, + 'nxn_pairid_filtered': 2, 'oct_aabb': 3, 'oct_child': 2, 'oct_coeff': 2, @@ -935,6 +803,7 @@ _NDIM = { 'sensor_acc_adr': 1, 'sensor_adr': 1, 'sensor_adr_to_contact_adr': 1, + 'sensor_collision_start_adr': 1, 'sensor_contact_adr': 1, 'sensor_cutoff': 1, 'sensor_datatype': 1, @@ -964,7 +833,6 @@ _NDIM = { 'site_size': 2, 'site_type': 1, 'stat__meaninertia': 0, - 'subtree_mass': 2, 'taxel_sensorid': 1, 'taxel_vertadr': 1, 'ten_wrapadr_site': 1, @@ -1038,13 +906,9 @@ _BATCH_DIM = { 'Data': { 'act': True, 'act_dot': True, - 'act_dot_rk': True, - 'act_t0': True, - 'act_vel_integration': True, 'actuator_force': True, 'actuator_length': True, 'actuator_moment': True, - 'actuator_trntype_body_ncon': True, 'actuator_velocity': True, 'cacc': True, 'cam_xmat': True, @@ -1063,11 +927,13 @@ _BATCH_DIM = { 'contact__frame': False, 'contact__friction': False, 'contact__geom': False, + 'contact__geomcollisionid': False, 'contact__includemargin': False, 'contact__pos': False, 'contact__solimp': False, 'contact__solref': False, 'contact__solreffriction': False, + 'contact__type': False, 'contact__worldid': False, 'crb': True, 'ctrl': True, @@ -1083,7 +949,6 @@ _BATCH_DIM = { 'efc__cholesky_L_tmp': True, 'efc__cholesky_y_tmp': True, 'efc__cost': True, - 'efc__cost_candidate': True, 'efc__done': True, 'efc__force': True, 'efc__frictionloss': True, @@ -1107,42 +972,16 @@ _BATCH_DIM = { 'efc__type': True, 'efc__vel': True, 'energy': True, - 'energy_vel_mul_m_skip': True, - 'epa_face': False, - 'epa_horizon': False, - 'epa_index': False, - 'epa_map': False, - 'epa_norm2': False, - 'epa_pr': False, - 'epa_vert': False, - 'epa_vert1': False, - 'epa_vert2': False, - 'epa_vert_index1': False, - 'epa_vert_index2': False, 'eq_active': True, 'flexedge_length': True, 'flexedge_velocity': True, 'flexvert_xpos': True, - 'fluid_applied': True, - 'geom_skip': False, 'geom_xmat': True, 'geom_xpos': True, - 'inverse_mul_m_skip': True, 'light_xdir': True, 'light_xpos': True, 'mocap_pos': True, 'mocap_quat': True, - 'multiccd_clipped': False, - 'multiccd_endvert': False, - 'multiccd_face1': False, - 'multiccd_face2': False, - 'multiccd_idx1': False, - 'multiccd_idx2': False, - 'multiccd_n1': False, - 'multiccd_n2': False, - 'multiccd_pdist': False, - 'multiccd_pnormal': False, - 'multiccd_polygon': False, 'nacon': False, 'naconmax': False, 'ncollision': False, @@ -1158,15 +997,9 @@ _BATCH_DIM = { 'nsolving': False, 'nworld': False, 'qLD': True, - 'qLD_integration': True, 'qLDiagInv': True, - 'qLDiagInv_integration': True, 'qM': True, - 'qM_integration': True, 'qacc': True, - 'qacc_discrete': True, - 'qacc_integration': True, - 'qacc_rk': True, 'qacc_smooth': True, 'qacc_warmstart': True, 'qfrc_actuator': True, @@ -1176,33 +1009,12 @@ _BATCH_DIM = { 'qfrc_damper': True, 'qfrc_fluid': True, 'qfrc_gravcomp': True, - 'qfrc_integration': True, 'qfrc_inverse': True, 'qfrc_passive': True, 'qfrc_smooth': True, 'qfrc_spring': True, 'qpos': True, - 'qpos_t0': True, 'qvel': True, - 'qvel_rk': True, - 'qvel_t0': True, - 'ray_bodyexclude': False, - 'ray_dist': True, - 'ray_geomid': True, - 'sap_cumulative_sum': True, - 'sap_projection_lower': True, - 'sap_projection_upper': True, - 'sap_range': True, - 'sap_segment_index': True, - 'sap_sort_index': True, - 'sensor_contact_criteria': True, - 'sensor_contact_direction': True, - 'sensor_contact_matchid': True, - 'sensor_contact_nmatch': True, - 'sensor_rangefinder_dist': True, - 'sensor_rangefinder_geomid': True, - 'sensor_rangefinder_pnt': True, - 'sensor_rangefinder_vec': True, 'sensordata': True, 'site_xmat': True, 'site_xpos': True, @@ -1212,15 +1024,11 @@ _BATCH_DIM = { 'subtree_com': True, 'subtree_linvel': True, 'ten_J': True, - 'ten_Jdot': True, - 'ten_actfrc': True, - 'ten_bias_coef': True, 'ten_length': True, 'ten_velocity': True, 'ten_wrapadr': True, 'ten_wrapnum': True, 'time': True, - 'wrap_geom_xpos': True, 'wrap_obj': True, 'wrap_xpos': True, 'xanchor': True, @@ -1268,11 +1076,11 @@ _BATCH_DIM = { 'block_dim__energy_vel_kinetic': False, 'block_dim__euler_dense': False, 'block_dim__mul_m_dense': False, - 'block_dim__qderiv_actuator_passive_actuation': False, - 'block_dim__qderiv_actuator_passive_no_actuation': False, 'block_dim__ray': False, 'block_dim__segmented_sort': False, 'block_dim__tendon_velocity': False, + 'block_dim__update_gradient_JTDAJ_dense': False, + 'block_dim__update_gradient_JTDAJ_sparse': False, 'block_dim__update_gradient_cholesky': False, 'body_conaffinity': False, 'body_contype': False, @@ -1350,7 +1158,7 @@ _BATCH_DIM = { 'flex_vertbodyid': False, 'flex_vertnum': False, 'flexedge_length0': False, - 'geom_aabb': False, + 'geom_aabb': True, 'geom_bodyid': False, 'geom_conaffinity': False, 'geom_condim': False, @@ -1435,6 +1243,7 @@ _BATCH_DIM = { 'nC': False, 'nM': False, 'na': False, + 'nacttrnbody': False, 'nbody': False, 'ncam': False, 'neq': False, @@ -1450,8 +1259,9 @@ _BATCH_DIM = { 'nhfielddata': False, 'njnt': False, 'nlight': False, - 'nlsp': False, 'nmat': False, + 'nmaxmeshdeg': False, + 'nmaxpolygon': False, 'nmeshface': False, 'nmeshgraph': False, 'nmeshpoly': False, @@ -1461,7 +1271,10 @@ _BATCH_DIM = { 'nmocap': False, 'npair': False, 'nq': False, + 'nrangefinder': False, 'nsensor': False, + 'nsensorcollision': False, + 'nsensorcontact': False, 'nsensordata': False, 'nsensortaxel': False, 'nsite': False, @@ -1530,6 +1343,7 @@ _BATCH_DIM = { 'sensor_acc_adr': False, 'sensor_adr': False, 'sensor_adr_to_contact_adr': False, + 'sensor_collision_start_adr': False, 'sensor_contact_adr': False, 'sensor_cutoff': False, 'sensor_datatype': False, @@ -1559,7 +1373,6 @@ _BATCH_DIM = { 'site_size': False, 'site_type': False, 'stat__meaninertia': False, - 'subtree_mass': True, 'taxel_sensorid': False, 'taxel_vertadr': False, 'ten_wrapadr_site': False,