Import google-deepmind/mujoco_warp from GitHub.

PiperOrigin-RevId: 830506734
Change-Id: I857bf59a34a46bd61fc24953fd6f10b2157febb7
This commit is contained in:
Taylor Howell
2025-11-10 10:34:01 -08:00
committed by Copybara-Service
parent f8cef22992
commit bfcceb583e
15 changed files with 371 additions and 1390 deletions
+199 -145
View File
@@ -13,23 +13,22 @@
# limitations under the License.
# ==============================================================================
from typing import Tuple
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
from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_filter
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import contact_params
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import geom_collision_pair
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact
from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame
from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAX_EPAFACES
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAX_EPAHORIZON
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXCONPAIR
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import Data
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model
@@ -82,14 +81,120 @@ def _check_convex_collision_pairs():
assert _check_convex_collision_pairs(), "_CONVEX_COLLISION_PAIRS is in invalid order."
@wp.func
def _hfield_filter(
# Model:
geom_dataid: wp.array(dtype=int),
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),
# Data in:
geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.array2d(dtype=wp.mat33),
# In:
worldid: int,
g1: int,
g2: int,
) -> Tuple[bool, float, float, float, float, float, float]:
"""Filter for height field collisions.
See MuJoCo mjc_ConvexHField.
"""
# 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)
pos2 = geom_xpos_in[worldid, g2]
pos = mat1T @ (pos2 - pos1)
r2 = geom_rbound[rbound_id, g2]
# TODO(team): margin?
margin = wp.max(geom_margin[margin_id, g1], geom_margin[margin_id, g2])
# box-sphere test: horizontal plane
for i in range(2):
if (size1[i] < pos[i] - r2 - margin) or (-size1[i] > pos[i] + r2 + margin):
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
# box-sphere test: vertical direction
if size1[2] < pos[2] - r2 - margin: # up
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
if -size1[3] > pos[2] + r2 + margin: # down
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
mat2 = geom_xmat_in[worldid, g2]
mat = mat1T @ mat2
# aabb for geom in height field frame
xmax = -MJ_MAXVAL
ymax = -MJ_MAXVAL
zmax = -MJ_MAXVAL
xmin = MJ_MAXVAL
ymin = MJ_MAXVAL
zmin = MJ_MAXVAL
aabb_id = worldid % geom_aabb.shape[0]
center2 = geom_aabb[aabb_id, g2, 0]
size2 = geom_aabb[aabb_id, g2, 1]
pos += mat1T @ center2
sign = wp.vec2(-1.0, 1.0)
for i in range(2):
for j in range(2):
for k in range(2):
corner_local = wp.vec3(sign[i] * size2[0], sign[j] * size2[1], sign[k] * size2[2])
corner_hf = mat @ corner_local
if corner_hf[0] > xmax:
xmax = corner_hf[0]
if corner_hf[1] > ymax:
ymax = corner_hf[1]
if corner_hf[2] > zmax:
zmax = corner_hf[2]
if corner_hf[0] < xmin:
xmin = corner_hf[0]
if corner_hf[1] < ymin:
ymin = corner_hf[1]
if corner_hf[2] < zmin:
zmin = corner_hf[2]
xmax += pos[0]
xmin += pos[0]
ymax += pos[1]
ymin += pos[1]
zmax += pos[2]
zmin += pos[2]
# box-box test
if (
(xmin - margin > size1[0])
or (xmax + margin < -size1[0])
or (ymin - margin > size1[1])
or (ymax + margin < -size1[1])
or (zmin - margin > size1[2])
or (zmax + margin < -size1[3])
):
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
else:
return False, xmin, xmax, ymin, ymax, zmin, zmax
@cache_kernel
def ccd_kernel_builder(
legacy_gjk: bool,
geomtype1: int,
geomtype2: int,
ccd_iterations: int,
epa_exact_neg_distance: bool,
depth_extension: float,
is_hfield: bool,
use_multiccd: bool,
):
@@ -155,108 +260,81 @@ def ccd_kernel_builder(
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):
simplex, normal = gjk_legacy(
ccd_iterations,
geom1,
geom2,
geomtype1,
geomtype2,
)
depth, normal = epa_legacy(
ccd_iterations, geom1, geom2, geomtype1, geomtype2, depth_extension, epa_exact_neg_distance, simplex, normal
)
dist = -depth
if dist >= 0.0 or depth < -depth_extension:
return 0
sphere = GeomType.SPHERE
ellipsoid = GeomType.ELLIPSOID
g1 = geoms[0]
g2 = geoms[1]
if geom_type[g1] == sphere or geom_type[g1] == ellipsoid or geom_type[g2] == sphere or geom_type[g2] == ellipsoid:
ncontact, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 1, 2, 1.0e-5)
else:
ncontact, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 4, 8, 1.0e-1)
frame = make_frame(normal)
points = mat3c()
witness1 = mat3c()
witness2 = mat3c()
geom1.margin = margin
geom2.margin = margin
if pairid[1] >= 0:
# if collision sensor, set large cutoff to work with various sensor cutoff values
cutoff = 1.0e32
else:
points = mat3c()
witness1 = mat3c()
witness2 = mat3c()
geom1.margin = margin
geom2.margin = margin
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,
geomtype1,
geomtype2,
x1,
x2,
epa_vert_in[tid],
epa_vert1_in[tid],
epa_vert2_in[tid],
epa_vert_index1_in[tid],
epa_vert_index2_in[tid],
epa_face_in[tid],
epa_pr_in[tid],
epa_norm2_in[tid],
epa_index_in[tid],
epa_map_in[tid],
epa_horizon_in[tid],
)
cutoff = 0.0
dist, ncontact, w1, w2, idx = ccd(
opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]],
cutoff,
ccd_iterations,
geom1,
geom2,
geomtype1,
geomtype2,
x1,
x2,
epa_vert_in[tid],
epa_vert1_in[tid],
epa_vert2_in[tid],
epa_vert_index1_in[tid],
epa_vert_index2_in[tid],
epa_face_in[tid],
epa_pr_in[tid],
epa_norm2_in[tid],
epa_index_in[tid],
epa_map_in[tid],
epa_horizon_in[tid],
)
if dist >= 0.0 and pairid[1] == -1:
return 0
if dist >= 0.0 and pairid[1] == -1:
return 0
witness1[0] = w1
witness2[0] = w2
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,
)
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)
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:
@@ -406,7 +484,7 @@ def ccd_kernel_builder(
# height field filter
if wp.static(is_hfield):
no_hf_collision, xmin, xmax, ymin, ymax, zmin, zmax = hfield_filter(
no_hf_collision, xmin, xmax, ymin, ymax, zmin, zmax = _hfield_filter(
geom_dataid, geom_aabb, geom_rbound, geom_margin, hfield_size, geom_xpos_in, geom_xmat_in, worldid, g1, g2
)
if no_hf_collision:
@@ -434,13 +512,10 @@ 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[geom_size_id, g1],
geom1, geom2 = geom_collision_pair(
geom_type,
geom_dataid,
geom_size,
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
@@ -455,35 +530,16 @@ def ccd_kernel_builder(
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in[worldid, g1],
geom_xmat_in[worldid, g1],
)
geom2_dataid = geom_dataid[g2]
geom2 = geom(
geomtype2,
geom2_dataid,
geom_size[geom_size_id, g2],
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
mesh_vert,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in[worldid, g2],
geom_xmat_in[worldid, g2],
geom_xpos_in,
geom_xmat_in,
geoms,
worldid,
)
# see MuJoCo mjc_ConvexHField
if wp.static(is_hfield):
geom1_dataid = geom_dataid[g1]
# height field subgrid
nrow = hfield_nrow[geom1_dataid]
ncol = hfield_ncol[geom1_dataid]
@@ -546,11 +602,10 @@ def ccd_kernel_builder(
# prism center
x1 = geom1.pos
if wp.static(not legacy_gjk):
x1_ = wp.vec3(0.0, 0.0, 0.0)
for i in range(6):
x1_ += prism[i]
x1 += geom1.rot @ (x1_ / 6.0)
x1_ = wp.vec3(0.0, 0.0, 0.0)
for i in range(6):
x1_ += prism[i]
x1 += geom1.rot @ (x1_ / 6.0)
ncontact = eval_ccd_write_contact(
opt_ccd_tolerance,
@@ -690,7 +745,6 @@ def convex_narrowphase(m: Model, d: Data):
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
@@ -749,7 +803,7 @@ def convex_narrowphase(m: Model, d: Data):
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, use_multiccd),
ccd_kernel_builder(g1, g2, m.opt.ccd_iterations, g1 == GeomType.HFIELD, use_multiccd),
dim=d.naconmax,
inputs=[
m.opt.ccd_tolerance,
@@ -1,715 +0,0 @@
# Copyright 2025 The Newton Developers
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import Geom
from mujoco.mjx.third_party.mujoco_warp._src.math import gjk_normalize
from mujoco.mjx.third_party.mujoco_warp._src.math import orthonormal
from mujoco.mjx.third_party.mujoco_warp._src.math import orthonormal_to_z
from mujoco.mjx.third_party.mujoco_warp._src.support import all_same
from mujoco.mjx.third_party.mujoco_warp._src.support import any_different
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MINVAL
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
# TODO(team): improve compile time to enable backward pass
wp.set_module_options({"enable_backward": False})
FLOAT_MIN = -1e30
FLOAT_MAX = 1e30
EPS_BEST_COUNT = 12
MULTI_CONTACT_COUNT = 4
MULTI_POLYGON_COUNT = 8
matc3 = wp.types.matrix(shape=(EPS_BEST_COUNT, 3), dtype=float)
vecc3 = wp.types.vector(EPS_BEST_COUNT * 3, dtype=float)
# Matrix definition for the `tris` scratch space which is used to store the
# triangles of the polytope. Note that the first dimension is 2, as we need
# to store the previous and current polytope. But since Warp doesn't support
# 3D matrices yet, we use 2 * 3 * EPS_BEST_COUNT as the first dimension.
TRIS_DIM = 3 * EPS_BEST_COUNT
mat2c3 = wp.types.matrix(shape=(2 * TRIS_DIM, 3), dtype=float)
mat3p = wp.types.matrix(shape=(MULTI_POLYGON_COUNT, 3), dtype=float)
mat3c = wp.types.matrix(shape=(MULTI_CONTACT_COUNT, 3), dtype=float)
mat43 = wp.types.matrix(shape=(4, 3), dtype=float)
vec6 = wp.types.vector(6, dtype=int)
VECI1 = vec6(0, 0, 0, 1, 1, 2)
VECI2 = vec6(1, 2, 3, 2, 3, 3)
@wp.func
def _gjk_support_geom(geom: Geom, geomtype: int, dir: wp.vec3):
local_dir = wp.transpose(geom.rot) @ dir
if geomtype == GeomType.SPHERE:
support_pt = geom.pos + geom.size[0] * dir
elif geomtype == GeomType.BOX:
res = wp.cw_mul(wp.sign(local_dir), geom.size)
support_pt = geom.rot @ res + geom.pos
elif geomtype == GeomType.CAPSULE:
res = local_dir * geom.size[0]
# add cylinder contribution
res[2] += wp.sign(local_dir[2]) * geom.size[1]
support_pt = geom.rot @ res + geom.pos
elif geomtype == GeomType.ELLIPSOID:
res = wp.cw_mul(local_dir, geom.size)
res = wp.normalize(res)
# transform to ellipsoid
res = wp.cw_mul(res, geom.size)
support_pt = geom.rot @ res + geom.pos
elif geomtype == GeomType.CYLINDER:
res = wp.vec3(0.0, 0.0, 0.0)
# set result in XY plane: support on circle
d = wp.sqrt(wp.dot(local_dir, local_dir))
if d > MJ_MINVAL:
scl = geom.size[0] / d
res[0] = local_dir[0] * scl
res[1] = local_dir[1] * scl
# set result in Z direction
res[2] = wp.sign(local_dir[2]) * geom.size[1]
support_pt = geom.rot @ res + geom.pos
elif geomtype == GeomType.MESH:
max_dist = float(FLOAT_MIN)
if geom.graphadr == -1 or geom.vertnum < 10:
# exhaustive search over all vertices
for i in range(geom.vertnum):
vert = geom.vert[geom.vertadr + i]
dist = wp.dot(vert, local_dir)
if dist > max_dist:
max_dist = dist
support_pt = vert
else:
numvert = geom.graph[geom.graphadr]
vert_edgeadr = geom.graphadr + 2
vert_globalid = geom.graphadr + 2 + numvert
edge_localid = geom.graphadr + 2 + 2 * numvert
# hillclimb until no change
prev = int(-1)
imax = int(0)
while True:
prev = int(imax)
i = int(geom.graph[vert_edgeadr + imax])
while geom.graph[edge_localid + i] >= 0:
subidx = geom.graph[edge_localid + i]
idx = geom.graph[vert_globalid + subidx]
dist = wp.dot(local_dir, geom.vert[geom.vertadr + idx])
if dist > max_dist:
max_dist = dist
imax = int(subidx)
i += int(1)
if imax == prev:
break
imax = geom.graph[vert_globalid + imax]
support_pt = geom.vert[geom.vertadr + imax]
support_pt = geom.rot @ support_pt + geom.pos
elif geomtype == GeomType.HFIELD:
max_dist = float(FLOAT_MIN)
for i in range(6):
vert = geom.hfprism[i]
dist = wp.dot(vert, local_dir)
if dist > max_dist:
max_dist = dist
support_pt = vert
support_pt = geom.rot @ support_pt + geom.pos
return wp.dot(support_pt, dir), support_pt
@wp.func
def _gjk_support(
# In:
geom1: Geom,
geom2: Geom,
geomtype1: int,
geomtype2: int,
dir: wp.vec3,
):
# Returns the distance between support points on two geoms, and the support point.
# Negative distance means objects are not intersecting along direction `dir`.
# Positive distance means objects are intersecting along the given direction `dir`.
dist1, s1 = _gjk_support_geom(geom1, geomtype1, dir)
dist2, s2 = _gjk_support_geom(geom2, geomtype2, -dir)
support_pt = s1 - s2
return dist1 + dist2, support_pt
@wp.func
def _expand_polytope(count: int, prev_count: int, dists: vecc3, tris: mat2c3, p: matc3):
# expand polytope greedily
for j in range(count):
best = int(0)
dd = dists[0]
for i in range(1, 3 * prev_count):
if dists[i] < dd:
dd = dists[i]
best = i
dists[best] = float(wp.static(2 * FLOAT_MAX))
parent_index = best // 3
child_index = best % 3
# fill in the new triangle at the next index
tris[TRIS_DIM + j * 3 + 0] = tris[parent_index * 3 + child_index]
tris[TRIS_DIM + j * 3 + 1] = tris[parent_index * 3 + ((child_index + 1) % 3)]
tris[TRIS_DIM + j * 3 + 2] = p[parent_index]
for r in range(wp.static(EPS_BEST_COUNT * 3)):
# swap triangles
swap = tris[TRIS_DIM + r]
tris[TRIS_DIM + r] = tris[r]
tris[r] = swap
return dists, tris
@wp.func
def gjk_legacy(
# In:
gjk_iterations: int,
geom1: Geom,
geom2: Geom,
geomtype1: int,
geomtype2: int,
):
dir = wp.vec3(0.0, 0.0, 1.0)
dir_n = -dir
depth = float(FLOAT_MAX)
dist_max, simplex0 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir)
dist_min, simplex1 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir_n)
if dist_max < dist_min:
depth = dist_max
normal = dir
else:
depth = dist_min
normal = dir_n
sd = wp.normalize(simplex0 - simplex1)
dir = orthonormal_to_z(sd)
dist_max, simplex3 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir)
# Initialize a 2-simplex with simplex[2]==simplex[1]. This ensures the
# correct winding order for face normals defined below. Face 0 and face 3
# are degenerate, and face 1 and 2 have opposing normals.
simplex = mat43()
simplex[0] = simplex0
simplex[1] = simplex1
simplex[2] = simplex[1]
simplex[3] = simplex3
if dist_max < depth:
depth = dist_max
normal = dir
if dist_min < depth:
depth = dist_min
normal = dir_n
plane = mat43()
for _ in range(gjk_iterations):
# winding orders: plane[0] ccw, plane[1] cw, plane[2] ccw, plane[3] cw
plane[0] = wp.cross(simplex[3] - simplex[2], simplex[1] - simplex[2])
plane[1] = wp.cross(simplex[3] - simplex[0], simplex[2] - simplex[0])
plane[2] = wp.cross(simplex[3] - simplex[1], simplex[0] - simplex[1])
plane[3] = wp.cross(simplex[2] - simplex[0], simplex[1] - simplex[0])
# Compute distance of each face halfspace to the origin. If dplane<0, then the
# origin is outside the halfspace. If dplane>0 then the origin is inside
# the halfspace defined by the face plane.
dplane = wp.vec4(float(FLOAT_MAX))
plane0, p0 = gjk_normalize(plane[0])
plane1, p1 = gjk_normalize(plane[1])
plane2, p2 = gjk_normalize(plane[2])
plane3, p3 = gjk_normalize(plane[3])
plane[0] = plane0
plane[1] = plane1
plane[2] = plane2
plane[3] = plane3
if p0:
dplane[0] = wp.dot(plane[0], simplex[2])
if p1:
dplane[1] = wp.dot(plane[1], simplex[0])
if p2:
dplane[2] = wp.dot(plane[2], simplex[1])
if p3:
dplane[3] = wp.dot(plane[3], simplex[0])
# pick plane normal with minimum distance to the origin
i1 = wp.where(dplane[0] < dplane[1], 0, 1)
i2 = wp.where(dplane[2] < dplane[3], 2, 3)
index = wp.where(dplane[i1] < dplane[i2], i1, i2)
if dplane[index] > 0.0:
# origin is inside the simplex, objects are intersecting
break
# add new support point to the simplex
dist, simplex_i = _gjk_support(geom1, geom2, geomtype1, geomtype2, plane[index])
simplex[index] = simplex_i
if dist < depth:
depth = dist
normal = plane[index]
# preserve winding order of the simplex faces
index1 = (index + 1) & 3
index2 = (index + 2) & 3
swap = simplex[index1]
simplex[index1] = simplex[index2]
simplex[index2] = swap
if dist < 0.0:
break # objects are likely non-intersecting
return simplex, normal
@wp.func
def epa_legacy(
# In:
epa_iterations: int,
geom1: Geom,
geom2: Geom,
geomtype1: int,
geomtype2: int,
depth_extension: float,
epa_exact_neg_distance: bool,
simplex: mat43,
normal: wp.vec3,
):
# get the support, if depth < 0: objects do not intersect
depth, simplex0 = _gjk_support(geom1, geom2, geomtype1, geomtype2, normal)
simplex[0] = simplex0
if depth < -depth_extension:
# Objects are not intersecting, and we do not obtain the closest points as
# specified by depth_extension.
return FLOAT_MAX, wp.vec3(wp.nan, wp.nan, wp.nan)
if epa_exact_neg_distance:
# Check closest points to all edges of the simplex, rather than just the
# face normals. This gives the exact depth/normal for the non-intersecting
# case.
for i in range(6):
i1 = VECI1[i]
i2 = VECI2[i]
si1 = simplex[i1]
si2 = simplex[i2]
if si1[0] != si2[0] or si1[1] != si2[1] or si1[2] != si2[2]:
v = si1 - si2
alpha = wp.dot(si1, v) / wp.dot(v, v)
# p0 is the closest segment point to the origin
p0 = wp.clamp(alpha, 0.0, 1.0) * v - si1
p0, pf = gjk_normalize(p0)
if pf:
depth2, _ = _gjk_support(geom1, geom2, geomtype1, geomtype2, p0)
if depth2 < depth:
depth = depth2
normal = p0
# supporting points for each triangle
p = matc3()
# distance to the origin for candidate triangles
dists = vecc3()
tris = mat2c3()
tris[0] = simplex[2]
tris[1] = simplex[1]
tris[2] = simplex[3]
tris[3] = simplex[0]
tris[4] = simplex[2]
tris[5] = simplex[3]
tris[6] = simplex[1]
tris[7] = simplex[0]
tris[8] = simplex[3]
tris[9] = simplex[0]
tris[10] = simplex[1]
tris[11] = simplex[2]
# Calculate the total number of iterations to avoid nested loop
# This is a hack to reduce compile time
count = int(4)
it = int(0)
for _ in range(epa_iterations):
it += count
count = wp.min(count * 3, EPS_BEST_COUNT)
count = int(4)
i = int(0)
for _ in range(it):
# Loop through all triangles, and obtain distances to the origin for each
# new triangle candidate.
ti = 3 * i
n = wp.cross(tris[ti + 2] - tris[ti + 0], tris[ti + 1] - tris[ti + 0])
n, nf = gjk_normalize(n)
if not nf:
for j in range(3):
dists[i * 3 + j] = wp.static(float(2 * FLOAT_MAX))
continue
dist, pi = _gjk_support(geom1, geom2, geomtype1, geomtype2, n)
p[i] = pi
if dist < depth:
depth = dist
normal = n
# iterate over edges and get distance using support point
for j in range(3):
if epa_exact_neg_distance:
# obtain closest point between new triangle edge and origin
tqj = tris[ti + j]
if (p[i, 0] != tqj[0]) or (p[i, 1] != tqj[1]) or (p[i, 2] != tqj[2]):
v = p[i] - tris[ti + j]
alpha = wp.dot(p[i], v) / wp.dot(v, v)
p0 = wp.clamp(alpha, 0.0, 1.0) * v - p[i]
p0, pf = gjk_normalize(p0)
if pf:
dist2, v = _gjk_support(geom1, geom2, geomtype1, geomtype2, p0)
if dist2 < depth:
depth = dist2
normal = p0
plane = wp.cross(p[i] - tris[ti + j], tris[ti + ((j + 1) % 3)] - tris[ti + j])
plane, pf = gjk_normalize(plane)
if pf:
dd = wp.dot(plane, tris[ti + j])
else:
dd = float(FLOAT_MAX)
if (dd < 0 and depth >= 0) or (
tris[ti + ((j + 2) % 3)][0] == p[i][0]
and tris[ti + ((j + 2) % 3)][1] == p[i][1]
and tris[ti + ((j + 2) % 3)][2] == p[i][2]
):
dists[i * 3 + j] = float(FLOAT_MAX)
else:
dists[i * 3 + j] = dd
if i == count - 1:
prev_count = count
count = wp.min(count * 3, EPS_BEST_COUNT)
dists, tris = _expand_polytope(count, prev_count, dists, tris, p)
i = int(0)
else:
i += 1
return depth, normal
@wp.func
def multicontact_legacy(
# In:
geom1: Geom,
geom2: Geom,
geomtype1: int,
geomtype2: int,
depth_extension: float,
depth: float,
normal: wp.vec3,
ncontact: int,
npolygon: int,
perturbation_angle: float,
):
# Calculates multiple contact points given the normal from EPA.
# 1. Calculates the polygon on each shape by tiling the normal
# "perturbation_angle" (radians) in the orthogonal component of the normal.
# The "perturbation_angle" can be changed to depend on the depth of the
# contact, in a future version.
# 2. The normal is tilted "npolygon" times in the directions evenly
# spaced in the orthogonal component of the normal.
# (works well for >= 6, default is 8).
# 3. The intersection between these two polygons is calculated in 2D space
# (complement to the normal). If they intersect, extreme points in both
# directions are found. This can be modified to the extremes in the
# direction of eigenvectors of the variance of points of each polygon. If
# they do not intersect, the closest points of both polygons are found.
assert ncontact <= MULTI_CONTACT_COUNT
assert npolygon <= MULTI_POLYGON_COUNT
if depth < -depth_extension:
return 0, mat3c()
dir = orthonormal(normal)
dir2 = wp.cross(normal, dir)
angle = perturbation_angle
c = wp.cos(angle)
s = wp.sin(angle)
tc = 1.0 - c
v1 = mat3p()
v2 = mat3p()
contact_points = mat3c()
# Obtain points on the polygon determined by the support and tilt angle,
# in the basis of the contact frame.
v1count = int(0)
v2count = int(0)
angle_ratio = wp.static(2.0 * wp.pi) / float(npolygon)
for i in range(npolygon):
angle = angle_ratio * float(i)
axis = wp.cos(angle) * dir + wp.sin(angle) * dir2
# Axis-angle rotation matrix. See
# https://en.wikipedia.org/wiki/Rotation_matrix#Rotation_matrix_from_axis_and_angle
mat0 = c + axis[0] * axis[0] * tc
mat5 = c + axis[1] * axis[1] * tc
mat10 = c + axis[2] * axis[2] * tc
t1 = axis[0] * axis[1] * tc
t2 = axis[2] * s
mat4 = t1 + t2
mat1 = t1 - t2
t1 = axis[0] * axis[2] * tc
t2 = axis[1] * s
mat8 = t1 - t2
mat2 = t1 + t2
t1 = axis[1] * axis[2] * tc
t2 = axis[0] * s
mat9 = t1 + t2
mat6 = t1 - t2
n = wp.vec3(
mat0 * normal[0] + mat1 * normal[1] + mat2 * normal[2],
mat4 * normal[0] + mat5 * normal[1] + mat6 * normal[2],
mat8 * normal[0] + mat9 * normal[1] + mat10 * normal[2],
)
_, p = _gjk_support_geom(geom1, geomtype1, n)
v1[v1count] = wp.vec3(wp.dot(p, dir), wp.dot(p, dir2), wp.dot(p, normal))
if i == 0:
v1count += 1
elif any_different(v1[v1count], v1[v1count - 1]):
v1count += 1
n = -n
_, p = _gjk_support_geom(geom2, geomtype2, n)
v2[v2count] = wp.vec3(wp.dot(p, dir), wp.dot(p, dir2), wp.dot(p, normal))
if i == 0:
v2count += 1
elif any_different(v2[v2count], v2[v2count - 1]):
v2count += 1
# remove duplicate vertices on the array boundary
if v1count > 1 and all_same(v1[v1count - 1], v1[0]):
v1count -= 1
if v2count > 1 and all_same(v2[v2count - 1], v2[0]):
v2count -= 1
# find an intersecting polygon between v1 and v2 in the 2D plane
out = mat43()
candCount = int(0)
if v2count > 1:
for i in range(v1count):
m1a = v1[i]
is_in = bool(True)
# check if point m1a is inside the v2 polygon on the 2D plane
for j in range(v2count):
j2 = (j + 1) % v2count
# Checks that orientation of the triangle (v2[j], v2[j2], m1a) is
# counter-clockwise. If so, point m1a is inside the v2 polygon.
is_in = is_in and ((v2[j2][0] - v2[j][0]) * (m1a[1] - v2[j][1]) - (v2[j2][1] - v2[j][1]) * (m1a[0] - v2[j][0]) >= 0.0)
if not is_in:
break
if is_in:
if not candCount or m1a[0] < out[0, 0]:
out[0] = m1a
if not candCount or m1a[0] > out[1, 0]:
out[1] = m1a
if not candCount or m1a[1] < out[2, 1]:
out[2] = m1a
if not candCount or m1a[1] > out[3, 1]:
out[3] = m1a
candCount += 1
if v1count > 1:
for i in range(v2count):
m1a = v2[i]
is_in = bool(True)
for j in range(v1count):
j2 = (j + 1) % v1count
is_in = is_in and (v1[j2][0] - v1[j][0]) * (m1a[1] - v1[j][1]) - (v1[j2][1] - v1[j][1]) * (m1a[0] - v1[j][0]) >= 0.0
if not is_in:
break
if is_in:
if not candCount or m1a[0] < out[0, 0]:
out[0] = m1a
if not candCount or m1a[0] > out[1, 0]:
out[1] = m1a
if not candCount or m1a[1] < out[2, 1]:
out[2] = m1a
if not candCount or m1a[1] > out[3, 1]:
out[3] = m1a
candCount += 1
if v1count > 1 and v2count > 1:
# Check all edge pairs, and store line segment intersections if they are
# on the edge of the boundary.
for i in range(v1count):
for j in range(v2count):
m1a = v1[i]
m1b = v1[(i + 1) % v1count]
m2a = v2[j]
m2b = v2[(j + 1) % v2count]
det = (m2a[1] - m2b[1]) * (m1b[0] - m1a[0]) - (m1a[1] - m1b[1]) * (m2b[0] - m2a[0])
if wp.abs(det) > 1e-12:
a11 = (m2a[1] - m2b[1]) / det
a12 = (m2b[0] - m2a[0]) / det
a21 = (m1a[1] - m1b[1]) / det
a22 = (m1b[0] - m1a[0]) / det
b1 = m2a[0] - m1a[0]
b2 = m2a[1] - m1a[1]
alpha = a11 * b1 + a12 * b2
beta = a21 * b1 + a22 * b2
if alpha >= 0.0 and alpha <= 1.0 and beta >= 0.0 and beta <= 1.0:
m0 = wp.vec3(
m1a[0] + alpha * (m1b[0] - m1a[0]),
m1a[1] + alpha * (m1b[1] - m1a[1]),
(m1a[2] + alpha * (m1b[2] - m1a[2]) + m2a[2] + beta * (m2b[2] - m2a[2])) * 0.5,
)
if not candCount or m0[0] < out[0, 0]:
out[0] = m0
if not candCount or m0[0] > out[1, 0]:
out[1] = m0
if not candCount or m0[1] < out[2, 1]:
out[2] = m0
if not candCount or m0[1] > out[3, 1]:
out[3] = m0
candCount += 1
var_rx = wp.vec3(0.0)
contact_count = int(0)
if candCount > 0:
# Polygon intersection was found.
# TODO(btaba): replace the above routine with the manifold point routine
# from MJX. Deduplicate the points properly.
last_pt = wp.vec3(FLOAT_MAX, FLOAT_MAX, FLOAT_MAX)
for k in range(ncontact):
pt = out[k, 0] * dir + out[k, 1] * dir2 + out[k, 2] * normal
# skip contact points that are too close
if wp.length(pt - last_pt) <= 1e-6:
continue
contact_points[contact_count] = pt
last_pt = pt
contact_count += 1
else:
# Polygon intersection was not found. Loop through all vertex pairs and
# calculate an approximate contact point.
minDist = float(0.0)
for i in range(v1count):
for j in range(v2count):
# Find the closest vertex pair. Calculate a contact point var_rx as the
# midpoint between the closest vertex pair.
m1 = v1[i]
m2 = v2[j]
dd = (m1[0] - m2[0]) * (m1[0] - m2[0]) + (m1[1] - m2[1]) * (m1[1] - m2[1])
if (i == 0 and j == 0) or (dd < minDist):
minDist = dd
var_rx = ((m1[0] + m2[0]) * dir + (m1[1] + m2[1]) * dir2 + (m1[2] + m2[2]) * normal) * 0.5
# Check for a closer point between a point on v2 and an edge on v1.
m1b = v1[(i + 1) % v1count]
m2b = v2[(j + 1) % v2count]
if v1count > 1:
dd = (m1b[0] - m1[0]) * (m1b[0] - m1[0]) + (m1b[1] - m1[1]) * (m1b[1] - m1[1])
t = ((m2[1] - m1[1]) * (m1b[0] - m1[0]) - (m2[0] - m1[0]) * (m1b[1] - m1[1])) / dd
dx = m2[0] + (m1b[1] - m1[1]) * t
dy = m2[1] - (m1b[0] - m1[0]) * t
dist = (dx - m2[0]) * (dx - m2[0]) + (dy - m2[1]) * (dy - m2[1])
if (
(dist < minDist)
and ((dx - m1[0]) * (m1b[0] - m1[0]) + (dy - m1[1]) * (m1b[1] - m1[1]) >= 0)
and ((dx - m1b[0]) * (m1[0] - m1b[0]) + (dy - m1b[1]) * (m1[1] - m1b[1]) >= 0)
):
alpha = wp.sqrt(((dx - m1[0]) * (dx - m1[0]) + (dy - m1[1]) * (dy - m1[1])) / dd)
minDist = dist
w = ((1.0 - alpha) * m1 + alpha * m1b + m2) * 0.5
var_rx = w[0] * dir + w[1] * dir2 + w[2] * normal
# check for a closer point between a point on v1 and an edge on v2
if v2count > 1:
dd = (m2b[0] - m2[0]) * (m2b[0] - m2[0]) + (m2b[1] - m2[1]) * (m2b[1] - m2[1])
t = ((m1[1] - m2[1]) * (m2b[0] - m2[0]) - (m1[0] - m2[0]) * (m2b[1] - m2[1])) / dd
dx = m1[0] + (m2b[1] - m2[1]) * t
dy = m1[1] - (m2b[0] - m2[0]) * t
dist = (dx - m1[0]) * (dx - m1[0]) + (dy - m1[1]) * (dy - m1[1])
if (
dist < minDist
and (dx - m2[0]) * (m2b[0] - m2[0]) + (dy - m2[1]) * (m2b[1] - m2[1]) >= 0
and (dx - m2b[0]) * (m2[0] - m2b[0]) + (dy - m2b[1]) * (m2[1] - m2b[1]) >= 0
):
alpha = wp.sqrt(((dx - m2[0]) * (dx - m2[0]) + (dy - m2[1]) * (dy - m2[1])) / dd)
minDist = dist
w = (m1 + (1.0 - alpha) * m2 + alpha * m2b) * 0.5
var_rx = w[0] * dir + w[1] * dir2 + w[2] * normal
for k in range(ncontact):
contact_points[k] = var_rx
contact_count = 1
return contact_count, contact_points
@@ -1,129 +0,0 @@
# Copyright 2025 The Newton Developers
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
from typing import Tuple
import warp as wp
from mujoco.mjx.third_party.mujoco_warp._src.types import MJ_MAXVAL
@wp.func
def hfield_filter(
# Model:
geom_dataid: wp.array(dtype=int),
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),
# Data in:
geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.array2d(dtype=wp.mat33),
# In:
worldid: int,
g1: int,
g2: int,
) -> Tuple[bool, float, float, float, float, float, float]:
"""Filter for height field collisions.
See MuJoCo mjc_ConvexHField.
"""
# 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)
pos2 = geom_xpos_in[worldid, g2]
pos = mat1T @ (pos2 - pos1)
r2 = geom_rbound[rbound_id, g2]
# TODO(team): margin?
margin = wp.max(geom_margin[margin_id, g1], geom_margin[margin_id, g2])
# box-sphere test: horizontal plane
for i in range(2):
if (size1[i] < pos[i] - r2 - margin) or (-size1[i] > pos[i] + r2 + margin):
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
# box-sphere test: vertical direction
if size1[2] < pos[2] - r2 - margin: # up
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
if -size1[3] > pos[2] + r2 + margin: # down
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
mat2 = geom_xmat_in[worldid, g2]
mat = mat1T @ mat2
# aabb for geom in height field frame
xmax = -MJ_MAXVAL
ymax = -MJ_MAXVAL
zmax = -MJ_MAXVAL
xmin = MJ_MAXVAL
ymin = MJ_MAXVAL
zmin = MJ_MAXVAL
aabb_id = worldid % geom_aabb.shape[0]
center2 = geom_aabb[aabb_id, g2, 0]
size2 = geom_aabb[aabb_id, g2, 1]
pos += mat1T @ center2
sign = wp.vec2(-1.0, 1.0)
for i in range(2):
for j in range(2):
for k in range(2):
corner_local = wp.vec3(sign[i] * size2[0], sign[j] * size2[1], sign[k] * size2[2])
corner_hf = mat @ corner_local
if corner_hf[0] > xmax:
xmax = corner_hf[0]
if corner_hf[1] > ymax:
ymax = corner_hf[1]
if corner_hf[2] > zmax:
zmax = corner_hf[2]
if corner_hf[0] < xmin:
xmin = corner_hf[0]
if corner_hf[1] < ymin:
ymin = corner_hf[1]
if corner_hf[2] < zmin:
zmin = corner_hf[2]
xmax += pos[0]
xmin += pos[0]
ymax += pos[1]
ymin += pos[1]
zmax += pos[2]
zmin += pos[2]
# box-box test
if (
(xmin - margin > size1[0])
or (xmax + margin < -size1[0])
or (ymin - margin > size1[1])
or (ymax + margin < -size1[1])
or (zmin - margin > size1[2])
or (zmax + margin < -size1[3])
):
return True, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf, wp.inf
else:
return False, xmin, xmax, ymin, ymax, zmin, zmax
@@ -79,12 +79,11 @@ class Geom:
@wp.func
def geom(
# kernel_analyzer: off
def geom_collision_pair(
# Model:
geom_type: int,
geom_dataid: int,
geom_size: wp.vec3,
geom_type: wp.array(dtype=int),
geom_dataid: wp.array(dtype=int),
geom_size: wp.array2d(dtype=wp.vec3),
mesh_vertadr: wp.array(dtype=int),
mesh_vertnum: wp.array(dtype=int),
mesh_graphadr: wp.array(dtype=int),
@@ -100,44 +99,73 @@ def geom(
mesh_polymapnum: wp.array(dtype=int),
mesh_polymap: wp.array(dtype=int),
# Data in:
geom_xpos_in: wp.vec3,
geom_xmat_in: wp.mat33,
# kernel_analyzer: on
) -> Geom:
geom = Geom()
geom.pos = geom_xpos_in
geom.rot = geom_xmat_in
geom.size = geom_size
geom.normal = wp.vec3(geom_xmat_in[0, 2], geom_xmat_in[1, 2], geom_xmat_in[2, 2]) # plane
geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.array2d(dtype=wp.mat33),
# In:
geoms: wp.vec2i,
worldid: int,
) -> Tuple[Geom, Geom]:
geom1 = Geom()
geom2 = Geom()
if geom_type == GeomType.MESH:
if geom_dataid >= 0:
geom.vertadr = mesh_vertadr[geom_dataid]
geom.vertnum = mesh_vertnum[geom_dataid]
geom.graphadr = mesh_graphadr[geom_dataid]
geom.mesh_polynum = mesh_polynum[geom_dataid]
geom.mesh_polyadr = mesh_polyadr[geom_dataid]
else:
geom.vertadr = -1
geom.vertnum = -1
geom.graphadr = -1
geom.mesh_polynum = -1
geom.mesh_polyadr = -1
g1 = geoms[0]
g2 = geoms[1]
geom_type1 = geom_type[g1]
geom_type2 = geom_type[g2]
geom.vert = mesh_vert
geom.graph = mesh_graph
geom.mesh_polynormal = mesh_polynormal
geom.mesh_polyvertadr = mesh_polyvertadr
geom.mesh_polyvertnum = mesh_polyvertnum
geom.mesh_polyvert = mesh_polyvert
geom.mesh_polymapadr = mesh_polymapadr
geom.mesh_polymapnum = mesh_polymapnum
geom.mesh_polymap = mesh_polymap
geom1.pos = geom_xpos_in[worldid, g1]
geom1.rot = geom_xmat_in[worldid, g1]
geom1.size = geom_size[worldid % geom_size.shape[0], g1]
geom1.normal = wp.vec3(geom1.rot[0, 2], geom1.rot[1, 2], geom1.rot[2, 2]) # plane
geom.index = -1
geom.margin = 0.0
geom2.pos = geom_xpos_in[worldid, g2]
geom2.rot = geom_xmat_in[worldid, g2]
geom2.size = geom_size[worldid % geom_size.shape[0], g2]
geom2.normal = wp.vec3(geom2.rot[0, 2], geom2.rot[1, 2], geom2.rot[2, 2]) # plane
return geom
if geom_type1 == GeomType.MESH:
dataid = geom_dataid[g1]
geom1.vertadr = wp.where(dataid >= 0, mesh_vertadr[dataid], -1)
geom1.vertnum = wp.where(dataid >= 0, mesh_vertnum[dataid], -1)
geom1.graphadr = wp.where(dataid >= 0, mesh_graphadr[dataid], -1)
geom1.mesh_polynum = wp.where(dataid >= 0, mesh_polynum[dataid], -1)
geom1.mesh_polyadr = wp.where(dataid >= 0, mesh_polyadr[dataid], -1)
geom1.vert = mesh_vert
geom1.graph = mesh_graph
geom1.mesh_polynormal = mesh_polynormal
geom1.mesh_polyvertadr = mesh_polyvertadr
geom1.mesh_polyvertnum = mesh_polyvertnum
geom1.mesh_polyvert = mesh_polyvert
geom1.mesh_polymapadr = mesh_polymapadr
geom1.mesh_polymapnum = mesh_polymapnum
geom1.mesh_polymap = mesh_polymap
if geom_type2 == GeomType.MESH:
dataid = geom_dataid[g2]
geom2.vertadr = wp.where(dataid >= 0, mesh_vertadr[dataid], -1)
geom2.vertnum = wp.where(dataid >= 0, mesh_vertnum[dataid], -1)
geom2.graphadr = wp.where(dataid >= 0, mesh_graphadr[dataid], -1)
geom2.mesh_polynum = wp.where(dataid >= 0, mesh_polynum[dataid], -1)
geom2.mesh_polyadr = wp.where(dataid >= 0, mesh_polyadr[dataid], -1)
geom2.vert = mesh_vert
geom2.graph = mesh_graph
geom2.mesh_polynormal = mesh_polynormal
geom2.mesh_polyvertadr = mesh_polyvertadr
geom2.mesh_polyvertnum = mesh_polyvertnum
geom2.mesh_polyvert = mesh_polyvert
geom2.mesh_polymapadr = mesh_polymapadr
geom2.mesh_polymapnum = mesh_polymapnum
geom2.mesh_polymap = mesh_polymap
geom1.index = -1
geom1.margin = 0.0
geom2.index = -1
geom2.margin = 0.0
return geom1, geom2
@wp.func
@@ -1575,11 +1603,6 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_
mesh_polymapadr: wp.array(dtype=int),
mesh_polymapnum: wp.array(dtype=int),
mesh_polymap: wp.array(dtype=int),
hfield_size: wp.array(dtype=wp.vec4),
hfield_nrow: wp.array(dtype=int),
hfield_ncol: wp.array(dtype=int),
hfield_adr: wp.array(dtype=int),
hfield_data: wp.array(dtype=float),
pair_dim: wp.array(dtype=int),
pair_solref: wp.array2d(dtype=wp.vec2),
pair_solreffriction: wp.array2d(dtype=wp.vec2),
@@ -1617,12 +1640,6 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_
return
geoms = collision_pair_in[tid]
g1 = geoms[0]
g2 = geoms[1]
type1 = geom_type[g1]
type2 = geom_type[g2]
worldid = collision_worldid_in[tid]
_, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
@@ -1647,12 +1664,10 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_
worldid,
)
geom1_dataid = geom_dataid[g1]
geom1 = geom(
type1,
geom1_dataid,
geom_size[worldid % geom_size.shape[0], g1],
geom1, geom2 = geom_collision_pair(
geom_type,
geom_dataid,
geom_size,
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
@@ -1667,37 +1682,17 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in[worldid, g1],
geom_xmat_in[worldid, g1],
)
geom2_dataid = geom_dataid[g2]
geom2 = geom(
type2,
geom2_dataid,
geom_size[worldid % geom_size.shape[0], g2],
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
mesh_vert,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in[worldid, g2],
geom_xmat_in[worldid, g2],
geom_xpos_in,
geom_xmat_in,
geoms,
worldid,
)
for i in range(wp.static(len(primitive_collisions_func))):
collision_type1 = wp.static(primitive_collisions_types[i][0])
collision_type2 = wp.static(primitive_collisions_types[i][1])
type1 = geom_type[geoms[0]]
type2 = geom_type[geoms[1]]
if collision_type1 == type1 and collision_type2 == type2:
wp.static(primitive_collisions_func[i])(
naconmax_in,
@@ -1792,11 +1787,6 @@ def primitive_narrowphase(m: Model, d: Data):
m.mesh_polymapadr,
m.mesh_polymapnum,
m.mesh_polymap,
m.hfield_size,
m.hfield_nrow,
m.hfield_ncol,
m.hfield_adr,
m.hfield_data,
m.pair_dim,
m.pair_solref,
m.pair_solreffriction,
@@ -803,8 +803,8 @@ def box_box(
if i != n:
points[n] = points[i]
points[n, 2] *= 0.5
depth[n] = points[n, 2]
points[n, 2] *= 0.5
n += 1
# Set up contact frame
+25 -57
View File
@@ -18,7 +18,7 @@ from typing import Tuple
import warp as wp
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 geom_collision_pair
from mujoco.mjx.third_party.mujoco_warp._src.collision_primitive import write_contact
from mujoco.mjx.third_party.mujoco_warp._src.math import make_frame
from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_mesh
@@ -650,11 +650,6 @@ def _sdf_narrowphase(
mesh_polymapadr: wp.array(dtype=int),
mesh_polymapnum: wp.array(dtype=int),
mesh_polymap: wp.array(dtype=int),
hfield_size: wp.array(dtype=wp.vec4),
hfield_nrow: wp.array(dtype=int),
hfield_ncol: wp.array(dtype=int),
hfield_adr: wp.array(dtype=int),
hfield_data: wp.array(dtype=float),
pair_dim: wp.array(dtype=int),
pair_solref: wp.array2d(dtype=wp.vec2),
pair_solreffriction: wp.array2d(dtype=wp.vec2),
@@ -725,56 +720,34 @@ def _sdf_narrowphase(
worldid,
)
geom_size_id = worldid % geom_size.shape[0]
aabb_id = worldid % geom_aabb.shape[0]
geom1, geom2 = geom_collision_pair(
geom_type,
geom_dataid,
geom_size,
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
mesh_vert,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in,
geom_xmat_in,
geoms,
worldid,
)
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[geom_size_id, g1],
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
mesh_vert,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in[worldid, g1],
geom_xmat_in[worldid, g1],
)
geom2_dataid = geom_dataid[g2]
geom2 = geom(
type2,
geom2_dataid,
geom_size[geom_size_id, g2],
mesh_vertadr,
mesh_vertnum,
mesh_graphadr,
mesh_vert,
mesh_graph,
mesh_polynum,
mesh_polyadr,
mesh_polynormal,
mesh_polyvertadr,
mesh_polyvertnum,
mesh_polyvert,
mesh_polymapadr,
mesh_polymapnum,
mesh_polymap,
geom_xpos_in[worldid, g2],
geom_xmat_in[worldid, g2],
)
g1_plugin = geom_plugin_index[g1]
g2_plugin = geom_plugin_index[g2]
@@ -923,11 +896,6 @@ def sdf_narrowphase(m: Model, d: Data):
m.mesh_polymapadr,
m.mesh_polymapnum,
m.mesh_polymap,
m.hfield_size,
m.hfield_nrow,
m.hfield_ncol,
m.hfield_adr,
m.hfield_data,
m.pair_dim,
m.pair_solref,
m.pair_solreffriction,
+47 -135
View File
@@ -46,14 +46,6 @@ from mujoco.mjx.third_party.mujoco_warp._src.warp_util import kernel as nested_k
wp.set_module_options({"enable_backward": False})
# RK4 tableau
_RK4_A = [
[0.5, 0.0, 0.0],
[0.0, 0.5, 0.0],
[0.0, 0.0, 1.0],
]
_RK4_B = [1.0 / 6.0, 1.0 / 3.0, 1.0 / 3.0, 1.0 / 6.0]
@wp.kernel
def _next_position(
@@ -105,12 +97,7 @@ def _next_position(
qpos_next[qpos_adr + 6] = qpos_quat_new[3]
elif jnttype == JointType.BALL:
qpos_quat = wp.quat(
qpos[qpos_adr + 0],
qpos[qpos_adr + 1],
qpos[qpos_adr + 2],
qpos[qpos_adr + 3],
)
qpos_quat = wp.quat(qpos[qpos_adr + 0], qpos[qpos_adr + 1], qpos[qpos_adr + 2], qpos[qpos_adr + 3])
qvel_ang = wp.vec3(qvel[dof_adr], qvel[dof_adr + 1], qvel[dof_adr + 2]) * qvel_scale_in
qpos_quat_new = math.quat_integrate(qpos_quat, qvel_ang, timestep)
@@ -242,79 +229,45 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None)
# TODO(team): can we assume static timesteps?
# advance activations
if m.na:
wp.launch(
_next_activation,
dim=(d.nworld, m.na),
inputs=[
m.opt.timestep,
m.actuator_dyntype,
m.actuator_actlimited,
m.actuator_dynprm,
m.actuator_actrange,
d.act,
d.act_dot,
1.0,
True,
],
outputs=[
d.act,
],
)
wp.launch(
_next_activation,
dim=(d.nworld, m.na),
inputs=[
m.opt.timestep,
m.actuator_dyntype,
m.actuator_actlimited,
m.actuator_dynprm,
m.actuator_actrange,
d.act,
d.act_dot,
1.0,
True,
],
outputs=[d.act],
)
wp.launch(
_next_velocity,
dim=(d.nworld, m.nv),
inputs=[
m.opt.timestep,
d.qvel,
qacc,
1.0,
],
outputs=[
d.qvel,
],
inputs=[m.opt.timestep, d.qvel, qacc, 1.0],
outputs=[d.qvel],
)
# advance positions with qvel if given, d.qvel otherwise (semi-implicit)
if qvel is not None:
qvel_in = qvel
else:
qvel_in = d.qvel
qvel_in = qvel or d.qvel
wp.launch(
_next_position,
dim=(d.nworld, m.njnt),
inputs=[
m.opt.timestep,
m.jnt_type,
m.jnt_qposadr,
m.jnt_dofadr,
d.qpos,
qvel_in,
1.0,
],
outputs=[
d.qpos,
],
inputs=[m.opt.timestep, m.jnt_type, m.jnt_qposadr, m.jnt_dofadr, d.qpos, qvel_in, 1.0],
outputs=[d.qpos],
)
wp.launch(
_next_time,
dim=(d.nworld,),
inputs=[
m.opt.timestep,
d.nefc,
d.time,
d.nworld,
d.naconmax,
d.njmax,
d.nacon,
d.ncollision,
],
outputs=[
d.time,
],
dim=d.nworld,
inputs=[m.opt.timestep, d.nefc, d.time, d.nworld, d.naconmax, d.njmax, d.nacon, d.ncollision],
outputs=[d.time],
)
wp.copy(d.qacc_warmstart, d.qacc)
@@ -491,6 +444,10 @@ def _rk_accumulate(
@event_scope
def rungekutta4(m: Model, d: Data):
"""Runge-Kutta explicit order 4 integrator."""
# RK4 tableau
A = [0.5, 0.5, 1.0] # diagonal only
B = [1.0 / 6.0, 1.0 / 3.0, 1.0 / 3.0, 1.0 / 6.0]
qpos_t0 = wp.clone(d.qpos)
qvel_t0 = wp.clone(d.qvel)
qvel_rk = wp.zeros((d.nworld, m.nv), dtype=float)
@@ -503,12 +460,10 @@ def rungekutta4(m: Model, d: Data):
act_t0 = None
act_dot_rk = None
A, B = _RK4_A, _RK4_B
_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]
a, b = float(A[i]), B[i + 1]
_rk_perturb_state(m, d, a, qpos_t0, qvel_t0, act_t0)
forward(m, d)
_rk_accumulate(m, d, b, qvel_rk, qacc_rk, act_dot_rk)
@@ -565,8 +520,8 @@ def fwd_position(m: Model, d: Data, factorize: bool = True):
smooth.transmission(m, d)
@cache_kernel
def _create_actuator_velocity_kernel(NV: int):
# TODO(team): sparse actuator_moment version
def _actuator_velocity(m: Model, d: Data):
@nested_kernel(module="unique", enable_backward=False)
def actuator_velocity(
# Data in:
@@ -576,36 +531,22 @@ def _create_actuator_velocity_kernel(NV: int):
actuator_velocity_out: wp.array2d(dtype=float),
):
worldid, actid = wp.tid()
moment_tile = wp.tile_load(actuator_moment_in[worldid, actid], shape=NV)
qvel_tile = wp.tile_load(qvel_in[worldid], shape=NV)
moment_tile = wp.tile_load(actuator_moment_in[worldid, actid], shape=wp.static(m.nv))
qvel_tile = wp.tile_load(qvel_in[worldid], shape=wp.static(m.nv))
moment_qvel_tile = wp.tile_map(wp.mul, moment_tile, qvel_tile)
actuator_velocity_tile = wp.tile_reduce(wp.add, moment_qvel_tile)
actuator_velocity_out[worldid, actid] = actuator_velocity_tile[0]
return actuator_velocity
# TODO(team): sparse actuator_moment version
def _actuator_velocity(m: Model, d: Data):
NV = m.nv
wp.launch_tiled(
_create_actuator_velocity_kernel(NV),
actuator_velocity,
dim=(d.nworld, m.nu),
inputs=[
d.qvel,
d.actuator_moment,
],
outputs=[
d.actuator_velocity,
],
inputs=[d.qvel, d.actuator_moment],
outputs=[d.actuator_velocity],
block_dim=m.block_dim.actuator_velocity,
)
def _tendon_velocity(m: Model, d: Data):
NV = m.nv
@nested_kernel(module="unique", enable_backward=False)
def tendon_velocity(
# Data in:
@@ -615,8 +556,8 @@ def _tendon_velocity(m: Model, d: Data):
ten_velocity_out: wp.array2d(dtype=float),
):
worldid, tenid = wp.tid()
ten_J_tile = wp.tile_load(ten_J_in[worldid, tenid], shape=NV)
qvel_tile = wp.tile_load(qvel_in[worldid], shape=NV)
ten_J_tile = wp.tile_load(ten_J_in[worldid, tenid], shape=wp.static(m.nv))
qvel_tile = wp.tile_load(qvel_in[worldid], shape=wp.static(m.nv))
ten_J_qvel_tile = wp.tile_map(wp.mul, ten_J_tile, qvel_tile)
ten_velocity_tile = wp.tile_reduce(wp.add, ten_J_qvel_tile)
ten_velocity_out[worldid, tenid] = ten_velocity_tile[0]
@@ -624,13 +565,8 @@ def _tendon_velocity(m: Model, d: Data):
wp.launch_tiled(
tendon_velocity,
dim=(d.nworld, m.ntendon),
inputs=[
d.qvel,
d.ten_J,
],
outputs=[
d.ten_velocity,
],
inputs=[d.qvel, d.ten_J],
outputs=[d.ten_velocity],
block_dim=m.block_dim.tendon_velocity,
)
@@ -715,16 +651,14 @@ 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:
act = act_in[worldid, act_last]
ctrl_act = _next_act(
opt_timestep[opt_timestep_id],
opt_timestep[worldid % opt_timestep.shape[0]],
dyntype,
dynprm,
actuator_actrange[actuator_actrange_id, uid],
actuator_actrange[worldid % actuator_actrange.shape[0], uid],
act,
act_dot,
1.0,
@@ -764,8 +698,6 @@ def _actuator_force(
force = gain * ctrl_act + bias
# TODO(team): tendon total force clamping
if actuator_forcelimited[uid]:
forcerange = actuator_forcerange[worldid % actuator_forcerange.shape[0], uid]
force = wp.clamp(force, forcerange[0], forcerange[1])
@@ -958,15 +890,8 @@ def fwd_acceleration(m: Model, d: Data, factorize: bool = False):
wp.launch(
_qfrc_smooth,
dim=(d.nworld, m.nv),
inputs=[
d.qfrc_applied,
d.qfrc_bias,
d.qfrc_passive,
d.qfrc_actuator,
],
outputs=[
d.qfrc_smooth,
],
inputs=[d.qfrc_applied, d.qfrc_bias, d.qfrc_passive, d.qfrc_actuator],
outputs=[d.qfrc_smooth],
)
xfrc_accumulate(m, d, d.qfrc_smooth)
@@ -976,15 +901,6 @@ def fwd_acceleration(m: Model, d: Data, factorize: bool = False):
smooth.solve_m(m, d, d.qacc_smooth, d.qfrc_smooth)
@wp.kernel
def _zero_energy(
# Data out:
energy_out: wp.array(dtype=wp.vec2),
):
tid = wp.tid()
energy_out[tid] = wp.vec2(0.0, 0.0)
@event_scope
def forward(m: Model, d: Data):
"""Forward dynamics."""
@@ -997,11 +913,7 @@ def forward(m: Model, d: Data):
if m.sensor_e_potential == 0: # not computed by sensor
sensor.energy_pos(m, d)
else:
wp.launch(
_zero_energy,
dim=d.nworld,
inputs=[d.energy],
)
d.energy.zero_()
fwd_velocity(m, d)
sensor.sensor_vel(m, d)
@@ -1048,7 +960,7 @@ def step1(m: Model, d: Data):
if m.sensor_e_potential == 0: # not computed by sensor
sensor.energy_pos(m, d)
else:
wp.launch(_zero_energy, dim=d.nworld, inputs=[d.energy])
d.energy.zero_()
fwd_velocity(m, d)
sensor.sensor_vel(m, d)
+6 -7
View File
@@ -589,7 +589,6 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
sdf_initpoints=mjm.opt.sdf_initpoints,
sdf_iterations=mjm.opt.sdf_iterations,
run_collision_detection=True,
legacy_gjk=False,
contact_sensor_maxmatch=64,
),
stat=types.Statistic(
@@ -975,7 +974,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
return m
def _get_padded_sizes(nv: int, njmax: int, nworld: int, is_sparse: bool, tile_size: int):
def _get_padded_sizes(nv: int, njmax: int, is_sparse: bool, tile_size: int):
# if dense - we just pad to the next multiple of 4 for nv, to get the fast load path.
# we pad to the next multiple of tile_size for njmax to avoid out of bounds accesses.
# if sparse - we pad to the next multiple of tile_size for njmax, and nv.
@@ -1006,7 +1005,7 @@ def make_data(
mjm: The model containing kinematic and dynamic information (host).
nworld: Number of worlds.
nconmax: Number of contacts to allocate per world. Contacts exist in large
heterogenous arrays: one world may have more than nconmax contacts.
heterogeneous arrays: one world may have more than nconmax contacts.
njmax: Number of constraints to allocate per world. Constraint arrays are
batched by world: no world may have more than njmax constraints.
naconmax: Number of contacts to allocate for all worlds. Overrides nconmax.
@@ -1047,7 +1046,7 @@ def make_data(
else:
tile_size = types.TILE_SIZE_JTDAJ_DENSE
njmax_padded, nv_padded = _get_padded_sizes(mjm.nv, njmax, nworld, mujoco.mj_isSparse(mjm), tile_size)
njmax_padded, nv_padded = _get_padded_sizes(mjm.nv, njmax, mujoco.mj_isSparse(mjm), tile_size)
# static geoms (attached to the world) have their poses calculated once during make_data instead
# of during each physics step. this speeds up scenes with many static geoms (e.g. terrains)
@@ -1319,7 +1318,7 @@ def put_data(
else:
tile_size = types.TILE_SIZE_JTDAJ_DENSE
njmax_padded, nv_padded = _get_padded_sizes(mjm.nv, njmax, nworld, mujoco.mj_isSparse(mjm), tile_size)
njmax_padded, nv_padded = _get_padded_sizes(mjm.nv, njmax, mujoco.mj_isSparse(mjm), tile_size)
efc_type_fill = np.zeros((nworld, njmax))
efc_id_fill = np.zeros((nworld, njmax))
@@ -1542,8 +1541,8 @@ def get_data_into(
nl = d.nl.numpy()[0]
# efc indexing
# mujoco expects contigious efc ordering for contacts
# this ordering is not guarenteed with mujoco warp, we enforce order here
# mujoco expects contiguous efc ordering for contacts
# this ordering is not guaranteed with mujoco warp, we enforce order here
if nacon > 0:
efc_idx_efl = np.arange(ne + nf + nl)
+2 -20
View File
@@ -455,7 +455,6 @@ def _clock(time_in: wp.array(dtype=float), worldid: int) -> float:
@wp.kernel
def _sensor_pos(
# Model:
ngeom: int,
opt_magnetic: wp.array(dtype=wp.vec3),
body_geomnum: wp.array(dtype=int),
body_geomadr: wp.array(dtype=int),
@@ -504,14 +503,6 @@ def _sensor_pos(
subtree_com_in: wp.array2d(dtype=wp.vec3),
ten_length_in: wp.array2d(dtype=float),
actuator_length_in: wp.array2d(dtype=float),
contact_dist_in: wp.array(dtype=float),
contact_pos_in: wp.array(dtype=wp.vec3),
contact_frame_in: wp.array(dtype=wp.mat33),
contact_geom_in: wp.array(dtype=wp.vec2i),
contact_worldid_in: wp.array(dtype=int),
contact_type_in: wp.array(dtype=int),
nacon_in: wp.array(dtype=int),
collision_pairid_in: wp.array(dtype=wp.vec2i),
# In:
rangefinder_dist_in: wp.array2d(dtype=float),
sensor_collision_in: wp.array4d(dtype=float),
@@ -826,7 +817,6 @@ def sensor_pos(m: Model, d: Data):
_sensor_pos,
dim=(d.nworld, m.sensor_pos_adr.size),
inputs=[
m.ngeom,
m.opt.magnetic,
m.body_geomnum,
m.body_geomadr,
@@ -874,14 +864,6 @@ def sensor_pos(m: Model, d: Data):
d.subtree_com,
d.ten_length,
d.actuator_length,
d.contact.dist,
d.contact.pos,
d.contact.frame,
d.contact.geom,
d.contact.worldid,
d.contact.type,
d.nacon,
d.collision_pairid,
rangefinder_dist,
sensor_collision,
],
@@ -2793,7 +2775,7 @@ def _energy_pos_passive_tendon(
def energy_pos(m: Model, d: Data):
"""Position-dependent energy (potential)."""
wp.launch(_energy_pos_zero, dim=(d.nworld,), outputs=[d.energy])
wp.launch(_energy_pos_zero, dim=d.nworld, outputs=[d.energy])
# init potential energy: -sum_i(body_i.mass * dot(gravity, body_i.pos))
if not m.opt.disableflags & DisableBit.GRAVITY:
@@ -2868,7 +2850,7 @@ def energy_vel(m: Model, d: Data):
wp.launch_tiled(
_energy_vel_kinetic(m.nv),
dim=(d.nworld,),
dim=d.nworld,
inputs=[d.qvel, d.efc.mv],
outputs=[d.energy],
block_dim=m.block_dim.energy_vel_kinetic,
+9 -73
View File
@@ -127,12 +127,7 @@ def _kinematics_level(
xaxis = math.rot_vec_quat(jnt_axis_, xquat)
if jnt_type_ == JointType.BALL:
qloc = wp.quat(
qpos[qadr + 0],
qpos[qadr + 1],
qpos[qadr + 2],
qpos[qadr + 3],
)
qloc = wp.quat(qpos[qadr + 0], qpos[qadr + 1], qpos[qadr + 2], qpos[qadr + 3])
qloc = wp.normalize(qloc)
xquat = math.mul_quat(xquat, qloc)
# correct for off-center rotation
@@ -1797,14 +1792,7 @@ def _transmission(
if jnt_typ == JointType.FREE:
actuator_length_out[worldid, actid] = 0.0
if trntype == TrnType.JOINTINPARENT:
quat = wp.normalize(
wp.quat(
qpos[qadr + 3],
qpos[qadr + 4],
qpos[qadr + 5],
qpos[qadr + 6],
)
)
quat = wp.normalize(wp.quat(qpos[qadr + 3], qpos[qadr + 4], qpos[qadr + 5], qpos[qadr + 6]))
quat_neg = math.quat_inv(quat)
gearaxis = math.rot_vec_quat(wp.spatial_bottom(gear), quat_neg)
actuator_moment_out[worldid, actid, vadr + 0] = gear[0]
@@ -1875,30 +1863,14 @@ def _transmission(
# get Jacobians of axis(jacA) and vec(jac)
# mj_jacPointAxis
jacp, jacr = support.jac(
body_parentid,
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
site_xpos_idslider,
site_bodyid[idslider],
i,
worldid,
body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_idslider, site_bodyid[idslider], i, worldid
)
jacS = jacp
jacA = wp.cross(jacr, axis)
# mj_jacSite
jac, _ = support.jac(
body_parentid,
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
site_xpos_id,
site_bodyid[id],
i,
worldid,
body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_id, site_bodyid[id], i, worldid
)
jac -= jacS
@@ -2023,28 +1995,12 @@ def _transmission(
# TODO(team): parallelize
for i in range(nv):
jacp, jacr = support.jac(
body_parentid,
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
site_xpos,
site_bodyid[siteid],
i,
worldid,
body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos, site_bodyid[siteid], i, worldid
)
# jacref: global Jacobian of reference site
jacpref, jacrref = support.jac(
body_parentid,
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
ref_xpos,
site_bodyid[refid],
i,
worldid,
body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, ref_xpos, site_bodyid[refid], i, worldid
)
jacpdif = jacp - jacpref
@@ -2156,28 +2112,8 @@ def _transmission_body_moment(
normal = wp.vec3(contact_frame[0, 0], contact_frame[0, 1], contact_frame[0, 2])
# get Jacobian difference
jacp1, _ = support.jac(
body_parentid,
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
contact_pos,
b1,
dofid,
worldid,
)
jacp2, _ = support.jac(
body_parentid,
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
contact_pos,
b2,
dofid,
worldid,
)
jacp1, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b1, dofid, worldid)
jacp2, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b2, dofid, worldid)
jacdif = jacp2 - jacp1
# project Jacobian along the normal of the contact frame
@@ -3118,7 +3054,7 @@ def tendon(m: Model, d: Data):
if spatial_site or spatial_geom:
wp.launch(
_spatial_tendon_wrap,
dim=(d.nworld,),
dim=d.nworld,
inputs=[m.ntendon, m.tendon_adr, m.tendon_num, m.wrap_type, m.wrap_objid, d.site_xpos, wrap_geom_xpos],
outputs=[d.ten_wrapadr, d.ten_wrapnum, d.wrap_obj, d.wrap_xpos],
)
+5 -5
View File
@@ -458,7 +458,7 @@ def _linesearch_iterative(m: types.Model, d: types.Data):
"""Iterative linesearch."""
wp.launch(
linesearch_iterative,
dim=(d.nworld,),
dim=d.nworld,
inputs=[
m.nv,
m.opt.impratio,
@@ -1811,7 +1811,7 @@ def _update_gradient(m: types.Model, d: types.Data):
if m.nv < 32:
wp.launch_tiled(
update_gradient_cholesky(m.nv),
dim=(d.nworld,),
dim=d.nworld,
inputs=[d.efc.grad, d.efc.h, d.efc.done],
outputs=[d.efc.Mgrad],
block_dim=m.block_dim.update_gradient_cholesky,
@@ -1819,7 +1819,7 @@ def _update_gradient(m: types.Model, d: types.Data):
else:
wp.launch_tiled(
update_gradient_cholesky_blocked(16),
dim=(d.nworld,),
dim=d.nworld,
inputs=[
d.efc.grad.reshape(shape=(d.nworld, m.nv, 1)),
d.efc.h,
@@ -1982,7 +1982,7 @@ def _solver_iteration(
if m.opt.solver == types.SolverType.CG:
wp.launch(
solve_beta,
dim=(d.nworld,),
dim=d.nworld,
inputs=[m.nv, d.efc.grad, d.efc.Mgrad, d.efc.prev_grad, d.efc.prev_Mgrad, d.efc.done],
outputs=[d.efc.beta],
)
@@ -1998,7 +1998,7 @@ def _solver_iteration(
wp.launch(
solve_done,
dim=(d.nworld,),
dim=d.nworld,
inputs=[
m.nv,
m.opt.tolerance,
-2
View File
@@ -637,7 +637,6 @@ class Option:
run_collision_detection: if False, skips collision detection and allows user-populated
contacts during the physics step (as opposed to DisableBit.CONTACT which explicitly
zeros out the contacts at each step)
legacy_gjk: run legacy gjk algorithm
contact_sensor_maxmatch: max number of contacts considered by contact sensor matching criteria
contacts matched after this value is exceded will be ignored
"""
@@ -671,7 +670,6 @@ class Option:
broadphase_filter: int
graph_conditional: bool
run_collision_detection: bool
legacy_gjk: bool
contact_sensor_maxmatch: int
-3
View File
@@ -109,7 +109,6 @@ def _collision_shim(
opt__ccd_iterations: int,
opt__ccd_tolerance: wp.array(dtype=float),
opt__disableflags: int,
opt__legacy_gjk: bool,
opt__sdf_initpoints: int,
opt__sdf_iterations: int,
# Data
@@ -192,7 +191,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.legacy_gjk = opt__legacy_gjk
_m.opt.sdf_initpoints = opt__sdf_initpoints
_m.opt.sdf_iterations = opt__sdf_iterations
_m.pair_dim = pair_dim
@@ -344,7 +342,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.legacy_gjk,
m.opt._impl.sdf_initpoints,
m.opt._impl.sdf_iterations,
d._impl.naconmax,
-6
View File
@@ -344,7 +344,6 @@ def _forward_shim(
opt__impratio: wp.array(dtype=float),
opt__is_sparse: bool,
opt__iterations: int,
opt__legacy_gjk: bool,
opt__ls_iterations: int,
opt__ls_parallel: bool,
opt__ls_parallel_min_step: float,
@@ -714,7 +713,6 @@ def _forward_shim(
_m.opt.impratio = opt__impratio
_m.opt.is_sparse = opt__is_sparse
_m.opt.iterations = opt__iterations
_m.opt.legacy_gjk = opt__legacy_gjk
_m.opt.ls_iterations = opt__ls_iterations
_m.opt.ls_parallel = opt__ls_parallel
_m.opt.ls_parallel_min_step = opt__ls_parallel_min_step
@@ -1519,7 +1517,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m.opt.impratio,
m.opt._impl.is_sparse,
m.opt.iterations,
m.opt._impl.legacy_gjk,
m.opt.ls_iterations,
m.opt._impl.ls_parallel,
m.opt._impl.ls_parallel_min_step,
@@ -2138,7 +2135,6 @@ def _step_shim(
opt__integrator: int,
opt__is_sparse: bool,
opt__iterations: int,
opt__legacy_gjk: bool,
opt__ls_iterations: int,
opt__ls_parallel: bool,
opt__ls_parallel_min_step: float,
@@ -2510,7 +2506,6 @@ def _step_shim(
_m.opt.integrator = opt__integrator
_m.opt.is_sparse = opt__is_sparse
_m.opt.iterations = opt__iterations
_m.opt.legacy_gjk = opt__legacy_gjk
_m.opt.ls_iterations = opt__ls_iterations
_m.opt.ls_parallel = opt__ls_parallel
_m.opt.ls_parallel_min_step = opt__ls_parallel_min_step
@@ -3317,7 +3312,6 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m.opt.integrator,
m.opt._impl.is_sparse,
m.opt.iterations,
m.opt._impl.legacy_gjk,
m.opt.ls_iterations,
m.opt._impl.ls_parallel,
m.opt._impl.ls_parallel_min_step,
-5
View File
@@ -94,7 +94,6 @@ class OptionWarp(PyTreeNode):
graph_conditional: bool
has_fluid: bool
is_sparse: bool
legacy_gjk: bool
ls_parallel: bool
ls_parallel_min_step: float
run_collision_detection: bool
@@ -765,7 +764,6 @@ _NDIM = {
'opt__integrator': 0,
'opt__is_sparse': 0,
'opt__iterations': 0,
'opt__legacy_gjk': 0,
'opt__ls_iterations': 0,
'opt__ls_parallel': 0,
'opt__ls_parallel_min_step': 0,
@@ -885,7 +883,6 @@ _NDIM = {
'integrator': 0,
'is_sparse': 0,
'iterations': 0,
'legacy_gjk': 0,
'ls_iterations': 0,
'ls_parallel': 0,
'ls_parallel_min_step': 0,
@@ -1305,7 +1302,6 @@ _BATCH_DIM = {
'opt__integrator': False,
'opt__is_sparse': False,
'opt__iterations': False,
'opt__legacy_gjk': False,
'opt__ls_iterations': False,
'opt__ls_parallel': False,
'opt__ls_parallel_min_step': False,
@@ -1425,7 +1421,6 @@ _BATCH_DIM = {
'integrator': False,
'is_sparse': False,
'iterations': False,
'legacy_gjk': False,
'ls_iterations': False,
'ls_parallel': False,
'ls_parallel_min_step': False,