diff --git a/mjx/mujoco/mjx/_src/collision_convex.py b/mjx/mujoco/mjx/_src/collision_convex.py index 62d357f9..21197c15 100644 --- a/mjx/mujoco/mjx/_src/collision_convex.py +++ b/mjx/mujoco/mjx/_src/collision_convex.py @@ -173,7 +173,7 @@ def _manifold_points( bp = b - poly dist_bp = jp.abs(bp.dot(bc)) + dist_mask dist_ap = jp.abs(ap.dot(ac)) + dist_mask - d_idx = jp.concatenate([dist_bp, dist_ap]).argmax() % poly.shape[0] + d_idx = (dist_bp + dist_ap).argmax() % poly.shape[0] return jp.array([a_idx, b_idx, c_idx, d_idx]) @@ -514,6 +514,7 @@ def plane_convex(plane: GeomInfo, convex: GeomInfo) -> Contact: frame = jp.stack([math.make_frame(n)] * 4, axis=0) unique = jp.tril(idx == idx[:, None]).sum(axis=1) == 1 dist = jp.where(unique, -support[idx], 1) + pos = pos - 0.5 * dist[:, None] * n return dist, pos, frame diff --git a/mjx/mujoco/mjx/_src/collision_driver_test.py b/mjx/mujoco/mjx/_src/collision_driver_test.py index 5b58f7f8..f0bca5e9 100644 --- a/mjx/mujoco/mjx/_src/collision_driver_test.py +++ b/mjx/mujoco/mjx/_src/collision_driver_test.py @@ -265,7 +265,7 @@ class ConvexTest(absltest.TestCase): - + @@ -281,7 +281,35 @@ class ConvexTest(absltest.TestCase): # extract the contact points with penetration c = jax.tree_map(lambda x: jp.take(x, jp.array([0, 1]), axis=0), dx.contact) for field in dataclasses.fields(Contact): - _assert_attr_eq(c, d.contact, field.name, 'box_plane', 1e-2) + _assert_attr_eq(c, d.contact, field.name, 'box_plane', 1e-5) + + _FLAT_BOX_PLANE = """ + + + + + + + + + + """ + + def test_flat_box_plane(self): + """Tests box collision with a plane.""" + with jax.disable_jit(): + d, dx = _collide(self._FLAT_BOX_PLANE) + + np.testing.assert_array_less(dx.contact.dist, 0) + + # sort positions for comparison + idx = np.lexsort((dx.contact.pos[:, 0], dx.contact.pos[:, 1])) + dx = dx.tree_replace({'contact.pos': dx.contact.pos[idx]}) + idx = np.lexsort((d.contact.pos[:, 0], d.contact.pos[:, 1])) + d.contact.pos[:] = d.contact.pos[idx] + + for field in dataclasses.fields(Contact): + _assert_attr_eq(dx.contact, d.contact, field.name, 'flat_box_plane', 1e-5) _BOX_BOX = """