Add ray filtering based on geom transparency and filter params.
PiperOrigin-RevId: 597316447 Change-Id: If9619320f0024c486bb7abe2c63e9a1ba1a631f0
This commit is contained in:
committed by
Copybara-Service
parent
8a9b041388
commit
85ab6183bc
@@ -111,6 +111,8 @@ def put_model(m: mujoco.MjModel, device=None) -> types.Model:
|
||||
for f in types.Model.fields()
|
||||
if f.type in (int, bytes, np.ndarray)
|
||||
}
|
||||
static_fields['geom_rgba'] = static_fields['geom_rgba'].reshape((-1, 4))
|
||||
static_fields['mat_rgba'] = static_fields['mat_rgba'].reshape((-1, 4))
|
||||
|
||||
device_fields = {
|
||||
f.name: copy.copy(getattr(m, f.name)) # copy because device_put is async
|
||||
|
||||
+43
-15
@@ -14,7 +14,7 @@
|
||||
# ==============================================================================
|
||||
"""Functions for ray interesection testing."""
|
||||
|
||||
from typing import Tuple
|
||||
from typing import Sequence, Tuple
|
||||
|
||||
import jax
|
||||
from jax import numpy as jp
|
||||
@@ -149,32 +149,60 @@ _RAY_FUNC = {
|
||||
|
||||
|
||||
def ray(
|
||||
m: Model, d: Data, pnt: jax.Array, vec: jax.Array
|
||||
m: Model,
|
||||
d: Data,
|
||||
pnt: jax.Array,
|
||||
vec: jax.Array,
|
||||
geomgroup: Sequence[int] = (),
|
||||
flg_static: bool = True,
|
||||
bodyexclude: int = -1,
|
||||
) -> Tuple[jax.Array, jax.Array]:
|
||||
"""Returns the geom id and distance at which a ray intersects with a geom."""
|
||||
"""Returns the geom id and distance at which a ray intersects with a geom.
|
||||
|
||||
ids = []
|
||||
dists = []
|
||||
Args:
|
||||
m: MJX model
|
||||
d: MJX data
|
||||
pnt: ray origin point (3,)
|
||||
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
|
||||
|
||||
Returns:
|
||||
dist: distance from ray origin to geom surface (or -1.0 for no intersection)
|
||||
id: id of intersected geom (or -1 for no intersection)
|
||||
"""
|
||||
|
||||
dists, ids = [], []
|
||||
geom_filter = m.geom_bodyid != bodyexclude
|
||||
geom_filter &= (m.geom_matid != -1) | (m.geom_rgba[:, 3] != 0)
|
||||
geom_filter &= (m.geom_matid == -1) | (m.mat_rgba[m.geom_matid, 3] != 0)
|
||||
geom_filter &= flg_static | (m.body_weldid[m.geom_bodyid] != 0)
|
||||
if geomgroup:
|
||||
geomgroup = np.array(geomgroup, dtype=bool)
|
||||
geom_filter &= geomgroup[np.clip(m.geom_group, 0, mujoco.mjNGROUP)]
|
||||
|
||||
# map ray to local geom frames
|
||||
geom_pnts = jax.vmap(lambda x, y: x.T @ (pnt - y))(d.geom_xmat, d.geom_xpos)
|
||||
geom_vecs = jax.vmap(lambda x: x.T @ vec)(d.geom_xmat)
|
||||
|
||||
for geom_type, fn in _RAY_FUNC.items():
|
||||
if not np.any(m.geom_type == geom_type):
|
||||
id_, = np.nonzero(geom_filter & (m.geom_type == geom_type))
|
||||
|
||||
if id_.size == 0:
|
||||
continue
|
||||
|
||||
geom_ids = jp.array(np.nonzero(m.geom_type == geom_type)[0])
|
||||
geom_dists = jax.vmap(fn)(
|
||||
m.geom_size[geom_ids], geom_pnts[geom_ids], geom_vecs[geom_ids]
|
||||
)
|
||||
ids.append(geom_ids)
|
||||
dists.append(geom_dists)
|
||||
size, pnt, vec = m.geom_size[id_], geom_pnts[id_], geom_vecs[id_]
|
||||
dist = jax.vmap(fn)(size, pnt, vec)
|
||||
dists, ids = dists + [dist], ids + [id_]
|
||||
|
||||
if not ids:
|
||||
return jp.array(-1), jp.array(-1.0)
|
||||
|
||||
ids = jp.concatenate(ids)
|
||||
dists = jp.concatenate(dists)
|
||||
ids = jp.concatenate(ids)
|
||||
min_id = jp.argmin(dists)
|
||||
id_ = jp.where(jp.isinf(dists[min_id]), -1, ids[min_id])
|
||||
dist = jp.where(jp.isinf(dists[min_id]), -1, dists[min_id])
|
||||
id_ = jp.where(jp.isinf(dists[min_id]), -1, ids[min_id])
|
||||
|
||||
return id_, dist
|
||||
return dist, id_
|
||||
|
||||
@@ -43,9 +43,9 @@ class RayTest(absltest.TestCase):
|
||||
mx, dx = mjx.put_model(m), mjx.put_data(m, d)
|
||||
|
||||
pnt, vec = jp.array([12.146, 1.865, 3.895]), jp.array([0, 0, -1.0])
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([-1]), 'geom_id')
|
||||
_assert_eq(dist, jp.array([-1]), 'dist')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, -1, 'geom_id')
|
||||
_assert_eq(dist, -1, 'dist')
|
||||
|
||||
def test_ray_plane(self):
|
||||
"""Tests MJX ray<>plane matches MuJoCo."""
|
||||
@@ -57,17 +57,17 @@ class RayTest(absltest.TestCase):
|
||||
# looking down at a slight angle
|
||||
pnt, vec = jp.array([2, 1, 3.0]), jp.array([0.1, 0.2, -1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([0]), 'geom_id')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, 0, 'geom_id')
|
||||
pnt, vec, unused = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32)
|
||||
mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused)
|
||||
_assert_eq(dist, mj_dist, 'dist')
|
||||
|
||||
# looking on wrong side of plane
|
||||
pnt = jp.array([0, 0, -0.5])
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([-1]), 'geom_id')
|
||||
_assert_eq(dist, jp.array([-1]), 'dist')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, -1, 'geom_id')
|
||||
_assert_eq(dist, -1, 'dist')
|
||||
|
||||
def test_ray_sphere(self):
|
||||
"""Tests MJX ray<>sphere matches MuJoCo."""
|
||||
@@ -79,8 +79,8 @@ class RayTest(absltest.TestCase):
|
||||
# looking down at sphere at a slight angle
|
||||
pnt, vec = jp.array([0, 0, 1.6]), jp.array([0.1, 0.2, -1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([1]), 'geom_id')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, 1, 'geom_id')
|
||||
pnt, vec, unused = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32)
|
||||
mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused)
|
||||
_assert_eq(dist, mj_dist, 'dist')
|
||||
@@ -95,8 +95,8 @@ class RayTest(absltest.TestCase):
|
||||
# looking down at capsule at a slight angle
|
||||
pnt, vec = jp.array([0.5, 1, 1.6]), jp.array([0, 0.05, -1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([2]), 'geom_id')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, 2, 'geom_id')
|
||||
pnt, vec, unused = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32)
|
||||
mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused)
|
||||
_assert_eq(dist, mj_dist, 'dist')
|
||||
@@ -104,8 +104,8 @@ class RayTest(absltest.TestCase):
|
||||
# looking up at capsule from below
|
||||
pnt, vec = jp.array([-0.5, 1, 0.05]), jp.array([0, 0.05, 1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([2]), 'geom_id')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, 2, 'geom_id')
|
||||
pnt, vec, unused = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32)
|
||||
mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused)
|
||||
_assert_eq(dist, mj_dist, 'dist')
|
||||
@@ -113,8 +113,8 @@ class RayTest(absltest.TestCase):
|
||||
# looking at cylinder of capsule from the side
|
||||
pnt, vec = jp.array([0, 1, 0.75]), jp.array([1, 0, 0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([2]), 'geom_id')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, 2, 'geom_id')
|
||||
pnt, vec, unused = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32)
|
||||
mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused)
|
||||
_assert_eq(dist, mj_dist, 'dist')
|
||||
@@ -129,8 +129,8 @@ class RayTest(absltest.TestCase):
|
||||
# looking down at box at a slight angle
|
||||
pnt, vec = jp.array([1, 0, 1.6]), jp.array([0, 0.05, -1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([3]), 'geom_id')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, 3, 'geom_id')
|
||||
pnt, vec, unused = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32)
|
||||
mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused)
|
||||
_assert_eq(dist, mj_dist, 'dist')
|
||||
@@ -138,12 +138,83 @@ class RayTest(absltest.TestCase):
|
||||
# looking up at box from below
|
||||
pnt, vec = jp.array([1, 0, 0.05]), jp.array([0, 0.05, 1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomid, dist = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, jp.array([3]), 'geom_id')
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, 3, 'geom_id')
|
||||
pnt, vec, unused = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32)
|
||||
mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused)
|
||||
_assert_eq(dist, mj_dist, 'dist')
|
||||
|
||||
def test_ray_geomgroup(self):
|
||||
"""Tests ray geomgroup filter."""
|
||||
m = test_util.load_test_file('ray.xml')
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_forward(m, d)
|
||||
mx, dx = mjx.put_model(m), mjx.put_data(m, d)
|
||||
ray_fn = jax.jit(mjx.ray, static_argnums=(4,))
|
||||
|
||||
# hits plane with geom_group[0] = 1
|
||||
pnt, vec = jp.array([2, 1, 3.0]), jp.array([0.1, 0.2, -1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomgroup = (1, 0, 0, 0, 0, 0)
|
||||
dist, geomid = ray_fn(mx, dx, pnt, vec, geomgroup)
|
||||
_assert_eq(geomid, 0, 'geom_id')
|
||||
pnt, vec, unused = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32)
|
||||
mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused)
|
||||
_assert_eq(dist, mj_dist, 'dist')
|
||||
|
||||
# nothing hit with geom_group[0] = 0
|
||||
pnt, vec = jp.array([2, 1, 3.0]), jp.array([0.1, 0.2, -1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
geomgroup = (0, 0, 0, 0, 0, 0)
|
||||
dist, geomid = ray_fn(mx, dx, pnt, vec, geomgroup)
|
||||
_assert_eq(geomid, -1, 'geom_id')
|
||||
_assert_eq(dist, -1, 'dist')
|
||||
|
||||
def test_ray_flg_static(self):
|
||||
"""Tests ray flg_static filter."""
|
||||
m = test_util.load_test_file('ray.xml')
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_forward(m, d)
|
||||
mx, dx = mjx.put_model(m), mjx.put_data(m, d)
|
||||
ray_fn = jax.jit(mjx.ray, static_argnames=('flg_static',))
|
||||
|
||||
# nothing hit with flg_static = False
|
||||
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, flg_static=False)
|
||||
_assert_eq(geomid, -1, 'geom_id')
|
||||
_assert_eq(dist, -1, 'dist')
|
||||
|
||||
def test_ray_bodyexclude(self):
|
||||
"""Tests ray bodyexclude filter."""
|
||||
m = test_util.load_test_file('ray.xml')
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_forward(m, d)
|
||||
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)
|
||||
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)
|
||||
_assert_eq(geomid, -1, 'geom_id')
|
||||
_assert_eq(dist, -1, 'dist')
|
||||
|
||||
def test_ray_invisible(self):
|
||||
"""Tests ray doesn't hit transparent geoms."""
|
||||
m = test_util.load_test_file('ray.xml')
|
||||
# nothing hit with transparent geoms:
|
||||
m.geom_rgba = 0
|
||||
d = mujoco.MjData(m)
|
||||
mujoco.mj_forward(m, d)
|
||||
mx, dx = mjx.put_model(m), mjx.put_data(m, d)
|
||||
|
||||
pnt, vec = jp.array([2, 1, 3.0]), jp.array([0.1, 0.2, -1.0])
|
||||
vec /= jp.linalg.norm(vec)
|
||||
dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec)
|
||||
_assert_eq(geomid, -1, 'geom_id')
|
||||
_assert_eq(dist, -1, 'dist')
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
absltest.main()
|
||||
|
||||
@@ -264,6 +264,7 @@ class Model(PyTreeNode):
|
||||
ngeom: number of geoms
|
||||
nsite: number of sites
|
||||
nmesh: number of meshes
|
||||
nmat: number of materials
|
||||
npair: number of predefined geom pairs
|
||||
nexclude: number of excluded geom pairs
|
||||
neq: number of equality constraints
|
||||
@@ -320,6 +321,8 @@ class Model(PyTreeNode):
|
||||
geom_conaffinity: geom contact affinity (ngeom,)
|
||||
geom_condim: contact dimensionality (1, 3, 4, 6) (ngeom,)
|
||||
geom_bodyid: id of geom's body (ngeom,)
|
||||
geom_group: group for visibility (ngeom,)
|
||||
geom_matid: material id for rendering (ngeom,)
|
||||
geom_priority: geom contact priority (ngeom,)
|
||||
geom_solmix: mixing coef for solref/imp in geom pair (ngeom,)
|
||||
geom_solref: constraint solver reference: contact (ngeom, mjNREF)
|
||||
@@ -330,9 +333,11 @@ class Model(PyTreeNode):
|
||||
geom_friction: friction for (slide, spin, roll) (ngeom, 3)
|
||||
geom_margin: include in solver if dist<margin-gap (ngeom,)
|
||||
geom_gap: include in solver if dist<margin-gap (ngeom,)
|
||||
geom_rgba: rgba when material is omitted (ngeom, 4)
|
||||
site_bodyid: id of site's body (nsite,)
|
||||
site_pos: local position offset rel. to body (nsite, 3)
|
||||
site_quat: local orientation offset rel. to body (nsite, 4)
|
||||
mat_rgba: rgba (nmat, 4)
|
||||
geom_convex_face: vertex face data, MJX only (ngeom,)
|
||||
geom_convex_vert: vertex data, MJX only (ngeom,)
|
||||
geom_convex_edge: unique edge data, MJX only (ngeom,)
|
||||
@@ -385,6 +390,7 @@ class Model(PyTreeNode):
|
||||
ngeom: int
|
||||
nsite: int
|
||||
nmesh: int
|
||||
nmat: int
|
||||
npair: int
|
||||
nexclude: int
|
||||
neq: int
|
||||
@@ -441,6 +447,8 @@ class Model(PyTreeNode):
|
||||
geom_conaffinity: np.ndarray
|
||||
geom_condim: np.ndarray
|
||||
geom_bodyid: np.ndarray
|
||||
geom_group: np.ndarray
|
||||
geom_matid: np.ndarray
|
||||
geom_priority: np.ndarray
|
||||
geom_solmix: jax.Array
|
||||
geom_solref: jax.Array
|
||||
@@ -451,9 +459,11 @@ class Model(PyTreeNode):
|
||||
geom_friction: jax.Array
|
||||
geom_margin: jax.Array
|
||||
geom_gap: jax.Array
|
||||
geom_rgba: np.ndarray
|
||||
site_bodyid: np.ndarray
|
||||
site_pos: jax.Array
|
||||
site_quat: jax.Array
|
||||
mat_rgba: np.ndarray
|
||||
pair_dim: np.ndarray
|
||||
pair_geom1: np.ndarray
|
||||
pair_geom2: np.ndarray
|
||||
|
||||
Reference in New Issue
Block a user