Add mj_multiRay function to perform multiple ray intersection with the same source.

PiperOrigin-RevId: 528684535
Change-Id: I6e1747a4a0910e6061d6154e19ad8723f56c22e3
This commit is contained in:
Alessio Quaglino
2023-05-01 23:02:12 -07:00
committed by Copybara-Service
parent 4d2bf636ea
commit 2ad82d5998
11 changed files with 440 additions and 21 deletions
+211 -8
View File
@@ -20,6 +20,7 @@
#include <mujoco/mjdata.h>
#include <mujoco/mjmodel.h>
#include <mujoco/mjvisualize.h>
#include "engine/engine_io.h"
#include "engine/engine_macro.h"
#include "engine/engine_util_blas.h"
#include "engine/engine_util_errmem.h"
@@ -46,6 +47,20 @@ static void ray_map(const mjtNum* pos, const mjtNum* mat, const mjtNum* pnt, con
// map to azimuth angle in spherical coordinates
static mjtNum longitude(const mjtNum vec[3]) {
return mju_atan2(vec[1], vec[0]);
}
// map to elevation angle in spherical coordinates
static mjtNum latitude(const mjtNum vec[3]) {
return mju_atan2(mju_sqrt(vec[0]*vec[0] + vec[1]*vec[1]), vec[2]);
}
// eliminate geom
static int ray_eliminate(const mjModel* m, const mjData* d, int geomid,
const mjtByte* geomgroup, mjtByte flg_static, int bodyexclude) {
@@ -207,13 +222,13 @@ static mjtNum ray_plane(const mjtNum* pos, const mjtNum* mat, const mjtNum* size
// sphere
static mjtNum ray_sphere(const mjtNum* pos, const mjtNum* mat, const mjtNum* size,
static mjtNum ray_sphere(const mjtNum* pos, const mjtNum* mat, mjtNum dist_sqr,
const mjtNum* pnt, const mjtNum* vec) {
// (x*vec+pnt-pos)'*(x*vec+pnt-pos) = size[0]*size[0]
mjtNum dif[3] = {pnt[0]-pos[0], pnt[1]-pos[1], pnt[2]-pos[2]};
mjtNum a = vec[0]*vec[0] + vec[1]*vec[1] + vec[2]*vec[2];
mjtNum b = vec[0]*dif[0] + vec[1]*dif[1] + vec[2]*dif[2];
mjtNum c = dif[0]*dif[0] + dif[1]*dif[1] + dif[2]*dif[2] - size[0]*size[0];
mjtNum c = dif[0]*dif[0] + dif[1]*dif[1] + dif[2]*dif[2] - dist_sqr;
// solve a*x^2 + 2*b*x + c = 0
mjtNum xx[2];
@@ -227,7 +242,7 @@ static mjtNum ray_capsule(const mjtNum* pos, const mjtNum* mat, const mjtNum* si
const mjtNum* pnt, const mjtNum* vec) {
// bounding sphere test
mjtNum ssz = size[0] + size[1];
if (ray_sphere(pos, NULL, &ssz, pnt, vec)<0) {
if (ray_sphere(pos, NULL, ssz*ssz, pnt, vec)<0) {
return -1;
}
@@ -315,8 +330,8 @@ static mjtNum ray_ellipsoid(const mjtNum* pos, const mjtNum* mat, const mjtNum*
static mjtNum ray_cylinder(const mjtNum* pos, const mjtNum* mat, const mjtNum* size,
const mjtNum* pnt, const mjtNum* vec) {
// bounding sphere test
mjtNum ssz = mju_sqrt(size[0]*size[0] + size[1]*size[1]);
if (ray_sphere(pos, NULL, &ssz, pnt, vec)<0) {
mjtNum ssz = size[0]*size[0] + size[1]*size[1];
if (ray_sphere(pos, NULL, ssz, pnt, vec)<0) {
return -1;
}
@@ -382,8 +397,8 @@ static mjtNum ray_box(const mjtNum* pos, const mjtNum* mat, const mjtNum* size,
}
// bounding sphere test
mjtNum ssz = mju_sqrt(size[0]*size[0] + size[1]*size[1] + size[2]*size[2]);
if (ray_sphere(pos, NULL, &ssz, pnt, vec)<0) {
mjtNum ssz = size[0]*size[0] + size[1]*size[1] + size[2]*size[2];
if (ray_sphere(pos, NULL, ssz, pnt, vec)<0) {
return -1;
}
@@ -659,7 +674,7 @@ mjtNum mju_rayGeom(const mjtNum* pos, const mjtNum* mat, const mjtNum* size,
return ray_plane(pos, mat, size, pnt, vec);
case mjGEOM_SPHERE:
return ray_sphere(pos, mat, size, pnt, vec);
return ray_sphere(pos, mat, size[0]*size[0], pnt, vec);
case mjGEOM_CAPSULE:
return ray_capsule(pos, mat, size, pnt, vec);
@@ -774,6 +789,28 @@ mjtNum mju_raySkin(int nface, int nvert, const int* face, const float* vert,
// return 1 if point is inside object-aligned bounding box, 0 otherwise
static int point_in_box(const mjtNum aabb[6], const mjtNum xpos[3],
const mjtNum xmat[9], const mjtNum pnt[3]) {
mjtNum point[3];
// compute point in local coordinates of the box
mju_sub3(point, pnt, xpos);
mju_rotVecMatT(point, point, xmat);
mju_subFrom3(point, aabb);
// check intersections
for (int j=0; j<3; j++) { // directions
if (mju_abs(point[j]) > aabb[3+j]) {
return 0;
}
}
return 1;
}
//---------------------------- main entry point ---------------------------------------------------
// intersect ray (pnt+x*vec, x>=0) with visible geoms, except geoms on bodyexclude
@@ -819,3 +856,169 @@ mjtNum mj_ray(const mjModel* m, const mjData* d, const mjtNum* pnt, const mjtNum
return dist;
}
// Initializes spherical bounding angles (geom_ba) and flag vector for a given source
void mju_multiRayPrepare(const mjModel* m, const mjData* d, const mjtNum pnt[3],
const mjtNum* ray_xmat, const mjtByte* geomgroup, mjtByte flg_static,
int bodyexclude, mjtNum* geom_ba, int* geom_eliminate) {
if (ray_xmat) {
mju_error("ray_xmat is currently unused, should be NULL");
}
if (geom_eliminate) {
// compute eliminate flag for all geoms
for (int geomid=0; geomid<m->ngeom; geomid++)
geom_eliminate[geomid] = ray_eliminate(m, d, geomid, geomgroup, flg_static, bodyexclude);
}
for (int b=0; b<m->nbody; b++) {
// skip precomputation if no bounding volume is available
if (m->body_bvhadr[b] == -1) {
continue;
}
// loop over child geoms, compute bounding angles
for (int i=0; i<m->body_geomnum[b]; i++) {
int g = i + m->body_geomadr[b];
mjtNum AABB[4] = {mjMAXVAL, mjMAXVAL, -mjMAXVAL, -mjMAXVAL};
mjtNum* aabb = m->geom_aabb + 6*g;
mjtNum* xpos = d->geom_xpos + 3*g;
mjtNum* xmat = d->geom_xmat + 9*g;
if (point_in_box(aabb, xpos, xmat, pnt)) {
(geom_ba+4*g)[0] = -mjPI;
(geom_ba+4*g)[1] = -mjPI/2;
(geom_ba+4*g)[2] = mjPI;
(geom_ba+4*g)[3] = mjPI/2;
continue;
}
// loop over box vertices, compute spherical aperture
for (int v=0; v<8; v++) {
mjtNum vert[3], box[3];
vert[0] = (v&1 ? aabb[0]+aabb[3] : aabb[0]-aabb[3]);
vert[1] = (v&2 ? aabb[1]+aabb[4] : aabb[1]-aabb[4]);
vert[2] = (v&4 ? aabb[2]+aabb[5] : aabb[2]-aabb[5]);
// rotate to the world frame
mju_rotVecMat(box, vert, xmat);
mju_addTo3(box, xpos);
// spherical coordinates
mju_sub3(vert, box, pnt);
mjtNum azimuth = longitude(vert);
mjtNum elevation = latitude(vert);
// update bounds
AABB[0] = mju_min(AABB[0], azimuth);
AABB[1] = mju_min(AABB[1], elevation);
AABB[2] = mju_max(AABB[2], azimuth);
AABB[3] = mju_max(AABB[3], elevation);
}
if (AABB[2]-AABB[0] > mjPI) {
AABB[0] = -mjPI;
AABB[2] = mjPI;
}
if (AABB[3]-AABB[1] > mjPI) { // SHOULD NOT OCCUR
mju_error("mj_ray: discontinuity in azimuth angle");
}
mju_copy(geom_ba+4*g, AABB, 4);
}
}
}
// Performs single ray intersection
static mjtNum mju_singleRay(const mjModel* m, mjData* d, const mjtNum pnt[3], const mjtNum vec[3],
int* ray_eliminate, mjtNum* geom_ba, int geomid[1]) {
mjtNum dist, newdist;
// check vector length
if (mju_norm3(vec)<mjMINVAL) {
mju_error("mj_ray: vector length is too small");
}
// clear result
dist = -1;
*geomid = -1;
// get ray spherical coordinates
mjtNum azimuth = longitude(vec);
mjtNum elevation = latitude(vec);
// loop over bodies not eliminated by bodyexclude
for (int b=0; b<m->nbody; b++) {
// exclude body using bounding sphere test
if (m->body_bvhadr[b] != -1) {
mjtNum* pos = m->bvh_aabb + 6*m->body_bvhadr[b];
mjtNum* size = pos + 3;
mjtNum ssz = size[0]*size[0] + size[1]*size[1] + size[2]*size[2];
if (ray_sphere(pos, NULL, ssz, pnt, vec)<0) {
continue;
}
}
// loop over geoms if bounding sphere test fails
for (int g=0; g<m->body_geomnum[b]; g++) {
int i = m->body_geomadr[b] + g;
if (ray_eliminate[i]) {
continue;
}
// exclude geom using bounding angles
if (m->body_bvhadr[b] != -1) {
if (azimuth<(geom_ba+4*i)[0] || elevation<(geom_ba+4*i)[1] ||
azimuth>(geom_ba+4*i)[2] || elevation>(geom_ba+4*i)[3]) {
continue;
}
}
// handle mesh and hfield separately
if (m->geom_type[i]==mjGEOM_MESH) {
newdist = mj_rayMesh(m, d, i, pnt, vec);
} else if (m->geom_type[i]==mjGEOM_HFIELD) {
newdist = mj_rayHfield(m, d, i, pnt, vec);
}
// otherwise general dispatch
else {
newdist = mju_rayGeom(d->geom_xpos+3*i, d->geom_xmat+9*i,
m->geom_size+3*i, pnt, vec, m->geom_type[i]);
}
// update if closer intersection found
if (newdist>=0 && (newdist<dist || dist<0)) {
dist = newdist;
*geomid = i;
}
}
}
return dist;
}
// Performs multiple ray intersections with the precomputes bv and flags
void mj_multiRay(const mjModel* m, mjData* d, const mjtNum pnt[3], const mjtNum* vec,
const mjtByte* geomgroup, mjtByte flg_static, int bodyexclude,
int* geomid, mjtNum* dist, int nray) {
mjMARKSTACK;
// allocate source
mjtNum* geom_ba = mj_stackAlloc(d, 4*m->ngeom);
int* geom_eliminate = mj_stackAllocInt(d, m->ngeom);
// initialize source
mju_multiRayPrepare(m, d, pnt, NULL, geomgroup, flg_static, bodyexclude, geom_ba, geom_eliminate);
// loop over rays
for (int i=0; i<nray; i++) {
dist[i] = mju_singleRay(m, d, pnt, vec+3*i, geom_eliminate, geom_ba, geomid+i);
}
mjFREESTACK;
}
+10
View File
@@ -23,6 +23,16 @@
extern "C" {
#endif
MJAPI void mju_multiRayPrepare(const mjModel* m, const mjData* d, const mjtNum pnt[3],
const mjtNum* ray_xmat, const mjtByte* geomgroup, mjtByte flg_static,
int bodyexclude, mjtNum* geom_ba, int* geom_eliminate);
// Intersect multiple rays emanating from a single source
// Similar semantics to mj_ray, but vec is an array of (nray x 3) directions.
MJAPI void mj_multiRay(const mjModel* m, mjData* d, const mjtNum pnt[3], const mjtNum* vec,
const mjtByte* geomgroup, mjtByte flg_static, int bodyexclude,
int* geomid, mjtNum* dist, int nray);
// intersect ray (pnt+x*vec, x>=0) with visible geoms, except geoms on bodyexclude
// return geomid and distance (x) to nearest surface, or -1 if no intersection
// geomgroup, flg_static are as in mjvOption; geomgroup==NULL skips group exclusion