Add ray filtering based on geom transparency and filter params.

PiperOrigin-RevId: 597316447
Change-Id: If9619320f0024c486bb7abe2c63e9a1ba1a631f0
This commit is contained in:
Erik Frey
2024-01-10 11:57:53 -08:00
committed by Copybara-Service
parent 8a9b041388
commit 85ab6183bc
4 changed files with 146 additions and 35 deletions
+2
View File
@@ -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
View File
@@ -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_
+91 -20
View File
@@ -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()
+10
View File
@@ -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