diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py
index 871652a3..2601e3ae 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_convex.py
@@ -104,7 +104,7 @@ def _max_contacts_height_field(
@cache_kernel
def ccd_kernel_builder(
- default_gjk: bool,
+ legacy_gjk: bool,
geomtype1: int,
geomtype2: int,
gjk_iterations: int,
@@ -286,10 +286,16 @@ def ccd_kernel_builder(
hftri_index,
)
- points = mat3c()
-
# TODO(kbayes): remove legacy GJK once multicontact can be enabled
- if default_gjk:
+ if wp.static(legacy_gjk):
+ # find prism center for height field
+ if geomtype1 == int(GeomType.HFIELD.value):
+ x1 = wp.vec3(0.0, 0.0, 0.0)
+ for i in range(6):
+ x1 += hfield_prism_vertex(geom1.hfprism, i)
+ x1 = geom1.pos + geom1.rot @ (x1 / 6.0)
+ geom1.pos = x1
+
simplex, normal = gjk_legacy(
gjk_iterations,
geom1,
@@ -303,24 +309,28 @@ def ccd_kernel_builder(
)
dist = -depth
- if (dist - margin) >= 0.0 or depth != depth:
+ if dist >= 0.0 or depth < -depth_extension:
+ count = 0
return
sphere = int(GeomType.SPHERE.value)
ellipsoid = int(GeomType.ELLIPSOID.value)
if geom_type[g1] == sphere or geom_type[g1] == ellipsoid or geom_type[g2] == sphere or geom_type[g2] == ellipsoid:
count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 1, 2, 1.0e-5)
else:
- count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 4, 8, 1.0e-3)
+ count, points = multicontact_legacy(geom1, geom2, geomtype1, geomtype2, depth_extension, depth, normal, 4, 8, 1.0e-1)
+ frame = make_frame(normal)
else:
+ points = mat3c()
+
x1 = geom1.pos
x2 = geom2.pos
# find prism center for height field
if geomtype1 == int(GeomType.HFIELD.value):
- x1 = wp.vec3(0.0, 0.0, 0.0)
+ x1_ = wp.vec3(0.0, 0.0, 0.0)
for i in range(6):
- x1 += hfield_prism_vertex(geom1.hfprism, i)
- x1 = x1 / 6.0
+ x1_ += hfield_prism_vertex(geom1.hfprism, i)
+ x1 += geom1.rot @ (x1_ / 6.0)
dist, count, witness1, witness2 = ccd(
False,
@@ -411,13 +421,13 @@ def convex_narrowphase(m: Model, d: Data):
if m.geom_pair_type_count[upper_trid_index(len(GeomType), geom_pair[0], geom_pair[1])]:
wp.launch(
ccd_kernel_builder(
- False,
+ m.opt.legacy_gjk,
geom_pair[0],
geom_pair[1],
m.opt.gjk_iterations,
m.opt.epa_iterations,
- False,
- 0.1,
+ True,
+ 1e9,
),
dim=d.nconmax,
inputs=[
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py
index 089d5856..969e638e 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk.py
@@ -113,7 +113,7 @@ def _support_margin(geom: Geom, geomtype: int, dir: wp.vec3):
local_dir = wp.transpose(geom.rot) @ dir
res = wp.vec3()
res[2] = wp.where(local_dir[2] >= 0, geom.size[1], -geom.size[1])
- sp.point = res
+ sp.point = geom.rot @ res + geom.pos
return sp
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py
index e8313be6..0b93ab15 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_gjk_legacy.py
@@ -19,6 +19,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.collision_hfield import hfield_pris
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
@@ -203,8 +204,8 @@ def gjk_legacy(
depth = dist_min
normal = dir_n
- sd = simplex0 - simplex1
- dir = orthonormal(sd)
+ sd = wp.normalize(simplex0 - simplex1)
+ dir = orthonormal_to_z(sd)
dist_max, simplex3 = _gjk_support(geom1, geom2, geomtype1, geomtype2, dir)
@@ -304,14 +305,15 @@ def epa_legacy(
normal: wp.vec3,
):
# get the support, if depth < 0: objects do not intersect
- depth, _ = _gjk_support(geom1, geom2, geomtype1, geomtype2, normal)
+ 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 wp.nan, wp.vec3(wp.nan, wp.nan, wp.nan)
+ return FLOAT_MAX, wp.vec3(wp.nan, wp.nan, wp.nan)
- if wp.static(epa_exact_neg_distance):
+ 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.
@@ -364,7 +366,7 @@ def epa_legacy(
# This is a hack to reduce compile time
count = int(4)
it = int(0)
- for _ in range(wp.static(epa_iterations)):
+ for _ in range(epa_iterations):
it += count
count = wp.min(count * 3, EPS_BEST_COUNT)
@@ -391,7 +393,7 @@ def epa_legacy(
# iterate over edges and get distance using support point
for j in range(3):
- if wp.static(epa_exact_neg_distance):
+ if epa_exact_neg_distance:
# obtain closest point between new triangle edge and origin
tqj = tris[ti + j]
@@ -663,7 +665,7 @@ def multicontact_legacy(
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:
+ 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
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py
index 8397700c..0458fd6b 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_hfield.py
@@ -17,74 +17,49 @@ from typing import Tuple
import warp as wp
+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
@wp.func
-def _hfield_overlap_range(
- # Model:
- geom_dataid: wp.array(dtype=int),
- geom_rbound: wp.array2d(dtype=float),
- geom_margin: wp.array2d(dtype=float),
- hfield_nrow: wp.array(dtype=int),
- hfield_ncol: wp.array(dtype=int),
- 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),
+def _hfield_subgrid(
# In:
- hfieldid: int,
- geomid: int,
- worldid: int,
+ nrow: int,
+ ncol: int,
+ size: wp.vec4,
+ xmax: float,
+ xmin: float,
+ ymax: float,
+ ymin: float,
) -> Tuple[int, int, int, int]:
- """Returns min/max grid coordinates of height field cells overlapped by a geom's bounds.
+ """Returns height field subgrid that overlaps with geom AABB.
Args:
- geom_dataid: Array of geom data IDs
- geom_rbound: Array of geom bounding radii
- geom_margin: Array of geom margins
- hfield_nrow: Array of heightfield rows
- hfield_ncol: Array of heightfield columns
- hfield_size: Array of heightfield sizes
- geom_xpos_in: Array of geom positions
- geom_xmat_in: Array of geom orientation matrices
- hfieldid: Index of the height field geom
- geomid: Index of the other geom
- worldid: Current world index
+ nrow: height field number of rows
+ ncol: height field number of columns
+ size: height field size
+ xmax: geom maximum x position
+ xmin: geom minimum x position
+ ymax: geom maximum y position
+ ymin: geom minimum y position
Returns:
- min_i, min_j, max_i, max_j: Grid coordinate bounds
+ grid coordinate bounds
"""
- # get height field dimensions
- dataid = geom_dataid[hfieldid]
- nrow = hfield_nrow[dataid]
- ncol = hfield_ncol[dataid]
- size = hfield_size[dataid] # (x, y, z_top, z_bottom)
- # get positions and transforms
- hf_pos = geom_xpos_in[worldid, hfieldid]
- hf_mat = geom_xmat_in[worldid, hfieldid]
- geom_pos = geom_xpos_in[worldid, geomid]
+ # grid resolution
+ x_scale = 0.5 * float(ncol - 1) / size[0]
+ y_scale = 0.5 * float(nrow - 1) / size[1]
- # transform geom_pos to height field local space
- local_pos = wp.transpose(hf_mat) @ (geom_pos - hf_pos)
+ # subgrid
+ cmin = wp.max(0, int(wp.floor((xmin + size[0]) * x_scale)))
+ cmax = wp.min(ncol - 1, int(wp.ceil((xmax + size[0]) * x_scale)))
+ rmin = wp.max(0, int(wp.floor((ymin + size[1]) * y_scale)))
+ rmax = wp.min(nrow - 1, int(wp.ceil((ymax + size[1]) * y_scale)))
- # get bounding radius of other geometry (including margin)
- bound_radius = geom_rbound[worldid, geomid] + geom_margin[worldid, geomid]
-
- # calculate grid resolution
- x_scale = 2.0 * size[0] / float(ncol - 1)
- y_scale = 2.0 * size[1] / float(nrow - 1)
-
- # calculate min/max grid coordinates that could contain the object
- min_i = wp.max(0, int((local_pos[0] - bound_radius + size[0]) / x_scale))
- max_i = wp.min(ncol - 2, int((local_pos[0] + bound_radius + size[0]) / x_scale) + 1)
- min_j = wp.max(0, int((local_pos[1] - bound_radius + size[1]) / y_scale))
- max_j = wp.min(nrow - 2, int((local_pos[1] + bound_radius + size[1]) / y_scale) + 1)
-
- return min_i, min_j, max_i, max_j
+ return cmin, rmin, cmax, rmax
@wp.func
@@ -100,20 +75,20 @@ def hfield_triangle_prism(
hfieldid: int,
hftri_index: int,
) -> wp.mat33:
- """Returns the vertices of a triangular prism for a heightfield triangle.
+ """Returns triangular prism vertex information in compressed representation.
Args:
- geom_dataid: Array of geometry data IDs
- hfield_adr: Array of heightfield addresses
- hfield_nrow: Array of heightfield rows
- hfield_ncol: Array of heightfield columns
- hfield_size: Array of heightfield sizes
- hfield_data: Array of heightfield data
- hfieldid: Index of the height field geometry
- hftri_index: Index of the triangle in the heightfield
+ geom_dataid: geom data ids
+ hfield_adr: address for height field
+ hfield_nrow: height field number of rows
+ hfield_ncol: height field number of columns
+ hfield_size: height field sizes
+ hfield_data: height field data
+ hfieldid: height field geom id
+ hftri_index: height field triangle index
Returns:
- 3x3 matrix containing the vertices of the triangular prism
+ triangular prism vertex information (compressed)
"""
# https://mujoco.readthedocs.io/en/stable/XMLreference.html#asset-hfield
@@ -160,21 +135,14 @@ def hfield_triangle_prism(
z10 = z10 * z_top
z11 = z11 * z_top
- # set bottom z-value
- z_bottom = -size[3]
+ x2 = wp.where(hftri_index % 2, 1.0, 0.0)
+ y2 = wp.where(hftri_index % 2, z10, z01)
+ z22 = -size[3]
# compress 6 prism vertices into 3x3 matrix, see hfield_prism_vertex for details
- return wp.mat33(
- x0,
- y0,
- z00,
- x1,
- y1,
- z11,
- wp.where(hftri_index % 2, 1.0, 0.0),
- wp.where(hftri_index % 2, z10, z01),
- z_bottom,
- )
+ return wp.mat33(x0, y0, z00,
+ x1, y1, z11,
+ x2, y2, z22) # fmt: off
@wp.func
@@ -194,10 +162,10 @@ def hfield_prism_vertex(prism: wp.mat33, vert_index: int) -> wp.vec3:
Args:
prism: 3x3 compressed representation of a triangular prism
- vert_index: Index of vertex to extract (0-5)
+ vert_index: index of vertex to extract (0-5)
Returns:
- The 3D coordinates of the requested vertex
+ 3D coordinates of the requested vertex
"""
if vert_index == 0 or vert_index == 1:
return prism[vert_index] # first two vertices stored directly
@@ -223,6 +191,7 @@ def _hfield_midphase(
# Model:
geom_type: wp.array(dtype=int),
geom_dataid: wp.array(dtype=int),
+ geom_aabb: wp.array2d(dtype=wp.vec3),
geom_rbound: wp.array2d(dtype=float),
geom_margin: wp.array2d(dtype=float),
hfield_nrow: wp.array(dtype=int),
@@ -250,130 +219,171 @@ def _hfield_midphase(
one for each potentially colliding triangle.
Args:
- geom_type: Array of geometry types
- geom_dataid: Array of geometry data IDs
- geom_rbound: Array of geometry bounding radii
- geom_margin: Array of geometry margins
- hfield_nrow: Array of heightfield rows
- hfield_ncol: Array of heightfield columns
- hfield_size: Array of heightfield sizes
- nconmax_in: Max number of collisions
- geom_xpos_in: Array of geometry positions
- geom_xmat_in: Array of geometry orientation matrices
- collision_pair_in: Array of collision pairs
- collision_hftri_index_in: Array of heightfield triangle indices, -1 for heightfield
- pairs
- collision_pairid_in: Array of collision pair IDs
- collision_worldid_in: Array of collision world IDs
-
- collision_pair_out: Output array of collision pairs
- collision_hftri_index_out: Output array of heightfield triangle indices
- collision_pairid_out: Output array of collision pair IDs
- collision_worldid_out: Output array of collision world IDs
- ncollision_out: Output counter for number of collisions
+ geom_type: geom type
+ geom_dataid: geom data id
+ geom_rbound: geom bounding sphere radius
+ geom_margin: geom margin
+ hfield_nrow: height field number of rows
+ hfield_ncol: height field number of columns
+ hfield_size: height field size
+ nconmax_in: maximum number of contacts
+ geom_xpos_in: geom position
+ geom_xmat_in: geom orientation
+ collision_pair_in: collision pair
+ collision_hftri_index_in: triangle indices, -1 for height field pair
+ collision_pairid_in: collision pair id from broadphase
+ collision_worldid_in: collision world id from broadphase
+ collision_pair_out: collision pair from midphase
+ collision_hftri_index_out: triangle indices from midphase
+ collision_pairid_out: collision pair id from midphase
+ collision_worldid_out: collision world id from midphase
+ ncollision_out: number of collisions from broadphase and midphase
"""
pairid = wp.tid()
- # only process pairs that are marked for heightfield collision (-1)
- # the buffer is cleared at the start of each frame in collision_driver.py
+ # only process pairs that are marked for height field collision (-1)
if collision_hftri_index_in[pairid] != -1:
return
- # get the collision pair info
- pair = collision_pair_in[pairid]
+ # collision pair info
worldid = collision_worldid_in[pairid]
pair_id = collision_pairid_in[pairid]
- # identify which geom is the heightfield
+ pair = collision_pair_in[pairid]
g1 = pair[0]
g2 = pair[1]
hfieldid = g1
geomid = g2
- # if the first geom is not a heightfield, swap them
- # in theory, shouldn't happen as _add_geom_pair already sorted the pair
+ # SHOULD NOT OCCUR: if the first geom is not a heightfield, swap
if geom_type[g1] != int(GeomType.HFIELD.value):
hfieldid = g2
geomid = g1
- # get min/max grid coordinates for overlap region
- min_i, min_j, max_i, max_j = _hfield_overlap_range(
- geom_dataid,
- geom_rbound,
- geom_margin,
- hfield_nrow,
- hfield_ncol,
- hfield_size,
- geom_xpos_in,
- geom_xmat_in,
- hfieldid,
- geomid,
- worldid,
- )
+ # height field info
+ hfdataid = geom_dataid[hfieldid]
+ size1 = hfield_size[hfdataid]
+ pos1 = geom_xpos_in[worldid, hfieldid]
+ mat1 = geom_xmat_in[worldid, hfieldid]
+ mat1T = wp.transpose(mat1)
- # get hfield dimensions for triangle index calculation
- dataid = geom_dataid[hfieldid]
- ncol = hfield_ncol[dataid]
+ # geom info
+ pos2 = geom_xpos_in[worldid, geomid]
+ pos = mat1T @ (pos2 - pos1)
+ r2 = geom_rbound[worldid, geomid]
- # loop through grid cells and add pairs for all triangles
- for j in range(min_j, max_j + 1):
- for i in range(min_i, max_i + 1):
- # each grid cell contains two triangles
- base_idx = ((j * (ncol - 1)) + i) * 2
+ # TODO(team): margin?
+ margin = wp.max(geom_margin[worldid, hfieldid], geom_margin[worldid, geomid])
+ # 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
+
+ # box-sphere test: vertical direction
+ if size1[2] < pos[2] - r2 - margin: # up
+ return
+
+ if -size1[3] > pos[2] + r2 + margin: # down
+ return
+
+ mat2 = geom_xmat_in[worldid, geomid]
+ 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
+
+ center2 = geom_aabb[geomid, 0]
+ size2 = geom_aabb[geomid, 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
+
+ # height field subgrid
+ nrow = hfield_nrow[hfieldid]
+ ncol = hfield_ncol[hfieldid]
+ size = hfield_size[hfieldid]
+ cmin, rmin, cmax, rmax = _hfield_subgrid(nrow, ncol, size, xmax, xmin, ymax, ymin)
+
+ # loop over subgrid triangles
+ for r in range(rmin, rmax):
+ for c in range(cmin, cmax):
# add both triangles from this cell
- for t in range(2):
- if i == 0 and j == 0 and t == 0:
+ for i in range(2):
+ if r == rmin and c == cmin and i == 0:
# reuse the initial pair for the 1st triangle
new_pairid = pairid
else:
- # for the rest create a new pair
+ # create a new pair
new_pairid = wp.atomic_add(ncollision_out, 0, 1)
if new_pairid >= nconmax_in:
return
collision_pair_out[new_pairid] = pair
- collision_hftri_index_out[new_pairid] = base_idx + t
+ collision_hftri_index_out[new_pairid] = 2 * (r * (ncol - 1) + c) + i
collision_pairid_out[new_pairid] = pair_id
collision_worldid_out[new_pairid] = worldid
def hfield_midphase(m: Model, d: Data):
- """Midphase collision detection for heightfield triangles with other geoms.
+ """Midphase collision detection for height field triangles with other geoms.
- Processes collision pairs from the broadphase where one geom is a heightfield and expands
+ Processes collision pairs from the broadphase where one geom is a height field and expands
them into multiple collision pairs, one for each potentially colliding triangle. The
function directly writes to the same collision buffers used by _add_geom_pair.
-
- Args:
- m: Model containing geometry and heightfield data
- - geom_type: Array of geometry types
- - geom_dataid: Array of geometry data IDs
- - hfield_nrow: Array of heightfield rows
- - hfield_ncol: Array of heightfield columns
- - hfield_size: Array of heightfield sizes
- - geom_rbound: Array of geometry bounding radii
- - geom_margin: Array of geometry margins
- d: Data containing current state and collision information
- - nconmax: Maximum number of contacts
- - geom_xpos: Array of geometry positions
- - geom_xmat: Array of geometry orientation matrices
- - collision_pair: Array of collision pairs
- - collision_hftri_index: Array of heightfield triangle indices
- - collision_pairid: Array of collision pair IDs
- - collision_worldid: Array of collision world IDs
- - ncollision: Number of collisions
"""
- # launch the midphase kernel to expand height field collision pairs
- # write directly to the same buffers that _add_geom_pair writes to
wp.launch(
kernel=_hfield_midphase,
- dim=d.nconmax, # launch threads to process all potential pairs
+ dim=d.nconmax,
inputs=[
m.geom_type,
m.geom_dataid,
+ m.geom_aabb,
m.geom_rbound,
m.geom_margin,
m.hfield_nrow,
@@ -387,11 +397,5 @@ def hfield_midphase(m: Model, d: Data):
d.collision_pairid,
d.collision_worldid,
],
- outputs=[
- d.collision_pair,
- d.collision_hftri_index,
- d.collision_pairid,
- d.collision_worldid,
- d.ncollision,
- ],
+ outputs=[d.collision_pair, d.collision_hftri_index, d.collision_pairid, d.collision_worldid, d.ncollision],
)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py
index 6ab7236e..9289ccf1 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/collision_primitive.py
@@ -919,7 +919,7 @@ def plane_convex(
a_dist = wp.float32(-_HUGE_VAL)
for i in range(convex.vertnum):
support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n)
- dist = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ dist = wp.where(support > threshold, support, -_HUGE_VAL)
if dist > a_dist:
indices[0] = i
a_dist = dist
@@ -942,7 +942,8 @@ def plane_convex(
for i in range(convex.vertnum):
support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n)
dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
- dist = wp.length_sq(ab - convex.vert[convex.vertadr + i]) + dist_mask
+ ap = a - convex.vert[convex.vertadr + i]
+ dist = wp.abs(wp.dot(ap, ab)) + dist_mask
if dist > c_dist:
indices[2] = i
c_dist = dist
@@ -955,8 +956,8 @@ def plane_convex(
for i in range(convex.vertnum):
support = wp.dot(plane_pos - convex.vert[convex.vertadr + i], n)
dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
- ap = ac - convex.vert[convex.vertadr + i]
- bp = bc - convex.vert[convex.vertadr + i]
+ ap = a - convex.vert[convex.vertadr + i]
+ bp = b - convex.vert[convex.vertadr + i]
dist_ap = wp.abs(wp.dot(ap, ac)) + dist_mask
dist_bp = wp.abs(wp.dot(bp, bc)) + dist_mask
if dist_ap + dist_bp > d_dist:
@@ -993,10 +994,6 @@ def plane_convex(
threshold = wp.max(0.0, max_support - 1e-3)
a_dist = wp.float32(-_HUGE_VAL)
- # hillclimb until no change
- prev = int(-1)
- imax = int(0)
-
while True:
prev = int(imax)
i = int(convex.graph[vert_edgeadr + imax])
@@ -1004,7 +1001,7 @@ def plane_convex(
subidx = convex.graph[edge_localid + i]
idx = convex.graph[vert_globalid + subidx]
support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
- dist = wp.where(support > threshold, 0.0, -_HUGE_VAL)
+ dist = wp.where(support > threshold, support, -_HUGE_VAL)
if dist > a_dist:
a_dist = dist
imax = int(subidx)
@@ -1017,10 +1014,6 @@ def plane_convex(
# Find point b (furthest from a)
b_dist = wp.float32(-_HUGE_VAL)
- # hillclimb until no change
- prev = int(-1)
- imax = int(0)
-
while True:
prev = int(imax)
i = int(convex.graph[vert_edgeadr + imax])
@@ -1043,10 +1036,6 @@ def plane_convex(
# Find point c (furthest along axis orthogonal to a-b)
ab = wp.cross(n, a - b)
c_dist = wp.float32(-_HUGE_VAL)
- # hillclimb until no change
- prev = int(-1)
- imax = int(0)
-
while True:
prev = int(imax)
i = int(convex.graph[vert_edgeadr + imax])
@@ -1055,7 +1044,8 @@ def plane_convex(
idx = convex.graph[vert_globalid + subidx]
support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
- dist = wp.length_sq(ab - convex.vert[convex.vertadr + idx]) + dist_mask
+ ap = a - convex.vert[convex.vertadr + i]
+ dist = wp.abs(wp.dot(ap, ab)) + dist_mask
if dist > c_dist:
c_dist = dist
imax = int(subidx)
@@ -1070,10 +1060,6 @@ def plane_convex(
ac = wp.cross(n, a - c)
bc = wp.cross(n, b - c)
d_dist = wp.float32(-_HUGE_VAL)
- # hillclimb until no change
- prev = int(-1)
- imax = int(0)
-
while True:
prev = int(imax)
i = int(convex.graph[vert_edgeadr + imax])
@@ -1082,8 +1068,8 @@ def plane_convex(
idx = convex.graph[vert_globalid + subidx]
support = wp.dot(plane_pos - convex.vert[convex.vertadr + idx], n)
dist_mask = wp.where(support > threshold, 0.0, -_HUGE_VAL)
- ap = ac - convex.vert[convex.vertadr + idx]
- bp = bc - convex.vert[convex.vertadr + idx]
+ ap = a - convex.vert[convex.vertadr + idx]
+ bp = b - convex.vert[convex.vertadr + idx]
dist_ap = wp.abs(wp.dot(ap, ac)) + dist_mask
dist_bp = wp.abs(wp.dot(bp, bc)) + dist_mask
if dist_ap + dist_bp > d_dist:
@@ -2549,7 +2535,7 @@ def box_box(
v = (y * ax - x * ay) * C
if nl == 0:
- if (u < 0 or u > 0) and (v < 0 or v > 1):
+ if (u < 0 or u > 1) and (v < 0 or v > 1):
continue
elif u < 0 or v < 0 or u > 1 or v > 1:
continue
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py
index 31309249..09a0a512 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward.py
@@ -312,6 +312,8 @@ def _advance(m: Model, d: Data, qacc: wp.array, qvel: Optional[wp.array] = None)
],
)
+ wp.copy(d.qacc_warmstart, d.qacc)
+
@wp.kernel
def _euler_damp_qfrc_sparse(
@@ -983,8 +985,8 @@ def forward(m: Model, d: Data):
energy = m.opt.enableflags & EnableBit.ENERGY
fwd_position(m, d, factorize=False)
+ d.sensordata.zero_()
sensor.sensor_pos(m, d)
-
if energy:
if m.sensor_e_potential == 0: # not computed by sensor
sensor.energy_pos(m, d)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py
index 73087c3c..890d3b1c 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/forward_test.py
@@ -421,6 +421,7 @@ class ForwardTest(parameterized.TestCase):
"qfrc_actuator",
"qfrc_smooth",
"qacc",
+ "qacc_warmstart",
"qvel",
"qpos",
"efc_force",
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py
index c9742982..2441f8db 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io.py
@@ -396,13 +396,12 @@ def put_model(mjm: mujoco.MjModel) -> types.Model:
# contact sensor
sensor_adr_to_contact_adr = np.clip(np.cumsum(mjm.sensor_type == mujoco.mjtSensor.mjSENS_CONTACT) - 1, a_min=0, a_max=None)
- # TODO(team): improve heuristic for selecting broadphase routine
- if mjm.ngeom > 1000:
- broadphase = types.BroadphaseType.SAP_SEGMENTED
- elif mjm.ngeom > 100:
+ if nxn_geom_pair_filtered.shape[0] < 250_000:
+ broadphase = types.BroadphaseType.NXN
+ elif mjm.ngeom < 1000:
broadphase = types.BroadphaseType.SAP_TILE
else:
- broadphase = types.BroadphaseType.NXN
+ broadphase = types.BroadphaseType.SAP_SEGMENTED
condim = np.concatenate((mjm.geom_condim, mjm.pair_dim))
condim_max = np.max(condim) if len(condim) > 0 else 0
@@ -473,6 +472,7 @@ 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,
),
stat=types.Statistic(
meaninertia=mjm.stat.meaninertia,
@@ -1016,7 +1016,6 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
cost=wp.zeros((nworld,), dtype=float),
prev_cost=wp.zeros((nworld,), dtype=float),
state=wp.zeros((nworld, njmax), dtype=int),
- gtol=wp.zeros((nworld,), dtype=float),
mv=wp.zeros((nworld, mjm.nv), dtype=float),
jv=wp.zeros((nworld, njmax), dtype=float),
quad=wp.zeros((nworld, njmax), dtype=wp.vec3f),
@@ -1028,18 +1027,6 @@ def make_data(mjm: mujoco.MjModel, nworld: int = 1, nconmax: int = -1, njmax: in
beta=wp.zeros((nworld,), dtype=float),
done=wp.zeros((nworld,), dtype=bool),
# linesearch
- ls_done=wp.zeros((nworld,), dtype=bool),
- p0=wp.zeros((nworld,), dtype=wp.vec3),
- lo=wp.zeros((nworld,), dtype=wp.vec3),
- lo_alpha=wp.zeros((nworld,), dtype=float),
- hi=wp.zeros((nworld,), dtype=wp.vec3),
- hi_alpha=wp.zeros((nworld,), dtype=float),
- lo_next=wp.zeros((nworld,), dtype=wp.vec3),
- lo_next_alpha=wp.zeros((nworld,), dtype=float),
- hi_next=wp.zeros((nworld,), dtype=wp.vec3),
- hi_next_alpha=wp.zeros((nworld,), dtype=float),
- mid=wp.zeros((nworld,), dtype=wp.vec3),
- mid_alpha=wp.zeros((nworld,), dtype=float),
cost_candidate=wp.zeros((nworld, mjm.opt.ls_iterations), dtype=float),
),
# RK4
@@ -1395,7 +1382,6 @@ def put_data(
cost=wp.empty(shape=(nworld,), dtype=float),
prev_cost=wp.empty(shape=(nworld,), dtype=float),
state=wp.empty(shape=(nworld, njmax), dtype=int),
- gtol=wp.empty(shape=(nworld,), dtype=float),
mv=wp.empty(shape=(nworld, mjm.nv), dtype=float),
jv=wp.empty(shape=(nworld, njmax), dtype=float),
quad=wp.empty(shape=(nworld, njmax), dtype=wp.vec3f),
@@ -1406,18 +1392,6 @@ def put_data(
prev_Mgrad=wp.empty(shape=(nworld, mjm.nv), dtype=float),
beta=wp.empty(shape=(nworld,), dtype=float),
done=wp.empty(shape=(nworld,), dtype=bool),
- ls_done=wp.zeros(shape=(nworld,), dtype=bool),
- p0=wp.empty(shape=(nworld,), dtype=wp.vec3),
- lo=wp.empty(shape=(nworld,), dtype=wp.vec3),
- lo_alpha=wp.empty(shape=(nworld,), dtype=float),
- hi=wp.empty(shape=(nworld,), dtype=wp.vec3),
- hi_alpha=wp.empty(shape=(nworld,), dtype=float),
- lo_next=wp.empty(shape=(nworld,), dtype=wp.vec3),
- lo_next_alpha=wp.empty(shape=(nworld,), dtype=float),
- hi_next=wp.empty(shape=(nworld,), dtype=wp.vec3),
- hi_next_alpha=wp.empty(shape=(nworld,), dtype=float),
- mid=wp.empty(shape=(nworld,), dtype=wp.vec3),
- mid_alpha=wp.empty(shape=(nworld,), dtype=float),
cost_candidate=wp.empty(shape=(nworld, mjm.opt.ls_iterations), dtype=float),
),
# TODO(team): skip allocation if integrator != RK4
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py
index 2d5932d8..1560550e 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/io_test.py
@@ -33,7 +33,7 @@ from mujoco.mjx.third_party.mujoco_warp._src.io import MAX_WORLDS
_IO_TEST_MODELS = (
"pendula.xml",
"collision_sdf/tactile.xml",
- "flex/cloth.xml",
+ "flex/floppy.xml",
"actuation/tendon_force_limit.xml",
"hfield/hfield.xml",
)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py
index 68c1e29b..4254f2cb 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/math.py
@@ -190,6 +190,16 @@ def orthonormal(normal: wp.vec3) -> wp.vec3:
return dir
+@wp.func
+def orthonormal_to_z(normal: wp.vec3) -> wp.vec3:
+ if wp.abs(normal[0]) < wp.abs(normal[1]):
+ dir = wp.vec3(1.0 - normal[0] * normal[0], -normal[0] * normal[1], -normal[0] * normal[2])
+ else:
+ dir = wp.vec3(-normal[1] * normal[0], 1.0 - normal[1] * normal[1], -normal[1] * normal[2])
+ dir, _ = gjk_normalize(dir)
+ return dir
+
+
@wp.func
def gjk_normalize(a: wp.vec3):
norm = wp.length(a)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py
index 3b82a825..4b14a8b9 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/sensor.py
@@ -226,18 +226,6 @@ def _ball_quat(jnt_qposadr: wp.array(dtype=int), qpos_in: wp.array2d(dtype=float
return wp.normalize(quat)
-@wp.kernel
-def _limit_pos_zero(
- # Model:
- sensor_adr: wp.array(dtype=int),
- sensor_limitpos_adr: wp.array(dtype=int),
- # Data out:
- sensordata_out: wp.array2d(dtype=float),
-):
- worldid, limitposid = wp.tid()
- sensordata_out[worldid, sensor_adr[sensor_limitpos_adr[limitposid]]] = 0.0
-
-
@wp.kernel
def _limit_pos(
# Model:
@@ -694,13 +682,6 @@ def sensor_pos(m: Model, d: Data):
)
# jointlimitpos and tendonlimitpos
- wp.launch(
- _limit_pos_zero,
- dim=(d.nworld, m.sensor_limitpos_adr.size),
- inputs=[m.sensor_adr, m.sensor_limitpos_adr],
- outputs=[d.sensordata],
- )
-
wp.launch(
_limit_pos,
dim=(d.nworld, d.njmax, m.sensor_limitpos_adr.size),
@@ -788,18 +769,6 @@ def _ball_ang_vel(jnt_dofadr: wp.array(dtype=int), qvel_in: wp.array2d(dtype=flo
return wp.vec3(qvel_in[worldid, adr + 0], qvel_in[worldid, adr + 1], qvel_in[worldid, adr + 2])
-@wp.kernel
-def _limit_vel_zero(
- # Model:
- sensor_adr: wp.array(dtype=int),
- sensor_limitvel_adr: wp.array(dtype=int),
- # Data out:
- sensordata_out: wp.array2d(dtype=float),
-):
- worldid, limitvelid = wp.tid()
- sensordata_out[worldid, sensor_adr[sensor_limitvel_adr[limitvelid]]] = 0.0
-
-
@wp.kernel
def _limit_vel(
# Model:
@@ -1250,13 +1219,6 @@ def sensor_vel(m: Model, d: Data):
outputs=[d.sensordata],
)
- wp.launch(
- _limit_vel_zero,
- dim=(d.nworld, m.sensor_limitvel_adr.size),
- inputs=[m.sensor_adr, m.sensor_limitvel_adr],
- outputs=[d.sensordata],
- )
-
wp.launch(
_limit_vel,
dim=(d.nworld, d.njmax, m.sensor_limitvel_adr.size),
@@ -1367,20 +1329,6 @@ def _joint_actuator_force(
return qfrc_actuator_in[worldid, jnt_dofadr[objid]]
-@wp.kernel
-def _tendon_actuator_force_zero(
- # Model:
- sensor_adr: wp.array(dtype=int),
- sensor_tendonactfrc_adr: wp.array(dtype=int),
- # Data out:
- sensordata_out: wp.array2d(dtype=float),
-):
- worldid, tenactfrcid = wp.tid()
- sensorid = sensor_tendonactfrc_adr[tenactfrcid]
- adr = sensor_adr[sensorid]
- sensordata_out[worldid, adr] = 0.0
-
-
@wp.kernel
def _tendon_actuator_force(
# Model:
@@ -1422,18 +1370,6 @@ def _tendon_actuator_force_cutoff(
_write_scalar(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, val, sensordata_out[worldid])
-@wp.kernel
-def _limit_frc_zero(
- # Model:
- sensor_adr: wp.array(dtype=int),
- sensor_limitfrc_adr: wp.array(dtype=int),
- # Data out:
- sensordata_out: wp.array2d(dtype=float),
-):
- worldid, limitfrcid = wp.tid()
- sensordata_out[worldid, sensor_adr[sensor_limitfrc_adr[limitfrcid]]] = 0.0
-
-
@wp.kernel
def _limit_frc(
# Model:
@@ -1760,20 +1696,6 @@ def _sensor_acc(
_write_vector(sensor_datatype, sensor_adr, sensor_cutoff, sensorid, 3, vec3, out)
-@wp.kernel
-def _sensor_touch_zero(
- # Model:
- sensor_adr: wp.array(dtype=int),
- sensor_touch_adr: wp.array(dtype=int),
- # Data out:
- sensordata_out: wp.array2d(dtype=float),
-):
- worldid, sensortouchadrid = wp.tid()
- sensorid = sensor_touch_adr[sensortouchadrid]
- adr = sensor_adr[sensorid]
- sensordata_out[worldid, adr] = 0.0
-
-
@wp.kernel
def _sensor_touch(
# Model:
@@ -1855,24 +1777,6 @@ def _sensor_touch(
wp.atomic_add(sensordata_out[worldid], adr, normalforce)
-@wp.kernel
-def _sensor_tactile_zero(
- # Model:
- sensor_type: wp.array(dtype=int),
- sensor_dim: wp.array(dtype=int),
- sensor_adr: wp.array(dtype=int),
- # Data out:
- sensordata_out: wp.array2d(dtype=float),
-):
- worldid, sensorid = wp.tid()
-
- if sensor_type[sensorid] != int(SensorType.TACTILE.value):
- return
-
- for i in range(sensor_dim[sensorid]):
- sensordata_out[worldid, sensor_adr[sensorid] + i] = 0.0
-
-
@wp.func
def _transform_spatial(vec: wp.spatial_vector, dif: wp.vec3) -> wp.vec3:
return wp.spatial_bottom(vec) - wp.cross(dif, wp.spatial_top(vec))
@@ -2093,18 +1997,6 @@ def sensor_acc(m: Model, d: Data):
if m.opt.disableflags & DisableBit.SENSOR:
return
- wp.launch(
- _sensor_touch_zero,
- dim=(d.nworld, m.sensor_touch_adr.size),
- inputs=[
- m.sensor_adr,
- m.sensor_touch_adr,
- ],
- outputs=[
- d.sensordata,
- ],
- )
-
wp.launch(
_sensor_touch,
dim=(d.nconmax, m.sensor_touch_adr.size),
@@ -2133,19 +2025,6 @@ def sensor_acc(m: Model, d: Data):
],
)
- wp.launch(
- _sensor_tactile_zero,
- dim=(d.nworld, m.nsensor),
- inputs=[
- m.sensor_type,
- m.sensor_dim,
- m.sensor_adr,
- ],
- outputs=[
- d.sensordata,
- ],
- )
-
wp.launch(
_sensor_tactile,
dim=(d.nconmax, m.nsensortaxel),
@@ -2280,18 +2159,6 @@ def sensor_acc(m: Model, d: Data):
outputs=[d.sensordata],
)
- wp.launch(
- _tendon_actuator_force_zero,
- dim=(d.nworld, m.sensor_tendonactfrc_adr.size),
- inputs=[
- m.sensor_adr,
- m.sensor_tendonactfrc_adr,
- ],
- outputs=[
- d.sensordata,
- ],
- )
-
wp.launch(
_tendon_actuator_force,
dim=(d.nworld, m.sensor_tendonactfrc_adr.size, m.nu),
@@ -2321,13 +2188,6 @@ def sensor_acc(m: Model, d: Data):
outputs=[d.sensordata],
)
- wp.launch(
- _limit_frc_zero,
- dim=(d.nworld, m.sensor_limitfrc_adr.size),
- inputs=[m.sensor_adr, m.sensor_limitfrc_adr],
- outputs=[d.sensordata],
- )
-
wp.launch(
_limit_frc,
dim=(d.nworld, d.njmax, m.sensor_limitfrc_adr.size),
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py
index ea5b10cf..b5b55024 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/smooth.py
@@ -347,10 +347,10 @@ def _subtree_com_init(
# Data in:
xipos_in: wp.array2d(dtype=wp.vec3),
# Data out:
- xipos_out: wp.array2d(dtype=wp.vec3),
+ subtree_com_out: wp.array2d(dtype=wp.vec3),
):
worldid, bodyid = wp.tid()
- xipos_out[worldid, bodyid] = xipos_in[worldid, bodyid] * body_mass[worldid, bodyid]
+ subtree_com_out[worldid, bodyid] = xipos_in[worldid, bodyid] * body_mass[worldid, bodyid]
@wp.kernel
@@ -486,7 +486,7 @@ def com_pos(m: Model, d: Data):
Accumulates the mass-weighted positions up the kinematic tree, divides by total mass, and
computes composite inertias and motion degrees of freedom in the subtree CoM frame.
"""
- wp.launch(_subtree_com_init, dim=(d.nworld, m.nbody), inputs=[m.body_mass, d.xipos, d.subtree_com])
+ wp.launch(_subtree_com_init, dim=(d.nworld, m.nbody), inputs=[m.body_mass, d.xipos], outputs=[d.subtree_com])
for i in reversed(range(len(m.body_tree))):
body_tree = m.body_tree[i]
@@ -514,37 +514,18 @@ def com_pos(m: Model, d: Data):
@wp.kernel
def _cam_local_to_global(
- # Model:
- cam_bodyid: wp.array(dtype=int),
- cam_pos: wp.array2d(dtype=wp.vec3),
- cam_quat: wp.array2d(dtype=wp.quat),
- # Data in:
- xpos_in: wp.array2d(dtype=wp.vec3),
- xquat_in: wp.array2d(dtype=wp.quat),
- # Data out:
- cam_xpos_out: wp.array2d(dtype=wp.vec3),
- cam_xmat_out: wp.array2d(dtype=wp.mat33),
-):
- """Fixed cameras."""
- worldid, camid = wp.tid()
- bodyid = cam_bodyid[camid]
- xpos = xpos_in[worldid, bodyid]
- xquat = xquat_in[worldid, bodyid]
- cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat)
- cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[worldid, camid]))
-
-
-@wp.kernel
-def _cam_fn(
# Model:
cam_mode: wp.array(dtype=int),
cam_bodyid: wp.array(dtype=int),
cam_targetbodyid: wp.array(dtype=int),
+ cam_pos: wp.array2d(dtype=wp.vec3),
+ cam_quat: wp.array2d(dtype=wp.quat),
cam_poscom0: wp.array2d(dtype=wp.vec3),
cam_pos0: wp.array2d(dtype=wp.vec3),
cam_mat0: wp.array2d(dtype=wp.mat33),
# Data in:
xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
subtree_com_in: wp.array2d(dtype=wp.vec3),
# Data out:
cam_xpos_out: wp.array2d(dtype=wp.vec3),
@@ -556,7 +537,11 @@ def _cam_fn(
)
invalid_target = is_target_cam and (cam_targetbodyid[camid] < 0)
if invalid_target:
- return
+ bodyid = cam_bodyid[camid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat)
+ cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[worldid, camid]))
elif cam_mode[camid] == wp.static(CamLightType.TRACK.value):
cam_xmat_out[worldid, camid] = cam_mat0[worldid, camid]
body_xpos = xpos_in[worldid, cam_bodyid[camid]]
@@ -567,6 +552,10 @@ def _cam_fn(
elif cam_mode[camid] == wp.static(CamLightType.TARGETBODY.value) or cam_mode[camid] == wp.static(
CamLightType.TARGETBODYCOM.value
):
+ bodyid = cam_bodyid[camid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat)
pos = xpos_in[worldid, cam_targetbodyid[camid]]
if cam_mode[camid] == wp.static(CamLightType.TARGETBODYCOM.value):
pos = subtree_com_in[worldid, cam_targetbodyid[camid]]
@@ -582,42 +571,28 @@ def _cam_fn(
mat_1[2], mat_2[2], mat_3[2]
)
# fmt: on
+ else:
+ bodyid = cam_bodyid[camid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ cam_xpos_out[worldid, camid] = xpos + math.rot_vec_quat(cam_pos[worldid, camid], xquat)
+ cam_xmat_out[worldid, camid] = math.quat_to_mat(math.mul_quat(xquat, cam_quat[worldid, camid]))
@wp.kernel
def _light_local_to_global(
- # Model:
- light_bodyid: wp.array(dtype=int),
- light_pos: wp.array2d(dtype=wp.vec3),
- light_dir: wp.array2d(dtype=wp.vec3),
- # Data in:
- xpos_in: wp.array2d(dtype=wp.vec3),
- xquat_in: wp.array2d(dtype=wp.quat),
- # Data out:
- light_xpos_out: wp.array2d(dtype=wp.vec3),
- light_xdir_out: wp.array2d(dtype=wp.vec3),
-):
- """Fixed lights."""
- worldid, lightid = wp.tid()
- bodyid = light_bodyid[lightid]
- xpos = xpos_in[worldid, bodyid]
- xquat = xquat_in[worldid, bodyid]
- light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat)
- light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[worldid, lightid], xquat)
-
-
-@wp.kernel
-def _light_fn(
# Model:
light_mode: wp.array(dtype=int),
light_bodyid: wp.array(dtype=int),
light_targetbodyid: wp.array(dtype=int),
+ light_pos: wp.array2d(dtype=wp.vec3),
+ light_dir: wp.array2d(dtype=wp.vec3),
light_poscom0: wp.array2d(dtype=wp.vec3),
light_pos0: wp.array2d(dtype=wp.vec3),
light_dir0: wp.array2d(dtype=wp.vec3),
# Data in:
xpos_in: wp.array2d(dtype=wp.vec3),
- light_xpos_in: wp.array2d(dtype=wp.vec3),
+ xquat_in: wp.array2d(dtype=wp.quat),
subtree_com_in: wp.array2d(dtype=wp.vec3),
# Data out:
light_xpos_out: wp.array2d(dtype=wp.vec3),
@@ -629,6 +604,11 @@ def _light_fn(
)
invalid_target = is_target_light and (light_targetbodyid[lightid] < 0)
if invalid_target:
+ bodyid = light_bodyid[lightid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat)
+ light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[worldid, lightid], xquat)
return
elif light_mode[lightid] == wp.static(CamLightType.TRACK.value):
light_xdir_out[worldid, lightid] = light_dir0[worldid, lightid]
@@ -640,10 +620,21 @@ def _light_fn(
elif light_mode[lightid] == wp.static(CamLightType.TARGETBODY.value) or light_mode[lightid] == wp.static(
CamLightType.TARGETBODYCOM.value
):
+ bodyid = light_bodyid[lightid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat)
pos = xpos_in[worldid, light_targetbodyid[lightid]]
if light_mode[lightid] == wp.static(CamLightType.TARGETBODYCOM.value):
pos = subtree_com_in[worldid, light_targetbodyid[lightid]]
- light_xdir_out[worldid, lightid] = pos - light_xpos_in[worldid, lightid]
+ light_xdir_out[worldid, lightid] = pos - light_xpos_out[worldid, lightid]
+ else:
+ bodyid = light_bodyid[lightid]
+ xpos = xpos_in[worldid, bodyid]
+ xquat = xquat_in[worldid, bodyid]
+ light_xpos_out[worldid, lightid] = xpos + math.rot_vec_quat(light_pos[worldid, lightid], xquat)
+ light_xdir_out[worldid, lightid] = math.rot_vec_quat(light_dir[worldid, lightid], xquat)
+
light_xdir_out[worldid, lightid] = wp.normalize(light_xdir_out[worldid, lightid])
@@ -658,33 +649,35 @@ def camlight(m: Model, d: Data):
wp.launch(
_cam_local_to_global,
dim=(d.nworld, m.ncam),
- inputs=[m.cam_bodyid, m.cam_pos, m.cam_quat, d.xpos, d.xquat],
- outputs=[d.cam_xpos, d.cam_xmat],
- )
- wp.launch(
- _cam_fn,
- dim=(d.nworld, m.ncam),
- inputs=[m.cam_mode, m.cam_bodyid, m.cam_targetbodyid, m.cam_poscom0, m.cam_pos0, m.cam_mat0, d.xpos, d.subtree_com],
+ inputs=[
+ m.cam_mode,
+ m.cam_bodyid,
+ m.cam_targetbodyid,
+ m.cam_pos,
+ m.cam_quat,
+ m.cam_poscom0,
+ m.cam_pos0,
+ m.cam_mat0,
+ d.xpos,
+ d.xquat,
+ d.subtree_com,
+ ],
outputs=[d.cam_xpos, d.cam_xmat],
)
wp.launch(
_light_local_to_global,
dim=(d.nworld, m.nlight),
- inputs=[m.light_bodyid, m.light_pos, m.light_dir, d.xpos, d.xquat],
- outputs=[d.light_xpos, d.light_xdir],
- )
- wp.launch(
- _light_fn,
- dim=(d.nworld, m.nlight),
inputs=[
m.light_mode,
m.light_bodyid,
m.light_targetbodyid,
+ m.light_pos,
+ m.light_dir,
m.light_poscom0,
m.light_pos0,
m.light_dir0,
d.xpos,
- d.light_xpos,
+ d.xquat,
d.subtree_com,
],
outputs=[d.light_xpos, d.light_xdir],
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py
index 7d208219..529f4939 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver.py
@@ -15,6 +15,7 @@
from math import ceil
from math import sqrt
+from typing import Tuple
import warp as wp
@@ -56,556 +57,416 @@ def _eval_pt(quad: wp.vec3, alpha: float) -> wp.vec3:
@wp.func
-def _eval(
- # Model:
- opt_impratio: wp.array(dtype=float),
+def _eval_frictionloss(
+ # In:
+ x: float,
+ f: float,
+ rf: float,
+ Jaref: float,
+ jv: float,
+ quad: wp.vec3,
+) -> wp.vec3:
+ # -bound < x < bound : quadratic
+ if (-rf < x) and (x < rf):
+ return quad
+ # x < -bound: linear negative
+ elif x <= -rf:
+ return wp.vec3(f * (-0.5 * rf - Jaref), -f * jv, 0.0)
+ # bound < x : linear positive
+ else:
+ return wp.vec3(f * (-0.5 * rf + Jaref), f * jv, 0.0)
+
+
+@wp.func
+def _eval_elliptic(
+ # In:
+ impratio_invsqrt: float,
+ friction: types.vec5,
+ quad: wp.vec3,
+ quad1: wp.vec3,
+ quad2: wp.vec3,
+ alpha: float,
+) -> wp.vec3:
+ mu = friction[0] * impratio_invsqrt
+
+ u0 = quad1[0]
+ v0 = quad1[1]
+ uu = quad1[2]
+ uv = quad2[0]
+ vv = quad2[1]
+ dm = quad2[2]
+
+ # compute N, Tsqr
+ N = u0 + alpha * v0
+ Tsqr = uu + alpha * (2.0 * uv + alpha * vv)
+
+ # no tangential force: top or bottom zone
+ if Tsqr <= 0.0:
+ # bottom zone: quadratic cost
+ if N < 0.0:
+ return _eval_pt(quad, alpha)
+
+ # top zone: nothing to do
+ # otherwise regular processing
+ else:
+ # tangential force
+ T = wp.sqrt(Tsqr)
+
+ # N >= mu * T : top zone
+ if N >= mu * T:
+ # nothing to do
+ pass
+ # mu * N + T <= 0 : bottom zone
+ elif mu * N + T <= 0.0:
+ return _eval_pt(quad, alpha)
+
+ # otherwise middle zone
+ else:
+ # derivatives
+ N1 = v0
+ T1 = (uv + alpha * vv) / T
+ T2 = vv / T - (uv + alpha * vv) * T1 / (T * T)
+
+ # add to cost
+ cost = wp.vec3(
+ 0.5 * dm * (N - mu * T) * (N - mu * T),
+ dm * (N - mu * T) * (N1 - mu * T1),
+ dm * ((N1 - mu * T1) * (N1 - mu * T1) + (N - mu * T) * (-mu * T2)),
+ )
+
+ return cost
+
+ return wp.vec3(0.0, 0.0, 0.0)
+
+
+@wp.func
+def _eval_init(
# Data in:
- ncon_in: wp.array(dtype=int),
- ne_in: wp.array(dtype=int),
- nf_in: wp.array(dtype=int),
contact_friction_in: wp.array(dtype=types.vec5),
contact_efc_address_in: wp.array2d(dtype=int),
- efc_type_in: wp.array2d(dtype=int),
- efc_id_in: wp.array2d(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_frictionloss_in: wp.array2d(dtype=float),
- efc_Jaref_in: wp.array2d(dtype=float),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
# In:
- worldid: int,
- efcid: int,
+ ne_clip: int,
+ nef_clip: int,
+ nefc_clip: int,
+ impratio_invsqrt: float,
+ type_in: wp.array(dtype=int),
+ id_in: wp.array(dtype=int),
+ D_in: wp.array(dtype=float),
+ frictionloss_in: wp.array(dtype=float),
+ Jaref_in: wp.array(dtype=float),
+ jv_in: wp.array(dtype=float),
+ quad_in: wp.array(dtype=wp.vec3),
alpha: float,
- # Out:
- out: wp.array(dtype=wp.vec3),
-):
- ne = ne_in[worldid]
- nf = nf_in[worldid]
+) -> wp.vec3:
+ lo = wp.vec3(0.0, 0.0, 0.0)
+ for efcid in range(ne_clip):
+ quad = quad_in[efcid]
+ lo += _eval_pt(quad, alpha)
+
+ for efcid in range(ne_clip, nef_clip):
+ D = D_in[efcid]
+ f = frictionloss_in[efcid]
+ Jaref = Jaref_in[efcid]
+ jv = jv_in[efcid]
- # equality
- if efcid < ne:
- wp.atomic_add(out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
- # friction
- elif efcid < ne + nf:
# search point, friction loss, bound (rf)
- start = efc_Jaref_in[worldid, efcid]
- dir = efc_jv_in[worldid, efcid]
- x = start + alpha * dir
- f = efc_frictionloss_in[worldid, efcid]
- rf = math.safe_div(f, efc_D_in[worldid, efcid])
+ x = Jaref + alpha * jv
+ rf = math.safe_div(f, D)
- # -bound < x < bound : quadratic
- if (-rf < x) and (x < rf):
- quad = efc_quad_in[worldid, efcid]
- # x < -bound: linear negative
- elif x <= -rf:
- quad = wp.vec3(f * (-0.5 * rf - start), -f * dir, 0.0)
- # bound < x : linear positive
+ quad_f = _eval_frictionloss(x, f, rf, Jaref, jv, quad_in[efcid])
+ lo += _eval_pt(quad_f, alpha)
+
+ for efcid in range(nef_clip, nefc_clip):
+ if type_in[efcid] == int(types.ConstraintType.CONTACT_ELLIPTIC.value):
+ conid = id_in[efcid]
+
+ efcid0 = contact_efc_address_in[conid, 0]
+ if efcid != efcid0:
+ continue
+
+ efcid1 = contact_efc_address_in[conid, 1]
+ efcid2 = contact_efc_address_in[conid, 2]
+ efc_quad0 = quad_in[efcid0]
+ efc_quad1 = quad_in[efcid1]
+ efc_quad2 = quad_in[efcid2]
+ friction = contact_friction_in[conid]
+
+ lo += _eval_elliptic(impratio_invsqrt, friction, efc_quad0, efc_quad1, efc_quad2, alpha)
else:
- quad = wp.vec3(f * (-0.5 * rf + start), f * dir, 0.0)
+ Jaref = Jaref_in[efcid]
+ jv = jv_in[efcid]
+ quad = quad_in[efcid]
- wp.atomic_add(out, worldid, _eval_pt(quad, alpha))
- # elliptic friction cone contact
- elif efc_type_in[worldid, efcid] == int(types.ConstraintType.CONTACT_ELLIPTIC.value):
- # extract contact info
- conid = efc_id_in[worldid, efcid]
+ x = Jaref + alpha * jv
+ res = _eval_pt(quad, alpha)
+ lo += res * float(x < 0.0)
- if conid >= ncon_in[0]:
- return
+ return lo
- efcid0 = contact_efc_address_in[conid, 0]
- if efcid != efcid0:
- return
+@wp.func
+def _eval(
+ # Data in:
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ # In:
+ ne_clip: int,
+ nef_clip: int,
+ nefc_clip: int,
+ impratio_invsqrt: float,
+ type_in: wp.array(dtype=int),
+ id_in: wp.array(dtype=int),
+ D_in: wp.array(dtype=float),
+ frictionloss_in: wp.array(dtype=float),
+ Jaref_in: wp.array(dtype=float),
+ jv_in: wp.array(dtype=float),
+ quad_in: wp.array(dtype=wp.vec3),
+ lo_alpha: float,
+ hi_alpha: float,
+ mid_alpha: float,
+) -> Tuple[wp.vec3, wp.vec3, wp.vec3]:
+ lo = wp.vec3(0.0, 0.0, 0.0)
+ hi = wp.vec3(0.0, 0.0, 0.0)
+ mid = wp.vec3(0.0, 0.0, 0.0)
+ for efcid in range(ne_clip):
+ quad = quad_in[efcid]
+ lo += _eval_pt(quad, lo_alpha)
+ hi += _eval_pt(quad, hi_alpha)
+ mid += _eval_pt(quad, mid_alpha)
- friction = contact_friction_in[conid]
- mu = friction[0] / wp.sqrt(opt_impratio[worldid])
+ for efcid in range(ne_clip, nef_clip):
+ quad = quad_in[efcid]
+ D = D_in[efcid]
+ f = frictionloss_in[efcid]
+ Jaref = Jaref_in[efcid]
+ jv = jv_in[efcid]
- # unpack quad
- efcid1 = contact_efc_address_in[conid, 1]
- efcid2 = contact_efc_address_in[conid, 2]
- u0 = efc_quad_in[worldid, efcid1][0]
- v0 = efc_quad_in[worldid, efcid1][1]
- uu = efc_quad_in[worldid, efcid1][2]
- uv = efc_quad_in[worldid, efcid2][0]
- vv = efc_quad_in[worldid, efcid2][1]
- dm = efc_quad_in[worldid, efcid2][2]
+ # search point, friction loss, bound (rf)
+ rf = math.safe_div(f, D)
+ x_lo = Jaref + lo_alpha * jv
+ x_hi = Jaref + hi_alpha * jv
+ x_mid = Jaref + mid_alpha * jv
- # compute N, Tsqr
- N = u0 + alpha * v0
- Tsqr = uu + alpha * (2.0 * uv + alpha * vv)
+ quad_f = _eval_frictionloss(x_lo, f, rf, Jaref, jv, quad)
+ lo += _eval_pt(quad_f, lo_alpha)
+ quad_f = _eval_frictionloss(x_hi, f, rf, Jaref, jv, quad)
+ hi += _eval_pt(quad_f, hi_alpha)
+ quad_f = _eval_frictionloss(x_mid, f, rf, Jaref, jv, quad)
+ mid += _eval_pt(quad_f, mid_alpha)
- # no tangential force: top or bottom zone
- if Tsqr <= 0.0:
- # bottom zone: quadratic cost
- if N < 0.0:
- wp.atomic_add(out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+ for efcid in range(nef_clip, nefc_clip):
+ if type_in[efcid] == int(types.ConstraintType.CONTACT_ELLIPTIC.value):
+ conid = id_in[efcid]
- # top zone: nothing to do
- # otherwise regular processing
+ efcid0 = contact_efc_address_in[conid, 0]
+ if efcid != efcid0:
+ continue
+
+ efcid1 = contact_efc_address_in[conid, 1]
+ efcid2 = contact_efc_address_in[conid, 2]
+ efc_quad0 = quad_in[efcid0]
+ efc_quad1 = quad_in[efcid1]
+ efc_quad2 = quad_in[efcid2]
+ friction = contact_friction_in[conid]
+
+ lo += _eval_elliptic(impratio_invsqrt, friction, efc_quad0, efc_quad1, efc_quad2, lo_alpha)
+ hi += _eval_elliptic(impratio_invsqrt, friction, efc_quad0, efc_quad1, efc_quad2, hi_alpha)
+ mid += _eval_elliptic(impratio_invsqrt, friction, efc_quad0, efc_quad1, efc_quad2, mid_alpha)
else:
- # tangential force
- T = wp.sqrt(Tsqr)
+ Jaref = Jaref_in[efcid]
+ jv = jv_in[efcid]
+ quad = quad_in[efcid]
- # N >= mu * T : top zone
- if N >= mu * T:
- # nothing to do
- pass
- # mu * N + T <= 0 : bottom zone
- elif mu * N + T <= 0.0:
- wp.atomic_add(out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+ x_lo = Jaref + lo_alpha * jv
+ x_hi = Jaref + hi_alpha * jv
+ x_mid = Jaref + mid_alpha * jv
+ lo += _eval_pt(quad, lo_alpha) * float(x_lo < 0.0)
+ hi += _eval_pt(quad, hi_alpha) * float(x_hi < 0.0)
+ mid += _eval_pt(quad, mid_alpha) * float(x_mid < 0.0)
- # otherwise middle zone
- else:
- # derivatives
- N1 = v0
- T1 = (uv + alpha * vv) / T
- T2 = vv / T - (uv + alpha * vv) * T1 / (T * T)
-
- # add to cost
- cost = wp.vec3(
- 0.5 * dm * (N - mu * T) * (N - mu * T),
- dm * (N - mu * T) * (N1 - mu * T1),
- dm * ((N1 - mu * T1) * (N1 - mu * T1) + (N - mu * T) * (-mu * T2)),
- )
-
- wp.atomic_add(out, worldid, cost)
- else:
- # search point
- x = efc_Jaref_in[worldid, efcid] + alpha * efc_jv_in[worldid, efcid]
-
- # active
- if x < 0.0:
- wp.atomic_add(out, worldid, _eval_pt(efc_quad_in[worldid, efcid], alpha))
+ return lo, hi, mid
@wp.kernel
-def linesearch_iterative_init_gtol_p0_gauss(
+def linesearch_iterative(
# Model:
nv: int,
+ opt_impratio: wp.array(dtype=float),
opt_tolerance: wp.array(dtype=float),
opt_ls_tolerance: wp.array(dtype=float),
+ opt_ls_iterations: int,
stat_meaninertia: float,
# Data in:
+ njmax_in: int,
+ ne_in: wp.array(dtype=int),
+ nf_in: wp.array(dtype=int),
+ nefc_in: wp.array(dtype=int),
+ contact_friction_in: wp.array(dtype=types.vec5),
+ contact_efc_address_in: wp.array2d(dtype=int),
+ efc_type_in: wp.array2d(dtype=int),
+ efc_id_in: wp.array2d(dtype=int),
+ efc_D_in: wp.array2d(dtype=float),
+ efc_frictionloss_in: wp.array2d(dtype=float),
+ efc_Jaref_in: wp.array2d(dtype=float),
efc_search_dot_in: wp.array(dtype=float),
- efc_quad_gauss_in: wp.array(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- # Data out:
- efc_gtol_out: wp.array(dtype=float),
- efc_p0_out: wp.array(dtype=wp.vec3),
-):
- worldid = wp.tid()
-
- if efc_done_in[worldid]:
- return
-
- tolerance = opt_tolerance[worldid]
- ls_tolerance = opt_ls_tolerance[worldid]
- snorm = wp.math.sqrt(efc_search_dot_in[worldid])
- scale = stat_meaninertia * wp.float(wp.max(1, nv))
- efc_gtol_out[worldid] = tolerance * ls_tolerance * snorm * scale
-
- quad = efc_quad_gauss_in[worldid]
- efc_p0_out[worldid] = wp.vec3(quad[0], quad[1], 2.0 * quad[2])
-
-
-@wp.kernel
-def linesearch_iterative_init_p0(
- # Model:
- opt_impratio: wp.array(dtype=float),
- # Data in:
- ncon_in: wp.array(dtype=int),
- ne_in: wp.array(dtype=int),
- nf_in: wp.array(dtype=int),
- nefc_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_efc_address_in: wp.array2d(dtype=int),
- efc_type_in: wp.array2d(dtype=int),
- efc_id_in: wp.array2d(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_frictionloss_in: wp.array2d(dtype=float),
- efc_Jaref_in: wp.array2d(dtype=float),
efc_jv_in: wp.array2d(dtype=float),
efc_quad_in: wp.array2d(dtype=wp.vec3),
- # Data out:
- efc_p0_out: wp.array(dtype=wp.vec3),
-):
- worldid, efcid = wp.tid()
-
- if efcid >= nefc_in[worldid]:
- return
-
- _eval(
- opt_impratio,
- ncon_in,
- ne_in,
- nf_in,
- contact_friction_in,
- contact_efc_address_in,
- efc_type_in,
- efc_id_in,
- efc_D_in,
- efc_frictionloss_in,
- efc_Jaref_in,
- efc_jv_in,
- efc_quad_in,
- worldid,
- efcid,
- 0.0,
- efc_p0_out,
- )
-
-
-@wp.kernel
-def linesearch_iterative_init_lo_gauss(
- # Data in:
efc_quad_gauss_in: wp.array(dtype=wp.vec3),
efc_done_in: wp.array(dtype=bool),
- efc_p0_in: wp.array(dtype=wp.vec3),
- # Data out:
- efc_lo_out: wp.array(dtype=wp.vec3),
- efc_lo_alpha_out: wp.array(dtype=float),
-):
- worldid = wp.tid()
-
- if efc_done_in[worldid]:
- return
-
- p0 = efc_p0_in[worldid]
- alpha = -math.safe_div(p0[1], p0[2])
- efc_lo_out[worldid] = _eval_pt(efc_quad_gauss_in[worldid], alpha)
- efc_lo_alpha_out[worldid] = alpha
-
-
-@wp.kernel
-def linesearch_iterative_init_lo(
- # Model:
- opt_impratio: wp.array(dtype=float),
- # Data in:
- ncon_in: wp.array(dtype=int),
- ne_in: wp.array(dtype=int),
- nf_in: wp.array(dtype=int),
- nefc_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_efc_address_in: wp.array2d(dtype=int),
- efc_type_in: wp.array2d(dtype=int),
- efc_id_in: wp.array2d(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_frictionloss_in: wp.array2d(dtype=float),
- efc_Jaref_in: wp.array2d(dtype=float),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_lo_alpha_in: wp.array(dtype=float),
- # Data out:
- efc_lo_out: wp.array(dtype=wp.vec3),
-):
- worldid, efcid = wp.tid()
-
- if efcid >= nefc_in[worldid]:
- return
-
- _eval(
- opt_impratio,
- ncon_in,
- ne_in,
- nf_in,
- contact_friction_in,
- contact_efc_address_in,
- efc_type_in,
- efc_id_in,
- efc_D_in,
- efc_frictionloss_in,
- efc_Jaref_in,
- efc_jv_in,
- efc_quad_in,
- worldid,
- efcid,
- efc_lo_alpha_in[worldid],
- efc_lo_out,
- )
-
-
-@wp.kernel
-def linesearch_iterative_init_bounds(
- # Data in:
- efc_done_in: wp.array(dtype=bool),
- efc_p0_in: wp.array(dtype=wp.vec3),
- efc_lo_in: wp.array(dtype=wp.vec3),
- efc_lo_alpha_in: wp.array(dtype=float),
- # Data out:
- efc_lo_out: wp.array(dtype=wp.vec3),
- efc_lo_alpha_out: wp.array(dtype=float),
- efc_hi_out: wp.array(dtype=wp.vec3),
- efc_hi_alpha_out: wp.array(dtype=float),
-):
- worldid = wp.tid()
-
- if efc_done_in[worldid]:
- return
-
- p0 = efc_p0_in[worldid]
- lo = efc_lo_in[worldid]
- lo_alpha = efc_lo_alpha_in[worldid]
- lo_less = lo[1] < p0[1]
-
- efc_lo_out[worldid] = wp.where(lo_less, lo, p0)
- efc_lo_alpha_out[worldid] = wp.where(lo_less, lo_alpha, 0.0)
- efc_hi_out[worldid] = wp.where(lo_less, p0, lo)
- efc_hi_alpha_out[worldid] = wp.where(lo_less, 0.0, lo_alpha)
-
-
-@wp.kernel
-def linesearch_iterative_next_alpha_gauss(
- # Data in:
- efc_quad_gauss_in: wp.array(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_ls_done_in: wp.array(dtype=bool),
- efc_lo_in: wp.array(dtype=wp.vec3),
- efc_lo_alpha_in: wp.array(dtype=float),
- efc_hi_in: wp.array(dtype=wp.vec3),
- efc_hi_alpha_in: wp.array(dtype=float),
- # Data out:
- efc_lo_next_out: wp.array(dtype=wp.vec3),
- efc_lo_next_alpha_out: wp.array(dtype=float),
- efc_hi_next_out: wp.array(dtype=wp.vec3),
- efc_hi_next_alpha_out: wp.array(dtype=float),
- efc_mid_out: wp.array(dtype=wp.vec3),
- efc_mid_alpha_out: wp.array(dtype=float),
-):
- worldid = wp.tid()
-
- if efc_ls_done_in[worldid]:
- return
-
- if efc_done_in[worldid]:
- return
-
- quad = efc_quad_gauss_in[worldid]
-
- lo = efc_lo_in[worldid]
- lo_alpha = efc_lo_alpha_in[worldid]
- lo_next_alpha = lo_alpha - math.safe_div(lo[1], lo[2])
- efc_lo_next_out[worldid] = _eval_pt(quad, lo_next_alpha)
- efc_lo_next_alpha_out[worldid] = lo_next_alpha
-
- hi = efc_hi_in[worldid]
- hi_alpha = efc_hi_alpha_in[worldid]
- hi_next_alpha = hi_alpha - math.safe_div(hi[1], hi[2])
- efc_hi_next_out[worldid] = _eval_pt(quad, hi_next_alpha)
- efc_hi_next_alpha_out[worldid] = hi_next_alpha
-
- mid_alpha = 0.5 * (lo_alpha + hi_alpha)
- efc_mid_out[worldid] = _eval_pt(quad, mid_alpha)
- efc_mid_alpha_out[worldid] = mid_alpha
-
-
-@wp.kernel
-def linesearch_iterative_next_quad(
- # Model:
- opt_impratio: wp.array(dtype=float),
- # Data in:
- ncon_in: wp.array(dtype=int),
- ne_in: wp.array(dtype=int),
- nf_in: wp.array(dtype=int),
- nefc_in: wp.array(dtype=int),
- contact_friction_in: wp.array(dtype=types.vec5),
- contact_efc_address_in: wp.array2d(dtype=int),
- efc_type_in: wp.array2d(dtype=int),
- efc_id_in: wp.array2d(dtype=int),
- efc_D_in: wp.array2d(dtype=float),
- efc_frictionloss_in: wp.array2d(dtype=float),
- efc_Jaref_in: wp.array2d(dtype=float),
- efc_jv_in: wp.array2d(dtype=float),
- efc_quad_in: wp.array2d(dtype=wp.vec3),
- efc_done_in: wp.array(dtype=bool),
- efc_ls_done_in: wp.array(dtype=bool),
- efc_lo_next_alpha_in: wp.array(dtype=float),
- efc_hi_next_alpha_in: wp.array(dtype=float),
- efc_mid_alpha_in: wp.array(dtype=float),
- # Data out:
- efc_lo_next_out: wp.array(dtype=wp.vec3),
- efc_hi_next_out: wp.array(dtype=wp.vec3),
- efc_mid_out: wp.array(dtype=wp.vec3),
-):
- worldid, efcid = wp.tid()
-
- if efcid >= nefc_in[worldid]:
- return
-
- if efc_done_in[worldid]:
- return
-
- if efc_ls_done_in[worldid]:
- return
-
- # lo_next
- _eval(
- opt_impratio,
- ncon_in,
- ne_in,
- nf_in,
- contact_friction_in,
- contact_efc_address_in,
- efc_type_in,
- efc_id_in,
- efc_D_in,
- efc_frictionloss_in,
- efc_Jaref_in,
- efc_jv_in,
- efc_quad_in,
- worldid,
- efcid,
- efc_lo_next_alpha_in[worldid],
- efc_lo_next_out,
- )
-
- # hi_next
- _eval(
- opt_impratio,
- ncon_in,
- ne_in,
- nf_in,
- contact_friction_in,
- contact_efc_address_in,
- efc_type_in,
- efc_id_in,
- efc_D_in,
- efc_frictionloss_in,
- efc_Jaref_in,
- efc_jv_in,
- efc_quad_in,
- worldid,
- efcid,
- efc_hi_next_alpha_in[worldid],
- efc_hi_next_out,
- )
-
- # mid
- _eval(
- opt_impratio,
- ncon_in,
- ne_in,
- nf_in,
- contact_friction_in,
- contact_efc_address_in,
- efc_type_in,
- efc_id_in,
- efc_D_in,
- efc_frictionloss_in,
- efc_Jaref_in,
- efc_jv_in,
- efc_quad_in,
- worldid,
- efcid,
- efc_mid_alpha_in[worldid],
- efc_mid_out,
- )
-
-
-@wp.kernel
-def linesearch_iterative_swap(
- # Data in:
- efc_gtol_in: wp.array(dtype=float),
- efc_done_in: wp.array(dtype=bool),
- efc_ls_done_in: wp.array(dtype=bool),
- efc_p0_in: wp.array(dtype=wp.vec3),
- efc_lo_in: wp.array(dtype=wp.vec3),
- efc_lo_alpha_in: wp.array(dtype=float),
- efc_hi_in: wp.array(dtype=wp.vec3),
- efc_hi_alpha_in: wp.array(dtype=float),
- efc_lo_next_in: wp.array(dtype=wp.vec3),
- efc_lo_next_alpha_in: wp.array(dtype=float),
- efc_hi_next_in: wp.array(dtype=wp.vec3),
- efc_hi_next_alpha_in: wp.array(dtype=float),
- efc_mid_in: wp.array(dtype=wp.vec3),
- efc_mid_alpha_in: wp.array(dtype=float),
# Data out:
efc_alpha_out: wp.array(dtype=float),
- efc_ls_done_out: wp.array(dtype=bool),
- efc_lo_out: wp.array(dtype=wp.vec3),
- efc_lo_alpha_out: wp.array(dtype=float),
- efc_hi_out: wp.array(dtype=wp.vec3),
- efc_hi_alpha_out: wp.array(dtype=float),
):
worldid = wp.tid()
if efc_done_in[worldid]:
return
- if efc_ls_done_in[worldid]:
- return
+ impratio = opt_impratio[worldid]
+ efc_type = efc_type_in[worldid]
+ efc_id = efc_id_in[worldid]
+ efc_D = efc_D_in[worldid]
+ efc_frictionloss = efc_frictionloss_in[worldid]
+ efc_Jaref = efc_Jaref_in[worldid]
+ efc_jv = efc_jv_in[worldid]
+ efc_quad = efc_quad_in[worldid]
+ efc_quad_gauss = efc_quad_gauss_in[worldid]
+ tolerance = opt_tolerance[worldid]
+ ls_tolerance = opt_ls_tolerance[worldid]
+ ne_clip = min(njmax_in, ne_in[worldid])
+ nef_clip = min(njmax_in, ne_clip + nf_in[worldid])
+ nefc_clip = min(njmax_in, nefc_in[worldid])
+ impratio_invsqrt = 1.0 / wp.sqrt(impratio)
- lo = efc_lo_in[worldid]
- lo_alpha = efc_lo_alpha_in[worldid]
- hi = efc_hi_in[worldid]
- hi_alpha = efc_hi_alpha_in[worldid]
- lo_next = efc_lo_next_in[worldid]
- lo_next_alpha = efc_lo_next_alpha_in[worldid]
- hi_next = efc_hi_next_in[worldid]
- hi_next_alpha = efc_hi_next_alpha_in[worldid]
- mid = efc_mid_in[worldid]
- mid_alpha = efc_mid_alpha_in[worldid]
+ # Calculate p0
+ snorm = wp.math.sqrt(efc_search_dot_in[worldid])
+ scale = stat_meaninertia * wp.float(wp.max(1, nv))
+ gtol = tolerance * ls_tolerance * snorm * scale
+ p0 = wp.vec3(efc_quad_gauss[0], efc_quad_gauss[1], 2.0 * efc_quad_gauss[2])
+ p0 += _eval_init(
+ contact_friction_in,
+ contact_efc_address_in,
+ ne_clip,
+ nef_clip,
+ nefc_clip,
+ impratio_invsqrt,
+ efc_type,
+ efc_id,
+ efc_D,
+ efc_frictionloss,
+ efc_Jaref,
+ efc_jv,
+ efc_quad,
+ 0.0,
+ )
- # swap lo:
- swap_lo_lo_next = _in_bracket(lo, lo_next)
- lo = wp.where(swap_lo_lo_next, lo_next, lo)
- lo_alpha = wp.where(swap_lo_lo_next, lo_next_alpha, lo_alpha)
- swap_lo_mid = _in_bracket(lo, mid)
- lo = wp.where(swap_lo_mid, mid, lo)
- lo_alpha = wp.where(swap_lo_mid, mid_alpha, lo_alpha)
- swap_lo_hi_next = _in_bracket(lo, hi_next)
- lo = wp.where(swap_lo_hi_next, hi_next, lo)
- lo_alpha = wp.where(swap_lo_hi_next, hi_next_alpha, lo_alpha)
- efc_lo_out[worldid] = lo
- efc_lo_alpha_out[worldid] = lo_alpha
- swap_lo = swap_lo_lo_next or swap_lo_mid or swap_lo_hi_next
+ # Calculate lo bound
+ lo_alpha_in = -math.safe_div(p0[1], p0[2])
+ lo_in = _eval_pt(efc_quad_gauss, lo_alpha_in)
+ lo_in += _eval_init(
+ contact_friction_in,
+ contact_efc_address_in,
+ ne_clip,
+ nef_clip,
+ nefc_clip,
+ impratio_invsqrt,
+ efc_type,
+ efc_id,
+ efc_D,
+ efc_frictionloss,
+ efc_Jaref,
+ efc_jv,
+ efc_quad,
+ lo_alpha_in,
+ )
- # swap hi:
- swap_hi_hi_next = _in_bracket(hi, hi_next)
- hi = wp.where(swap_hi_hi_next, hi_next, hi)
- hi_alpha = wp.where(swap_hi_hi_next, hi_next_alpha, hi_alpha)
- swap_hi_mid = _in_bracket(hi, mid)
- hi = wp.where(swap_hi_mid, mid, hi)
- hi_alpha = wp.where(swap_hi_mid, mid_alpha, hi_alpha)
- swap_hi_lo_next = _in_bracket(hi, lo_next)
- hi = wp.where(swap_hi_lo_next, lo_next, hi)
- hi_alpha = wp.where(swap_hi_lo_next, lo_next_alpha, hi_alpha)
- efc_hi_out[worldid] = hi
- efc_hi_alpha_out[worldid] = hi_alpha
- swap_hi = swap_hi_hi_next or swap_hi_mid or swap_hi_lo_next
+ # Initialize bounds
+ lo_less = lo_in[1] < p0[1]
+ lo = wp.where(lo_less, lo_in, p0)
+ lo_alpha = wp.where(lo_less, lo_alpha_in, 0.0)
+ hi = wp.where(lo_less, p0, lo_in)
+ hi_alpha = wp.where(lo_less, 0.0, lo_alpha_in)
- # if we did not adjust the interval, we are done
- # also done if either low or hi slope is nearly flat
- gtol = efc_gtol_in[worldid]
- efc_ls_done_out[worldid] = (not swap_lo and not swap_hi) or (lo[1] < 0 and lo[1] > -gtol) or (hi[1] > 0 and hi[1] < gtol)
+ # Launch main linesearch iterative loop
+ alpha = float(0.0)
+ for _ in range(opt_ls_iterations):
+ lo_next_alpha = lo_alpha - math.safe_div(lo[1], lo[2])
+ hi_next_alpha = hi_alpha - math.safe_div(hi[1], hi[2])
+ mid_alpha = 0.5 * (lo_alpha + hi_alpha)
+
+ lo_next, hi_next, mid = _eval(
+ contact_friction_in,
+ contact_efc_address_in,
+ ne_clip,
+ nef_clip,
+ nefc_clip,
+ impratio_invsqrt,
+ efc_type,
+ efc_id,
+ efc_D,
+ efc_frictionloss,
+ efc_Jaref,
+ efc_jv,
+ efc_quad,
+ lo_next_alpha,
+ hi_next_alpha,
+ mid_alpha,
+ )
+ lo_next += _eval_pt(efc_quad_gauss, lo_next_alpha)
+ hi_next += _eval_pt(efc_quad_gauss, hi_next_alpha)
+ mid += _eval_pt(efc_quad_gauss, mid_alpha)
+
+ # swap lo:
+ swap_lo_lo_next = _in_bracket(lo, lo_next)
+ lo = wp.where(swap_lo_lo_next, lo_next, lo)
+ lo_alpha = wp.where(swap_lo_lo_next, lo_next_alpha, lo_alpha)
+ swap_lo_mid = _in_bracket(lo, mid)
+ lo = wp.where(swap_lo_mid, mid, lo)
+ lo_alpha = wp.where(swap_lo_mid, mid_alpha, lo_alpha)
+ swap_lo_hi_next = _in_bracket(lo, hi_next)
+ lo = wp.where(swap_lo_hi_next, hi_next, lo)
+ lo_alpha = wp.where(swap_lo_hi_next, hi_next_alpha, lo_alpha)
+ swap_lo = swap_lo_lo_next or swap_lo_mid or swap_lo_hi_next
+
+ # swap hi:
+ swap_hi_hi_next = _in_bracket(hi, hi_next)
+ hi = wp.where(swap_hi_hi_next, hi_next, hi)
+ hi_alpha = wp.where(swap_hi_hi_next, hi_next_alpha, hi_alpha)
+ swap_hi_mid = _in_bracket(hi, mid)
+ hi = wp.where(swap_hi_mid, mid, hi)
+ hi_alpha = wp.where(swap_hi_mid, mid_alpha, hi_alpha)
+ swap_hi_lo_next = _in_bracket(hi, lo_next)
+ hi = wp.where(swap_hi_lo_next, lo_next, hi)
+ hi_alpha = wp.where(swap_hi_lo_next, lo_next_alpha, hi_alpha)
+ swap_hi = swap_hi_hi_next or swap_hi_mid or swap_hi_lo_next
+
+ # if we did not adjust the interval, we are done
+ # also done if either low or hi slope is nearly flat
+ ls_done = (not swap_lo and not swap_hi) or (lo[1] < 0 and lo[1] > -gtol) or (hi[1] > 0 and hi[1] < gtol)
+
+ # update alpha if we have an improvement
+ improved = lo[0] < p0[0] or hi[0] < p0[0]
+ lo_better = lo[0] < hi[0]
+ alpha = wp.where(improved and lo_better, lo_alpha, alpha)
+ alpha = wp.where(improved and not lo_better, hi_alpha, alpha)
+ if ls_done:
+ break
- # update alpha if we have an improvement
- p0 = efc_p0_in[worldid]
- alpha = 0.0
- improved = lo[0] < p0[0] or hi[0] < p0[0]
- lo_better = lo[0] < hi[0]
- alpha = wp.where(improved and lo_better, lo_alpha, alpha)
- alpha = wp.where(improved and not lo_better, hi_alpha, alpha)
efc_alpha_out[worldid] = alpha
def _linesearch_iterative(m: types.Model, d: types.Data):
"""Iterative linesearch."""
- d.efc.ls_done.zero_()
-
wp.launch(
- linesearch_iterative_init_gtol_p0_gauss,
+ linesearch_iterative,
dim=(d.nworld,),
- inputs=[m.nv, m.opt.tolerance, m.opt.ls_tolerance, m.stat.meaninertia, d.efc.search_dot, d.efc.quad_gauss, d.efc.done],
- outputs=[d.efc.gtol, d.efc.p0],
- )
-
- wp.launch(
- linesearch_iterative_init_p0,
- dim=(d.nworld, d.njmax),
inputs=[
+ m.nv,
m.opt.impratio,
- d.ncon,
+ m.opt.tolerance,
+ m.opt.ls_tolerance,
+ m.opt.ls_iterations,
+ m.stat.meaninertia,
+ d.njmax,
d.ne,
d.nf,
d.nefc,
@@ -616,109 +477,15 @@ def _linesearch_iterative(m: types.Model, d: types.Data):
d.efc.D,
d.efc.frictionloss,
d.efc.Jaref,
+ d.efc.search_dot,
d.efc.jv,
d.efc.quad,
+ d.efc.quad_gauss,
+ d.efc.done,
],
- outputs=[d.efc.p0],
+ outputs=[d.efc.alpha],
)
- wp.launch(
- linesearch_iterative_init_lo_gauss,
- dim=(d.nworld,),
- inputs=[d.efc.quad_gauss, d.efc.done, d.efc.p0],
- outputs=[d.efc.lo, d.efc.lo_alpha],
- )
- wp.launch(
- linesearch_iterative_init_lo,
- dim=(d.nworld, d.njmax),
- inputs=[
- m.opt.impratio,
- d.ncon,
- d.ne,
- d.nf,
- d.nefc,
- d.contact.friction,
- d.contact.efc_address,
- d.efc.type,
- d.efc.id,
- d.efc.D,
- d.efc.frictionloss,
- d.efc.Jaref,
- d.efc.jv,
- d.efc.quad,
- d.efc.lo_alpha,
- ],
- outputs=[d.efc.lo],
- )
-
- # set the lo/hi interval bounds
- wp.launch(
- linesearch_iterative_init_bounds,
- dim=(d.nworld,),
- inputs=[d.efc.done, d.efc.p0, d.efc.lo, d.efc.lo_alpha],
- outputs=[d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha],
- )
-
- for _ in range(m.opt.ls_iterations):
- # NOTE: we always launch ls_iterations kernels, but the kernels may early exit if done
- # is true. this preserves cudagraph requirements (no dynamic kernel launching) at the
- # expense of extra launches
- wp.launch(
- linesearch_iterative_next_alpha_gauss,
- dim=(d.nworld,),
- inputs=[d.efc.quad_gauss, d.efc.done, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha],
- outputs=[d.efc.lo_next, d.efc.lo_next_alpha, d.efc.hi_next, d.efc.hi_next_alpha, d.efc.mid, d.efc.mid_alpha],
- )
-
- wp.launch(
- linesearch_iterative_next_quad,
- dim=(d.nworld, d.njmax),
- inputs=[
- m.opt.impratio,
- d.ncon,
- d.ne,
- d.nf,
- d.nefc,
- d.contact.friction,
- d.contact.efc_address,
- d.efc.type,
- d.efc.id,
- d.efc.D,
- d.efc.frictionloss,
- d.efc.Jaref,
- d.efc.jv,
- d.efc.quad,
- d.efc.done,
- d.efc.ls_done,
- d.efc.lo_next_alpha,
- d.efc.hi_next_alpha,
- d.efc.mid_alpha,
- ],
- outputs=[d.efc.lo_next, d.efc.hi_next, d.efc.mid],
- )
-
- wp.launch(
- linesearch_iterative_swap,
- dim=(d.nworld,),
- inputs=[
- d.efc.gtol,
- d.efc.done,
- d.efc.ls_done,
- d.efc.p0,
- d.efc.lo,
- d.efc.lo_alpha,
- d.efc.hi,
- d.efc.hi_alpha,
- d.efc.lo_next,
- d.efc.lo_next_alpha,
- d.efc.hi_next,
- d.efc.hi_next_alpha,
- d.efc.mid,
- d.efc.mid_alpha,
- ],
- outputs=[d.efc.alpha, d.efc.ls_done, d.efc.lo, d.efc.lo_alpha, d.efc.hi, d.efc.hi_alpha],
- )
-
@wp.func
def _log_scale(min_value: float, max_value: float, num_values: int, i: int) -> float:
@@ -1417,21 +1184,6 @@ def update_constraint_efc(
efc_state_out[worldid, efcid] = int(types.ConstraintState.CONE.value)
-@wp.kernel
-def update_constraint_zero_qfrc_constraint(
- # Data in:
- efc_done_in: wp.array(dtype=bool),
- # Data out:
- qfrc_constraint_out: wp.array2d(dtype=float),
-):
- worldid, dofid = wp.tid()
-
- if efc_done_in[worldid]:
- return
-
- qfrc_constraint_out[worldid, dofid] = 0.0
-
-
@wp.kernel
def update_constraint_init_qfrc_constraint(
# Data in:
@@ -1454,7 +1206,7 @@ def update_constraint_init_qfrc_constraint(
force = efc_force_in[worldid, efcid]
sum_qfrc += efc_J * force
- qfrc_constraint_out[worldid, dofid] += sum_qfrc
+ qfrc_constraint_out[worldid, dofid] = sum_qfrc
@cache_kernel
@@ -1530,13 +1282,6 @@ def _update_constraint(m: types.Model, d: types.Data):
)
# qfrc_constraint = efc_J.T @ efc_force
- wp.launch(
- update_constraint_zero_qfrc_constraint,
- dim=(d.nworld, m.nv),
- inputs=[d.efc.done],
- outputs=[d.qfrc_constraint],
- )
-
wp.launch(
update_constraint_init_qfrc_constraint,
dim=(d.nworld, m.nv),
@@ -2023,7 +1768,7 @@ def _update_gradient(m: types.Model, d: types.Data):
)
else:
wp.launch_tiled(
- update_gradient_cholesky_blocked(32),
+ update_gradient_cholesky_blocked(16),
dim=(d.nworld,),
inputs=[
d.efc.grad.reshape(shape=(d.nworld, m.nv, 1)),
@@ -2242,16 +1987,11 @@ def create_context(m: types.Model, d: types.Data, grad: bool = True):
_update_gradient(m, d)
-def _copy_acc(m: types.Model, d: types.Data):
- wp.copy(d.qacc, d.qacc_smooth)
- wp.copy(d.qacc_warmstart, d.qacc_smooth)
- d.solver_niter.fill_(0)
-
-
@event_scope
def solve(m: types.Model, d: types.Data):
if d.njmax == 0:
- _copy_acc(m, d)
+ wp.copy(d.qacc, d.qacc_smooth)
+ d.solver_niter.fill_(0)
else:
_solve(m, d)
@@ -2295,5 +2035,3 @@ def _solve(m: types.Model, d: types.Data):
# It should be removed when JAX becomes compatible.
for _ in range(m.opt.iterations):
_solver_iteration(m, d)
-
- wp.copy(d.qacc_warmstart, d.qacc)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py
index fa192929..eb424372 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/solver_test.py
@@ -254,9 +254,7 @@ class SolverTest(parameterized.TestCase):
ls_parallel=ls_parallel,
)
- qacc_warmstart = mjd.qacc_warmstart.copy()
mujoco.mj_forward(mjm, mjd)
- mjd.qacc_warmstart = qacc_warmstart
d.qacc.zero_()
d.qfrc_constraint.zero_()
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py
index 29d2693f..7ebd0108 100644
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py
+++ b/mjx/mujoco/mjx/third_party/mujoco_warp/_src/types.py
@@ -55,7 +55,7 @@ class BlockDim:
cholesky_solve: int = 256
cholesky_factorize_solve: int = 256
# solver
- update_gradient_cholesky: int = 256
+ update_gradient_cholesky: int = 64
# support
mul_m_dense: int = 256
@@ -570,6 +570,7 @@ 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
"""
timestep: wp.array(dtype=float)
@@ -600,6 +601,7 @@ class Option:
sdf_initpoints: int
sdf_iterations: int
run_collision_detection: bool # warp only
+ legacy_gjk: bool
@dataclasses.dataclass
@@ -639,7 +641,6 @@ class Constraint:
cost: constraint + Gauss cost (nworld,)
prev_cost: cost from previous iter (nworld,)
state: constraint state (nworld, njmax)
- gtol: linesearch termination tolerance (nworld,)
mv: qM @ search (nworld, nv)
jv: efc_J @ search (nworld, njmax)
quad: quadratic cost coefficients (nworld, njmax, 3)
@@ -650,18 +651,6 @@ class Constraint:
prev_Mgrad: previous Mgrad (nworld, nv)
beta: polak-ribiere beta (nworld,)
done: solver done (nworld,)
- ls_done: linesearch done (nworld,)
- p0: initial point (nworld, 3)
- lo: low point bounding the line search interval (nworld, 3)
- lo_alpha: alpha for low point (nworld,)
- hi: high point bounding the line search interval (nworld, 3)
- hi_alpha: alpha for high point (nworld,)
- lo_next: next low point (nworld, 3)
- lo_next_alpha: alpha for next low point (nworld,)
- hi_next: next high point (nworld, 3)
- hi_next_alpha: alpha for next high point (nworld,)
- mid: loss at mid_alpha (nworld, 3)
- mid_alpha: midpoint between lo_alpha and hi_alpha (nworld,)
cost_candidate: costs associated with step sizes (nworld, nlsp)
"""
@@ -688,7 +677,6 @@ class Constraint:
cost: wp.array(dtype=float)
prev_cost: wp.array(dtype=float)
state: wp.array2d(dtype=int)
- gtol: wp.array(dtype=float)
mv: wp.array2d(dtype=float)
jv: wp.array2d(dtype=float)
quad: wp.array2d(dtype=wp.vec3)
@@ -700,18 +688,6 @@ class Constraint:
beta: wp.array(dtype=float)
done: wp.array(dtype=bool)
# linesearch
- ls_done: wp.array(dtype=bool)
- p0: wp.array(dtype=wp.vec3)
- lo: wp.array(dtype=wp.vec3)
- lo_alpha: wp.array(dtype=float)
- hi: wp.array(dtype=wp.vec3)
- hi_alpha: wp.array(dtype=float)
- lo_next: wp.array(dtype=wp.vec3)
- lo_next_alpha: wp.array(dtype=float)
- hi_next: wp.array(dtype=wp.vec3)
- hi_next_alpha: wp.array(dtype=float)
- mid: wp.array(dtype=wp.vec3)
- mid_alpha: wp.array(dtype=float)
cost_candidate: wp.array2d(dtype=float)
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml
deleted file mode 100644
index 89e0adfa..00000000
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/cloth.xml
+++ /dev/null
@@ -1,45 +0,0 @@
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
diff --git a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml b/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml
deleted file mode 100644
index 9237b943..00000000
--- a/mjx/mujoco/mjx/third_party/mujoco_warp/test_data/flex/mannequin.xml
+++ /dev/null
@@ -1,174 +0,0 @@
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
-
diff --git a/mjx/mujoco/mjx/warp/collision_driver.py b/mjx/mujoco/mjx/warp/collision_driver.py
index 798d9f8b..ca4f200a 100644
--- a/mjx/mujoco/mjx/warp/collision_driver.py
+++ b/mjx/mujoco/mjx/warp/collision_driver.py
@@ -42,6 +42,7 @@ _e = mjwarp.Constraint(
**{f.name: None for f in dataclasses.fields(mjwarp.Constraint) if f.init}
)
+
@ffi.format_args_for_warp
def _collision_shim(
# Model
@@ -105,6 +106,7 @@ def _collision_shim(
opt__epa_iterations: int,
opt__gjk_iterations: int,
opt__graph_conditional: bool,
+ opt__legacy_gjk: bool,
opt__sdf_initpoints: int,
opt__sdf_iterations: int,
# Data
@@ -201,6 +203,7 @@ def _collision_shim(
_m.opt.epa_iterations = opt__epa_iterations
_m.opt.gjk_iterations = opt__gjk_iterations
_m.opt.graph_conditional = opt__graph_conditional
+ _m.opt.legacy_gjk = opt__legacy_gjk
_m.opt.sdf_initpoints = opt__sdf_initpoints
_m.opt.sdf_iterations = opt__sdf_iterations
_m.pair_dim = pair_dim
@@ -400,6 +403,7 @@ def _collision_jax_impl(m: types.Model, d: types.Data):
m.opt._impl.epa_iterations,
m.opt._impl.gjk_iterations,
m.opt._impl.graph_conditional,
+ m.opt._impl.legacy_gjk,
m.opt._impl.sdf_initpoints,
m.opt._impl.sdf_iterations,
d._impl.nconmax,
diff --git a/mjx/mujoco/mjx/warp/forward.py b/mjx/mujoco/mjx/warp/forward.py
index d2437fdc..9aab9aba 100644
--- a/mjx/mujoco/mjx/warp/forward.py
+++ b/mjx/mujoco/mjx/warp/forward.py
@@ -227,7 +227,6 @@ def _forward_shim(
nlsp: int,
nmeshface: int,
nmocap: int,
- nsensor: int,
nsensortaxel: int,
nsite: int,
ntendon: int,
@@ -331,6 +330,7 @@ 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,
@@ -496,24 +496,11 @@ def _forward_shim(
efc__gauss: wp.array(dtype=float),
efc__grad: wp.array2d(dtype=float),
efc__grad_dot: wp.array(dtype=float),
- efc__gtol: wp.array(dtype=float),
efc__h: wp.array3d(dtype=float),
- efc__hi: wp.array(dtype=wp.vec3),
- efc__hi_alpha: wp.array(dtype=float),
- efc__hi_next: wp.array(dtype=wp.vec3),
- efc__hi_next_alpha: wp.array(dtype=float),
efc__id: wp.array2d(dtype=int),
efc__jv: wp.array2d(dtype=float),
- efc__lo: wp.array(dtype=wp.vec3),
- efc__lo_alpha: wp.array(dtype=float),
- efc__lo_next: wp.array(dtype=wp.vec3),
- efc__lo_next_alpha: wp.array(dtype=float),
- efc__ls_done: wp.array(dtype=bool),
efc__margin: wp.array2d(dtype=float),
- efc__mid: wp.array(dtype=wp.vec3),
- efc__mid_alpha: wp.array(dtype=float),
efc__mv: wp.array2d(dtype=float),
- efc__p0: wp.array(dtype=wp.vec3),
efc__pos: wp.array2d(dtype=float),
efc__prev_Mgrad: wp.array2d(dtype=float),
efc__prev_cost: wp.array(dtype=float),
@@ -710,7 +697,6 @@ def _forward_shim(
_m.nlsp = nlsp
_m.nmeshface = nmeshface
_m.nmocap = nmocap
- _m.nsensor = nsensor
_m.nsensortaxel = nsensortaxel
_m.nsite = nsite
_m.ntendon = ntendon
@@ -733,6 +719,7 @@ 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
@@ -880,24 +867,11 @@ def _forward_shim(
_d.efc.gauss = efc__gauss
_d.efc.grad = efc__grad
_d.efc.grad_dot = efc__grad_dot
- _d.efc.gtol = efc__gtol
_d.efc.h = efc__h
- _d.efc.hi = efc__hi
- _d.efc.hi_alpha = efc__hi_alpha
- _d.efc.hi_next = efc__hi_next
- _d.efc.hi_next_alpha = efc__hi_next_alpha
_d.efc.id = efc__id
_d.efc.jv = efc__jv
- _d.efc.lo = efc__lo
- _d.efc.lo_alpha = efc__lo_alpha
- _d.efc.lo_next = efc__lo_next
- _d.efc.lo_next_alpha = efc__lo_next_alpha
- _d.efc.ls_done = efc__ls_done
_d.efc.margin = efc__margin
- _d.efc.mid = efc__mid
- _d.efc.mid_alpha = efc__mid_alpha
_d.efc.mv = efc__mv
- _d.efc.p0 = efc__p0
_d.efc.pos = efc__pos
_d.efc.prev_Mgrad = efc__prev_Mgrad
_d.efc.prev_cost = efc__prev_cost
@@ -1161,24 +1135,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'efc__gauss': d._impl.efc__gauss.shape,
'efc__grad': d._impl.efc__grad.shape,
'efc__grad_dot': d._impl.efc__grad_dot.shape,
- 'efc__gtol': d._impl.efc__gtol.shape,
'efc__h': d._impl.efc__h.shape,
- 'efc__hi': d._impl.efc__hi.shape,
- 'efc__hi_alpha': d._impl.efc__hi_alpha.shape,
- 'efc__hi_next': d._impl.efc__hi_next.shape,
- 'efc__hi_next_alpha': d._impl.efc__hi_next_alpha.shape,
'efc__id': d._impl.efc__id.shape,
'efc__jv': d._impl.efc__jv.shape,
- 'efc__lo': d._impl.efc__lo.shape,
- 'efc__lo_alpha': d._impl.efc__lo_alpha.shape,
- 'efc__lo_next': d._impl.efc__lo_next.shape,
- 'efc__lo_next_alpha': d._impl.efc__lo_next_alpha.shape,
- 'efc__ls_done': d._impl.efc__ls_done.shape,
'efc__margin': d._impl.efc__margin.shape,
- 'efc__mid': d._impl.efc__mid.shape,
- 'efc__mid_alpha': d._impl.efc__mid_alpha.shape,
'efc__mv': d._impl.efc__mv.shape,
- 'efc__p0': d._impl.efc__p0.shape,
'efc__pos': d._impl.efc__pos.shape,
'efc__prev_Mgrad': d._impl.efc__prev_Mgrad.shape,
'efc__prev_cost': d._impl.efc__prev_cost.shape,
@@ -1193,7 +1154,7 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
}
jf = ffi.jax_callable_variadic_tuple(
_forward_shim,
- num_outputs=177,
+ num_outputs=164,
output_dims=output_dims,
vmap_method=None,
in_out_argnames={
@@ -1345,24 +1306,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'efc__gauss',
'efc__grad',
'efc__grad_dot',
- 'efc__gtol',
'efc__h',
- 'efc__hi',
- 'efc__hi_alpha',
- 'efc__hi_next',
- 'efc__hi_next_alpha',
'efc__id',
'efc__jv',
- 'efc__lo',
- 'efc__lo_alpha',
- 'efc__lo_next',
- 'efc__lo_next_alpha',
- 'efc__ls_done',
'efc__margin',
- 'efc__mid',
- 'efc__mid_alpha',
'efc__mv',
- 'efc__p0',
'efc__pos',
'efc__prev_Mgrad',
'efc__prev_cost',
@@ -1558,7 +1506,6 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
m._impl.nlsp,
m.nmeshface,
m.nmocap,
- m.nsensor,
m._impl.nsensortaxel,
m.nsite,
m.ntendon,
@@ -1662,6 +1609,7 @@ 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,
@@ -1826,24 +1774,11 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
d._impl.efc__gauss,
d._impl.efc__grad,
d._impl.efc__grad_dot,
- d._impl.efc__gtol,
d._impl.efc__h,
- d._impl.efc__hi,
- d._impl.efc__hi_alpha,
- d._impl.efc__hi_next,
- d._impl.efc__hi_next_alpha,
d._impl.efc__id,
d._impl.efc__jv,
- d._impl.efc__lo,
- d._impl.efc__lo_alpha,
- d._impl.efc__lo_next,
- d._impl.efc__lo_next_alpha,
- d._impl.efc__ls_done,
d._impl.efc__margin,
- d._impl.efc__mid,
- d._impl.efc__mid_alpha,
d._impl.efc__mv,
- d._impl.efc__p0,
d._impl.efc__pos,
d._impl.efc__prev_Mgrad,
d._impl.efc__prev_cost,
@@ -2005,35 +1940,22 @@ def _forward_jax_impl(m: types.Model, d: types.Data):
'_impl.efc__gauss': out[145],
'_impl.efc__grad': out[146],
'_impl.efc__grad_dot': out[147],
- '_impl.efc__gtol': out[148],
- '_impl.efc__h': out[149],
- '_impl.efc__hi': out[150],
- '_impl.efc__hi_alpha': out[151],
- '_impl.efc__hi_next': out[152],
- '_impl.efc__hi_next_alpha': out[153],
- '_impl.efc__id': out[154],
- '_impl.efc__jv': out[155],
- '_impl.efc__lo': out[156],
- '_impl.efc__lo_alpha': out[157],
- '_impl.efc__lo_next': out[158],
- '_impl.efc__lo_next_alpha': out[159],
- '_impl.efc__ls_done': out[160],
- '_impl.efc__margin': out[161],
- '_impl.efc__mid': out[162],
- '_impl.efc__mid_alpha': out[163],
- '_impl.efc__mv': out[164],
- '_impl.efc__p0': out[165],
- '_impl.efc__pos': out[166],
- '_impl.efc__prev_Mgrad': out[167],
- '_impl.efc__prev_cost': out[168],
- '_impl.efc__prev_grad': out[169],
- '_impl.efc__quad': out[170],
- '_impl.efc__quad_gauss': out[171],
- '_impl.efc__search': out[172],
- '_impl.efc__search_dot': out[173],
- '_impl.efc__state': out[174],
- '_impl.efc__type': out[175],
- '_impl.efc__vel': out[176],
+ '_impl.efc__h': out[148],
+ '_impl.efc__id': out[149],
+ '_impl.efc__jv': out[150],
+ '_impl.efc__margin': out[151],
+ '_impl.efc__mv': out[152],
+ '_impl.efc__pos': out[153],
+ '_impl.efc__prev_Mgrad': out[154],
+ '_impl.efc__prev_cost': out[155],
+ '_impl.efc__prev_grad': out[156],
+ '_impl.efc__quad': out[157],
+ '_impl.efc__quad_gauss': out[158],
+ '_impl.efc__search': out[159],
+ '_impl.efc__search_dot': out[160],
+ '_impl.efc__state': out[161],
+ '_impl.efc__type': out[162],
+ '_impl.efc__vel': out[163],
})
return d
@@ -2254,7 +2176,6 @@ def _step_shim(
nlsp: int,
nmeshface: int,
nmocap: int,
- nsensor: int,
nsensortaxel: int,
nsite: int,
ntendon: int,
@@ -2359,6 +2280,7 @@ 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,
@@ -2536,24 +2458,11 @@ def _step_shim(
efc__gauss: wp.array(dtype=float),
efc__grad: wp.array2d(dtype=float),
efc__grad_dot: wp.array(dtype=float),
- efc__gtol: wp.array(dtype=float),
efc__h: wp.array3d(dtype=float),
- efc__hi: wp.array(dtype=wp.vec3),
- efc__hi_alpha: wp.array(dtype=float),
- efc__hi_next: wp.array(dtype=wp.vec3),
- efc__hi_next_alpha: wp.array(dtype=float),
efc__id: wp.array2d(dtype=int),
efc__jv: wp.array2d(dtype=float),
- efc__lo: wp.array(dtype=wp.vec3),
- efc__lo_alpha: wp.array(dtype=float),
- efc__lo_next: wp.array(dtype=wp.vec3),
- efc__lo_next_alpha: wp.array(dtype=float),
- efc__ls_done: wp.array(dtype=bool),
efc__margin: wp.array2d(dtype=float),
- efc__mid: wp.array(dtype=wp.vec3),
- efc__mid_alpha: wp.array(dtype=float),
efc__mv: wp.array2d(dtype=float),
- efc__p0: wp.array(dtype=wp.vec3),
efc__pos: wp.array2d(dtype=float),
efc__prev_Mgrad: wp.array2d(dtype=float),
efc__prev_cost: wp.array(dtype=float),
@@ -2751,7 +2660,6 @@ def _step_shim(
_m.nlsp = nlsp
_m.nmeshface = nmeshface
_m.nmocap = nmocap
- _m.nsensor = nsensor
_m.nsensortaxel = nsensortaxel
_m.nsite = nsite
_m.ntendon = ntendon
@@ -2775,6 +2683,7 @@ 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
@@ -2924,24 +2833,11 @@ def _step_shim(
_d.efc.gauss = efc__gauss
_d.efc.grad = efc__grad
_d.efc.grad_dot = efc__grad_dot
- _d.efc.gtol = efc__gtol
_d.efc.h = efc__h
- _d.efc.hi = efc__hi
- _d.efc.hi_alpha = efc__hi_alpha
- _d.efc.hi_next = efc__hi_next
- _d.efc.hi_next_alpha = efc__hi_next_alpha
_d.efc.id = efc__id
_d.efc.jv = efc__jv
- _d.efc.lo = efc__lo
- _d.efc.lo_alpha = efc__lo_alpha
- _d.efc.lo_next = efc__lo_next
- _d.efc.lo_next_alpha = efc__lo_next_alpha
- _d.efc.ls_done = efc__ls_done
_d.efc.margin = efc__margin
- _d.efc.mid = efc__mid
- _d.efc.mid_alpha = efc__mid_alpha
_d.efc.mv = efc__mv
- _d.efc.p0 = efc__p0
_d.efc.pos = efc__pos
_d.efc.prev_Mgrad = efc__prev_Mgrad
_d.efc.prev_cost = efc__prev_cost
@@ -3227,24 +3123,11 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'efc__gauss': d._impl.efc__gauss.shape,
'efc__grad': d._impl.efc__grad.shape,
'efc__grad_dot': d._impl.efc__grad_dot.shape,
- 'efc__gtol': d._impl.efc__gtol.shape,
'efc__h': d._impl.efc__h.shape,
- 'efc__hi': d._impl.efc__hi.shape,
- 'efc__hi_alpha': d._impl.efc__hi_alpha.shape,
- 'efc__hi_next': d._impl.efc__hi_next.shape,
- 'efc__hi_next_alpha': d._impl.efc__hi_next_alpha.shape,
'efc__id': d._impl.efc__id.shape,
'efc__jv': d._impl.efc__jv.shape,
- 'efc__lo': d._impl.efc__lo.shape,
- 'efc__lo_alpha': d._impl.efc__lo_alpha.shape,
- 'efc__lo_next': d._impl.efc__lo_next.shape,
- 'efc__lo_next_alpha': d._impl.efc__lo_next_alpha.shape,
- 'efc__ls_done': d._impl.efc__ls_done.shape,
'efc__margin': d._impl.efc__margin.shape,
- 'efc__mid': d._impl.efc__mid.shape,
- 'efc__mid_alpha': d._impl.efc__mid_alpha.shape,
'efc__mv': d._impl.efc__mv.shape,
- 'efc__p0': d._impl.efc__p0.shape,
'efc__pos': d._impl.efc__pos.shape,
'efc__prev_Mgrad': d._impl.efc__prev_Mgrad.shape,
'efc__prev_cost': d._impl.efc__prev_cost.shape,
@@ -3259,7 +3142,7 @@ def _step_jax_impl(m: types.Model, d: types.Data):
}
jf = ffi.jax_callable_variadic_tuple(
_step_shim,
- num_outputs=189,
+ num_outputs=176,
output_dims=output_dims,
vmap_method=None,
in_out_argnames={
@@ -3423,24 +3306,11 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'efc__gauss',
'efc__grad',
'efc__grad_dot',
- 'efc__gtol',
'efc__h',
- 'efc__hi',
- 'efc__hi_alpha',
- 'efc__hi_next',
- 'efc__hi_next_alpha',
'efc__id',
'efc__jv',
- 'efc__lo',
- 'efc__lo_alpha',
- 'efc__lo_next',
- 'efc__lo_next_alpha',
- 'efc__ls_done',
'efc__margin',
- 'efc__mid',
- 'efc__mid_alpha',
'efc__mv',
- 'efc__p0',
'efc__pos',
'efc__prev_Mgrad',
'efc__prev_cost',
@@ -3637,7 +3507,6 @@ def _step_jax_impl(m: types.Model, d: types.Data):
m._impl.nlsp,
m.nmeshface,
m.nmocap,
- m.nsensor,
m._impl.nsensortaxel,
m.nsite,
m.ntendon,
@@ -3742,6 +3611,7 @@ 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,
@@ -3918,24 +3788,11 @@ def _step_jax_impl(m: types.Model, d: types.Data):
d._impl.efc__gauss,
d._impl.efc__grad,
d._impl.efc__grad_dot,
- d._impl.efc__gtol,
d._impl.efc__h,
- d._impl.efc__hi,
- d._impl.efc__hi_alpha,
- d._impl.efc__hi_next,
- d._impl.efc__hi_next_alpha,
d._impl.efc__id,
d._impl.efc__jv,
- d._impl.efc__lo,
- d._impl.efc__lo_alpha,
- d._impl.efc__lo_next,
- d._impl.efc__lo_next_alpha,
- d._impl.efc__ls_done,
d._impl.efc__margin,
- d._impl.efc__mid,
- d._impl.efc__mid_alpha,
d._impl.efc__mv,
- d._impl.efc__p0,
d._impl.efc__pos,
d._impl.efc__prev_Mgrad,
d._impl.efc__prev_cost,
@@ -4109,35 +3966,22 @@ def _step_jax_impl(m: types.Model, d: types.Data):
'_impl.efc__gauss': out[157],
'_impl.efc__grad': out[158],
'_impl.efc__grad_dot': out[159],
- '_impl.efc__gtol': out[160],
- '_impl.efc__h': out[161],
- '_impl.efc__hi': out[162],
- '_impl.efc__hi_alpha': out[163],
- '_impl.efc__hi_next': out[164],
- '_impl.efc__hi_next_alpha': out[165],
- '_impl.efc__id': out[166],
- '_impl.efc__jv': out[167],
- '_impl.efc__lo': out[168],
- '_impl.efc__lo_alpha': out[169],
- '_impl.efc__lo_next': out[170],
- '_impl.efc__lo_next_alpha': out[171],
- '_impl.efc__ls_done': out[172],
- '_impl.efc__margin': out[173],
- '_impl.efc__mid': out[174],
- '_impl.efc__mid_alpha': out[175],
- '_impl.efc__mv': out[176],
- '_impl.efc__p0': out[177],
- '_impl.efc__pos': out[178],
- '_impl.efc__prev_Mgrad': out[179],
- '_impl.efc__prev_cost': out[180],
- '_impl.efc__prev_grad': out[181],
- '_impl.efc__quad': out[182],
- '_impl.efc__quad_gauss': out[183],
- '_impl.efc__search': out[184],
- '_impl.efc__search_dot': out[185],
- '_impl.efc__state': out[186],
- '_impl.efc__type': out[187],
- '_impl.efc__vel': out[188],
+ '_impl.efc__h': out[160],
+ '_impl.efc__id': out[161],
+ '_impl.efc__jv': out[162],
+ '_impl.efc__margin': out[163],
+ '_impl.efc__mv': out[164],
+ '_impl.efc__pos': out[165],
+ '_impl.efc__prev_Mgrad': out[166],
+ '_impl.efc__prev_cost': out[167],
+ '_impl.efc__prev_grad': out[168],
+ '_impl.efc__quad': out[169],
+ '_impl.efc__quad_gauss': out[170],
+ '_impl.efc__search': out[171],
+ '_impl.efc__search_dot': out[172],
+ '_impl.efc__state': out[173],
+ '_impl.efc__type': out[174],
+ '_impl.efc__vel': out[175],
})
return d
diff --git a/mjx/mujoco/mjx/warp/types.py b/mjx/mujoco/mjx/warp/types.py
index 8b556e5f..3cb5f035 100644
--- a/mjx/mujoco/mjx/warp/types.py
+++ b/mjx/mujoco/mjx/warp/types.py
@@ -94,6 +94,7 @@ 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
@@ -256,24 +257,11 @@ class DataWarp(PyTreeNode):
efc__gauss: jax.Array
efc__grad: jax.Array
efc__grad_dot: jax.Array
- efc__gtol: jax.Array
efc__h: jax.Array
- efc__hi: jax.Array
- efc__hi_alpha: jax.Array
- efc__hi_next: jax.Array
- efc__hi_next_alpha: jax.Array
efc__id: jax.Array
efc__jv: jax.Array
- efc__lo: jax.Array
- efc__lo_alpha: jax.Array
- efc__lo_next: jax.Array
- efc__lo_next_alpha: jax.Array
- efc__ls_done: jax.Array
efc__margin: jax.Array
- efc__mid: jax.Array
- efc__mid_alpha: jax.Array
efc__mv: jax.Array
- efc__p0: jax.Array
efc__pos: jax.Array
efc__prev_Mgrad: jax.Array
efc__prev_cost: jax.Array
@@ -488,24 +476,11 @@ _NDIM = {
'efc__gauss': 1,
'efc__grad': 2,
'efc__grad_dot': 1,
- 'efc__gtol': 1,
'efc__h': 3,
- 'efc__hi': 2,
- 'efc__hi_alpha': 1,
- 'efc__hi_next': 2,
- 'efc__hi_next_alpha': 1,
'efc__id': 2,
'efc__jv': 2,
- 'efc__lo': 2,
- 'efc__lo_alpha': 1,
- 'efc__lo_next': 2,
- 'efc__lo_next_alpha': 1,
- 'efc__ls_done': 1,
'efc__margin': 2,
- 'efc__mid': 2,
- 'efc__mid_alpha': 1,
'efc__mv': 2,
- 'efc__p0': 2,
'efc__pos': 2,
'efc__prev_Mgrad': 2,
'efc__prev_cost': 1,
@@ -884,6 +859,7 @@ _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,
@@ -1002,6 +978,7 @@ _NDIM = {
'integrator': 0,
'is_sparse': 0,
'iterations': 0,
+ 'legacy_gjk': 0,
'ls_iterations': 0,
'ls_parallel': 0,
'ls_parallel_min_step': 0,
@@ -1075,24 +1052,11 @@ _BATCH_DIM = {
'efc__gauss': True,
'efc__grad': True,
'efc__grad_dot': True,
- 'efc__gtol': True,
'efc__h': True,
- 'efc__hi': True,
- 'efc__hi_alpha': True,
- 'efc__hi_next': True,
- 'efc__hi_next_alpha': True,
'efc__id': True,
'efc__jv': True,
- 'efc__lo': True,
- 'efc__lo_alpha': True,
- 'efc__lo_next': True,
- 'efc__lo_next_alpha': True,
- 'efc__ls_done': True,
'efc__margin': True,
- 'efc__mid': True,
- 'efc__mid_alpha': True,
'efc__mv': True,
- 'efc__p0': True,
'efc__pos': True,
'efc__prev_Mgrad': True,
'efc__prev_cost': True,
@@ -1471,6 +1435,7 @@ _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,
@@ -1589,6 +1554,7 @@ _BATCH_DIM = {
'integrator': False,
'is_sparse': False,
'iterations': False,
+ 'legacy_gjk': False,
'ls_iterations': False,
'ls_parallel': False,
'ls_parallel_min_step': False,