diff --git a/mjx/mujoco/mjx/_src/io.py b/mjx/mujoco/mjx/_src/io.py index d7ae4a1d..5b406185 100644 --- a/mjx/mujoco/mjx/_src/io.py +++ b/mjx/mujoco/mjx/_src/io.py @@ -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 diff --git a/mjx/mujoco/mjx/_src/ray.py b/mjx/mujoco/mjx/_src/ray.py index 582d6e4f..d093e926 100644 --- a/mjx/mujoco/mjx/_src/ray.py +++ b/mjx/mujoco/mjx/_src/ray.py @@ -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_ diff --git a/mjx/mujoco/mjx/_src/ray_test.py b/mjx/mujoco/mjx/_src/ray_test.py index 712aca63..0556c9a3 100644 --- a/mjx/mujoco/mjx/_src/ray_test.py +++ b/mjx/mujoco/mjx/_src/ray_test.py @@ -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() diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 595d7f90..79cecaf5 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -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