Breaking change: Add surface normal output to MuJoCo raycast functions.

PiperOrigin-RevId: 855781592
Change-Id: Id96b1ca7eaf722e260cc69d7706c28dc51f52d92
This commit is contained in:
Yuval Tassa
2026-01-13 10:21:31 -08:00
committed by Copybara-Service
parent 37762e3f70
commit 218226fc95
17 changed files with 457 additions and 326 deletions
+37 -84
View File
@@ -559,8 +559,8 @@ static mjtNum ray_box(const mjtNum pos[3], const mjtNum mat[9], const mjtNum siz
// intersect ray with hfield, compute normal if given
static mjtNum mj_rayHfieldNormal(const mjModel* m, const mjData* d, int geomid,
const mjtNum pnt[3], const mjtNum vec[3], mjtNum normal[3]) {
mjtNum mj_rayHfield(const mjModel* m, const mjData* d, int geomid,
const mjtNum pnt[3], const mjtNum vec[3], mjtNum normal[3]) {
// clear normal if given
if (normal) mju_zero3(normal);
@@ -738,13 +738,6 @@ static mjtNum mj_rayHfieldNormal(const mjModel* m, const mjData* d, int geomid,
}
// intersect ray with hfield
mjtNum mj_rayHfield(const mjModel* m, const mjData* d, int geomid,
const mjtNum pnt[3], const mjtNum vec[3]) {
return mj_rayHfieldNormal(m, d, geomid, pnt, vec, NULL);
}
// ray vs axis-aligned bounding box using slab method
// see Ericson, Real-time Collision Detection section 5.3.3.
int mju_raySlab(const mjtNum aabb[6], const mjtNum xpos[3],
@@ -889,8 +882,8 @@ mjtNum mju_rayTree(const mjModel* m, const mjData* d, int id, const mjtNum pnt[3
// intersect ray with signed distance field, compute normal if given
static mjtNum mj_raySdfNormal(const mjModel* m, const mjData* d, int g,
const mjtNum pnt[3], const mjtNum vec[3], mjtNum normal[3]) {
static mjtNum mj_raySdf(const mjModel* m, const mjData* d, int g,
const mjtNum pnt[3], const mjtNum vec[3], mjtNum normal[3]) {
if (normal) mju_zero3(normal);
mjtNum distance_total = 0;
@@ -956,8 +949,8 @@ static mjtNum mj_raySdfNormal(const mjModel* m, const mjData* d, int g,
}
// intersect ray with mesh, compute normal if given
static mjtNum mj_rayMeshNormal(const mjModel* m, const mjData* d, int id, const mjtNum pnt[3],
const mjtNum vec[3], mjtNum normal[3]) {
mjtNum mj_rayMesh(const mjModel* m, const mjData* d, int id, const mjtNum pnt[3],
const mjtNum vec[3], mjtNum normal[3]) {
// clear normal if given
if (normal) mju_zero3(normal);
@@ -975,17 +968,10 @@ static mjtNum mj_rayMeshNormal(const mjModel* m, const mjData* d, int id, const
}
// intersect ray with mesh
mjtNum mj_rayMesh(const mjModel* m, const mjData* d, int id, const mjtNum pnt[3],
const mjtNum vec[3]) {
return mj_rayMeshNormal(m, d, id, pnt, vec, NULL);
}
// intersect ray and find normal with primitive geom, no meshes or hfields, compute normal if given
mjtNum mju_rayGeomNormal(const mjtNum pos[3], const mjtNum mat[9], const mjtNum size[3],
const mjtNum pnt[3], const mjtNum vec[3], int geomtype,
mjtNum normal[3]) {
// intersect ray with primitive geom, no meshes or hfields, compute normal if given
mjtNum mju_rayGeom(const mjtNum pos[3], const mjtNum mat[9], const mjtNum size[3],
const mjtNum pnt[3], const mjtNum vec[3], int geomtype,
mjtNum normal[3]) {
switch ((mjtGeom) geomtype) {
case mjGEOM_PLANE:
return ray_plane(pos, mat, size, pnt, vec, normal);
@@ -1012,18 +998,11 @@ mjtNum mju_rayGeomNormal(const mjtNum pos[3], const mjtNum mat[9], const mjtNum
}
// intersect ray with primitive geom, no meshes or hfields
mjtNum mju_rayGeom(const mjtNum pos[3], const mjtNum mat[9], const mjtNum size[3],
const mjtNum pnt[3], const mjtNum vec[3], int geomtype) {
return mju_rayGeomNormal(pos, mat, size, pnt, vec, geomtype, NULL);
}
// intersect ray with flex, return nearest vertex id, compute normal if given
mjtNum mju_rayFlexNormal(const mjModel* m, const mjData* d, int flex_layer,
mjtByte flg_vert, mjtByte flg_edge, mjtByte flg_face,
mjtByte flg_skin, int flexid, const mjtNum pnt[3],
const mjtNum vec[3], int vertid[1], mjtNum normal[3]) {
mjtNum mj_rayFlex(const mjModel* m, const mjData* d, int flex_layer,
mjtByte flg_vert, mjtByte flg_edge, mjtByte flg_face,
mjtByte flg_skin, int flexid, const mjtNum pnt[3],
const mjtNum vec[3], int vertid[1], mjtNum normal[3]) {
int dim = m->flex_dim[flexid];
// clear normal if given
@@ -1102,8 +1081,8 @@ mjtNum mju_rayFlexNormal(const mjModel* m, const mjData* d, int flex_layer,
mju_quat2Mat(mat, quat);
// intersect ray with capsule
mjtNum sol = mju_rayGeomNormal(pos, mat, size, pnt, vec, mjGEOM_CAPSULE,
normal ? normal_local : NULL);
mjtNum sol = mju_rayGeom(pos, mat, size, pnt, vec, mjGEOM_CAPSULE,
normal ? normal_local : NULL);
// update
if (sol >= 0 && (x < 0 || sol < x)) {
@@ -1136,8 +1115,8 @@ mjtNum mju_rayFlexNormal(const mjModel* m, const mjData* d, int flex_layer,
size[0] = radius;
// intersect ray with sphere
mjtNum sol = mju_rayGeomNormal(vpos, NULL, size, pnt, vec, mjGEOM_SPHERE,
normal ? normal_local : NULL);
mjtNum sol = mju_rayGeom(vpos, NULL, size, pnt, vec, mjGEOM_SPHERE,
normal ? normal_local : NULL);
// update
if (sol >= 0 && (x < 0 || sol < x)) {
@@ -1208,17 +1187,6 @@ mjtNum mju_rayFlexNormal(const mjModel* m, const mjData* d, int flex_layer,
return x;
}
// intersect ray with flex, return nearest vertex id
mjtNum mju_rayFlex(const mjModel* m, const mjData* d, int flex_layer,
mjtByte flg_vert, mjtByte flg_edge, mjtByte flg_face,
mjtByte flg_skin, int flexid, const mjtNum pnt[3],
const mjtNum vec[3], int vertid[1]) {
return mju_rayFlexNormal(m, d, flex_layer, flg_vert, flg_edge, flg_face,
flg_skin, flexid, pnt, vec, vertid, NULL);
}
// intersect ray with skin, return nearest vertex id
mjtNum mju_raySkin(int nface, int nvert, const int* face, const float* vert,
const mjtNum pnt[3], const mjtNum vec[3], int vertid[1]) {
@@ -1337,9 +1305,9 @@ static int point_in_box(const mjtNum aabb[6], const mjtNum xpos[3],
// 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
mjtNum mj_rayNormal(const mjModel* m, const mjData* d, const mjtNum pnt[3], const mjtNum vec[3],
const mjtByte* geomgroup, mjtByte flg_static, int bodyexclude,
int geomid[1], mjtNum normal[3]) {
mjtNum mj_ray(const mjModel* m, const mjData* d, const mjtNum pnt[3], const mjtNum vec[3],
const mjtByte* geomgroup, mjtByte flg_static, int bodyexclude,
int geomid[1], mjtNum normal[3]) {
int ngeom = m->ngeom;
mjtNum dist, newdist;
mjtNum normal_local[3];
@@ -1360,13 +1328,13 @@ mjtNum mj_rayNormal(const mjModel* m, const mjData* d, const mjtNum pnt[3], cons
if (!ray_eliminate(m, d, i, geomgroup, flg_static, bodyexclude)) {
int type = m->geom_type[i];
if (type == mjGEOM_MESH) {
newdist = mj_rayMeshNormal(m, d, i, pnt, vec, p_normal);
newdist = mj_rayMesh(m, d, i, pnt, vec, p_normal);
} else if (type == mjGEOM_HFIELD) {
newdist = mj_rayHfieldNormal(m, d, i, pnt, vec, p_normal);
newdist = mj_rayHfield(m, d, i, pnt, vec, p_normal);
} else if (type == mjGEOM_SDF) {
newdist = mj_raySdfNormal(m, d, i, pnt, vec, p_normal);
newdist = mj_raySdf(m, d, i, pnt, vec, p_normal);
} else {
newdist = mju_rayGeomNormal(d->geom_xpos+3*i, d->geom_xmat+9*i,
newdist = mju_rayGeom(d->geom_xpos+3*i, d->geom_xmat+9*i,
m->geom_size+3*i, pnt, vec, type, p_normal);
}
@@ -1383,13 +1351,6 @@ mjtNum mj_rayNormal(const mjModel* m, const mjData* d, const mjtNum pnt[3], cons
}
// intersect ray
mjtNum mj_ray(const mjModel* m, const mjData* d, const mjtNum pnt[3], const mjtNum vec[3],
const mjtByte* geomgroup, mjtByte flg_static, int bodyexclude, int geomid[1]) {
return mj_rayNormal(m, d, pnt, vec, geomgroup, flg_static, bodyexclude, geomid, NULL);
}
// 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[9], const mjtByte* geomgroup, mjtByte flg_static,
@@ -1502,7 +1463,7 @@ static mjtNum mju_singleRay(const mjModel* m, mjData* d, const mjtNum pnt[3], co
// clear result
dist = -1;
*geomid = -1;
if (geomid) *geomid = -1;
if (normal) mju_zero3(normal);
// get ray spherical coordinates
@@ -1556,20 +1517,20 @@ static mjtNum mju_singleRay(const mjModel* m, mjData* d, const mjtNum pnt[3], co
// dispatch to type-specific ray function
int type = m->geom_type[i];
if (type == mjGEOM_MESH) {
newdist = mj_rayMeshNormal(m, d, i, pnt, vec, p_normal);
newdist = mj_rayMesh(m, d, i, pnt, vec, p_normal);
} else if (type == mjGEOM_HFIELD) {
newdist = mj_rayHfieldNormal(m, d, i, pnt, vec, p_normal);
newdist = mj_rayHfield(m, d, i, pnt, vec, p_normal);
} else if (type == mjGEOM_SDF) {
newdist = mj_raySdfNormal(m, d, i, pnt, vec, p_normal);
newdist = mj_raySdf(m, d, i, pnt, vec, p_normal);
} else {
newdist = mju_rayGeomNormal(d->geom_xpos+3*i, d->geom_xmat+9*i,
m->geom_size+3*i, pnt, vec, type, p_normal);
newdist = mju_rayGeom(d->geom_xpos+3*i, d->geom_xmat+9*i,
m->geom_size+3*i, pnt, vec, type, p_normal);
}
// update if closer intersection found
if (newdist >= 0 && (newdist < dist || dist < 0)) {
dist = newdist;
*geomid = i;
if (geomid) *geomid = i;
if (normal) mju_copy3(normal, normal_local);
}
}
@@ -1580,9 +1541,9 @@ static mjtNum mju_singleRay(const mjModel* m, mjData* d, const mjtNum pnt[3], co
// performs multiple ray intersections, compute normals if given
void mj_multiRayNormal(const mjModel* m, mjData* d, const mjtNum pnt[3], const mjtNum* vec,
const mjtByte* geomgroup, mjtByte flg_static, int bodyexclude,
int* geomid, mjtNum* dist, mjtNum* normal, int nray, mjtNum cutoff) {
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, mjtNum* normal, int nray, mjtNum cutoff) {
mj_markStack(d);
// allocate source
@@ -1598,7 +1559,8 @@ void mj_multiRayNormal(const mjModel* m, mjData* d, const mjtNum pnt[3], const m
if (mju_dot3(vec+3*i, vec+3*i) < mjMINVAL) {
dist[i] = -1;
} else {
dist[i] = mju_singleRay(m, d, pnt, vec+3*i, geom_eliminate, geom_ba, geomid+i,
int* p_geomid = geomid ? geomid + i : NULL;
dist[i] = mju_singleRay(m, d, pnt, vec+3*i, geom_eliminate, geom_ba, p_geomid,
normal ? normal+3*i : NULL);
}
}
@@ -1606,12 +1568,3 @@ void mj_multiRayNormal(const mjModel* m, mjData* d, const mjtNum pnt[3], const m
mj_freeStack(d);
}
// performs multiple ray intersections with the precomputed 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, mjtNum cutoff) {
mj_multiRayNormal(m, d, pnt, vec, geomgroup, flg_static, bodyexclude,
geomid, dist, NULL, nray, cutoff);
}