diff --git a/mjx/mujoco/mjx/_src/collision_driver.py b/mjx/mujoco/mjx/_src/collision_driver.py index bdcf5275..87ef925f 100644 --- a/mjx/mujoco/mjx/_src/collision_driver.py +++ b/mjx/mujoco/mjx/_src/collision_driver.py @@ -293,7 +293,7 @@ def _collide_geoms( size2 = jp.max(m.geom_size[jp.array(geom2)], axis=-1) # TODO(btaba): consider re-using collision info for (sphere, sphere) dists = jax.vmap(jp.linalg.norm)(g2.pos - g1.pos) - (size1 + size2) - _, idx = jax.lax.top_k(dists, k=n_pairs) + _, idx = jax.lax.top_k(-dists, k=n_pairs) g1, g2, params = jax.tree_map( lambda x, idx=idx: x[idx, ...], (g1, g2, params) ) diff --git a/mjx/mujoco/mjx/_src/collision_driver_test.py b/mjx/mujoco/mjx/_src/collision_driver_test.py index 3b7aff1d..c57f38fb 100644 --- a/mjx/mujoco/mjx/_src/collision_driver_test.py +++ b/mjx/mujoco/mjx/_src/collision_driver_test.py @@ -583,6 +583,7 @@ class TopKContactTest(absltest.TestCase): self.assertEqual(dx_all.contact.dist.shape, (6,)) self.assertEqual(dx_top_k.contact.dist.shape, (2,)) + self.assertTrue((dx_top_k.contact.dist < 0).all()) if __name__ == '__main__':