Import google-deepmind/mujoco_warp from GitHub.
PiperOrigin-RevId: 829017789 Change-Id: I2756e349297cd0ae53a63a38fb203e2fd2d5cb99
This commit is contained in:
committed by
Copybara-Service
parent
7dfd92098f
commit
fa062768f4
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
+169
-55
@@ -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
@@ -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
@@ -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
|
||||
|
||||
|
||||
+530
-427
File diff suppressed because it is too large
Load Diff
+107
-124
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
+754
-978
File diff suppressed because it is too large
Load Diff
@@ -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
@@ -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
@@ -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,
|
||||
|
||||
+251
-309
File diff suppressed because it is too large
Load Diff
+202
-202
File diff suppressed because it is too large
Load Diff
+195
-141
@@ -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
@@ -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,
|
||||
],
|
||||
)
|
||||
|
||||
+455
-499
File diff suppressed because it is too large
Load Diff
+26
-32
@@ -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)
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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}...")
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
+344
-873
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user