Allow mjx.ray to exclude multiple bodies.

PiperOrigin-RevId: 820236386
Change-Id: I41229fa52043d844d116b540d60053628d1a842b
This commit is contained in:
Andrea Gesmundo
2025-10-16 08:02:38 -07:00
committed by Copybara-Service
parent 2fc8306e3d
commit 2893c3344c
2 changed files with 17 additions and 8 deletions
+8 -4
View File
@@ -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)]
+9 -4
View File
@@ -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')