Allow mjx.ray to exclude multiple bodies.
PiperOrigin-RevId: 820236386 Change-Id: I41229fa52043d844d116b540d60053628d1a842b
This commit is contained in:
committed by
Copybara-Service
parent
2fc8306e3d
commit
2893c3344c
@@ -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)]
|
||||
|
||||
@@ -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')
|
||||
|
||||
|
||||
Reference in New Issue
Block a user