From a02fc405af8da2589b22cbc1ac1b4f5ea783e5ba Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Tue, 16 Jan 2024 16:28:19 -0800 Subject: [PATCH] Add naive ray mesh implementation. PiperOrigin-RevId: 598993152 Change-Id: I708505bf20a89d5a4840451a11b99771195003e6 --- doc/changelog.rst | 2 +- mjx/mujoco/mjx/_src/ray.py | 87 +++++++++++++++++++++++++++++--- mjx/mujoco/mjx/_src/ray_test.py | 37 ++++++++++++++ mjx/mujoco/mjx/_src/types.py | 14 +++++ mjx/mujoco/mjx/test_data/ray.xml | 2 + 5 files changed, 133 insertions(+), 9 deletions(-) diff --git a/doc/changelog.rst b/doc/changelog.rst index fb886273..eb9aeab0 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -15,7 +15,7 @@ MJX 2. Added :ref:`dyntype` ``filterexact``. 3. Added :at:`site` transmission. 4. Updated MJX colab tutorial with more stable quadruped environment. -5. Added ``mjx.ray`` which mirrors :ref:`mj_ray` for planes, spheres, capsules, and boxes. +5. Added ``mjx.ray`` which mirrors :ref:`mj_ray` for planes, spheres, capsules, boxes, and meshes. Bug fixes ^^^^^^^^^ diff --git a/mjx/mujoco/mjx/_src/ray.py b/mjx/mujoco/mjx/_src/ray.py index d093e926..a9413b21 100644 --- a/mjx/mujoco/mjx/_src/ray.py +++ b/mjx/mujoco/mjx/_src/ray.py @@ -19,6 +19,7 @@ from typing import Sequence, Tuple import jax from jax import numpy as jp import mujoco +from mujoco.mjx._src import math # pylint: disable=g-importing-member from mujoco.mjx._src.types import Data from mujoco.mjx._src.types import GeomType @@ -129,14 +130,79 @@ def _ray_box( return jp.min(jp.where(valid, x, jp.inf)) -def _ray_mesh( - size: jax.Array, +def _ray_triangle( + vert: jax.Array, pnt: jax.Array, vec: jax.Array, + b0: jax.Array, + b1: jax.Array, ) -> jax.Array: - """Returns the distance at which a ray intersects with a mesh.""" - del size, pnt, vec - raise NotImplementedError("ray <> mesh not implemented yet") + """Returns the distance at which a ray intersects with a triangle.""" + # project difference vectors in ray normal plane + planar = jp.dot(jp.array([b0, b1]), (vert - pnt).T) + + # determine if origin is inside planar projection of triangle + # A = (p0-p2, p1-p2), b = -p2, solve A*t = b + A = jp.array( # pylint: disable=invalid-name + [planar[:, 0] - planar[:, 2], planar[:, 1] - planar[:, 2]] + ).T.flatten() + b = -planar[:, 2] + det = A[0] * A[3] - A[1] * A[2] + valid = jp.abs(det) >= mujoco.mjMINVAL + + t0 = (A[3] * b[0] - A[1] * b[1]) / det + t1 = (-A[2] * b[0] + A[0] * b[1]) / det + valid &= (t0 >= 0) & (t1 >= 0) & (t0 + t1 <= 1) + + # intersect ray with plane of triangle + nrm = jp.cross(vert[0] - vert[2], vert[1] - vert[2]) + denom = jp.dot(vec, nrm) + valid &= jp.abs(denom) >= mujoco.mjMINVAL + + dist = jp.where(valid, -jp.dot(pnt - vert[2], nrm) / denom, jp.inf) + + return dist + + +def _ray_mesh( + m: Model, + geom_id: np.ndarray, + unused_size: jax.Array, + pnt: jax.Array, + vec: jax.Array, +) -> Tuple[jax.Array, jax.Array]: + """Returns the best distance and geom_id for ray mesh intersections.""" + data_id = m.geom_dataid[geom_id] + + ray_basis = lambda x: math.orthogonals(math.normalize(x)) + b0, b1 = jax.vmap(ray_basis)(vec) + + faceadr = np.append(m.mesh_faceadr, m.nmeshface) + vertadr = np.append(m.mesh_vertadr, m.nmeshvert) + + dists = [] + for i, id_ in enumerate(data_id): + face = m.mesh_face[faceadr[id_] : faceadr[id_ + 1]] + vert = m.mesh_vert[vertadr[id_] : vertadr[id_ + 1]] + dist = jax.vmap(_ray_triangle, in_axes=(0, None, None, None, None))( + vert[face], pnt[i], vec[i], b0[i], b1[i] + ) + dists.append(dist) + + # map the triangle id to data id + tri_id = np.append(0, (faceadr[data_id + 1] - faceadr[data_id]).cumsum()) + tri_data_id = np.zeros(tri_id[-1], dtype=np.int32) + tri_data_id[tri_id[:-1]] = 1 + tri_data_id = tri_data_id.cumsum() - 1 + + dists = jp.concatenate(dists) + min_id = jp.argmin(dists) + # Grab the best distance amongst all meshes, bypassing the argmin in `ray`. + # This avoids having to compute the best distance per mesh. + dist = dists[min_id, None] + id_ = jp.array(geom_id)[jp.array(tri_data_id)[min_id], None] + + return dist, id_ _RAY_FUNC = { @@ -144,7 +210,7 @@ _RAY_FUNC = { GeomType.SPHERE: _ray_sphere, GeomType.CAPSULE: _ray_capsule, GeomType.BOX: _ray_box, - # GeomType.MESH: _ray_mesh, + GeomType.MESH: _ray_mesh, } @@ -192,8 +258,13 @@ def ray( if id_.size == 0: continue - size, pnt, vec = m.geom_size[id_], geom_pnts[id_], geom_vecs[id_] - dist = jax.vmap(fn)(size, pnt, vec) + args = m.geom_size[id_], geom_pnts[id_], geom_vecs[id_] + + if geom_type == GeomType.MESH: + dist, id_ = fn(m, id_, *args) + else: + dist = jax.vmap(fn)(*args) + dists, ids = dists + [dist], ids + [id_] if not ids: diff --git a/mjx/mujoco/mjx/_src/ray_test.py b/mjx/mujoco/mjx/_src/ray_test.py index 0556c9a3..f89f0d87 100644 --- a/mjx/mujoco/mjx/_src/ray_test.py +++ b/mjx/mujoco/mjx/_src/ray_test.py @@ -144,6 +144,43 @@ class RayTest(absltest.TestCase): mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, unused) _assert_eq(dist, mj_dist, 'dist') + def test_ray_mesh(self): + """Tests MJX ray<>mesh matches MuJoCo.""" + 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) + + # look at the tetrahedron + pnt, vec = jp.array([2.0, 2.0, 2.0]), -jp.array([ + 1.0, + 1.0, + 1.0, + ]) + vec /= jp.linalg.norm(vec) + dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec) + _assert_eq(geomid, 4, 'geom_id') + + pnt, vec, geomid = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, geomid) + _assert_eq(geomid, 4, 'geom_id') + _assert_eq(dist, mj_dist, 'dist-tetrahedron') + + # look at the dodecahedron + pnt, vec = jp.array([4.0, 2.0, 2.0]), -jp.array([ + 2.0, + 1.0, + 1.0, + ]) + vec /= jp.linalg.norm(vec) + dist, geomid = jax.jit(mjx.ray)(mx, dx, pnt, vec) + _assert_eq(geomid, 5, 'geom_id') + + pnt, vec, geomid = np.array(pnt), np.array(vec), np.zeros(1, dtype=np.int32) + mj_dist = mujoco.mj_ray(m, d, pnt, vec, None, 1, -1, geomid) + _assert_eq(geomid, 5, 'geom_id') + _assert_eq(dist, mj_dist, 'dist-dodecahedron') + def test_ray_geomgroup(self): """Tests ray geomgroup filter.""" m = test_util.load_test_file('ray.xml') diff --git a/mjx/mujoco/mjx/_src/types.py b/mjx/mujoco/mjx/_src/types.py index 79cecaf5..1527b583 100644 --- a/mjx/mujoco/mjx/_src/types.py +++ b/mjx/mujoco/mjx/_src/types.py @@ -264,6 +264,8 @@ class Model(PyTreeNode): ngeom: number of geoms nsite: number of sites nmesh: number of meshes + nmeshvert: number of vertices in all meshes + nmeshface: number of triangular faces in all meshes nmat: number of materials npair: number of predefined geom pairs nexclude: number of excluded geom pairs @@ -321,6 +323,7 @@ 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_dataid: id of geom's mesh/hfield; -1: none (ngeom,) geom_group: group for visibility (ngeom,) geom_matid: material id for rendering (ngeom,) geom_priority: geom contact priority (ngeom,) @@ -338,6 +341,10 @@ class Model(PyTreeNode): 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) + mesh_vertadr: first vertex address (nmesh x 1) + mesh_faceadr: first face address (nmesh x 1) + mesh_vert: vertex positions for all meshes (nmeshvert, 3) + mesh_face: vertex face data (nmeshface, 3) 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,) @@ -390,6 +397,8 @@ class Model(PyTreeNode): ngeom: int nsite: int nmesh: int + nmeshvert: int + nmeshface: int nmat: int npair: int nexclude: int @@ -447,6 +456,7 @@ class Model(PyTreeNode): geom_conaffinity: np.ndarray geom_condim: np.ndarray geom_bodyid: np.ndarray + geom_dataid: np.ndarray geom_group: np.ndarray geom_matid: np.ndarray geom_priority: np.ndarray @@ -463,6 +473,10 @@ class Model(PyTreeNode): site_bodyid: np.ndarray site_pos: jax.Array site_quat: jax.Array + mesh_vertadr: np.ndarray + mesh_faceadr: np.ndarray + mesh_vert: np.ndarray + mesh_face: np.ndarray mat_rgba: np.ndarray pair_dim: np.ndarray pair_geom1: np.ndarray diff --git a/mjx/mujoco/mjx/test_data/ray.xml b/mjx/mujoco/mjx/test_data/ray.xml index a6424ec4..bbf42b1e 100644 --- a/mjx/mujoco/mjx/test_data/ray.xml +++ b/mjx/mujoco/mjx/test_data/ray.xml @@ -1,6 +1,7 @@ + @@ -12,5 +13,6 @@ +