From bf74d01d932ea3f424a6604a16c00823177f44f3 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Fri, 27 Feb 2026 14:34:25 -0800 Subject: [PATCH] Add collisions between flexes and SDF geoms. The flex-SDF algorithm uses the same logic as mesh-SDF. Previous to this change, the flex was instead colliding with the convex hull of the mesh generated automatically with marching cubes from the SDF. PiperOrigin-RevId: 876420302 Change-Id: I3bf6552b6b8e108ba4c4209bbcf10fe5f2b2fd90 --- doc/changelog.rst | 5 + model/flex/cloth_sdf.xml | 51 +++ src/engine/engine_collision_driver.c | 70 ++- src/engine/engine_collision_driver.h | 3 + src/engine/engine_collision_sdf.c | 517 +++++++++++++++++------ src/engine/engine_collision_sdf.h | 4 + test/engine/engine_collision_sdf_test.cc | 66 +++ 7 files changed, 588 insertions(+), 128 deletions(-) create mode 100644 model/flex/cloth_sdf.xml diff --git a/doc/changelog.rst b/doc/changelog.rst index 83d80fea..20f297ed 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -5,6 +5,11 @@ Changelog Upcoming version (not yet released) ----------------------------------- +General +^^^^^^^ + +- Flexes now support collisions with SDF geoms. + MJX ^^^ diff --git a/model/flex/cloth_sdf.xml b/model/flex/cloth_sdf.xml new file mode 100644 index 00000000..96697257 --- /dev/null +++ b/model/flex/cloth_sdf.xml @@ -0,0 +1,51 @@ + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + diff --git a/src/engine/engine_collision_driver.c b/src/engine/engine_collision_driver.c index ceaac7ef..27a8f516 100644 --- a/src/engine/engine_collision_driver.c +++ b/src/engine/engine_collision_driver.c @@ -484,6 +484,14 @@ void mj_collision(const mjModel* m, mjData* d) { continue; } + // SDF special processing + if (m->geom_type[g] == mjGEOM_SDF) { + int ncon_before = d->ncon; + mj_collideSdfFlex(m, d, g, f); + filterFlexContacts(d, ncon_before); + continue; + } + // collide geom with flex elements int ncon_before = d->ncon; int elemnum = m->flex_elemnum[f]; @@ -754,6 +762,15 @@ void mj_collideTree(const mjModel* m, mjData* d, int bf1, int bf2, } } + // for body:flex, if body has SDFs, call mj_collideSdfFlex directly + if (isbody1 && !isbody2) { + for (int i=m->body_geomadr[bf1]; i < m->body_geomadr[bf1]+m->body_geomnum[bf1]; i++) { + if (m->geom_type[i] == mjGEOM_SDF) { + mj_collideSdfFlex(m, d, i, f2); + } + } + } + // collide trees while (nstack) { // pop from stack @@ -820,8 +837,9 @@ void mj_collideTree(const mjModel* m, mjData* d, int bf1, int bf2, d->geom_xpos + 3*nodeid1, d->geom_xmat + 9*nodeid1, NULL, NULL, margin, NULL, NULL, &initialize)) { - // collide unless geom is plane (plane:flex handled separately) - if (m->geom_type[nodeid1] != mjGEOM_PLANE) { + // collide unless geom is plane or SDF (handled separately) + if (m->geom_type[nodeid1] != mjGEOM_PLANE && + m->geom_type[nodeid1] != mjGEOM_SDF) { mj_collideGeomElem(m, d, nodeid1, f2, nodeid2); } if (mark_active) { @@ -1800,6 +1818,54 @@ void mj_collidePlaneFlex(const mjModel* m, mjData* d, int g, int f) { } +// test an SDF geom and a flex for collision, add to contact list +void mj_collideSdfFlex(const mjModel* m, mjData* d, int g, int f) { + // only support dim==2 (triangular elements) + if (m->flex_dim[f] != 2) { + return; + } + + // prepare contact parameters (same for all contacts) + mjtNum margin = mj_assignMargin(m, m->geom_margin[g] + m->flex_margin[f]); + int condim; + mjtNum gap, solref[mjNREF], solimp[mjNIMP], friction[5]; + mjtNum solreffriction[mjNREF] = {0}; + mj_contactParam(m, &condim, &gap, solref, solimp, friction, g, -1, -1, f); + + // allocate temporary contact array on stack (zero-initialized) + mj_markStack(d); + mjContact* con = mjSTACKALLOC(d, mjMAXCONPAIR, mjContact); + memset(con, 0, mjMAXCONPAIR * sizeof(mjContact)); + + // call batched flex-SDF collision + int num = mjc_FlexSDF(m, d, con, g, f, margin); + + // add contacts to mjData + for (int i = 0; i < num; i++) { + // set contact ids + con[i].geom[0] = g; + con[i].geom[1] = -1; + con[i].flex[0] = -1; + con[i].flex[1] = f; + con[i].elem[0] = -1; + // con[i].elem[1] already set by mjc_FlexSDF with actual element index + con[i].vert[0] = -1; + con[i].vert[1] = -1; + + // set remaining contact parameters + mj_setContact(m, con + i, condim, margin-gap, solref, solreffriction, solimp, friction); + + // add to mjData, abort if too many contacts + if (mj_addContact(m, d, con + i)) { + mj_freeStack(d); + return; + } + } + + mj_freeStack(d); +} + + // test single triangle plane : vertex static int planeVertex(mjContact* con, const mjtNum* pos, mjtNum rad, int t0, int t1, int t2, int v) { diff --git a/src/engine/engine_collision_driver.h b/src/engine/engine_collision_driver.h index a27f84c7..4667d91d 100644 --- a/src/engine/engine_collision_driver.h +++ b/src/engine/engine_collision_driver.h @@ -55,6 +55,9 @@ void mj_collideGeoms(const mjModel* m, mjData* d, int g1, int g2); // test a plane geom and a flex for collision, add to contact list void mj_collidePlaneFlex(const mjModel* m, mjData* d, int g, int f); +// test an SDF geom and a flex for collision, add to contact list +void mj_collideSdfFlex(const mjModel* m, mjData* d, int g, int f); + // test for internal flex collisions, add to contact list void mj_collideFlexInternal(const mjModel* m, mjData* d, int f); diff --git a/src/engine/engine_collision_sdf.c b/src/engine/engine_collision_sdf.c index a49acecc..543c93d4 100644 --- a/src/engine/engine_collision_sdf.c +++ b/src/engine/engine_collision_sdf.c @@ -22,9 +22,7 @@ #include // IWYU pragma: keep #include #include "engine/engine_collision_primitive.h" -#include "engine/engine_memory.h" #include "engine/engine_plugin.h" -#include "engine/engine_sort.h" #include "engine/engine_util_blas.h" #include "engine/engine_util_errmem.h" #include "engine/engine_util_misc.h" @@ -526,9 +524,12 @@ 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, mjData* d) { + int cnt, const mjModel* m, const mjSDF* s, mjData* d, + int flipNormal) { // check if there is a collision if (dist > 0 || isknown(points, x, cnt)) { return cnt; @@ -539,7 +540,17 @@ static int addContact(mjtNum* points, mjContact* con, const mjtNum x[3], // compute normal in local coordinates mjtNum norm[3], vec[3]; mjc_gradient(m, d, s, norm, x); - mju_scl3(norm, norm, -1); + + // validate normal - skip if gradient is degenerate (zero or near-zero) + mjtNum norm_len = mju_normalize3(norm); + if (norm_len < mjMINVAL) { + return cnt; // degenerate gradient, skip this contact + } + + // normal direction: flipNormal=0 -> INTO SDF, flipNormal=1 -> OUT of SDF + if (!flipNormal) { + mju_scl3(norm, norm, -1); + } // construct contact con[cnt].dist = dist; @@ -720,6 +731,106 @@ static int boxIntersect(const mjtNum bvh[6], const mjtNum offset[3], return mjc_distance(m, d, s, candidate) < r; } +// Farthest point sampling (FPS) to select spatially diverse contacts +// Returns indices of selected candidates in nselected, fills selected_indices array +static int selectFPS(const mjtNum* candidate, const mjtNum* dist, int ncandidate, + int* selected_indices, int max_select) { + if (ncandidate <= 0) return 0; + + mjtByte selected[mjMAXCONPAIR] = {0}; + mjtNum min_dist2[mjMAXCONPAIR]; + for (int i = 0; i < ncandidate; i++) { + min_dist2[i] = mjMAXVAL; + } + + // start with deepest penetrating contact + int best = 0; + mjtNum bestval = -dist[0]; + for (int i = 1; i < ncandidate; i++) { + if (-dist[i] > bestval) { + bestval = -dist[i]; + best = i; + } + } + + // iteratively select contacts using FPS + int nselected = 0; + while (nselected < max_select && nselected < mjMAXCONPAIR && best >= 0) { + selected[best] = 1; + selected_indices[nselected] = best; + nselected++; + + const mjtNum* bestpos = candidate + 3*best; + + // find next farthest point + int nextbest = -1; + mjtNum nextbestdist = -1; + for (int i = 0; i < ncandidate; i++) { + if (selected[i]) continue; + + mjtNum dx = candidate[3*i+0] - bestpos[0]; + mjtNum dy = candidate[3*i+1] - bestpos[1]; + mjtNum dz = candidate[3*i+2] - bestpos[2]; + mjtNum d2 = dx*dx + dy*dy + dz*dz; + if (d2 < min_dist2[i]) { + min_dist2[i] = d2; + } + if (min_dist2[i] > nextbestdist) { + nextbestdist = min_dist2[i]; + nextbest = i; + } + } + best = nextbest; + } + + return nselected; +} + +// Process triangle corners against SDF using Halton sampling + Frank-Wolfe. +// Corners are assumed to already be in SDF local coordinates. +// Adds penetrating candidates to candidate/dist arrays. +static void processSdfCorners(const mjtNum corners[9], const mjModel* m, mjData* d, + const mjSDF* sdf, int nstartpts, + mjtNum* candidate, mjtNum* dist, int* ncandidate) { + mjtNum x[3], depth; + + // stricter culling using triangle circumsphere + if (!triangleIntersect(corners, m, sdf, d)) { + return; + } + + // sample multiple starting points using Halton sequence + for (int sp = 0; sp < nstartpts; sp++) { + if (*ncandidate >= mjMAXCONPAIR) break; + + // barycentric coordinates from Halton sequence + mjtNum u = mju_Halton(sp + 1, 2); + mjtNum v = mju_Halton(sp + 1, 3); + if (u + v > 1) { + u = 1 - u; + v = 1 - v; + } + mjtNum b0 = 1 - u - v; + mjtNum b1 = u; + mjtNum b2 = v; + + // starting point + x[0] = b0*corners[0] + b1*corners[3] + b2*corners[6]; + x[1] = b0*corners[1] + b1*corners[4] + b2*corners[7]; + x[2] = b0*corners[2] + b1*corners[5] + b2*corners[8]; + + depth = stepFrankWolfe(x, corners, 3, m, sdf, d); + + // store candidate if penetration + if (depth < 0) { + int nc = *ncandidate; + mju_copy3(candidate + 3*nc, x); + dist[nc] = depth; + (*ncandidate)++; + } + } +} + // Context for inline face processing during BVH traversal typedef struct { const mjModel* m; @@ -738,9 +849,8 @@ typedef struct { // process a single mesh face inline during BVH traversal static void processOneFace(int faceid, mjtByte* bvh_active, int node, MeshSDFContext* ctx) { - mjtNum corners[9], x[3], depth; + mjtNum corners[9]; const mjModel* m = ctx->m; - mjData* d = ctx->d; int face = ctx->faceadr + faceid; // transform triangle vertices to SDF local coordinates @@ -754,85 +864,48 @@ static void processOneFace(int faceid, mjtByte* bvh_active, int node, mju_addTo3(corners + 3*v, ctx->offset); } - // stricter culling using triangle circumsphere - if (!triangleIntersect(corners, m, ctx->sdf, d)) { - return; - } + // use shared helper for Halton sampling + Frank-Wolfe + processSdfCorners(corners, m, ctx->d, ctx->sdf, ctx->nstartpts, + ctx->candidate, ctx->dist, ctx->ncandidate); - if (bvh_active) bvh_active[node] = 1; - - // sample multiple starting points using Halton sequence - for (int sp = 0; sp < ctx->nstartpts; sp++) { - if (*(ctx->ncandidate) >= mjMAXCONPAIR) break; - - // barycentric coordinates from Halton sequence - mjtNum u = mju_Halton(sp + 1, 2); - mjtNum v = mju_Halton(sp + 1, 3); - if (u + v > 1) { - u = 1 - u; - v = 1 - v; - } - mjtNum b0 = 1 - u - v; - mjtNum b1 = u; - mjtNum b2 = v; - - // starting point - x[0] = b0*corners[0] + b1*corners[3] + b2*corners[6]; - x[1] = b0*corners[1] + b1*corners[4] + b2*corners[7]; - x[2] = b0*corners[2] + b1*corners[5] + b2*corners[8]; - - depth = stepFrankWolfe(x, corners, 3, m, ctx->sdf, d); - - // store candidate if penetration - if (depth < 0) { - int nc = *(ctx->ncandidate); - mju_copy3(ctx->candidate + 3*nc, x); - ctx->dist[nc] = depth; - (*(ctx->ncandidate))++; - } - } + if (bvh_active && *(ctx->ncandidate) > 0) bvh_active[node] = 1; } -// tree vs sdf binary search with inline face processing -static void collideBVHInline(const mjModel* m, mjData* d, int g, - MeshSDFContext* ctx) { - const int bvhadr = m->mesh_bvhadr[m->geom_dataid[g]]; +// Callback type for leaf node processing during BVH traversal +// Returns 1 if node should be marked active (for visualization), 0 otherwise +typedef int (*BVHLeafCallback)(int leaf_id, int node, void* ctx); - // mesh has no BVH tree - if (bvhadr < 0) { - return; - } - - const int* faceid = m->bvh_nodeid + bvhadr; - const mjtNum* bvh = m->bvh_aabb + 6*bvhadr; - const int* child = m->bvh_child + 2*bvhadr; - mjtByte* bvh_active = m->vis.global.bvactive ? d->bvh_active + bvhadr : NULL; - - mj_markStack(d); - int max_stack = m->mesh_bvhnum[m->geom_dataid[g]]; - struct CollideTreeArgs_ { - int node; - }; - typedef struct CollideTreeArgs_ CollideTreeArgs; - CollideTreeArgs* stack = mjSTACKALLOC(d, max_stack, CollideTreeArgs); +// Generic BVH traversal with callback for leaf processing +// bvh: pointer to BVH AABBs (6 floats per node: center xyz, half-size xyz) +// nodeid: leaf ID per node (-1 for intermediate nodes) +// child: child indices (2 per node) +// bvh_active: visualization array or NULL +// offset/rotation: transform for boxIntersect +// m, d, sdf: for boxIntersect +// callback: function called for each leaf node +// ctx: user context passed to callback +static void traverseBVH(const mjtNum* bvh, const int* nodeid, const int* child, + mjtByte* bvh_active, const mjtNum* offset, const mjtNum* rotation, + const mjModel* m, mjData* d, const mjSDF* sdf, + BVHLeafCallback callback, void* ctx) { + int stack[64]; int nstack = 0; - stack[nstack].node = 0; - nstack++; + stack[nstack++] = 0; while (nstack) { - nstack--; - int node = stack[nstack].node; + int node = stack[--nstack]; - // leaf node: process face inline - if (faceid[node] != -1) { - if (boxIntersect(bvh + 6*node, ctx->offset, ctx->rotation, m, ctx->sdf, d)) { - processOneFace(faceid[node], bvh_active, node, ctx); + // leaf node: call callback + if (nodeid[node] != -1) { + if (boxIntersect(bvh + 6*node, offset, rotation, m, sdf, d)) { + int active = callback(nodeid[node], node, ctx); + if (bvh_active && active) bvh_active[node] = 1; } continue; } // intermediate node: check bounding box - if (!boxIntersect(bvh + 6*node, ctx->offset, ctx->rotation, m, ctx->sdf, d)) { + if (!boxIntersect(bvh + 6*node, offset, rotation, m, sdf, d)) { continue; } @@ -841,16 +914,24 @@ static void collideBVHInline(const mjModel* m, mjData* d, int g, // push children for (int i = 0; i < 2; i++) { if (child[2*node+i] != -1) { - if (nstack >= max_stack) mjERROR("BVH stack depth exceeded."); - stack[nstack].node = child[2*node+i]; - nstack++; + if (nstack >= 64) { + mjERROR("BVH stack depth exceeded."); + } + stack[nstack++] = child[2*node+i]; } } } - - mj_freeStack(d); } +// mesh face callback for traverseBVH +static int meshFaceCallback(int face_id, int node, void* ctx) { + MeshSDFContext* mctx = (MeshSDFContext*)ctx; + int ncandidate_before = *(mctx->ncandidate); + processOneFace(face_id, NULL, node, mctx); + return *(mctx->ncandidate) > ncandidate_before; // mark active if candidates added +} + + //------------------------------ collision functions ----------------------------------------------- // collision between a height field and a signed distance field @@ -901,65 +982,36 @@ int mjc_MeshSDF(const mjModel* m, const mjData* d, mjContact* con, int g1, int g ctx.dist = dist; ctx.ncandidate = &ncandidate; - // binary tree search with inline face processing - collideBVHInline(m, (mjData*)d, g1, &ctx); + // BVH traversal for mesh faces + const int bvhadr = m->mesh_bvhadr[m->geom_dataid[g1]]; + if (bvhadr >= 0) { + const int* nodeid = m->bvh_nodeid + bvhadr; + const mjtNum* bvh = m->bvh_aabb + 6*bvhadr; + const int* child = m->bvh_child + 2*bvhadr; + mjtByte* bvh_active = m->vis.global.bvactive ? d->bvh_active + bvhadr : NULL; + + traverseBVH(bvh, nodeid, child, bvh_active, ctx.offset, ctx.rotation, + m, (mjData*)d, ctx.sdf, meshFaceCallback, &ctx); + } // 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, (mjData*)d); + dist[i], cnt, m, &sdf, (mjData*)d, 0); } return cnt; } // farthest point sampling (FPS) for spatial coverage - mjtByte selected[mjMAXCONPAIR] = {0}; - mjtNum min_dist2[mjMAXCONPAIR]; - for (int i = 0; i < ncandidate; i++) { - min_dist2[i] = mjMAXVAL; - } + int selected_indices[mjMAXCONPAIR]; + int nselected = selectFPS(candidate, dist, ncandidate, selected_indices, mjMAXCONPAIR); - // start with deepest penetrating contact - int best = 0; - mjtNum bestval = -dist[0]; - for (int i = 1; i < ncandidate; i++) { - if (-dist[i] > bestval) { - bestval = -dist[i]; - best = i; - } - } - - // iteratively select contacts using FPS - int nselected = 0; - while (nselected < mjMAXCONPAIR && best >= 0) { - selected[best] = 1; - mjtNum* bestpos = candidate + 3*best; - - // add the selected contact - cnt = addContact(points, con, bestpos, pos2, sdf_quat, - dist[best], cnt, m, &sdf, (mjData*)d); - nselected++; - - // find next farthest point - int nextbest = -1; - mjtNum nextbestdist = -1; - for (int i = 0; i < ncandidate; i++) { - if (selected[i]) continue; - - mjtNum dx = candidate[3*i+0] - bestpos[0]; - mjtNum dy = candidate[3*i+1] - bestpos[1]; - mjtNum dz = candidate[3*i+2] - bestpos[2]; - mjtNum d2 = dx*dx + dy*dy + dz*dz; - if (d2 < min_dist2[i]) { - min_dist2[i] = d2; - } - if (min_dist2[i] > nextbestdist) { - nextbestdist = min_dist2[i]; - nextbest = i; - } - } - best = nextbest; + // 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, (mjData*)d, 0); } return cnt; @@ -1092,7 +1144,7 @@ int mjc_SDF(const mjModel* m, const mjData* d, mjContact* con, int g1, int g2, m // 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, (mjData*)d); + cnt = addContact(contacts, con, x, pos2, quat2, dist, cnt, m, &sdf, (mjData*)d, 0); // SHOULD NOT OCCUR if (cnt > mjMAXCONPAIR) { @@ -1102,3 +1154,216 @@ int mjc_SDF(const mjModel* m, const mjData* d, mjContact* con, int g1, int g2, m return cnt; } + +// Context for flex element processing during BVH traversal +typedef struct { + const mjModel* m; + mjData* d; + const mjSDF* sdf; + const mjtNum* offset; + const mjtNum* rotation; + const mjtNum* vertxpos; + int f; + int dim; + int nstartpts; + mjtNum* candidate; + mjtNum* dist; + int* elem_id; + int* ncandidate; +} FlexSDFContext; + +// flex element callback for traverseBVH +static int flexElemCallback(int elem_idx, int node, void* ctx) { + FlexSDFContext* fctx = (FlexSDFContext*)ctx; + const mjModel* m = fctx->m; + int dim = fctx->dim; + + // get element vertex indices + const int* edata = m->flex_elem + m->flex_elemdataadr[fctx->f] + elem_idx*(dim+1); + + // get triangle corners in world coordinates + mjtNum world_corners[9]; + for (int v = 0; v < 3; v++) { + mju_copy3(world_corners + 3*v, fctx->vertxpos + 3*edata[v]); + } + + // transform corners to SDF local coordinates + mjtNum corners[9]; + for (int v = 0; v < 3; v++) { + mju_mulMatVec3(corners + 3*v, fctx->rotation, world_corners + 3*v); + mju_addTo3(corners + 3*v, fctx->offset); + } + + // process triangle (per-element candidates) + mjtNum elem_candidate[3*mjMAXCONPAIR]; + mjtNum elem_dist[mjMAXCONPAIR]; + int elem_ncandidate = 0; + processSdfCorners(corners, m, fctx->d, fctx->sdf, fctx->nstartpts, + elem_candidate, elem_dist, &elem_ncandidate); + + // keep only deepest penetrating candidate from this element + if (elem_ncandidate > 0) { + int best_i = 0; + for (int i = 1; i < elem_ncandidate; i++) { + if (elem_dist[i] < elem_dist[best_i]) { + best_i = i; + } + } + // add to global candidate list + int nc = *(fctx->ncandidate); + if (nc < mjMAXCONPAIR) { + mju_copy3(fctx->candidate + 3*nc, elem_candidate + 3*best_i); + fctx->dist[nc] = elem_dist[best_i]; + fctx->elem_id[nc] = elem_idx; + (*(fctx->ncandidate))++; + return 1; // mark active + } + } + return 0; +} + + +int mjc_FlexSDF(const mjModel* m, mjData* d, mjContact* con, + int g, int f, mjtNum margin) { + // g = SDF geom, f = flex + int dim = m->flex_dim[f]; + + // only support dim==2 (triangular elements) + if (dim != 2) { + return 0; + } + + mjtNum points[3*mjMAXCONPAIR], dist[mjMAXCONPAIR], candidate[3*mjMAXCONPAIR]; + int elem_id[mjMAXCONPAIR]; // track which element generated each candidate + int cnt = 0, ncandidate = 0; + + // get SDF info (once for entire flex) + int instance = m->geom_plugin[g]; + const mjpPlugin* sdf_ptr = instance == -1 ? NULL : mjc_getSDF(m, g); + instance = instance == -1 ? m->geom_dataid[g] : instance; + mjtGeom geomtype = mjGEOM_SDF; + + mjSDF sdf; + sdf.id = &instance; + sdf.type = mjSDFTYPE_SINGLE; + sdf.plugin = &sdf_ptr; + sdf.geomtype = &geomtype; + + // get SDF pose (once for entire flex) + mjtNum sdf_quat[4]; + const mjtNum* sdf_pos = d->geom_xpos + 3*g; + const mjtNum* sdf_mat = d->geom_xmat + 9*g; + mju_mat2Quat(sdf_quat, sdf_mat); + + // compute world-to-SDF transform (once for entire flex) + mjtNum rotation[9], offset[3]; + mjtNum world_origin[3] = {0, 0, 0}, world_quat[4] = {1, 0, 0, 0}; + mapPose(world_origin, world_quat, sdf_pos, sdf_quat, offset, rotation); + + // get flex element and vertex data + const mjtNum* vertxpos = d->flexvert_xpos + 3*m->flex_vertadr[f]; + int nstartpts = mju_max(1, m->opt.sdf_initpoints); + int elemnum = m->flex_elemnum[f]; + + // flex has no BVH tree: iterate all elements directly + const int bvhadr = m->flex_bvhadr[f]; + if (bvhadr < 0) { + mjtNum elem_candidate[3*mjMAXCONPAIR]; + mjtNum elem_dist[mjMAXCONPAIR]; + for (int e = 0; e < elemnum; e++) { + const int* edata = m->flex_elem + m->flex_elemdataadr[f] + e*(dim+1); + + // get triangle corners in world coordinates + mjtNum world_corners[9]; + for (int v = 0; v < 3; v++) { + mju_copy3(world_corners + 3*v, vertxpos + 3*edata[v]); + } + + // transform corners to SDF local coordinates + mjtNum corners[9]; + for (int v = 0; v < 3; v++) { + mju_mulMatVec3(corners + 3*v, rotation, world_corners + 3*v); + mju_addTo3(corners + 3*v, offset); + } + + // process triangle using shared helper (per-element candidates) + int elem_ncandidate = 0; + processSdfCorners(corners, m, d, &sdf, nstartpts, elem_candidate, elem_dist, &elem_ncandidate); + + // add contacts for this element immediately (per-element points tracking) + mjtNum elem_points[3*mjMAXCONPAIR]; + int elem_cnt = 0; + for (int i = 0; i < elem_ncandidate && cnt < mjMAXCONPAIR; i++) { + int old_elem_cnt = elem_cnt; + elem_cnt = addContact(elem_points, con + cnt, elem_candidate + 3*i, sdf_pos, sdf_quat, + elem_dist[i], elem_cnt, m, &sdf, d, 1); + // set element ID for successfully added contact + if (elem_cnt > old_elem_cnt) { + con[cnt].elem[1] = e; + cnt++; + } + } + } + return cnt; + } + + // have BVH: use tree traversal for efficiency + const int* elemid = m->bvh_nodeid + bvhadr; + // flexes use dynamic AABBs (vertices move during simulation) + const int nbvhstatic = m->nbvhstatic; + const mjtNum* bvh = d->bvh_aabb_dyn + 6*(bvhadr - nbvhstatic); + const int* child = m->bvh_child + 2*bvhadr; + mjtByte* bvh_active = m->vis.global.bvactive ? d->bvh_active + bvhadr : NULL; + + // set up context for flex element processing + FlexSDFContext fctx; + fctx.m = m; + fctx.d = d; + fctx.sdf = &sdf; + fctx.offset = offset; + fctx.rotation = rotation; + fctx.vertxpos = vertxpos; + fctx.f = f; + fctx.dim = dim; + fctx.nstartpts = nstartpts; + fctx.candidate = candidate; + fctx.dist = dist; + fctx.elem_id = elem_id; + fctx.ncandidate = &ncandidate; + + traverseBVH(bvh, elemid, child, bvh_active, offset, rotation, + m, d, &sdf, flexElemCallback, &fctx); + + + // if few candidates, add them all directly + if (ncandidate <= mjMAXCONPAIR) { + for (int i = 0; i < ncandidate; i++) { + int old_cnt = cnt; + cnt = addContact(points, con, candidate + 3*i, sdf_pos, sdf_quat, + dist[i], cnt, m, &sdf, d, 1); + // set element ID for successfully added contact + if (cnt > old_cnt) { + con[old_cnt].elem[1] = elem_id[i]; + } + } + return cnt; + } + + // farthest point sampling (FPS) for spatial coverage + int selected_indices[mjMAXCONPAIR]; + int nselected = selectFPS(candidate, dist, ncandidate, selected_indices, mjMAXCONPAIR); + + // add selected contacts + for (int i = 0; i < nselected; i++) { + int idx = selected_indices[i]; + int old_cnt = cnt; + cnt = addContact(points, con, candidate + 3*idx, sdf_pos, sdf_quat, + dist[idx], cnt, m, &sdf, d, 1); + // set element ID for successfully added contact + if (cnt > old_cnt) { + con[old_cnt].elem[1] = elem_id[idx]; + } + } + + return cnt; +} diff --git a/src/engine/engine_collision_sdf.h b/src/engine/engine_collision_sdf.h index c0ea7cdb..cf572450 100644 --- a/src/engine/engine_collision_sdf.h +++ b/src/engine/engine_collision_sdf.h @@ -44,6 +44,10 @@ int mjc_MeshSDF(const mjModel* m, const mjData* d, mjContact* con, int g1, int g // collision between two signed distance fields int mjc_SDF(const mjModel* m, const mjData* d, mjContact* con, int g1, int g2, mjtNum margin); +// collision between entire flex and SDF geom (batched processing) +int mjc_FlexSDF(const mjModel* m, mjData* d, mjContact* con, + int g, int f, mjtNum margin); + #ifdef __cplusplus } #endif diff --git a/test/engine/engine_collision_sdf_test.cc b/test/engine/engine_collision_sdf_test.cc index f42dc09f..a8f93c32 100644 --- a/test/engine/engine_collision_sdf_test.cc +++ b/test/engine/engine_collision_sdf_test.cc @@ -81,5 +81,71 @@ TEST_F(SdfTest, SdfPrimitive) { mj_deleteModel(model); } +static constexpr char kFlexSdfModel[] = R"( + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +)"; + +TEST_F(SdfTest, FlexSdfCollision) { + char error[1024]; + mjModel* model = LoadModelFromString(kFlexSdfModel, error, sizeof(error)); + ASSERT_THAT(model, NotNull()) << error; + mjData* data = mj_makeData(model); + ASSERT_THAT(data, NotNull()); + + // check there is at least one flex + ASSERT_GT(model->nflex, 0); + + // simulate for a few steps to let flex fall onto torus + for (int i = 0; i < 100; i++) { + mj_step(model, data); + } + + // should have contacts between flex and SDF + EXPECT_GT(data->ncon, 0) << "Expected flex-SDF contacts"; + + // check that contacts involve the flex + bool has_flex_contact = false; + for (int i = 0; i < data->ncon; i++) { + if (data->contact[i].flex[0] >= 0 || data->contact[i].flex[1] >= 0) { + has_flex_contact = true; + break; + } + } + EXPECT_TRUE(has_flex_contact) << "Expected at least one flex contact"; + + mj_deleteData(data); + mj_deleteModel(model); +} + } // namespace } // namespace mujoco +