Introduce mjPreContact, a minimal struct passed into the collision functions.

PiperOrigin-RevId: 918533795
Change-Id: I2b5af05c1479b25d5c2cfdc690321a26fce6ede6
This commit is contained in:
Kyle Bayes
2026-05-20 10:47:11 -07:00
committed by Copybara-Service
parent 53b3137a12
commit 7174d33f08
25 changed files with 709 additions and 406 deletions
+55 -31
View File
@@ -21,7 +21,6 @@
#include <mujoco/mjmodel.h>
#include <mujoco/mjsan.h> // IWYU pragma: keep
#include <mujoco/mjtnum.h>
#include "engine/engine_collision_primitive.h"
#include "engine/engine_plugin.h"
#include "engine/engine_util_blas.h"
#include "engine/engine_util_errmem.h"
@@ -543,10 +542,10 @@ static int isknown(const mjtNum* points, const mjtNum x[3], int cnt) {
// adds candidate point to result
// flipNormal: 0 = normal points INTO SDF (for mesh-SDF where SDF is g2)
// 1 = normal points OUT of SDF (for flex-SDF where SDF is g1)
static int addContact(mjtNum* points, mjContact* con, const mjtNum x[3],
const mjtNum pos2[3], const mjtNum quat2[4], mjtNum dist,
int cnt, const mjModel* m, const mjSDF* s, const mjData* d,
int flipNormal) {
static int addPreContact(mjtNum* points, mjPreContact* con, const mjtNum x[3],
const mjtNum pos2[3], const mjtNum quat2[4], mjtNum dist,
int cnt, const mjModel* m, const mjSDF* s, const mjData* d,
int flipNormal) {
// check if there is a collision
if (dist > 0 || isknown(points, x, cnt)) {
return cnt;
@@ -570,18 +569,37 @@ static int addContact(mjtNum* points, mjContact* con, const mjtNum x[3],
}
// construct contact
con[cnt].dist = dist;
mju_rotVecQuat(con[cnt].frame, norm, quat2);
mju_zero3(con[cnt].frame+3);
mju_makeFrame(con[cnt].frame);
mju_scl3(vec, con[cnt].frame, -con[cnt].dist/2);
mju_rotVecQuat(con[cnt].pos, x, quat2);
mju_addTo3(con[cnt].pos, pos2);
mju_addTo3(con[cnt].pos, vec);
con->dist = dist;
mju_rotVecQuat(con->normal, norm, quat2);
mju_scl3(vec, con->normal, - 0.5 * dist);
mju_rotVecQuat(con->pos, x, quat2);
mju_zero3(con->tangent);
mju_addTo3(con->pos, pos2);
mju_addTo3(con->pos, vec);
return cnt+1;
}
// adds candidate point to result
// flipNormal: 0 = normal points INTO SDF (for mesh-SDF where SDF is g2)
// 1 = normal points OUT of SDF (for flex-SDF where SDF is g1)
static int addContact(mjtNum* points, mjContact* con, const mjtNum x[3],
const mjtNum pos2[3], const mjtNum quat2[4], mjtNum dist,
int cnt, const mjModel* m, const mjSDF* s, const mjData* d,
int flipNormal) {
mjPreContact precon;
int ncon = addPreContact(points, &precon, x, pos2, quat2, dist, cnt, m, s, d,
flipNormal);
if (ncon > cnt) {
con[cnt].dist = precon.dist;
mju_copy3(con[cnt].pos, precon.pos);
mju_copy3(con[cnt].frame, precon.normal);
mju_zero3(con[cnt].frame + 3);
}
return ncon;
}
// finds minimum of Frank-Wolfe objective
static mjtNum stepFrankWolfe(mjtNum x[3], const mjtNum* corners, int ncorners,
const mjModel* m, const mjSDF* sdf, const mjData* d) {
@@ -952,18 +970,21 @@ static int meshFaceCallback(int face_id, int node, void* ctx) {
//------------------------------ collision functions -----------------------------------------------
// collision between a height field and a signed distance field
int mjc_HFieldSDF(const mjModel* m, mjData* d, mjContact* con, int g1, int g2, mjtNum margin) {
int mjc_HFieldSDF(const mjModel* m, mjData* d, mjPreContact* con, int g1, int g2, mjtNum margin) {
mju_warning("HField vs SDF collision not yet supported!");
return 0;
}
// collision between a mesh and a signed distance field
int mjc_MeshSDF(const mjModel* m, mjData* d, mjContact* con, int g1, int g2, mjtNum margin) {
mjGETINFO;
int mjc_MeshSDF(const mjModel* m, mjData* d, mjPreContact* con, int g1, int g2, mjtNum margin) {
const mjtNum* pos1 = d->geom_xpos + 3*g1;
const mjtNum* mat1 = d->geom_xmat + 9*g1;
const mjtNum* pos2 = d->geom_xpos + 3*g2;
const mjtNum* mat2 = d->geom_xmat + 9*g2;
mjtNum offset[3], rotation[9];
mjtNum points[3*mjMAXCONPAIR], dist[mjMAXCONPAIR], candidate[3*mjMAXCONPAIR];
mjtNum points[3*mjMAXCONPAIR], dist2[mjMAXCONPAIR], candidate[3*mjMAXCONPAIR];
int cnt = 0, ncandidate = 0;
// get sdf plugin
@@ -996,7 +1017,7 @@ int mjc_MeshSDF(const mjModel* m, mjData* d, mjContact* con, int g1, int g2, mjt
ctx.faceadr = m->mesh_faceadr[m->geom_dataid[g1]];
ctx.nstartpts = mju_max(1, m->opt.sdf_initpoints);
ctx.candidate = candidate;
ctx.dist = dist;
ctx.dist = dist2;
ctx.ncandidate = &ncandidate;
// BVH traversal for mesh faces
@@ -1014,34 +1035,37 @@ int mjc_MeshSDF(const mjModel* m, mjData* d, mjContact* con, int g1, int g2, mjt
// if few candidates, add them all directly
if (ncandidate <= mjMAXCONPAIR) {
for (int i = 0; i < ncandidate; i++) {
cnt = addContact(points, con, candidate + 3*i, pos2, sdf_quat,
dist[i], cnt, m, &sdf, d, 0);
cnt = addPreContact(points, con + cnt, candidate + 3*i, pos2, sdf_quat,
dist2[i], cnt, m, &sdf, d, 0);
}
return cnt;
}
// farthest point sampling (FPS) for spatial coverage
int selected_indices[mjMAXCONPAIR];
int nselected = selectFPS(candidate, dist, ncandidate, selected_indices, mjMAXCONPAIR);
int nselected = selectFPS(candidate, dist2, ncandidate, selected_indices, mjMAXCONPAIR);
// add selected contacts
for (int i = 0; i < nselected; i++) {
int idx = selected_indices[i];
cnt = addContact(points, con, candidate + 3*idx, pos2, sdf_quat,
dist[idx], cnt, m, &sdf, d, 0);
cnt = addPreContact(points, con + cnt, candidate + 3*idx, pos2, sdf_quat,
dist2[idx], cnt, m, &sdf, d, 0);
}
return cnt;
}
// collision between two SDFs
int mjc_SDF(const mjModel* m, mjData* d, mjContact* con, int g1, int g2, mjtNum margin) {
mjGETINFO;
size1 = m->geom_aabb + 6*g1;
size2 = m->geom_aabb + 6*g2;
int mjc_SDF(const mjModel* m, mjData* d, mjPreContact* con, int g1, int g2, mjtNum margin) {
const mjtNum* pos1 = d->geom_xpos + 3*g1;
const mjtNum* mat1 = d->geom_xmat + 9*g1;
const mjtNum* pos2 = d->geom_xpos + 3*g2;
const mjtNum* mat2 = d->geom_xmat + 9*g2;
const mjtNum* size1 = m->geom_aabb + 6*g1;
const mjtNum* size2 = m->geom_aabb + 6*g2;
int cnt = 0;
mjtNum x[3], y[3], dist, vec1[3], vec2[3];
mjtNum x[3], y[3], dist2, vec1[3], vec2[3];
mjtNum aabb1[6] = {mjMAXVAL, mjMAXVAL, mjMAXVAL, -mjMAXVAL, -mjMAXVAL, -mjMAXVAL};
mjtNum aabb2[6] = {mjMAXVAL, mjMAXVAL, mjMAXVAL, -mjMAXVAL, -mjMAXVAL, -mjMAXVAL};
mjtNum aabb[6] = {mjMAXVAL, mjMAXVAL, mjMAXVAL, -mjMAXVAL, -mjMAXVAL, -mjMAXVAL};
@@ -1153,15 +1177,15 @@ int mjc_SDF(const mjModel* m, mjData* d, mjContact* con, int g1, int g2, mjtNum
// gradient descent - we use a special function of the two SDF as objective
sdf.type = mjSDFTYPE_COLLISION;
dist = stepGradient(x, m, &sdf, d, m->opt.sdf_iterations);
dist2 = stepGradient(x, m, &sdf, d, m->opt.sdf_iterations);
// inexact SDFs can yield spurious collisions, filter them by projecting on the midsurface
sdf.type = mjSDFTYPE_INTERSECTION;
dist = stepGradient(x, m, &sdf, d, 1);
dist2 = stepGradient(x, m, &sdf, d, 1);
// contact point and normal - we use the midsurface where SDF1=SDF2 as zero level set
sdf.type = mjSDFTYPE_MIDSURFACE;
cnt = addContact(contacts, con, x, pos2, quat2, dist, cnt, m, &sdf, d, 0);
cnt = addPreContact(contacts, con + cnt, x, pos2, quat2, dist2, cnt, m, &sdf, d, 0);
// SHOULD NOT OCCUR
if (cnt > mjMAXCONPAIR) {