diff --git a/mjx/mujoco/mjx/_src/ray.py b/mjx/mujoco/mjx/_src/ray.py index 2e2055cf..dd2424af 100644 --- a/mjx/mujoco/mjx/_src/ray.py +++ b/mjx/mujoco/mjx/_src/ray.py @@ -237,7 +237,7 @@ def ray( vec: jax.Array, geomgroup: Sequence[int] = (), flg_static: bool = True, - bodyexclude: int = -1, + bodyexclude: Sequence[int] | int = -1, ) -> Tuple[jax.Array, jax.Array]: """Returns the geom id and distance at which a ray intersects with a geom. @@ -248,7 +248,7 @@ def ray( vec: ray direction (3,) geomgroup: group inclusion/exclusion mask, or empty to ignore flg_static: if True, allows rays to intersect with static geoms - bodyexclude: ignore geoms on specified body id + bodyexclude: ignore geoms on specified body id or sequence of body ids Returns: dist: distance from ray origin to geom surface (or -1.0 for no intersection) @@ -256,8 +256,12 @@ def ray( """ dists, ids = [], [] - geom_filter = m.geom_bodyid != bodyexclude - geom_filter &= flg_static | (m.body_weldid[m.geom_bodyid] != 0) + if not isinstance(bodyexclude, Sequence): + bodyexclude = [bodyexclude] + geom_filter = flg_static | (m.body_weldid[m.geom_bodyid] != 0) + # Loop through the body IDs to exclude and update the filter + for bodyid in bodyexclude: + geom_filter &= (m.geom_bodyid != bodyid) if geomgroup: geomgroup = np.array(geomgroup, dtype=bool) geom_filter &= geomgroup[np.clip(m.geom_group, 0, mujoco.mjNGROUP)] diff --git a/mjx/mujoco/mjx/_src/ray_test.py b/mjx/mujoco/mjx/_src/ray_test.py index 505878d3..ed0d6f19 100644 --- a/mjx/mujoco/mjx/_src/ray_test.py +++ b/mjx/mujoco/mjx/_src/ray_test.py @@ -15,6 +15,7 @@ """Tests for ray functions.""" from absl.testing import absltest +from absl.testing import parameterized import jax from jax import numpy as jp import mujoco @@ -33,7 +34,7 @@ def _assert_eq(a, b, name): np.testing.assert_allclose(a, b, err_msg=err_msg, atol=tol, rtol=tol) -class RayTest(absltest.TestCase): +class RayTest(parameterized.TestCase): def test_ray_nothing(self): """Tests that MJX ray returns -1 when nothing is hit.""" @@ -220,7 +221,11 @@ class RayTest(absltest.TestCase): _assert_eq(geomid, -1, 'geom_id') _assert_eq(dist, -1, 'dist') - def test_ray_bodyexclude(self): + @parameterized.named_parameters( + ('int', 0), + ('sequence', (0,)), + ) + def test_ray_bodyexclude(self, bodyexclude): """Tests ray bodyexclude filter.""" m = test_util.load_test_file('ray.xml') d = mujoco.MjData(m) @@ -228,10 +233,10 @@ class RayTest(absltest.TestCase): mx, dx = mjx.put_model(m), mjx.put_data(m, d) ray_fn = jax.jit(mjx.ray, static_argnames=('bodyexclude',)) - # nothing hit with bodyexclude = 0 (world body) + # The ray should hit the plane (geom 0, body 0), but body 0 is excluded. pnt, vec = jp.array([2, 1, 3.0]), jp.array([0.1, 0.2, -1.0]) vec /= jp.linalg.norm(vec) - dist, geomid = ray_fn(mx, dx, pnt, vec, bodyexclude=0) + dist, geomid = ray_fn(mx, dx, pnt, vec, bodyexclude=bodyexclude) _assert_eq(geomid, -1, 'geom_id') _assert_eq(dist, -1, 'dist')