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,