Import google-deepmind/mujoco_warp from GitHub.

PiperOrigin-RevId: 829017789
Change-Id: I2756e349297cd0ae53a63a38fb203e2fd2d5cb99
This commit is contained in:
Baruch Tabanpour
2025-11-06 10:41:32 -08:00
committed by Copybara-Service
parent 7dfd92098f
commit fa062768f4
34 changed files with 4292 additions and 4898 deletions
+2 -2
View File
@@ -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):
+1 -1
View File
@@ -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
+3
View File
@@ -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
+46 -32
View File
@@ -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()
@@ -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
@@ -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,
],
)
+130 -128
View File
@@ -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 `<exclude>` 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)
+19 -74
View File
@@ -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
@@ -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
File diff suppressed because it is too large Load Diff
@@ -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)
+56 -47
View File
@@ -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,
],
)
+172 -126
View File
@@ -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,
+28 -31
View File
@@ -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
+128 -118
View File
@@ -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),
+23 -11
View File
@@ -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)
File diff suppressed because it is too large Load Diff
-3
View File
@@ -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)
+41 -42
View File
@@ -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,
],
+57 -65
View File
@@ -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,
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+195 -141
View File
@@ -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)
+384 -93
View File
@@ -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,
],
)
File diff suppressed because it is too large Load Diff
+26 -32
View File
@@ -37,7 +37,7 @@ def is_intersect(p1: wp.vec2, p2: wp.vec2, p3: wp.vec2, p4: wp.vec2) -> 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)
+8 -9
View File
@@ -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.")
+17 -2
View File
@@ -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
+6 -3
View File
@@ -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}...")
+39 -191
View File
@@ -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
File diff suppressed because it is too large Load Diff
+33 -31
View File
@@ -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
+9 -8
View File
@@ -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_
+40 -227
View File
@@ -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,