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. # limitations under the License.
# ============================================================================== # ==============================================================================
from typing import Tuple
import warp as wp 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 ccd
from mujoco.mjx.third_party.mujoco_warp._src.collision_gjk import multicontact 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 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 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.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 make_frame
from mujoco.mjx.third_party.mujoco_warp._src.math import upper_trid_index 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_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_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_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 Data
from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType from mujoco.mjx.third_party.mujoco_warp._src.types import GeomType
from mujoco.mjx.third_party.mujoco_warp._src.types import Model from mujoco.mjx.third_party.mujoco_warp._src.types import Model
@@ -82,14 +81,120 @@ def _check_convex_collision_pairs():
assert _check_convex_collision_pairs(), "_CONVEX_COLLISION_PAIRS is in invalid order." 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 @cache_kernel
def ccd_kernel_builder( def ccd_kernel_builder(
legacy_gjk: bool,
geomtype1: int, geomtype1: int,
geomtype2: int, geomtype2: int,
ccd_iterations: int, ccd_iterations: int,
epa_exact_neg_distance: bool,
depth_extension: float,
is_hfield: bool, is_hfield: bool,
use_multiccd: bool, use_multiccd: bool,
): ):
@@ -155,108 +260,81 @@ def ccd_kernel_builder(
contact_geomcollisionid_out: wp.array(dtype=int), contact_geomcollisionid_out: wp.array(dtype=int),
nacon_out: wp.array(dtype=int), nacon_out: wp.array(dtype=int),
) -> int: ) -> int:
# TODO(kbayes): remove legacy GJK once multicontact can be enabled points = mat3c()
if wp.static(legacy_gjk): witness1 = mat3c()
simplex, normal = gjk_legacy( witness2 = mat3c()
ccd_iterations, geom1.margin = margin
geom1, geom2.margin = margin
geom2, if pairid[1] >= 0:
geomtype1, # if collision sensor, set large cutoff to work with various sensor cutoff values
geomtype2, cutoff = 1.0e32
)
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)
else: else:
points = mat3c() cutoff = 0.0
witness1 = mat3c() dist, ncontact, w1, w2, idx = ccd(
witness2 = mat3c() opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]],
geom1.margin = margin cutoff,
geom2.margin = margin ccd_iterations,
if pairid[1] >= 0: geom1,
# if collision sensor, set large cutoff to work with various sensor cutoff values geom2,
cutoff = 1.0e32 geomtype1,
else: geomtype2,
cutoff = 0.0 x1,
dist, ncontact, w1, w2, idx = ccd( x2,
opt_ccd_tolerance[worldid % opt_ccd_tolerance.shape[0]], epa_vert_in[tid],
cutoff, epa_vert1_in[tid],
ccd_iterations, epa_vert2_in[tid],
geom1, epa_vert_index1_in[tid],
geom2, epa_vert_index2_in[tid],
geomtype1, epa_face_in[tid],
geomtype2, epa_pr_in[tid],
x1, epa_norm2_in[tid],
x2, epa_index_in[tid],
epa_vert_in[tid], epa_map_in[tid],
epa_vert1_in[tid], epa_horizon_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: if dist >= 0.0 and pairid[1] == -1:
return 0 return 0
witness1[0] = w1 witness1[0] = w1
witness2[0] = w2 witness2[0] = w2
if wp.static(use_multiccd): if wp.static(use_multiccd):
if ( if (
geom1.margin == 0.0 geom1.margin == 0.0
and geom2.margin == 0.0 and geom2.margin == 0.0
and (geomtype1 == GeomType.BOX or (geomtype1 == GeomType.MESH and geom1.mesh_polyadr > -1)) 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)) and (geomtype2 == GeomType.BOX or (geomtype2 == GeomType.MESH and geom2.mesh_polyadr > -1))
): ):
ncontact, witness1, witness2 = multicontact( ncontact, witness1, witness2 = multicontact(
multiccd_polygon_in[tid], multiccd_polygon_in[tid],
multiccd_clipped_in[tid], multiccd_clipped_in[tid],
multiccd_pnormal_in[tid], multiccd_pnormal_in[tid],
multiccd_pdist_in[tid], multiccd_pdist_in[tid],
multiccd_idx1_in[tid], multiccd_idx1_in[tid],
multiccd_idx2_in[tid], multiccd_idx2_in[tid],
multiccd_n1_in[tid], multiccd_n1_in[tid],
multiccd_n2_in[tid], multiccd_n2_in[tid],
multiccd_endvert_in[tid], multiccd_endvert_in[tid],
multiccd_face1_in[tid], multiccd_face1_in[tid],
multiccd_face2_in[tid], multiccd_face2_in[tid],
epa_vert1_in[tid], epa_vert1_in[tid],
epa_vert2_in[tid], epa_vert2_in[tid],
epa_vert_index1_in[tid], epa_vert_index1_in[tid],
epa_vert_index2_in[tid], epa_vert_index2_in[tid],
epa_face_in[tid, idx], epa_face_in[tid, idx],
w1, w1,
w2, w2,
geom1, geom1,
geom2, geom2,
geomtype1, geomtype1,
geomtype2, geomtype2,
) )
for i in range(ncontact): for i in range(ncontact):
points[i] = 0.5 * (witness1[i] + witness2[i]) points[i] = 0.5 * (witness1[i] + witness2[i])
normal = witness1[0] - witness2[0] normal = witness1[0] - witness2[0]
frame = make_frame(normal) frame = make_frame(normal)
# flip if collision sensor # flip if collision sensor
if pairid[1] >= 0: if pairid[1] >= 0:
@@ -406,7 +484,7 @@ def ccd_kernel_builder(
# height field filter # height field filter
if wp.static(is_hfield): 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 geom_dataid, geom_aabb, geom_rbound, geom_margin, hfield_size, geom_xpos_in, geom_xmat_in, worldid, g1, g2
) )
if no_hf_collision: if no_hf_collision:
@@ -434,13 +512,10 @@ def ccd_kernel_builder(
worldid, worldid,
) )
geom_size_id = worldid % geom_size.shape[0] geom1, geom2 = geom_collision_pair(
geom_type,
geom1_dataid = geom_dataid[g1] geom_dataid,
geom1 = geom( geom_size,
geomtype1,
geom1_dataid,
geom_size[geom_size_id, g1],
mesh_vertadr, mesh_vertadr,
mesh_vertnum, mesh_vertnum,
mesh_graphadr, mesh_graphadr,
@@ -455,35 +530,16 @@ def ccd_kernel_builder(
mesh_polymapadr, mesh_polymapadr,
mesh_polymapnum, mesh_polymapnum,
mesh_polymap, mesh_polymap,
geom_xpos_in[worldid, g1], geom_xpos_in,
geom_xmat_in[worldid, g1], geom_xmat_in,
) geoms,
worldid,
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],
) )
# see MuJoCo mjc_ConvexHField # see MuJoCo mjc_ConvexHField
if wp.static(is_hfield): if wp.static(is_hfield):
geom1_dataid = geom_dataid[g1]
# height field subgrid # height field subgrid
nrow = hfield_nrow[geom1_dataid] nrow = hfield_nrow[geom1_dataid]
ncol = hfield_ncol[geom1_dataid] ncol = hfield_ncol[geom1_dataid]
@@ -546,11 +602,10 @@ def ccd_kernel_builder(
# prism center # prism center
x1 = geom1.pos x1 = geom1.pos
if wp.static(not legacy_gjk): x1_ = wp.vec3(0.0, 0.0, 0.0)
x1_ = wp.vec3(0.0, 0.0, 0.0) for i in range(6):
for i in range(6): x1_ += prism[i]
x1_ += prism[i] x1 += geom1.rot @ (x1_ / 6.0)
x1 += geom1.rot @ (x1_ / 6.0)
ncontact = eval_ccd_write_contact( ncontact = eval_ccd_write_contact(
opt_ccd_tolerance, 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 kernel for each type of convex collision pair present in the model, avoiding unnecessary
computations for non-existent pair types. 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): 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 return
@@ -749,7 +803,7 @@ def convex_narrowphase(m: Model, d: Data):
g2 = geom_pair[1].value g2 = geom_pair[1].value
if m.geom_pair_type_count[upper_trid_index(len(GeomType), g1, g2)]: if m.geom_pair_type_count[upper_trid_index(len(GeomType), g1, g2)]:
wp.launch( 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, dim=d.naconmax,
inputs=[ inputs=[
m.opt.ccd_tolerance, 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 @wp.func
def geom( def geom_collision_pair(
# kernel_analyzer: off
# Model: # Model:
geom_type: int, geom_type: wp.array(dtype=int),
geom_dataid: int, geom_dataid: wp.array(dtype=int),
geom_size: wp.vec3, geom_size: wp.array2d(dtype=wp.vec3),
mesh_vertadr: wp.array(dtype=int), mesh_vertadr: wp.array(dtype=int),
mesh_vertnum: wp.array(dtype=int), mesh_vertnum: wp.array(dtype=int),
mesh_graphadr: wp.array(dtype=int), mesh_graphadr: wp.array(dtype=int),
@@ -100,44 +99,73 @@ def geom(
mesh_polymapnum: wp.array(dtype=int), mesh_polymapnum: wp.array(dtype=int),
mesh_polymap: wp.array(dtype=int), mesh_polymap: wp.array(dtype=int),
# Data in: # Data in:
geom_xpos_in: wp.vec3, geom_xpos_in: wp.array2d(dtype=wp.vec3),
geom_xmat_in: wp.mat33, geom_xmat_in: wp.array2d(dtype=wp.mat33),
# kernel_analyzer: on # In:
) -> Geom: geoms: wp.vec2i,
geom = Geom() worldid: int,
geom.pos = geom_xpos_in ) -> Tuple[Geom, Geom]:
geom.rot = geom_xmat_in geom1 = Geom()
geom.size = geom_size geom2 = Geom()
geom.normal = wp.vec3(geom_xmat_in[0, 2], geom_xmat_in[1, 2], geom_xmat_in[2, 2]) # plane
if geom_type == GeomType.MESH: g1 = geoms[0]
if geom_dataid >= 0: g2 = geoms[1]
geom.vertadr = mesh_vertadr[geom_dataid] geom_type1 = geom_type[g1]
geom.vertnum = mesh_vertnum[geom_dataid] geom_type2 = geom_type[g2]
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
geom.vert = mesh_vert geom1.pos = geom_xpos_in[worldid, g1]
geom.graph = mesh_graph geom1.rot = geom_xmat_in[worldid, g1]
geom.mesh_polynormal = mesh_polynormal geom1.size = geom_size[worldid % geom_size.shape[0], g1]
geom.mesh_polyvertadr = mesh_polyvertadr geom1.normal = wp.vec3(geom1.rot[0, 2], geom1.rot[1, 2], geom1.rot[2, 2]) # plane
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
geom.index = -1 geom2.pos = geom_xpos_in[worldid, g2]
geom.margin = 0.0 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 @wp.func
@@ -1575,11 +1603,6 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_
mesh_polymapadr: wp.array(dtype=int), mesh_polymapadr: wp.array(dtype=int),
mesh_polymapnum: wp.array(dtype=int), mesh_polymapnum: wp.array(dtype=int),
mesh_polymap: 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_dim: wp.array(dtype=int),
pair_solref: wp.array2d(dtype=wp.vec2), pair_solref: wp.array2d(dtype=wp.vec2),
pair_solreffriction: 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 return
geoms = collision_pair_in[tid] geoms = collision_pair_in[tid]
g1 = geoms[0]
g2 = geoms[1]
type1 = geom_type[g1]
type2 = geom_type[g2]
worldid = collision_worldid_in[tid] worldid = collision_worldid_in[tid]
_, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params( _, margin, gap, condim, friction, solref, solreffriction, solimp = contact_params(
@@ -1647,12 +1664,10 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_
worldid, worldid,
) )
geom1_dataid = geom_dataid[g1] geom1, geom2 = geom_collision_pair(
geom_type,
geom1 = geom( geom_dataid,
type1, geom_size,
geom1_dataid,
geom_size[worldid % geom_size.shape[0], g1],
mesh_vertadr, mesh_vertadr,
mesh_vertnum, mesh_vertnum,
mesh_graphadr, mesh_graphadr,
@@ -1667,37 +1682,17 @@ def _create_narrowphase_kernel(primitive_collisions_types, primitive_collisions_
mesh_polymapadr, mesh_polymapadr,
mesh_polymapnum, mesh_polymapnum,
mesh_polymap, mesh_polymap,
geom_xpos_in[worldid, g1], geom_xpos_in,
geom_xmat_in[worldid, g1], geom_xmat_in,
) geoms,
worldid,
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],
) )
for i in range(wp.static(len(primitive_collisions_func))): for i in range(wp.static(len(primitive_collisions_func))):
collision_type1 = wp.static(primitive_collisions_types[i][0]) collision_type1 = wp.static(primitive_collisions_types[i][0])
collision_type2 = wp.static(primitive_collisions_types[i][1]) 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: if collision_type1 == type1 and collision_type2 == type2:
wp.static(primitive_collisions_func[i])( wp.static(primitive_collisions_func[i])(
naconmax_in, naconmax_in,
@@ -1792,11 +1787,6 @@ def primitive_narrowphase(m: Model, d: Data):
m.mesh_polymapadr, m.mesh_polymapadr,
m.mesh_polymapnum, m.mesh_polymapnum,
m.mesh_polymap, m.mesh_polymap,
m.hfield_size,
m.hfield_nrow,
m.hfield_ncol,
m.hfield_adr,
m.hfield_data,
m.pair_dim, m.pair_dim,
m.pair_solref, m.pair_solref,
m.pair_solreffriction, m.pair_solreffriction,
@@ -803,8 +803,8 @@ def box_box(
if i != n: if i != n:
points[n] = points[i] points[n] = points[i]
points[n, 2] *= 0.5
depth[n] = points[n, 2] depth[n] = points[n, 2]
points[n, 2] *= 0.5
n += 1 n += 1
# Set up contact frame # Set up contact frame
+25 -57
View File
@@ -18,7 +18,7 @@ from typing import Tuple
import warp as wp 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 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.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 make_frame
from mujoco.mjx.third_party.mujoco_warp._src.ray import ray_mesh 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_polymapadr: wp.array(dtype=int),
mesh_polymapnum: wp.array(dtype=int), mesh_polymapnum: wp.array(dtype=int),
mesh_polymap: 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_dim: wp.array(dtype=int),
pair_solref: wp.array2d(dtype=wp.vec2), pair_solref: wp.array2d(dtype=wp.vec2),
pair_solreffriction: wp.array2d(dtype=wp.vec2), pair_solreffriction: wp.array2d(dtype=wp.vec2),
@@ -725,56 +720,34 @@ def _sdf_narrowphase(
worldid, worldid,
) )
geom_size_id = worldid % geom_size.shape[0] geom1, geom2 = geom_collision_pair(
aabb_id = worldid % geom_aabb.shape[0] 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] g1 = geoms[0]
type1 = geom_type[g1] 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] g1_plugin = geom_plugin_index[g1]
g2_plugin = geom_plugin_index[g2] g2_plugin = geom_plugin_index[g2]
@@ -923,11 +896,6 @@ def sdf_narrowphase(m: Model, d: Data):
m.mesh_polymapadr, m.mesh_polymapadr,
m.mesh_polymapnum, m.mesh_polymapnum,
m.mesh_polymap, m.mesh_polymap,
m.hfield_size,
m.hfield_nrow,
m.hfield_ncol,
m.hfield_adr,
m.hfield_data,
m.pair_dim, m.pair_dim,
m.pair_solref, m.pair_solref,
m.pair_solreffriction, 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}) 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 @wp.kernel
def _next_position( def _next_position(
@@ -105,12 +97,7 @@ def _next_position(
qpos_next[qpos_adr + 6] = qpos_quat_new[3] qpos_next[qpos_adr + 6] = qpos_quat_new[3]
elif jnttype == JointType.BALL: elif jnttype == JointType.BALL:
qpos_quat = wp.quat( qpos_quat = wp.quat(qpos[qpos_adr + 0], qpos[qpos_adr + 1], qpos[qpos_adr + 2], qpos[qpos_adr + 3])
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 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) 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? # TODO(team): can we assume static timesteps?
# advance activations # advance activations
if m.na: wp.launch(
wp.launch( _next_activation,
_next_activation, dim=(d.nworld, m.na),
dim=(d.nworld, m.na), inputs=[
inputs=[ m.opt.timestep,
m.opt.timestep, m.actuator_dyntype,
m.actuator_dyntype, m.actuator_actlimited,
m.actuator_actlimited, m.actuator_dynprm,
m.actuator_dynprm, m.actuator_actrange,
m.actuator_actrange, d.act,
d.act, d.act_dot,
d.act_dot, 1.0,
1.0, True,
True, ],
], outputs=[d.act],
outputs=[ )
d.act,
],
)
wp.launch( wp.launch(
_next_velocity, _next_velocity,
dim=(d.nworld, m.nv), dim=(d.nworld, m.nv),
inputs=[ inputs=[m.opt.timestep, d.qvel, qacc, 1.0],
m.opt.timestep, outputs=[d.qvel],
d.qvel,
qacc,
1.0,
],
outputs=[
d.qvel,
],
) )
# advance positions with qvel if given, d.qvel otherwise (semi-implicit) # advance positions with qvel if given, d.qvel otherwise (semi-implicit)
if qvel is not None: qvel_in = qvel or d.qvel
qvel_in = qvel
else:
qvel_in = d.qvel
wp.launch( wp.launch(
_next_position, _next_position,
dim=(d.nworld, m.njnt), dim=(d.nworld, m.njnt),
inputs=[ inputs=[m.opt.timestep, m.jnt_type, m.jnt_qposadr, m.jnt_dofadr, d.qpos, qvel_in, 1.0],
m.opt.timestep, outputs=[d.qpos],
m.jnt_type,
m.jnt_qposadr,
m.jnt_dofadr,
d.qpos,
qvel_in,
1.0,
],
outputs=[
d.qpos,
],
) )
wp.launch( wp.launch(
_next_time, _next_time,
dim=(d.nworld,), dim=d.nworld,
inputs=[ inputs=[m.opt.timestep, d.nefc, d.time, d.nworld, d.naconmax, d.njmax, d.nacon, d.ncollision],
m.opt.timestep, outputs=[d.time],
d.nefc,
d.time,
d.nworld,
d.naconmax,
d.njmax,
d.nacon,
d.ncollision,
],
outputs=[
d.time,
],
) )
wp.copy(d.qacc_warmstart, d.qacc) wp.copy(d.qacc_warmstart, d.qacc)
@@ -491,6 +444,10 @@ def _rk_accumulate(
@event_scope @event_scope
def rungekutta4(m: Model, d: Data): def rungekutta4(m: Model, d: Data):
"""Runge-Kutta explicit order 4 integrator.""" """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) qpos_t0 = wp.clone(d.qpos)
qvel_t0 = wp.clone(d.qvel) qvel_t0 = wp.clone(d.qvel)
qvel_rk = wp.zeros((d.nworld, m.nv), dtype=float) qvel_rk = wp.zeros((d.nworld, m.nv), dtype=float)
@@ -503,12 +460,10 @@ def rungekutta4(m: Model, d: Data):
act_t0 = None act_t0 = None
act_dot_rk = None act_dot_rk = None
A, B = _RK4_A, _RK4_B
_rk_accumulate(m, d, B[0], qvel_rk, qacc_rk, act_dot_rk) _rk_accumulate(m, d, B[0], qvel_rk, qacc_rk, act_dot_rk)
for i in range(3): 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) _rk_perturb_state(m, d, a, qpos_t0, qvel_t0, act_t0)
forward(m, d) forward(m, d)
_rk_accumulate(m, d, b, qvel_rk, qacc_rk, act_dot_rk) _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) smooth.transmission(m, d)
@cache_kernel # TODO(team): sparse actuator_moment version
def _create_actuator_velocity_kernel(NV: int): def _actuator_velocity(m: Model, d: Data):
@nested_kernel(module="unique", enable_backward=False) @nested_kernel(module="unique", enable_backward=False)
def actuator_velocity( def actuator_velocity(
# Data in: # Data in:
@@ -576,36 +531,22 @@ def _create_actuator_velocity_kernel(NV: int):
actuator_velocity_out: wp.array2d(dtype=float), actuator_velocity_out: wp.array2d(dtype=float),
): ):
worldid, actid = wp.tid() worldid, actid = wp.tid()
moment_tile = wp.tile_load(actuator_moment_in[worldid, actid], 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=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) 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_tile = wp.tile_reduce(wp.add, moment_qvel_tile)
actuator_velocity_out[worldid, actid] = actuator_velocity_tile[0] 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( wp.launch_tiled(
_create_actuator_velocity_kernel(NV), actuator_velocity,
dim=(d.nworld, m.nu), dim=(d.nworld, m.nu),
inputs=[ inputs=[d.qvel, d.actuator_moment],
d.qvel, outputs=[d.actuator_velocity],
d.actuator_moment,
],
outputs=[
d.actuator_velocity,
],
block_dim=m.block_dim.actuator_velocity, block_dim=m.block_dim.actuator_velocity,
) )
def _tendon_velocity(m: Model, d: Data): def _tendon_velocity(m: Model, d: Data):
NV = m.nv
@nested_kernel(module="unique", enable_backward=False) @nested_kernel(module="unique", enable_backward=False)
def tendon_velocity( def tendon_velocity(
# Data in: # Data in:
@@ -615,8 +556,8 @@ def _tendon_velocity(m: Model, d: Data):
ten_velocity_out: wp.array2d(dtype=float), ten_velocity_out: wp.array2d(dtype=float),
): ):
worldid, tenid = wp.tid() worldid, tenid = wp.tid()
ten_J_tile = wp.tile_load(ten_J_in[worldid, tenid], 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=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_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_tile = wp.tile_reduce(wp.add, ten_J_qvel_tile)
ten_velocity_out[worldid, tenid] = ten_velocity_tile[0] ten_velocity_out[worldid, tenid] = ten_velocity_tile[0]
@@ -624,13 +565,8 @@ def _tendon_velocity(m: Model, d: Data):
wp.launch_tiled( wp.launch_tiled(
tendon_velocity, tendon_velocity,
dim=(d.nworld, m.ntendon), dim=(d.nworld, m.ntendon),
inputs=[ inputs=[d.qvel, d.ten_J],
d.qvel, outputs=[d.ten_velocity],
d.ten_J,
],
outputs=[
d.ten_velocity,
],
block_dim=m.block_dim.tendon_velocity, block_dim=m.block_dim.tendon_velocity,
) )
@@ -715,16 +651,14 @@ def _actuator_force(
act_dot_out[worldid, act_last] = act_dot act_dot_out[worldid, act_last] = act_dot
if actuator_actearly[uid]: 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: if dyntype == DynType.INTEGRATOR or dyntype == DynType.NONE:
act = act_in[worldid, act_last] act = act_in[worldid, act_last]
ctrl_act = _next_act( ctrl_act = _next_act(
opt_timestep[opt_timestep_id], opt_timestep[worldid % opt_timestep.shape[0]],
dyntype, dyntype,
dynprm, dynprm,
actuator_actrange[actuator_actrange_id, uid], actuator_actrange[worldid % actuator_actrange.shape[0], uid],
act, act,
act_dot, act_dot,
1.0, 1.0,
@@ -764,8 +698,6 @@ def _actuator_force(
force = gain * ctrl_act + bias force = gain * ctrl_act + bias
# TODO(team): tendon total force clamping
if actuator_forcelimited[uid]: if actuator_forcelimited[uid]:
forcerange = actuator_forcerange[worldid % actuator_forcerange.shape[0], uid] forcerange = actuator_forcerange[worldid % actuator_forcerange.shape[0], uid]
force = wp.clamp(force, forcerange[0], forcerange[1]) force = wp.clamp(force, forcerange[0], forcerange[1])
@@ -958,15 +890,8 @@ def fwd_acceleration(m: Model, d: Data, factorize: bool = False):
wp.launch( wp.launch(
_qfrc_smooth, _qfrc_smooth,
dim=(d.nworld, m.nv), dim=(d.nworld, m.nv),
inputs=[ inputs=[d.qfrc_applied, d.qfrc_bias, d.qfrc_passive, d.qfrc_actuator],
d.qfrc_applied, outputs=[d.qfrc_smooth],
d.qfrc_bias,
d.qfrc_passive,
d.qfrc_actuator,
],
outputs=[
d.qfrc_smooth,
],
) )
xfrc_accumulate(m, d, 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) 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 @event_scope
def forward(m: Model, d: Data): def forward(m: Model, d: Data):
"""Forward dynamics.""" """Forward dynamics."""
@@ -997,11 +913,7 @@ def forward(m: Model, d: Data):
if m.sensor_e_potential == 0: # not computed by sensor if m.sensor_e_potential == 0: # not computed by sensor
sensor.energy_pos(m, d) sensor.energy_pos(m, d)
else: else:
wp.launch( d.energy.zero_()
_zero_energy,
dim=d.nworld,
inputs=[d.energy],
)
fwd_velocity(m, d) fwd_velocity(m, d)
sensor.sensor_vel(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 if m.sensor_e_potential == 0: # not computed by sensor
sensor.energy_pos(m, d) sensor.energy_pos(m, d)
else: else:
wp.launch(_zero_energy, dim=d.nworld, inputs=[d.energy]) d.energy.zero_()
fwd_velocity(m, d) fwd_velocity(m, d)
sensor.sensor_vel(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_initpoints=mjm.opt.sdf_initpoints,
sdf_iterations=mjm.opt.sdf_iterations, sdf_iterations=mjm.opt.sdf_iterations,
run_collision_detection=True, run_collision_detection=True,
legacy_gjk=False,
contact_sensor_maxmatch=64, contact_sensor_maxmatch=64,
), ),
stat=types.Statistic( stat=types.Statistic(
@@ -975,7 +974,7 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
return m 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. # 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. # 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. # 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). mjm: The model containing kinematic and dynamic information (host).
nworld: Number of worlds. nworld: Number of worlds.
nconmax: Number of contacts to allocate per world. Contacts exist in large 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 njmax: Number of constraints to allocate per world. Constraint arrays are
batched by world: no world may have more than njmax constraints. batched by world: no world may have more than njmax constraints.
naconmax: Number of contacts to allocate for all worlds. Overrides nconmax. naconmax: Number of contacts to allocate for all worlds. Overrides nconmax.
@@ -1047,7 +1046,7 @@ def make_data(
else: else:
tile_size = types.TILE_SIZE_JTDAJ_DENSE 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 # 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) # of during each physics step. this speeds up scenes with many static geoms (e.g. terrains)
@@ -1319,7 +1318,7 @@ def put_data(
else: else:
tile_size = types.TILE_SIZE_JTDAJ_DENSE 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_type_fill = np.zeros((nworld, njmax))
efc_id_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] nl = d.nl.numpy()[0]
# efc indexing # efc indexing
# mujoco expects contigious efc ordering for contacts # mujoco expects contiguous efc ordering for contacts
# this ordering is not guarenteed with mujoco warp, we enforce order here # this ordering is not guaranteed with mujoco warp, we enforce order here
if nacon > 0: if nacon > 0:
efc_idx_efl = np.arange(ne + nf + nl) 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 @wp.kernel
def _sensor_pos( def _sensor_pos(
# Model: # Model:
ngeom: int,
opt_magnetic: wp.array(dtype=wp.vec3), opt_magnetic: wp.array(dtype=wp.vec3),
body_geomnum: wp.array(dtype=int), body_geomnum: wp.array(dtype=int),
body_geomadr: 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), subtree_com_in: wp.array2d(dtype=wp.vec3),
ten_length_in: wp.array2d(dtype=float), ten_length_in: wp.array2d(dtype=float),
actuator_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: # In:
rangefinder_dist_in: wp.array2d(dtype=float), rangefinder_dist_in: wp.array2d(dtype=float),
sensor_collision_in: wp.array4d(dtype=float), sensor_collision_in: wp.array4d(dtype=float),
@@ -826,7 +817,6 @@ def sensor_pos(m: Model, d: Data):
_sensor_pos, _sensor_pos,
dim=(d.nworld, m.sensor_pos_adr.size), dim=(d.nworld, m.sensor_pos_adr.size),
inputs=[ inputs=[
m.ngeom,
m.opt.magnetic, m.opt.magnetic,
m.body_geomnum, m.body_geomnum,
m.body_geomadr, m.body_geomadr,
@@ -874,14 +864,6 @@ def sensor_pos(m: Model, d: Data):
d.subtree_com, d.subtree_com,
d.ten_length, d.ten_length,
d.actuator_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, rangefinder_dist,
sensor_collision, sensor_collision,
], ],
@@ -2793,7 +2775,7 @@ def _energy_pos_passive_tendon(
def energy_pos(m: Model, d: Data): def energy_pos(m: Model, d: Data):
"""Position-dependent energy (potential).""" """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)) # init potential energy: -sum_i(body_i.mass * dot(gravity, body_i.pos))
if not m.opt.disableflags & DisableBit.GRAVITY: if not m.opt.disableflags & DisableBit.GRAVITY:
@@ -2868,7 +2850,7 @@ def energy_vel(m: Model, d: Data):
wp.launch_tiled( wp.launch_tiled(
_energy_vel_kinetic(m.nv), _energy_vel_kinetic(m.nv),
dim=(d.nworld,), dim=d.nworld,
inputs=[d.qvel, d.efc.mv], inputs=[d.qvel, d.efc.mv],
outputs=[d.energy], outputs=[d.energy],
block_dim=m.block_dim.energy_vel_kinetic, 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) xaxis = math.rot_vec_quat(jnt_axis_, xquat)
if jnt_type_ == JointType.BALL: if jnt_type_ == JointType.BALL:
qloc = wp.quat( qloc = wp.quat(qpos[qadr + 0], qpos[qadr + 1], qpos[qadr + 2], qpos[qadr + 3])
qpos[qadr + 0],
qpos[qadr + 1],
qpos[qadr + 2],
qpos[qadr + 3],
)
qloc = wp.normalize(qloc) qloc = wp.normalize(qloc)
xquat = math.mul_quat(xquat, qloc) xquat = math.mul_quat(xquat, qloc)
# correct for off-center rotation # correct for off-center rotation
@@ -1797,14 +1792,7 @@ def _transmission(
if jnt_typ == JointType.FREE: if jnt_typ == JointType.FREE:
actuator_length_out[worldid, actid] = 0.0 actuator_length_out[worldid, actid] = 0.0
if trntype == TrnType.JOINTINPARENT: if trntype == TrnType.JOINTINPARENT:
quat = wp.normalize( quat = wp.normalize(wp.quat(qpos[qadr + 3], qpos[qadr + 4], qpos[qadr + 5], qpos[qadr + 6]))
wp.quat(
qpos[qadr + 3],
qpos[qadr + 4],
qpos[qadr + 5],
qpos[qadr + 6],
)
)
quat_neg = math.quat_inv(quat) quat_neg = math.quat_inv(quat)
gearaxis = math.rot_vec_quat(wp.spatial_bottom(gear), quat_neg) gearaxis = math.rot_vec_quat(wp.spatial_bottom(gear), quat_neg)
actuator_moment_out[worldid, actid, vadr + 0] = gear[0] actuator_moment_out[worldid, actid, vadr + 0] = gear[0]
@@ -1875,30 +1863,14 @@ def _transmission(
# get Jacobians of axis(jacA) and vec(jac) # get Jacobians of axis(jacA) and vec(jac)
# mj_jacPointAxis # mj_jacPointAxis
jacp, jacr = support.jac( jacp, jacr = support.jac(
body_parentid, body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_idslider, site_bodyid[idslider], i, worldid
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
site_xpos_idslider,
site_bodyid[idslider],
i,
worldid,
) )
jacS = jacp jacS = jacp
jacA = wp.cross(jacr, axis) jacA = wp.cross(jacr, axis)
# mj_jacSite # mj_jacSite
jac, _ = support.jac( jac, _ = support.jac(
body_parentid, body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos_id, site_bodyid[id], i, worldid
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
site_xpos_id,
site_bodyid[id],
i,
worldid,
) )
jac -= jacS jac -= jacS
@@ -2023,28 +1995,12 @@ def _transmission(
# TODO(team): parallelize # TODO(team): parallelize
for i in range(nv): for i in range(nv):
jacp, jacr = support.jac( jacp, jacr = support.jac(
body_parentid, body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, site_xpos, site_bodyid[siteid], i, worldid
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
site_xpos,
site_bodyid[siteid],
i,
worldid,
) )
# jacref: global Jacobian of reference site # jacref: global Jacobian of reference site
jacpref, jacrref = support.jac( jacpref, jacrref = support.jac(
body_parentid, body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, ref_xpos, site_bodyid[refid], i, worldid
body_rootid,
dof_bodyid,
subtree_com_in,
cdof_in,
ref_xpos,
site_bodyid[refid],
i,
worldid,
) )
jacpdif = jacp - jacpref 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]) normal = wp.vec3(contact_frame[0, 0], contact_frame[0, 1], contact_frame[0, 2])
# get Jacobian difference # get Jacobian difference
jacp1, _ = support.jac( jacp1, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b1, dofid, worldid)
body_parentid, jacp2, _ = support.jac(body_parentid, body_rootid, dof_bodyid, subtree_com_in, cdof_in, contact_pos, b2, dofid, worldid)
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 jacdif = jacp2 - jacp1
# project Jacobian along the normal of the contact frame # 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: if spatial_site or spatial_geom:
wp.launch( wp.launch(
_spatial_tendon_wrap, _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], 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], 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.""" """Iterative linesearch."""
wp.launch( wp.launch(
linesearch_iterative, linesearch_iterative,
dim=(d.nworld,), dim=d.nworld,
inputs=[ inputs=[
m.nv, m.nv,
m.opt.impratio, m.opt.impratio,
@@ -1811,7 +1811,7 @@ def _update_gradient(m: types.Model, d: types.Data):
if m.nv < 32: if m.nv < 32:
wp.launch_tiled( wp.launch_tiled(
update_gradient_cholesky(m.nv), update_gradient_cholesky(m.nv),
dim=(d.nworld,), dim=d.nworld,
inputs=[d.efc.grad, d.efc.h, d.efc.done], inputs=[d.efc.grad, d.efc.h, d.efc.done],
outputs=[d.efc.Mgrad], outputs=[d.efc.Mgrad],
block_dim=m.block_dim.update_gradient_cholesky, block_dim=m.block_dim.update_gradient_cholesky,
@@ -1819,7 +1819,7 @@ def _update_gradient(m: types.Model, d: types.Data):
else: else:
wp.launch_tiled( wp.launch_tiled(
update_gradient_cholesky_blocked(16), update_gradient_cholesky_blocked(16),
dim=(d.nworld,), dim=d.nworld,
inputs=[ inputs=[
d.efc.grad.reshape(shape=(d.nworld, m.nv, 1)), d.efc.grad.reshape(shape=(d.nworld, m.nv, 1)),
d.efc.h, d.efc.h,
@@ -1982,7 +1982,7 @@ def _solver_iteration(
if m.opt.solver == types.SolverType.CG: if m.opt.solver == types.SolverType.CG:
wp.launch( wp.launch(
solve_beta, 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], inputs=[m.nv, d.efc.grad, d.efc.Mgrad, d.efc.prev_grad, d.efc.prev_Mgrad, d.efc.done],
outputs=[d.efc.beta], outputs=[d.efc.beta],
) )
@@ -1998,7 +1998,7 @@ def _solver_iteration(
wp.launch( wp.launch(
solve_done, solve_done,
dim=(d.nworld,), dim=d.nworld,
inputs=[ inputs=[
m.nv, m.nv,
m.opt.tolerance, m.opt.tolerance,
-2
View File
@@ -637,7 +637,6 @@ class Option:
run_collision_detection: if False, skips collision detection and allows user-populated run_collision_detection: if False, skips collision detection and allows user-populated
contacts during the physics step (as opposed to DisableBit.CONTACT which explicitly contacts during the physics step (as opposed to DisableBit.CONTACT which explicitly
zeros out the contacts at each step) 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 contact_sensor_maxmatch: max number of contacts considered by contact sensor matching criteria
contacts matched after this value is exceded will be ignored contacts matched after this value is exceded will be ignored
""" """
@@ -671,7 +670,6 @@ class Option:
broadphase_filter: int broadphase_filter: int
graph_conditional: bool graph_conditional: bool
run_collision_detection: bool run_collision_detection: bool
legacy_gjk: bool
contact_sensor_maxmatch: int contact_sensor_maxmatch: int
-3
View File
@@ -109,7 +109,6 @@ def _collision_shim(
opt__ccd_iterations: int, opt__ccd_iterations: int,
opt__ccd_tolerance: wp.array(dtype=float), opt__ccd_tolerance: wp.array(dtype=float),
opt__disableflags: int, opt__disableflags: int,
opt__legacy_gjk: bool,
opt__sdf_initpoints: int, opt__sdf_initpoints: int,
opt__sdf_iterations: int, opt__sdf_iterations: int,
# Data # Data
@@ -192,7 +191,6 @@ def _collision_shim(
_m.opt.ccd_iterations = opt__ccd_iterations _m.opt.ccd_iterations = opt__ccd_iterations
_m.opt.ccd_tolerance = opt__ccd_tolerance _m.opt.ccd_tolerance = opt__ccd_tolerance
_m.opt.disableflags = opt__disableflags _m.opt.disableflags = opt__disableflags
_m.opt.legacy_gjk = opt__legacy_gjk
_m.opt.sdf_initpoints = opt__sdf_initpoints _m.opt.sdf_initpoints = opt__sdf_initpoints
_m.opt.sdf_iterations = opt__sdf_iterations _m.opt.sdf_iterations = opt__sdf_iterations
_m.pair_dim = pair_dim _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_iterations,
m.opt._impl.ccd_tolerance, m.opt._impl.ccd_tolerance,
m.opt.disableflags, m.opt.disableflags,
m.opt._impl.legacy_gjk,
m.opt._impl.sdf_initpoints, m.opt._impl.sdf_initpoints,
m.opt._impl.sdf_iterations, m.opt._impl.sdf_iterations,
d._impl.naconmax, d._impl.naconmax,
-6
View File
@@ -344,7 +344,6 @@ def _forward_shim(
opt__impratio: wp.array(dtype=float), opt__impratio: wp.array(dtype=float),
opt__is_sparse: bool, opt__is_sparse: bool,
opt__iterations: int, opt__iterations: int,
opt__legacy_gjk: bool,
opt__ls_iterations: int, opt__ls_iterations: int,
opt__ls_parallel: bool, opt__ls_parallel: bool,
opt__ls_parallel_min_step: float, opt__ls_parallel_min_step: float,
@@ -714,7 +713,6 @@ def _forward_shim(
_m.opt.impratio = opt__impratio _m.opt.impratio = opt__impratio
_m.opt.is_sparse = opt__is_sparse _m.opt.is_sparse = opt__is_sparse
_m.opt.iterations = opt__iterations _m.opt.iterations = opt__iterations
_m.opt.legacy_gjk = opt__legacy_gjk
_m.opt.ls_iterations = opt__ls_iterations _m.opt.ls_iterations = opt__ls_iterations
_m.opt.ls_parallel = opt__ls_parallel _m.opt.ls_parallel = opt__ls_parallel
_m.opt.ls_parallel_min_step = opt__ls_parallel_min_step _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.impratio,
m.opt._impl.is_sparse, m.opt._impl.is_sparse,
m.opt.iterations, m.opt.iterations,
m.opt._impl.legacy_gjk,
m.opt.ls_iterations, m.opt.ls_iterations,
m.opt._impl.ls_parallel, m.opt._impl.ls_parallel,
m.opt._impl.ls_parallel_min_step, m.opt._impl.ls_parallel_min_step,
@@ -2138,7 +2135,6 @@ def _step_shim(
opt__integrator: int, opt__integrator: int,
opt__is_sparse: bool, opt__is_sparse: bool,
opt__iterations: int, opt__iterations: int,
opt__legacy_gjk: bool,
opt__ls_iterations: int, opt__ls_iterations: int,
opt__ls_parallel: bool, opt__ls_parallel: bool,
opt__ls_parallel_min_step: float, opt__ls_parallel_min_step: float,
@@ -2510,7 +2506,6 @@ def _step_shim(
_m.opt.integrator = opt__integrator _m.opt.integrator = opt__integrator
_m.opt.is_sparse = opt__is_sparse _m.opt.is_sparse = opt__is_sparse
_m.opt.iterations = opt__iterations _m.opt.iterations = opt__iterations
_m.opt.legacy_gjk = opt__legacy_gjk
_m.opt.ls_iterations = opt__ls_iterations _m.opt.ls_iterations = opt__ls_iterations
_m.opt.ls_parallel = opt__ls_parallel _m.opt.ls_parallel = opt__ls_parallel
_m.opt.ls_parallel_min_step = opt__ls_parallel_min_step _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.integrator,
m.opt._impl.is_sparse, m.opt._impl.is_sparse,
m.opt.iterations, m.opt.iterations,
m.opt._impl.legacy_gjk,
m.opt.ls_iterations, m.opt.ls_iterations,
m.opt._impl.ls_parallel, m.opt._impl.ls_parallel,
m.opt._impl.ls_parallel_min_step, m.opt._impl.ls_parallel_min_step,
-5
View File
@@ -94,7 +94,6 @@ class OptionWarp(PyTreeNode):
graph_conditional: bool graph_conditional: bool
has_fluid: bool has_fluid: bool
is_sparse: bool is_sparse: bool
legacy_gjk: bool
ls_parallel: bool ls_parallel: bool
ls_parallel_min_step: float ls_parallel_min_step: float
run_collision_detection: bool run_collision_detection: bool
@@ -765,7 +764,6 @@ _NDIM = {
'opt__integrator': 0, 'opt__integrator': 0,
'opt__is_sparse': 0, 'opt__is_sparse': 0,
'opt__iterations': 0, 'opt__iterations': 0,
'opt__legacy_gjk': 0,
'opt__ls_iterations': 0, 'opt__ls_iterations': 0,
'opt__ls_parallel': 0, 'opt__ls_parallel': 0,
'opt__ls_parallel_min_step': 0, 'opt__ls_parallel_min_step': 0,
@@ -885,7 +883,6 @@ _NDIM = {
'integrator': 0, 'integrator': 0,
'is_sparse': 0, 'is_sparse': 0,
'iterations': 0, 'iterations': 0,
'legacy_gjk': 0,
'ls_iterations': 0, 'ls_iterations': 0,
'ls_parallel': 0, 'ls_parallel': 0,
'ls_parallel_min_step': 0, 'ls_parallel_min_step': 0,
@@ -1305,7 +1302,6 @@ _BATCH_DIM = {
'opt__integrator': False, 'opt__integrator': False,
'opt__is_sparse': False, 'opt__is_sparse': False,
'opt__iterations': False, 'opt__iterations': False,
'opt__legacy_gjk': False,
'opt__ls_iterations': False, 'opt__ls_iterations': False,
'opt__ls_parallel': False, 'opt__ls_parallel': False,
'opt__ls_parallel_min_step': False, 'opt__ls_parallel_min_step': False,
@@ -1425,7 +1421,6 @@ _BATCH_DIM = {
'integrator': False, 'integrator': False,
'is_sparse': False, 'is_sparse': False,
'iterations': False, 'iterations': False,
'legacy_gjk': False,
'ls_iterations': False, 'ls_iterations': False,
'ls_parallel': False, 'ls_parallel': False,
'ls_parallel_min_step': False, 'ls_parallel_min_step': False,