Fix some more bugs in sphere/capsule-convex collisions.
PiperOrigin-RevId: 617646119 Change-Id: Ie0c99ccbe1dd2a18663d45a111e77fed4112e638
This commit is contained in:
committed by
Copybara-Service
parent
07da7336e8
commit
624becca52
@@ -215,46 +215,47 @@ def sphere_convex(sphere: GeomInfo, convex: GeomInfo) -> Contact:
|
||||
return jp.dot(pos - faces[0], normal)
|
||||
|
||||
support = get_support(faces, normals)
|
||||
has_separating_axis = jp.any(support >= 0)
|
||||
|
||||
# Pick the face with minimal penetration as long as it has support.
|
||||
support = jp.where(support >= 0, -1e12, support)
|
||||
# Pick the face with the best separating axis.
|
||||
best_idx = support.argmax()
|
||||
face = faces[best_idx]
|
||||
normal = normals[best_idx]
|
||||
face_normal = normals[best_idx]
|
||||
|
||||
# Get closest point between the polygon face and the sphere center point.
|
||||
# Project the sphere center point onto poly plane. If it's inside polygon
|
||||
# edge normals, then we're done.
|
||||
pt = _project_pt_onto_plane(sphere_pos, face[0], normal)
|
||||
# side planes, then we're done.
|
||||
pt = _project_pt_onto_plane(sphere_pos, face[0], face_normal)
|
||||
edge_p0 = jp.roll(face, 1, axis=0)
|
||||
edge_p1 = face
|
||||
edge_normals = jax.vmap(jp.cross, in_axes=[0, None])(
|
||||
side_normals = jax.vmap(jp.cross, in_axes=[0, None])(
|
||||
edge_p1 - edge_p0,
|
||||
normal,
|
||||
face_normal,
|
||||
)
|
||||
edge_dist = jax.vmap(
|
||||
lambda plane_pt, plane_norm: (pt - plane_pt).dot(plane_norm)
|
||||
)(edge_p0, edge_normals)
|
||||
inside = jp.all(edge_dist <= 0) # lte to handle degenerate edges
|
||||
)(edge_p0, side_normals)
|
||||
pt_on_face = jp.all(edge_dist <= 0) # lte to handle degenerate edges
|
||||
|
||||
# If the point is outside edge normals, project onto the closest edge plane
|
||||
# If the point is outside side planes, project onto the closest side plane
|
||||
# that the point is in front of.
|
||||
degenerate_edge = jp.all(edge_normals == 0, axis=1)
|
||||
degenerate_edge = jp.all(side_normals == 0, axis=1)
|
||||
behind = edge_dist < 0.0
|
||||
edge_dist = jp.where(degenerate_edge | behind, 1e12, edge_dist)
|
||||
idx = edge_dist.argmin()
|
||||
edge_pt = math.closest_segment_point(edge_p0[idx], edge_p1[idx], pt)
|
||||
|
||||
pt = jp.where(inside, pt, edge_pt)
|
||||
pt = jp.where(pt_on_face, pt, edge_pt)
|
||||
|
||||
# Get the normal, dist, and contact position.
|
||||
n, d = math.normalize_with_norm(pt - sphere_pos)
|
||||
n = jp.where(inside | (d < 1e-6), normal, n)
|
||||
pt_normal, d = math.normalize_with_norm(pt - sphere_pos)
|
||||
# Ensure normal points towards convex centroid.
|
||||
n *= jp.where(jp.dot(pt, n) < 0, 1, -1)
|
||||
inside = jp.dot(pt, pt_normal) > 0
|
||||
sign = jp.where(inside, -1, 1)
|
||||
n = jp.where(pt_on_face | (d < 1e-6), -face_normal, sign * pt_normal)
|
||||
d *= sign
|
||||
|
||||
spt = sphere_pos + n * sphere.size[0]
|
||||
dist = d - sphere.size[0]
|
||||
dist = jp.where(has_separating_axis, 1.0, d - sphere.size[0])
|
||||
pos = (pt + spt) * 0.5
|
||||
|
||||
# Go back to world frame.
|
||||
@@ -292,8 +293,7 @@ def capsule_convex(cap: GeomInfo, convex: GeomInfo) -> Contact:
|
||||
support = get_support(faces, normals)
|
||||
has_support = jp.all(support < 0)
|
||||
|
||||
# Pick the face with minimal penetration as long as it has support.
|
||||
support = jp.where(support >= 0, -1e12, support)
|
||||
# Pick the face with minimal penetration.
|
||||
best_idx = support.argmax()
|
||||
face = faces[best_idx]
|
||||
normal = normals[best_idx]
|
||||
@@ -316,44 +316,64 @@ def capsule_convex(cap: GeomInfo, convex: GeomInfo) -> Contact:
|
||||
# Create variables for the face contact.
|
||||
pos = (cap_pts_clipped + face_pts) * 0.5
|
||||
contact_normal = -jp.stack([normal] * 2, 0)
|
||||
penetration = jp.where(
|
||||
face_penetration = jp.where(
|
||||
mask & has_support, jp.dot(face_pts - cap_pts_clipped, normal), -1
|
||||
)
|
||||
|
||||
# Get a potential edge contact.
|
||||
edge_closest, cap_closest = jax.vmap(
|
||||
math.closest_segment_to_segment_points,
|
||||
in_axes=[0, 0, None, None],
|
||||
)(edge_p0, edge_p1, cap_pts[0], cap_pts[1])
|
||||
e_idx = ((edge_closest - cap_closest) ** 2).sum(axis=1).argmin()
|
||||
cap_closest_pt, edge_closest_pt = cap_closest[e_idx], edge_closest[e_idx]
|
||||
edge_dir = edge_closest_pt - cap_closest_pt
|
||||
degenerate_edge_dir = jp.sum(jp.square(edge_dir)) < 1e-6
|
||||
edge_dir, edge_dist = math.normalize_with_norm(edge_dir)
|
||||
face_edge_normals = convex.face_edge_normal[best_idx][e_idx]
|
||||
inside_edge_voronoi_front = ((face_edge_normals @ edge_dir) < 0).all()
|
||||
inside_edge_voronoi_back = ((face_edge_normals @ edge_dir) > 0).all()
|
||||
outside_edge_voronoi = ~inside_edge_voronoi_front & ~inside_edge_voronoi_back
|
||||
edge_axis = math.normalize(-edge_closest_pt) # approximate edge axis
|
||||
edge_axis = jp.where(
|
||||
~degenerate_edge_dir & inside_edge_voronoi_front,
|
||||
edge_dir, # shallow edge penetration
|
||||
edge_axis, # deep edge penetration
|
||||
)
|
||||
# Pick a potential shallow edge contact.
|
||||
def get_edge_axis(edge):
|
||||
edge_closest_pt, cap_closest_pt = math.closest_segment_to_segment_points(
|
||||
edge[0], edge[1], cap_pts[0], cap_pts[1]
|
||||
)
|
||||
edge_dir = edge_closest_pt - cap_closest_pt
|
||||
degenerate_edge_dir = jp.sum(jp.square(edge_dir)) < 1e-6
|
||||
edge_axis, edge_dist = math.normalize_with_norm(edge_dir)
|
||||
return (
|
||||
edge_dist,
|
||||
edge_axis,
|
||||
degenerate_edge_dir,
|
||||
edge_closest_pt,
|
||||
cap_closest_pt,
|
||||
)
|
||||
|
||||
edge = jp.take(convex.vert, convex.face_edge, axis=0).reshape(-1, 2, 3)
|
||||
face_edge_normal = convex.face_edge_normal.reshape((-1, 2, 3)) # pytype: disable=attribute-error
|
||||
|
||||
res = jax.vmap(get_edge_axis)(edge.reshape(-1, 2, 3))
|
||||
e_idx = jp.abs(res[0]).argmin()
|
||||
(
|
||||
edge_dist,
|
||||
edge_axis,
|
||||
degenerate_edge_dir,
|
||||
edge_closest_pt,
|
||||
cap_closest_pt,
|
||||
) = jax.tree_map(lambda x, i=e_idx: jp.take(x, i, axis=0), res)
|
||||
|
||||
face_edge_normals = face_edge_normal[e_idx]
|
||||
edge_voronoi_front = ((face_edge_normals @ edge_axis) < 0).all()
|
||||
shallow = ~degenerate_edge_dir & edge_voronoi_front
|
||||
edge_penetration = jp.where(shallow, cap.size[0] - edge_dist, -1)
|
||||
|
||||
# Determine edge contact position.
|
||||
edge_pos = (
|
||||
edge_closest_pt + (cap_closest_pt + edge_axis * cap.size[0])
|
||||
) * 0.5
|
||||
edge_penetration = cap.size[0] - edge_dist
|
||||
# prefer face contact if the edge axis is parallel to the face normal
|
||||
edge_dir_parallel_to_face = (
|
||||
jp.abs(edge_dir.dot(normal)) > 0.99
|
||||
jp.abs(edge_axis.dot(normal)) > 0.99
|
||||
) & ~degenerate_edge_dir
|
||||
min_face_penetration = face_penetration.min()
|
||||
has_edge_contact = (
|
||||
(edge_penetration > 0)
|
||||
# prefer edge contact if the edge is smaller than face penetration
|
||||
& jp.where(
|
||||
min_face_penetration > 0,
|
||||
edge_penetration < min_face_penetration,
|
||||
True,
|
||||
)
|
||||
# prefer face contact if the edge axis is parallel to the face normal
|
||||
& ~edge_dir_parallel_to_face
|
||||
& ~outside_edge_voronoi
|
||||
# make sure we have a shallow contact
|
||||
& edge_voronoi_front
|
||||
)
|
||||
|
||||
# Get the contact info.
|
||||
@@ -367,7 +387,7 @@ def capsule_convex(cap: GeomInfo, convex: GeomInfo) -> Contact:
|
||||
n = n @ convex.mat.T
|
||||
|
||||
dist = -jp.where(
|
||||
has_edge_contact, penetration.at[0].set(edge_penetration), penetration
|
||||
has_edge_contact, jp.array([edge_penetration, -1]), face_penetration
|
||||
)
|
||||
frame = jax.vmap(math.make_frame)(n)
|
||||
return dist, pos, frame
|
||||
|
||||
@@ -132,37 +132,70 @@ class SphereCollisionTest(parameterized.TestCase):
|
||||
</mujoco>
|
||||
"""
|
||||
|
||||
def test_sphere_convex(self):
|
||||
d, dx = _collide(self._SPHERE_CONVEX)
|
||||
def test_sphere_convex_face(self):
|
||||
# no contact
|
||||
xml = self._SPHERE_CONVEX.replace(
|
||||
'<body pos="0.52 0 0.52">', '<body pos="0.55 0 0.5">'
|
||||
)
|
||||
d, dx = _collide(xml)
|
||||
self.assertEmpty(d.contact.dist)
|
||||
self.assertGreater(dx.contact.dist, 0)
|
||||
|
||||
# face contact
|
||||
xml = self._SPHERE_CONVEX.replace(
|
||||
'<body pos="0.52 0 0.52">', '<body pos="0.51 0 0.25">'
|
||||
)
|
||||
d, dx = _collide(xml)
|
||||
for field in dataclasses.fields(Contact):
|
||||
_assert_attr_eq(dx.contact, d.contact, field.name, 'sphere_convex', 1e-4)
|
||||
_assert_attr_eq(dx.contact, d.contact, field.name, 'face', 1e-4)
|
||||
|
||||
# test deep penetration
|
||||
# deep face contact
|
||||
xml = self._SPHERE_CONVEX.replace(
|
||||
'<body pos="0.52 0 0.52">', '<body pos="0.48 0 0.47">'
|
||||
)
|
||||
d, dx = _collide(xml)
|
||||
self.assertTrue((dx.contact.dist < 0).all())
|
||||
self.assertTrue((d.contact.dist < 0).all())
|
||||
np.testing.assert_allclose(dx.contact.dist, [-0.07], atol=1e-5)
|
||||
np.testing.assert_array_almost_equal(dx.contact.pos, d.contact.pos)
|
||||
np.testing.assert_array_almost_equal(
|
||||
dx.contact.frame, d.contact.frame.reshape((-1, 3, 3))
|
||||
)
|
||||
|
||||
def test_sphere_convex_edge(self):
|
||||
# edge contact
|
||||
d, dx = _collide(self._SPHERE_CONVEX)
|
||||
for field in dataclasses.fields(Contact):
|
||||
_assert_attr_eq(dx.contact, d.contact, field.name, 'edge', 1e-4)
|
||||
|
||||
# deep edge penetration
|
||||
xml = self._SPHERE_CONVEX.replace(
|
||||
'<body pos="0.52 0 0.52">', '<body pos="0.49 0 0.49">'
|
||||
)
|
||||
d, dx = _collide(xml)
|
||||
|
||||
self.assertTrue((dx.contact.dist < 0).all())
|
||||
self.assertTrue((d.contact.dist < 0).all())
|
||||
np.testing.assert_allclose(dx.contact.dist, [-0.06], atol=1e-5)
|
||||
np.testing.assert_array_almost_equal(dx.contact.pos, d.contact.pos)
|
||||
np.testing.assert_array_almost_equal(
|
||||
dx.contact.frame, d.contact.frame.reshape((-1, 3, 3))
|
||||
)
|
||||
|
||||
# test sphere center on vertex
|
||||
# vertex contact
|
||||
xml = self._SPHERE_CONVEX.replace(
|
||||
'<body pos="0.52 0 0.52">', '<body pos="0.5 0.52 0.51">'
|
||||
)
|
||||
d, dx = _collide(xml)
|
||||
for field in dataclasses.fields(Contact):
|
||||
_assert_attr_eq(dx.contact, d.contact, field.name, 'vertex', 1e-4)
|
||||
|
||||
# sphere center on vertex
|
||||
xml = self._SPHERE_CONVEX.replace(
|
||||
'<body pos="0.52 0 0.52">', '<body pos="0.5 0 0.5">'
|
||||
)
|
||||
d, dx = _collide(xml)
|
||||
|
||||
self.assertTrue((dx.contact.dist < 0).all())
|
||||
self.assertTrue((d.contact.dist < 0).all())
|
||||
np.testing.assert_array_almost_equal(dx.contact.pos, d.contact.pos)
|
||||
np.testing.assert_array_almost_equal(
|
||||
dx.contact.frame, d.contact.frame.reshape((-1, 3, 3))
|
||||
)
|
||||
for field in dataclasses.fields(Contact):
|
||||
_assert_attr_eq(dx.contact, d.contact, field.name, 'vertex_center', 1e-4)
|
||||
|
||||
|
||||
class CapsuleCollisionTest(parameterized.TestCase):
|
||||
@@ -309,16 +342,17 @@ class CapsuleCollisionTest(parameterized.TestCase):
|
||||
def test_capsule_convex_edge_deep(self):
|
||||
"""Tests deep edge penetration."""
|
||||
xml = self._CAP_EDGE_BOX.replace(
|
||||
'<body pos="0.5 0 0.55"', '<body pos="0.5 0 0.48"'
|
||||
'<body pos="0.5 0 0.55"', '<body pos="0.5 0 0.42"'
|
||||
)
|
||||
_, dx = _collide(xml)
|
||||
|
||||
np.testing.assert_array_equal(dx.contact.dist < 0, np.array([True, False]))
|
||||
np.testing.assert_array_almost_equal(dx.contact.dist[0], np.array([-0.13]))
|
||||
np.testing.assert_array_almost_equal(
|
||||
dx.contact.pos[0], np.array([0.478, 0, 0.475]), decimal=3
|
||||
dx.contact.pos[0], np.array([0.5, 0, 0.435]), decimal=3
|
||||
)
|
||||
s2 = np.sqrt(2) / 2.0
|
||||
np.testing.assert_array_almost_equal(
|
||||
dx.contact.frame[0, 0], np.array([-s2, 0, -s2]), decimal=3
|
||||
dx.contact.frame[0, 0], np.array([0, 0, -1]), decimal=3
|
||||
)
|
||||
|
||||
def test_capsule_convex_edge_shallow_tip(self):
|
||||
|
||||
Reference in New Issue
Block a user