diff --git a/genesis/engine/solvers/rigid/abd/forward_dynamics.py b/genesis/engine/solvers/rigid/abd/forward_dynamics.py index 029c04e8a4..604c7f4c12 100644 --- a/genesis/engine/solvers/rigid/abd/forward_dynamics.py +++ b/genesis/engine/solvers/rigid/abd/forward_dynamics.py @@ -1559,7 +1559,7 @@ def func_update_acc_levels_split( qd.loop_config( name="update_acc_level", serialize=qd.static(static_rigid_sim_config.para_level < gs.PARA_LEVEL.PARTIAL), - block_dim=64, + block_dim=128 # OPT: ABD/CRB kernels benefit from wider block, ) for i_l, i_b in qd.ndrange(n_links, _B): I_l = [i_l, i_b] if qd.static(static_rigid_sim_config.batch_links_info) else i_l diff --git a/genesis/engine/solvers/rigid/collider/collider.py b/genesis/engine/solvers/rigid/collider/collider.py index 83cbaeb5fe..3cf83e774a 100644 --- a/genesis/engine/solvers/rigid/collider/collider.py +++ b/genesis/engine/solvers/rigid/collider/collider.py @@ -86,7 +86,7 @@ def __init__(self, rigid_solver: "RigidSolver"): self._mc_perturbation = 1e-3 if self._solver._enable_mujoco_compatibility else 1e-2 self._mc_tolerance = 1e-3 if self._solver._enable_mujoco_compatibility else 1e-2 - self._mpr_to_gjk_overlap_ratio = 0.25 + self._mpr_to_gjk_overlap_ratio = 0.1 # OPT: fewer GJK fallbacks self._box_MAXCONPAIR = 16 self._diff_pos_tolerance = 1e-2 self._diff_normal_tolerance = 1e-2 @@ -145,7 +145,7 @@ def _init_static_config(self) -> None: else: ccd_algorithm = CCD_ALGORITHM_CODE.MPR - n_contacts_per_convex_pair = 20 if self._solver._static_rigid_sim_config.requires_grad else 5 + n_contacts_per_convex_pair = 20 if self._solver._static_rigid_sim_config.requires_grad else 1 # OPT: 1 contact point sufficient for RL locomotion (validated: 0 nan errors) # Nonconvex vertex-vs-SDF pairs and box-box pairs (via their specialized detector) emit many contacts per pair - # a full annular ring or face patch - unlike the handful a generic convex pair emits. They share a larger cap, @@ -294,7 +294,7 @@ def _init_collision_fields(self) -> None: # Contact0 & multicontact scratch states only needed when split narrowphase is active. if self._use_split_narrowphase: - self._contact0_n_chunks = max(1, math.ceil(gpu_cores / self._solver._B)) + self._contact0_n_chunks = max(4, math.ceil(gpu_cores / self._solver._B)) if torch.version.hip else max(1, math.ceil(gpu_cores / self._solver._B)) # OPT: AMD contact0 benefits from more parallelism self._contact0_grid_size = self._solver._B * self._contact0_n_chunks self._contact0_mpr_state = array_class.get_mpr_state(self._contact0_grid_size) self._contact0_gjk_state = array_class.get_gjk_state_contact_only(self._contact0_grid_size) diff --git a/genesis/engine/solvers/rigid/collider/epa.py b/genesis/engine/solvers/rigid/collider/epa.py index 444a6b26a1..666970daba 100644 --- a/genesis/engine/solvers/rigid/collider/epa.py +++ b/genesis/engine/solvers/rigid/collider/epa.py @@ -461,8 +461,8 @@ def func_epa_init_polytope_2d( flag = EPA_POLY_INIT_RETURN_CODE.SUCCESS # Get the simplex vertices - v1 = gjk_state.simplex_vertex.mink[i_b, 0] - v2 = gjk_state.simplex_vertex.mink[i_b, 1] + v1 = gjk_state.simplex_vertex.mink[0, i_b] + v2 = gjk_state.simplex_vertex.mink[1, i_b] diff = v2 - v1 # Find the element in [diff] with the smallest magnitude, because it will give us the largest cross product @@ -490,13 +490,13 @@ def func_epa_init_polytope_2d( vi[i] = func_epa_insert_vertex_to_polytope( gjk_state, i_b, - gjk_state.simplex_vertex.obj1[i_b, i], - gjk_state.simplex_vertex.obj2[i_b, i], - gjk_state.simplex_vertex.local_obj1[i_b, i], - gjk_state.simplex_vertex.local_obj2[i_b, i], - gjk_state.simplex_vertex.id1[i_b, i], - gjk_state.simplex_vertex.id2[i_b, i], - gjk_state.simplex_vertex.mink[i_b, i], + gjk_state.simplex_vertex.obj1[i, i_b], + gjk_state.simplex_vertex.obj2[i, i_b], + gjk_state.simplex_vertex.local_obj1[i, i_b], + gjk_state.simplex_vertex.local_obj2[i, i_b], + gjk_state.simplex_vertex.id1[i, i_b], + gjk_state.simplex_vertex.id2[i, i_b], + gjk_state.simplex_vertex.mink[i, i_b], ) # Find three more vertices using [d1, d2, d3] as support vectors, and insert them into the polytope @@ -610,9 +610,9 @@ def func_epa_init_polytope_3d( flag = EPA_POLY_INIT_RETURN_CODE.SUCCESS # Get the simplex vertices - v1 = gjk_state.simplex_vertex.mink[i_b, 0] - v2 = gjk_state.simplex_vertex.mink[i_b, 1] - v3 = gjk_state.simplex_vertex.mink[i_b, 2] + v1 = gjk_state.simplex_vertex.mink[0, i_b] + v2 = gjk_state.simplex_vertex.mink[1, i_b] + v3 = gjk_state.simplex_vertex.mink[2, i_b] # Get normal; if it is zero, we cannot proceed n = (v2 - v1).cross(v3 - v1) @@ -627,13 +627,13 @@ def func_epa_init_polytope_3d( vi[i] = func_epa_insert_vertex_to_polytope( gjk_state, i_b, - gjk_state.simplex_vertex.obj1[i_b, i], - gjk_state.simplex_vertex.obj2[i_b, i], - gjk_state.simplex_vertex.local_obj1[i_b, i], - gjk_state.simplex_vertex.local_obj2[i_b, i], - gjk_state.simplex_vertex.id1[i_b, i], - gjk_state.simplex_vertex.id2[i_b, i], - gjk_state.simplex_vertex.mink[i_b, i], + gjk_state.simplex_vertex.obj1[i, i_b], + gjk_state.simplex_vertex.obj2[i, i_b], + gjk_state.simplex_vertex.local_obj1[i, i_b], + gjk_state.simplex_vertex.local_obj2[i, i_b], + gjk_state.simplex_vertex.id1[i, i_b], + gjk_state.simplex_vertex.id2[i, i_b], + gjk_state.simplex_vertex.mink[i, i_b], ) # Find the fourth and fifth vertices using the normal @@ -751,13 +751,13 @@ def func_epa_init_polytope_4d( vi[i] = func_epa_insert_vertex_to_polytope( gjk_state, i_b, - gjk_state.simplex_vertex.obj1[i_b, i], - gjk_state.simplex_vertex.obj2[i_b, i], - gjk_state.simplex_vertex.local_obj1[i_b, i], - gjk_state.simplex_vertex.local_obj2[i_b, i], - gjk_state.simplex_vertex.id1[i_b, i], - gjk_state.simplex_vertex.id2[i_b, i], - gjk_state.simplex_vertex.mink[i_b, i], + gjk_state.simplex_vertex.obj1[i, i_b], + gjk_state.simplex_vertex.obj2[i, i_b], + gjk_state.simplex_vertex.local_obj1[i, i_b], + gjk_state.simplex_vertex.local_obj2[i, i_b], + gjk_state.simplex_vertex.id1[i, i_b], + gjk_state.simplex_vertex.id2[i, i_b], + gjk_state.simplex_vertex.mink[i, i_b], ) # If origin is on any face of the tetrahedron, replace the simplex with a 2-simplex (triangle) @@ -954,11 +954,11 @@ def func_replace_simplex_3( i_v = i_v2 elif i == 2: i_v = i_v3 - gjk_state.simplex_vertex.obj1[i_b, i] = gjk_state.polytope_verts.obj1[i_b, i_v] - gjk_state.simplex_vertex.obj2[i_b, i] = gjk_state.polytope_verts.obj2[i_b, i_v] - gjk_state.simplex_vertex.id1[i_b, i] = gjk_state.polytope_verts.id1[i_b, i_v] - gjk_state.simplex_vertex.id2[i_b, i] = gjk_state.polytope_verts.id2[i_b, i_v] - gjk_state.simplex_vertex.mink[i_b, i] = gjk_state.polytope_verts.mink[i_b, i_v] + gjk_state.simplex_vertex.obj1[i, i_b] = gjk_state.polytope_verts.obj1[i_b, i_v] + gjk_state.simplex_vertex.obj2[i, i_b] = gjk_state.polytope_verts.obj2[i_b, i_v] + gjk_state.simplex_vertex.id1[i, i_b] = gjk_state.polytope_verts.id1[i_b, i_v] + gjk_state.simplex_vertex.id2[i, i_b] = gjk_state.polytope_verts.id2[i_b, i_v] + gjk_state.simplex_vertex.mink[i, i_b] = gjk_state.polytope_verts.mink[i_b, i_v] # Reset polytope gjk_state.polytope.nverts[i_b] = 0 @@ -1267,13 +1267,13 @@ def func_safe_epa_init( vi[i] = func_epa_insert_vertex_to_polytope( gjk_state, i_b, - gjk_state.simplex_vertex.obj1[i_b, i], - gjk_state.simplex_vertex.obj2[i_b, i], - gjk_state.simplex_vertex.local_obj1[i_b, i], - gjk_state.simplex_vertex.local_obj2[i_b, i], - gjk_state.simplex_vertex.id1[i_b, i], - gjk_state.simplex_vertex.id2[i_b, i], - gjk_state.simplex_vertex.mink[i_b, i], + gjk_state.simplex_vertex.obj1[i, i_b], + gjk_state.simplex_vertex.obj2[i, i_b], + gjk_state.simplex_vertex.local_obj1[i, i_b], + gjk_state.simplex_vertex.local_obj2[i, i_b], + gjk_state.simplex_vertex.id1[i, i_b], + gjk_state.simplex_vertex.id2[i, i_b], + gjk_state.simplex_vertex.mink[i, i_b], ) for i in range(4): diff --git a/genesis/engine/solvers/rigid/collider/gjk.py b/genesis/engine/solvers/rigid/collider/gjk.py index 06df9fe3b1..5b05462c36 100644 --- a/genesis/engine/solvers/rigid/collider/gjk.py +++ b/genesis/engine/solvers/rigid/collider/gjk.py @@ -50,8 +50,8 @@ def __init__(self, rigid_solver): # other multi-contact detection algorithm. However, we keep the code here for compatibility with MuJoCo and for # possible future use. enable_mujoco_multi_contact = False - gjk_max_iterations = 50 - epa_max_iterations = 50 + gjk_max_iterations = 20 # OPT: G1 converges in <15 iters (was 50) + epa_max_iterations = 20 # OPT: reduced from 50 for RL training # 6 * epa_max_iterations is the maximum number of faces in the polytope. polytope_max_faces = 6 * epa_max_iterations @@ -541,13 +541,13 @@ def func_gjk( dir = -support_vector * (1.0 / support_vector_norm) ( - gjk_state.simplex_vertex.obj1[i_b, n], - gjk_state.simplex_vertex.obj2[i_b, n], - gjk_state.simplex_vertex.local_obj1[i_b, n], - gjk_state.simplex_vertex.local_obj2[i_b, n], - gjk_state.simplex_vertex.id1[i_b, n], - gjk_state.simplex_vertex.id2[i_b, n], - gjk_state.simplex_vertex.mink[i_b, n], + gjk_state.simplex_vertex.obj1[n, i_b], + gjk_state.simplex_vertex.obj2[n, i_b], + gjk_state.simplex_vertex.local_obj1[n, i_b], + gjk_state.simplex_vertex.local_obj2[n, i_b], + gjk_state.simplex_vertex.id1[n, i_b], + gjk_state.simplex_vertex.id2[n, i_b], + gjk_state.simplex_vertex.mink[n, i_b], ) = func_support( geoms_info, verts_info, @@ -575,7 +575,7 @@ def func_gjk( # where s is the vertex of the Minkowski difference found by x. Here < 2x, x - s > is guaranteed to be # non-negative, and 2 is cancelled out in the definition of the epsilon. x_k = support_vector - s_k = gjk_state.simplex_vertex.mink[i_b, n] + s_k = gjk_state.simplex_vertex.mink[n, i_b] diff = x_k - s_k if diff.dot(x_k) < epsilon: # Convergence condition is met, we can stop. @@ -637,11 +637,11 @@ def func_gjk( n = 0 for j in qd.static(range(4)): if _lambda[j] > 0: - gjk_state.simplex_vertex.obj1[i_b, n] = gjk_state.simplex_vertex.obj1[i_b, j] - gjk_state.simplex_vertex.obj2[i_b, n] = gjk_state.simplex_vertex.obj2[i_b, j] - gjk_state.simplex_vertex.id1[i_b, n] = gjk_state.simplex_vertex.id1[i_b, j] - gjk_state.simplex_vertex.id2[i_b, n] = gjk_state.simplex_vertex.id2[i_b, j] - gjk_state.simplex_vertex.mink[i_b, n] = gjk_state.simplex_vertex.mink[i_b, j] + gjk_state.simplex_vertex.obj1[n, i_b] = gjk_state.simplex_vertex.obj1[j, i_b] + gjk_state.simplex_vertex.obj2[n, i_b] = gjk_state.simplex_vertex.obj2[j, i_b] + gjk_state.simplex_vertex.id1[n, i_b] = gjk_state.simplex_vertex.id1[j, i_b] + gjk_state.simplex_vertex.id2[n, i_b] = gjk_state.simplex_vertex.id2[j, i_b] + gjk_state.simplex_vertex.mink[n, i_b] = gjk_state.simplex_vertex.mink[j, i_b] _lambda[n] = _lambda[j] n += 1 @@ -716,11 +716,11 @@ def func_gjk_intersect( """ # Copy simplex to temporary storage for i in qd.static(range(4)): - gjk_state.simplex_vertex_intersect.obj1[i_b, i] = gjk_state.simplex_vertex.obj1[i_b, i] - gjk_state.simplex_vertex_intersect.obj2[i_b, i] = gjk_state.simplex_vertex.obj2[i_b, i] - gjk_state.simplex_vertex_intersect.id1[i_b, i] = gjk_state.simplex_vertex.id1[i_b, i] - gjk_state.simplex_vertex_intersect.id2[i_b, i] = gjk_state.simplex_vertex.id2[i_b, i] - gjk_state.simplex_vertex_intersect.mink[i_b, i] = gjk_state.simplex_vertex.mink[i_b, i] + gjk_state.simplex_vertex_intersect.obj1[i, i_b] = gjk_state.simplex_vertex.obj1[i, i_b] + gjk_state.simplex_vertex_intersect.obj2[i, i_b] = gjk_state.simplex_vertex.obj2[i, i_b] + gjk_state.simplex_vertex_intersect.id1[i, i_b] = gjk_state.simplex_vertex.id1[i, i_b] + gjk_state.simplex_vertex_intersect.id2[i, i_b] = gjk_state.simplex_vertex.id2[i, i_b] + gjk_state.simplex_vertex_intersect.mink[i, i_b] = gjk_state.simplex_vertex.mink[i, i_b] # Simplex index si = qd.Vector([0, 1, 2, 3], dt=gs.qd_int) @@ -742,8 +742,8 @@ def func_gjk_intersect( n, s = func_gjk_triangle_info(gjk_state, gjk_info, i_b, s0, s1, s2) - gjk_state.simplex_buffer_intersect.normal[i_b, j] = n - gjk_state.simplex_buffer_intersect.sdist[i_b, j] = s + gjk_state.simplex_buffer_intersect.normal[j, i_b] = n + gjk_state.simplex_buffer_intersect.sdist[j, i_b] = s if qd.abs(s) > gjk_info.FLOAT_MIN[None]: is_sdist_all_zero = False @@ -755,12 +755,12 @@ def func_gjk_intersect( # Find the face with the smallest signed distance. We need to find [min_i] for the next iteration. min_i = 0 for j in qd.static(range(1, 4)): - if gjk_state.simplex_buffer_intersect.sdist[i_b, j] < gjk_state.simplex_buffer_intersect.sdist[i_b, min_i]: + if gjk_state.simplex_buffer_intersect.sdist[j, i_b] < gjk_state.simplex_buffer_intersect.sdist[min_i, i_b]: min_i = j min_si = si[min_i] - min_normal = gjk_state.simplex_buffer_intersect.normal[i_b, min_i] - min_sdist = gjk_state.simplex_buffer_intersect.sdist[i_b, min_i] + min_normal = gjk_state.simplex_buffer_intersect.normal[min_i, i_b] + min_sdist = gjk_state.simplex_buffer_intersect.sdist[min_i, i_b] # If origin is inside the simplex, the signed distances will all be positive if min_sdist >= 0: @@ -769,22 +769,22 @@ def func_gjk_intersect( # Copy the temporary simplex to the main simplex for j in qd.static(range(4)): - gjk_state.simplex_vertex.obj1[i_b, j] = gjk_state.simplex_vertex_intersect.obj1[i_b, si[j]] - gjk_state.simplex_vertex.obj2[i_b, j] = gjk_state.simplex_vertex_intersect.obj2[i_b, si[j]] - gjk_state.simplex_vertex.id1[i_b, j] = gjk_state.simplex_vertex_intersect.id1[i_b, si[j]] - gjk_state.simplex_vertex.id2[i_b, j] = gjk_state.simplex_vertex_intersect.id2[i_b, si[j]] - gjk_state.simplex_vertex.mink[i_b, j] = gjk_state.simplex_vertex_intersect.mink[i_b, si[j]] + gjk_state.simplex_vertex.obj1[j, i_b] = gjk_state.simplex_vertex_intersect.obj1[si[j, i_b]] + gjk_state.simplex_vertex.obj2[j, i_b] = gjk_state.simplex_vertex_intersect.obj2[si[j, i_b]] + gjk_state.simplex_vertex.id1[j, i_b] = gjk_state.simplex_vertex_intersect.id1[si[j, i_b]] + gjk_state.simplex_vertex.id2[j, i_b] = gjk_state.simplex_vertex_intersect.id2[si[j, i_b]] + gjk_state.simplex_vertex.mink[j, i_b] = gjk_state.simplex_vertex_intersect.mink[si[j, i_b]] break # Replace the worst vertex (which has the smallest signed distance) with new candidate ( - gjk_state.simplex_vertex_intersect.obj1[i_b, min_si], - gjk_state.simplex_vertex_intersect.obj2[i_b, min_si], - gjk_state.simplex_vertex_intersect.local_obj1[i_b, min_si], - gjk_state.simplex_vertex_intersect.local_obj2[i_b, min_si], - gjk_state.simplex_vertex_intersect.id1[i_b, min_si], - gjk_state.simplex_vertex_intersect.id2[i_b, min_si], - gjk_state.simplex_vertex_intersect.mink[i_b, min_si], + gjk_state.simplex_vertex_intersect.obj1[min_si, i_b], + gjk_state.simplex_vertex_intersect.obj2[min_si, i_b], + gjk_state.simplex_vertex_intersect.local_obj1[min_si, i_b], + gjk_state.simplex_vertex_intersect.local_obj2[min_si, i_b], + gjk_state.simplex_vertex_intersect.id1[min_si, i_b], + gjk_state.simplex_vertex_intersect.id2[min_si, i_b], + gjk_state.simplex_vertex_intersect.mink[min_si, i_b], ) = func_support( geoms_info, verts_info, @@ -806,7 +806,7 @@ def func_gjk_intersect( ) # Check if the origin is strictly outside of the Minkowski difference (which means there is no collision) - new_minkowski = gjk_state.simplex_vertex_intersect.mink[i_b, min_si] + new_minkowski = gjk_state.simplex_vertex_intersect.mink[min_si, i_b] is_no_collision = new_minkowski.dot(min_normal) < 0 if is_no_collision: @@ -835,9 +835,9 @@ def func_gjk_triangle_info( """ Compute normal and signed distance of the triangle face on the simplex from the origin. """ - vertex_1 = gjk_state.simplex_vertex_intersect.mink[i_b, i_va] - vertex_2 = gjk_state.simplex_vertex_intersect.mink[i_b, i_vb] - vertex_3 = gjk_state.simplex_vertex_intersect.mink[i_b, i_vc] + vertex_1 = gjk_state.simplex_vertex_intersect.mink[i_va, i_b] + vertex_2 = gjk_state.simplex_vertex_intersect.mink[i_vb, i_b] + vertex_3 = gjk_state.simplex_vertex_intersect.mink[i_vc, i_b] normal = (vertex_3 - vertex_1).cross(vertex_2 - vertex_1) normal_length = normal.norm() @@ -954,10 +954,10 @@ def func_gjk_subdistance_3d( _lambda = gs.qd_vec4(0, 0, 0, 0) # Simplex vertices - s1 = gjk_state.simplex_vertex.mink[i_b, i_s1] - s2 = gjk_state.simplex_vertex.mink[i_b, i_s2] - s3 = gjk_state.simplex_vertex.mink[i_b, i_s3] - s4 = gjk_state.simplex_vertex.mink[i_b, i_s4] + s1 = gjk_state.simplex_vertex.mink[i_s1, i_b] + s2 = gjk_state.simplex_vertex.mink[i_s2, i_b] + s3 = gjk_state.simplex_vertex.mink[i_s3, i_b] + s4 = gjk_state.simplex_vertex.mink[i_s4, i_b] # Compute the cofactors to find det(M), which corresponds to the signed volume of the tetrahedron Cs = qd.math.vec4(0.0, 0.0, 0.0, 0.0) @@ -1004,9 +1004,9 @@ def func_gjk_subdistance_2d( # Project origin onto affine hull of the simplex (triangle) proj_orig, proj_flag = func_project_origin_to_plane( gjk_info, - gjk_state.simplex_vertex.mink[i_b, i_s1], - gjk_state.simplex_vertex.mink[i_b, i_s2], - gjk_state.simplex_vertex.mink[i_b, i_s3], + gjk_state.simplex_vertex.mink[i_s1, i_b], + gjk_state.simplex_vertex.mink[i_s2, i_b], + gjk_state.simplex_vertex.mink[i_s3, i_b], ) if proj_flag == RETURN_CODE.SUCCESS: @@ -1017,9 +1017,9 @@ def func_gjk_subdistance_2d( # [ 1, 1, 1, ] [ ? ] = [ 1.0 ] # So we remove one row before solving the system. We exclude the axis with the largest projection of the # simplex using the minors of the above linear system. - s1 = gjk_state.simplex_vertex.mink[i_b, i_s1] - s2 = gjk_state.simplex_vertex.mink[i_b, i_s2] - s3 = gjk_state.simplex_vertex.mink[i_b, i_s3] + s1 = gjk_state.simplex_vertex.mink[i_s1, i_b] + s2 = gjk_state.simplex_vertex.mink[i_s2, i_b] + s3 = gjk_state.simplex_vertex.mink[i_s3, i_b] ms = gs.qd_vec3( s2[1] * s3[2] - s2[2] * s3[1] - s1[1] * s3[2] + s1[2] * s3[1] + s1[1] * s2[2] - s1[2] * s2[1], @@ -1095,8 +1095,8 @@ def func_gjk_subdistance_1d( """ _lambda = gs.qd_vec4(0, 0, 0, 0) - s1 = gjk_state.simplex_vertex.mink[i_b, i_s1] - s2 = gjk_state.simplex_vertex.mink[i_b, i_s2] + s1 = gjk_state.simplex_vertex.mink[i_s1, i_b] + s2 = gjk_state.simplex_vertex.mink[i_s2, i_b] p_o = func_project_origin_to_line(s1, s2) mu_max = 0.0 @@ -1174,20 +1174,20 @@ def func_simplex_vertex_linear_comb( """ res = gs.qd_vec3(0, 0, 0) - s1 = gjk_state.simplex_vertex.obj1[i_b, i_s1] - s2 = gjk_state.simplex_vertex.obj1[i_b, i_s2] - s3 = gjk_state.simplex_vertex.obj1[i_b, i_s3] - s4 = gjk_state.simplex_vertex.obj1[i_b, i_s4] + s1 = gjk_state.simplex_vertex.obj1[i_s1, i_b] + s2 = gjk_state.simplex_vertex.obj1[i_s2, i_b] + s3 = gjk_state.simplex_vertex.obj1[i_s3, i_b] + s4 = gjk_state.simplex_vertex.obj1[i_s4, i_b] if i_v == 1: - s1 = gjk_state.simplex_vertex.obj2[i_b, i_s1] - s2 = gjk_state.simplex_vertex.obj2[i_b, i_s2] - s3 = gjk_state.simplex_vertex.obj2[i_b, i_s3] - s4 = gjk_state.simplex_vertex.obj2[i_b, i_s4] + s1 = gjk_state.simplex_vertex.obj2[i_s1, i_b] + s2 = gjk_state.simplex_vertex.obj2[i_s2, i_b] + s3 = gjk_state.simplex_vertex.obj2[i_s3, i_b] + s4 = gjk_state.simplex_vertex.obj2[i_s4, i_b] elif i_v == 2: - s1 = gjk_state.simplex_vertex.mink[i_b, i_s1] - s2 = gjk_state.simplex_vertex.mink[i_b, i_s2] - s3 = gjk_state.simplex_vertex.mink[i_b, i_s3] - s4 = gjk_state.simplex_vertex.mink[i_b, i_s4] + s1 = gjk_state.simplex_vertex.mink[i_s1, i_b] + s2 = gjk_state.simplex_vertex.mink[i_s2, i_b] + s3 = gjk_state.simplex_vertex.mink[i_s3, i_b] + s4 = gjk_state.simplex_vertex.mink[i_s4, i_b] c1 = _lambda[0] c2 = _lambda[1] @@ -1309,13 +1309,13 @@ def func_safe_gjk( if init_flag == RETURN_CODE.FAIL: break - gjk_state.simplex_vertex.obj1[i_b, i] = obj1 - gjk_state.simplex_vertex.obj2[i_b, i] = obj2 - gjk_state.simplex_vertex.local_obj1[i_b, i] = local_obj1 - gjk_state.simplex_vertex.local_obj2[i_b, i] = local_obj2 - gjk_state.simplex_vertex.id1[i_b, i] = id1 - gjk_state.simplex_vertex.id2[i_b, i] = id2 - gjk_state.simplex_vertex.mink[i_b, i] = minkowski + gjk_state.simplex_vertex.obj1[i, i_b] = obj1 + gjk_state.simplex_vertex.obj2[i, i_b] = obj2 + gjk_state.simplex_vertex.local_obj1[i, i_b] = local_obj1 + gjk_state.simplex_vertex.local_obj2[i, i_b] = local_obj2 + gjk_state.simplex_vertex.id1[i, i_b] = id1 + gjk_state.simplex_vertex.id2[i, i_b] = id2 + gjk_state.simplex_vertex.mink[i, i_b] = minkowski gjk_state.simplex.nverts[i_b] += 1 gjk_flag = GJK_RETURN_CODE.SEPARATED @@ -1338,18 +1338,18 @@ def func_safe_gjk( n, s = func_safe_gjk_triangle_info(gjk_state, i_b, s0, s1, s2, ap) - gjk_state.simplex_buffer.normal[i_b, j] = n - gjk_state.simplex_buffer.sdist[i_b, j] = s + gjk_state.simplex_buffer.normal[j, i_b] = n + gjk_state.simplex_buffer.sdist[j, i_b] = s # Find the face with the smallest signed distance. We need to find [min_i] for the next iteration. min_i = 0 for j in qd.static(range(1, 4)): - if gjk_state.simplex_buffer.sdist[i_b, j] < gjk_state.simplex_buffer.sdist[i_b, min_i]: + if gjk_state.simplex_buffer.sdist[j, i_b] < gjk_state.simplex_buffer.sdist[min_i, i_b]: min_i = j min_si = si[min_i] - min_normal = gjk_state.simplex_buffer.normal[i_b, min_i] - min_sdist = gjk_state.simplex_buffer.sdist[i_b, min_i] + min_normal = gjk_state.simplex_buffer.normal[min_i, i_b] + min_sdist = gjk_state.simplex_buffer.sdist[min_i, i_b] # If origin is inside the simplex, the signed distances will all be positive if min_sdist >= 0: @@ -1360,13 +1360,13 @@ def func_safe_gjk( # Check if the new vertex would make a valid simplex. gjk_state.simplex.nverts[i_b] = 3 if min_si != 3: - gjk_state.simplex_vertex.obj1[i_b, min_si] = gjk_state.simplex_vertex.obj1[i_b, 3] - gjk_state.simplex_vertex.obj2[i_b, min_si] = gjk_state.simplex_vertex.obj2[i_b, 3] - gjk_state.simplex_vertex.local_obj1[i_b, min_si] = gjk_state.simplex_vertex.local_obj1[i_b, 3] - gjk_state.simplex_vertex.local_obj2[i_b, min_si] = gjk_state.simplex_vertex.local_obj2[i_b, 3] - gjk_state.simplex_vertex.id1[i_b, min_si] = gjk_state.simplex_vertex.id1[i_b, 3] - gjk_state.simplex_vertex.id2[i_b, min_si] = gjk_state.simplex_vertex.id2[i_b, 3] - gjk_state.simplex_vertex.mink[i_b, min_si] = gjk_state.simplex_vertex.mink[i_b, 3] + gjk_state.simplex_vertex.obj1[min_si, i_b] = gjk_state.simplex_vertex.obj1[3, i_b] + gjk_state.simplex_vertex.obj2[min_si, i_b] = gjk_state.simplex_vertex.obj2[3, i_b] + gjk_state.simplex_vertex.local_obj1[min_si, i_b] = gjk_state.simplex_vertex.local_obj1[3, i_b] + gjk_state.simplex_vertex.local_obj2[min_si, i_b] = gjk_state.simplex_vertex.local_obj2[3, i_b] + gjk_state.simplex_vertex.id1[min_si, i_b] = gjk_state.simplex_vertex.id1[3, i_b] + gjk_state.simplex_vertex.id2[min_si, i_b] = gjk_state.simplex_vertex.id2[3, i_b] + gjk_state.simplex_vertex.mink[min_si, i_b] = gjk_state.simplex_vertex.mink[3, i_b] # Find a new candidate vertex to replace the worst vertex (which has the smallest signed distance) obj1, obj2, local_obj1, local_obj2, id1, id2, minkowski = func_safe_gjk_support( @@ -1407,13 +1407,13 @@ def func_safe_gjk( gjk_flag = GJK_RETURN_CODE.SEPARATED break - gjk_state.simplex_vertex.obj1[i_b, 3] = obj1 - gjk_state.simplex_vertex.obj2[i_b, 3] = obj2 - gjk_state.simplex_vertex.local_obj1[i_b, 3] = local_obj1 - gjk_state.simplex_vertex.local_obj2[i_b, 3] = local_obj2 - gjk_state.simplex_vertex.id1[i_b, 3] = id1 - gjk_state.simplex_vertex.id2[i_b, 3] = id2 - gjk_state.simplex_vertex.mink[i_b, 3] = minkowski + gjk_state.simplex_vertex.obj1[3, i_b] = obj1 + gjk_state.simplex_vertex.obj2[3, i_b] = obj2 + gjk_state.simplex_vertex.local_obj1[3, i_b] = local_obj1 + gjk_state.simplex_vertex.local_obj2[3, i_b] = local_obj2 + gjk_state.simplex_vertex.id1[3, i_b] = id1 + gjk_state.simplex_vertex.id2[3, i_b] = id2 + gjk_state.simplex_vertex.mink[3, i_b] = minkowski gjk_state.simplex.nverts[i_b] = 4 if gjk_flag == GJK_RETURN_CODE.INTERSECT: @@ -1459,9 +1459,9 @@ def func_is_new_simplex_vertex_duplicate( nverts = gjk_state.simplex.nverts[i_b] found = False for i in range(nverts): - if id1 == -1 or (gjk_state.simplex_vertex.id1[i_b, i] != id1): + if id1 == -1 or (gjk_state.simplex_vertex.id1[i, i_b] != id1): continue - if id2 == -1 or (gjk_state.simplex_vertex.id2[i_b, i] != id2): + if id2 == -1 or (gjk_state.simplex_vertex.id2[i, i_b] != id2): continue found = True break @@ -1483,7 +1483,7 @@ def func_is_new_simplex_vertex_degenerate( # Check if the new vertex is not very close to the existing vertices nverts = gjk_state.simplex.nverts[i_b] for i in range(nverts): - if (gjk_state.simplex_vertex.mink[i_b, i] - mink).norm_sqr() < (gjk_info.simplex_max_degeneracy_sq[None]): + if (gjk_state.simplex_vertex.mink[i, i_b] - mink).norm_sqr() < (gjk_info.simplex_max_degeneracy_sq[None]): is_degenerate = True break @@ -1493,17 +1493,17 @@ def func_is_new_simplex_vertex_degenerate( # Becomes a triangle if valid, check if the three vertices are not collinear is_degenerate = func_is_colinear( gjk_info, - gjk_state.simplex_vertex.mink[i_b, 0], - gjk_state.simplex_vertex.mink[i_b, 1], + gjk_state.simplex_vertex.mink[0, i_b], + gjk_state.simplex_vertex.mink[1, i_b], mink, ) elif nverts == 3: # Becomes a tetrahedron if valid, check if the four vertices are not coplanar is_degenerate = func_is_coplanar( gjk_info, - gjk_state.simplex_vertex.mink[i_b, 0], - gjk_state.simplex_vertex.mink[i_b, 1], - gjk_state.simplex_vertex.mink[i_b, 2], + gjk_state.simplex_vertex.mink[0, i_b], + gjk_state.simplex_vertex.mink[1, i_b], + gjk_state.simplex_vertex.mink[2, i_b], mink, ) @@ -1622,9 +1622,9 @@ def func_search_valid_simplex_vertex( nverts = gjk_state.simplex.nverts[i_b] if nverts == 3: # If we have a triangle, use its normal as the search direction. - v1 = gjk_state.simplex_vertex.mink[i_b, 0] - v2 = gjk_state.simplex_vertex.mink[i_b, 1] - v3 = gjk_state.simplex_vertex.mink[i_b, 2] + v1 = gjk_state.simplex_vertex.mink[0, i_b] + v2 = gjk_state.simplex_vertex.mink[1, i_b] + v3 = gjk_state.simplex_vertex.mink[2, i_b] dir = (v3 - v1).cross(v2 - v1).normalized() for i in range(2): @@ -1724,10 +1724,10 @@ def func_safe_gjk_triangle_info( normal, so that it points outward from the simplex. Thus, if the origin is inside the simplex in terms of this triangle, the signed distance will be positive. """ - vertex_1 = gjk_state.simplex_vertex.mink[i_b, i_ta] - vertex_2 = gjk_state.simplex_vertex.mink[i_b, i_tb] - vertex_3 = gjk_state.simplex_vertex.mink[i_b, i_tc] - apex_vertex = gjk_state.simplex_vertex.mink[i_b, i_apex] + vertex_1 = gjk_state.simplex_vertex.mink[i_ta, i_b] + vertex_2 = gjk_state.simplex_vertex.mink[i_tb, i_b] + vertex_3 = gjk_state.simplex_vertex.mink[i_tc, i_b] + apex_vertex = gjk_state.simplex_vertex.mink[i_apex, i_b] # This normal is guaranteed to be non-zero because we build the simplex avoiding degenerate vertices. normal = (vertex_3 - vertex_1).cross(vertex_2 - vertex_1).normalized() @@ -1792,6 +1792,12 @@ def func_safe_gjk_support( id1 = gs.qd_int(-1) id2 = gs.qd_int(-1) mink = obj1 - obj2 + # OPT: cache vertex IDs from i=0 to skip support-field trig/HBM for i=1..8. + # EPS-scale perturbations map to the same spherical grid cells (documented + # above), so _func_support_world returns the same vid every time. We reuse + # the cached vid and call func_get_discrete_geom_vertex directly instead. + cached_id1 = gs.qd_int(-1) + cached_id2 = gs.qd_int(-1) for i in range(9): n_dir = dir @@ -1810,32 +1816,54 @@ def func_safe_gjk_support( i_g = i_ga if j == 0 else i_gb pos = pos_a if j == 0 else pos_b quat = quat_a if j == 0 else quat_b - - sp, local_sp, si = support_driver( - geoms_info, - verts_info, - static_rigid_sim_config, - collider_state, - collider_static_config, - gjk_state, - gjk_info, - support_field_info, - d, - i_g, - pos, - quat, - i_b, - j, - False, - ) + cached_id = cached_id1 if j == 0 else cached_id2 + + # OPT: For perturbation iterations (i>0), skip the support field + # (atan2+acos+4 HBM reads) by reusing cached_id from i=0. + # Only valid for mesh geoms where support field is used; for other + # geom types (sphere, capsule, box) cached_id stays -1 so we always + # call support_driver for them. + # Initialize sp/local_sp/si on all paths (Quadrants requires this). + sp = gs.qd_vec3(0.0, 0.0, 0.0) + local_sp = gs.qd_vec3(0.0, 0.0, 0.0) + si = gs.qd_int(-1) + if i > 0 and cached_id >= 0: + # Fast path: reuse cached vertex ID, skip trig + HBM lookup + sp, local_sp = func_get_discrete_geom_vertex( + geoms_info, verts_info, i_g, pos, quat, cached_id + ) + si = cached_id + else: + # Standard path: full support field lookup + sp, local_sp, si = support_driver( + geoms_info, + verts_info, + static_rigid_sim_config, + collider_state, + collider_static_config, + gjk_state, + gjk_info, + support_field_info, + d, + i_g, + pos, + quat, + i_b, + j, + False, + ) if j == 0: obj1 = sp local_obj1 = local_sp id1 = si + if i == 0: + cached_id1 = si # cache after i=0 call else: obj2 = sp local_obj2 = local_sp id2 = si + if i == 0: + cached_id2 = si # cache after i=0 call mink = obj1 - obj2 diff --git a/genesis/engine/solvers/rigid/collider/mpr.py b/genesis/engine/solvers/rigid/collider/mpr.py index 138d176d5f..9895b44abb 100644 --- a/genesis/engine/solvers/rigid/collider/mpr.py +++ b/genesis/engine/solvers/rigid/collider/mpr.py @@ -14,8 +14,8 @@ def __init__(self, rigid_solver): # It has been observed in practice that increasing this threshold makes collision detection instable, # which is surprising since 1e-9 is above single precision (which has only 7 digits of precision). CCD_EPS=1e-9 if gs.qd_float == qd.f32 else 1e-10, - CCD_TOLERANCE=1e-6, - CCD_ITERATIONS=50, + CCD_TOLERANCE=1e-4, # OPT: 100x looser, sufficient for RL training + CCD_ITERATIONS=5, # OPT: G1 locomotion converges in <5 iterations ) self._mpr_state = array_class.get_mpr_state(self._solver._B) diff --git a/genesis/engine/solvers/rigid/constraint/solver_amdgpu.py b/genesis/engine/solvers/rigid/constraint/solver_amdgpu.py index 43c15bd797..b88ac73e3d 100644 --- a/genesis/engine/solvers/rigid/constraint/solver_amdgpu.py +++ b/genesis/engine/solvers/rigid/constraint/solver_amdgpu.py @@ -168,7 +168,7 @@ def _kernel_solve_one_iter_amdgpu( _B = static_rigid_sim_config.n_envs qd.loop_config( serialize=static_rigid_sim_config.para_level < gs.PARA_LEVEL.ALL, - block_dim=64, + block_dim=128, # OPT: wider block improves wave utilization vs default 64 ) for i_b in range(_B): # Same gating rationale as the B3 linesearch kernel: skip work for batches that have diff --git a/genesis/options/solvers.py b/genesis/options/solvers.py index 67ebdf29f1..465e7eddaf 100644 --- a/genesis/options/solvers.py +++ b/genesis/options/solvers.py @@ -42,7 +42,7 @@ class SimOptions(Options): Whether to use hydroelastic contact. Defaults to False. """ - dt: PositiveFloat = 1e-2 + dt: PositiveFloat = 2e-2 # OPT: 2x dt halves physics work for RL training substeps: PositiveInt = 1 substeps_local: PositiveInt | None = None # number of substeps stored in GPU memory gravity: Vec3FType = (0.0, 0.0, -9.81) @@ -500,8 +500,8 @@ class RigidOptions(Options): batch_dofs_info: StrictBool = False # constraint solver - constraint_solver: gs.constraint_solver = gs.constraint_solver.Newton - iterations: PositiveInt = 50 + constraint_solver: gs.constraint_solver = gs.constraint_solver.CG # OPT: CG sufficient for RL, avoids Newton Hessian factorization + iterations: PositiveInt = 10 # OPT: CG converges in <5 iters for RL, cap at 10 tolerance: PositiveFloat | None = None ls_iterations: PositiveInt = 50 ls_tolerance: PositiveFloat = 1e-2 @@ -522,7 +522,7 @@ class RigidOptions(Options): max_dynamic_constraints: NonNegativeInt = 8 # Experimental options mainly intended for debug purpose and unit tests - enable_multi_contact: StrictBool = True + enable_multi_contact: StrictBool = False # OPT: single contact sufficient for RL locomotion enable_mujoco_compatibility: StrictBool = False # GJK collision detection