diff --git a/mjx/mujoco/mjx/_src/collision_convex.py b/mjx/mujoco/mjx/_src/collision_convex.py index 90a709ab..870759ed 100644 --- a/mjx/mujoco/mjx/_src/collision_convex.py +++ b/mjx/mujoco/mjx/_src/collision_convex.py @@ -109,106 +109,6 @@ def _closest_segment_point_plane( return segment_point -def _closest_triangle_point( - p0: jax.Array, p1: jax.Array, p2: jax.Array, pt: jax.Array -) -> jax.Array: - """Gets the closest point between a triangle and a point in space. - - Args: - p0: triangle point - p1: triangle point - p2: triangle point - pt: point to test - - Returns: - closest point on the triangle w.r.t point pt - """ - # Parametrize the triangle s.t. a point inside the triangle is - # Q = p0 + u * e0 + v * e1, when 0 <= u <= 1, 0 <= v <= 1, and - # 0 <= u + v <= 1. Let e0 = (p1 - p0) and e1 = (p2 - p0). - # We analytically minimize the distance between the point pt and Q. - e0 = p1 - p0 - e1 = p2 - p0 - a = e0.dot(e0) - b = e0.dot(e1) - c = e1.dot(e1) - d = pt - p0 - # The determinant is 0 only if the angle between e1 and e0 is 0 - # (i.e. the triangle has overlapping lines). - det = a * c - b * b - u = (c * e0.dot(d) - b * e1.dot(d)) / det - v = (-b * e0.dot(d) + a * e1.dot(d)) / det - inside = (0 <= u) & (u <= 1) & (0 <= v) & (v <= 1) & (u + v <= 1) - closest_p = p0 + u * e0 + v * e1 - d0 = (closest_p - pt).dot(closest_p - pt) - - # If the closest point is outside the triangle, it must be on an edge, so we - # check each triangle edge for a closest point to the point pt. - closest_p1, d1 = math.closest_segment_point_and_dist(p0, p1, pt) - closest_p = jp.where((d0 < d1) & inside, closest_p, closest_p1) - min_d = jp.where((d0 < d1) & inside, d0, d1) - - closest_p2, d2 = math.closest_segment_point_and_dist(p1, p2, pt) - closest_p = jp.where(d2 < min_d, closest_p2, closest_p) - min_d = jp.minimum(min_d, d2) - - closest_p3, d3 = math.closest_segment_point_and_dist(p2, p0, pt) - closest_p = jp.where(d3 < min_d, closest_p3, closest_p) - - return closest_p - - -def _closest_segment_triangle_points( - a: jax.Array, - b: jax.Array, - p0: jax.Array, - p1: jax.Array, - p2: jax.Array, - triangle_normal: jax.Array, -) -> Tuple[jax.Array, jax.Array]: - """Gets the closest points between a line segment and triangle. - - Args: - a: first line segment point - b: second line segment point - p0: triangle point - p1: triangle point - p2: triangle point - triangle_normal: normal of triangle - - Returns: - closest point on the triangle w.r.t the line segment - """ - # The closest triangle point is either on the edge or within the triangle. - # First check triangle edges for the closest point. - # TODO(robotics-simulation): consider vmapping over closest point functions - seg_pt1, tri_pt1 = math.closest_segment_to_segment_points(a, b, p0, p1) - d1 = (seg_pt1 - tri_pt1).dot(seg_pt1 - tri_pt1) - seg_pt2, tri_pt2 = math.closest_segment_to_segment_points(a, b, p1, p2) - d2 = (seg_pt2 - tri_pt2).dot(seg_pt2 - tri_pt2) - seg_pt3, tri_pt3 = math.closest_segment_to_segment_points(a, b, p0, p2) - d3 = (seg_pt3 - tri_pt3).dot(seg_pt3 - tri_pt3) - - # Next, handle the case where the closest triangle point is inside the - # triangle. Either the line segment intersects the triangle or a segment - # endpoint is closest to a point inside the triangle. - seg_pt4 = _closest_segment_point_plane(a, b, p0, triangle_normal) - tri_pt4 = _closest_triangle_point(p0, p1, p2, seg_pt4) - d4 = (seg_pt4 - tri_pt4).dot(seg_pt4 - tri_pt4) - - # Get the point with minimum distance from the line segment point to the - # triangle point. - distance = jp.array([[d1, d2, d3, d4]]) - min_dist = jp.amin(distance) - mask = (distance == min_dist).T - seg_pt = jp.array([seg_pt1, seg_pt2, seg_pt3, seg_pt4]) * mask - tri_pt = jp.array([tri_pt1, tri_pt2, tri_pt3, tri_pt4]) * mask - seg_pt = jp.sum(seg_pt, axis=0) / jp.sum(mask) - tri_pt = jp.sum(tri_pt, axis=0) / jp.sum(mask) - - return seg_pt, tri_pt - - def _manifold_points( poly: jax.Array, poly_mask: jax.Array, poly_norm: jax.Array ) -> jax.Array: