From 490c1f4126c35a480d8f90ac2984c4746d76a835 Mon Sep 17 00:00:00 2001 From: Kyle Bayes Date: Tue, 3 Mar 2026 13:10:21 -0800 Subject: [PATCH] Refactor mj_collideGeoms. PiperOrigin-RevId: 878091736 Change-Id: I738c2e5a21036d859b7b163cb000dc516b63a878 --- src/engine/engine_collision_box.c | 79 +++++++++++++- src/engine/engine_collision_driver.c | 126 ++++------------------- src/engine/engine_collision_driver.h | 2 +- test/engine/engine_collision_box_test.cc | 4 +- 4 files changed, 98 insertions(+), 113 deletions(-) diff --git a/src/engine/engine_collision_box.c b/src/engine/engine_collision_box.c index b953f76f..5b10f821 100644 --- a/src/engine/engine_collision_box.c +++ b/src/engine/engine_collision_box.c @@ -17,6 +17,7 @@ #include "engine/engine_collision_primitive.h" #include "engine/engine_inline.h" #include "engine/engine_util_blas.h" +#include "engine/engine_util_misc.h" // hard-clamp vector to range [-limit(i), +limit(i)] static void mju_clampVec(mjtNum* vec, const mjtNum* limit, int n) @@ -600,8 +601,9 @@ int mjc_CapsuleBox(const mjModel* m, const mjData* d, mjContact* con, } -// box : box -int mjc_BoxBox(const mjModel* M, const mjData* D, mjContact* con, int g1, int g2, mjtNum margin) +// internal box : box +static inline +int _boxbox(const mjModel* M, const mjData* D, mjContact* con, int g1, int g2, mjtNum margin) { const mjtNum* pos1 = D->geom_xpos + 3 * g1; const mjtNum* mat1 = D->geom_xmat + 9 * g1; @@ -1337,3 +1339,76 @@ edgeedge: #undef rotaxis #undef rotmatx } + + +// box : box +int mjc_BoxBox(const mjModel* m, const mjData* d, mjContact* con, int g1, int g2, mjtNum margin) { + int num = _boxbox(m, d, con, g1, g2, margin); + + // use dim field to mark: -1: bad, 0: good + for (int i=0; i < num; i++) { + con[i].dim = 0; + } + + // get box info + const mjtNum* pos1 = d->geom_xpos + 3 * g1; + const mjtNum* mat1 = d->geom_xmat + 9 * g1; + const mjtNum* size1 = m->geom_size + 3 * g1; + const mjtNum* pos2 = d->geom_xpos + 3 * g2; + const mjtNum* mat2 = d->geom_xmat + 9 * g2; + const mjtNum* size2 = m->geom_size + 3 * g2; + + // find bad: contacts outside one of the boxes + for (int i=0; i < num; i++) { + // box sizes with margin + mjtNum sz1[3] = {size1[0] + margin, size1[1] + margin, size1[2] + margin}; + mjtNum sz2[3] = {size2[0] + margin, size2[1] + margin, size2[2] + margin}; + + // relative distance from surface (1%) outside of which box-box contacts are removed + static mjtNum kRemoveRatio = 1.01; + + // is the contact outside: 1, inside: -1, within the removal width: 0 + int out1 = mju_outsideBox(con[i].pos, pos1, mat1, sz1, kRemoveRatio); + int out2 = mju_outsideBox(con[i].pos, pos2, mat2, sz2, kRemoveRatio); + + // mark as bad if outside one box and not inside the other box + if ((out1 == 1 && out2 != -1) || (out2 == 1 && out1 != -1)) { + con[i].dim = -1; + } + } + + // find duplicates + for (int i=0; i < num-1; i++) { + if (con[i].dim == -1) { + continue; // already marked bad: skip + } + for (int j=i+1; j < num; j++) { + if (con[j].dim == -1) { + continue; // already marked bad: skip + } + if (con[i].pos[0] == con[j].pos[0] && + con[i].pos[1] == con[j].pos[1] && + con[i].pos[2] == con[j].pos[2]) { + con[i].dim = -1; + break; + } + } + } + + // consolidate good + int i = 0; + for (int j=0; j < num; j++) { + // good: maybe copy + if (con[j].dim == 0) { + // different: copy + if (i < j) { + con[i] = con[j]; + } + + // advance either way + i++; + } + } + + return i; +} diff --git a/src/engine/engine_collision_driver.c b/src/engine/engine_collision_driver.c index 27a8f516..cd8bb5e4 100644 --- a/src/engine/engine_collision_driver.c +++ b/src/engine/engine_collision_driver.c @@ -386,14 +386,12 @@ void mj_collision(const mjModel* m, mjData* d) { // merge predefined geom pairs int merged = 0; int startadr = pairadr; - if (npair) { - // test all predefined pairs for which pair_signature<=signature - while (pairadr < npair && m->pair_signature[pairadr] <= signature) { - if (m->pair_signature[pairadr] == signature) { - merged = 1; - } - mj_collideGeoms(m, d, pairadr++, -1); - } + // test all predefined pairs for which pair_signature <= signature + for (; pairadr < npair && m->pair_signature[pairadr] <= signature; pairadr++) { + merged = (m->pair_signature[pairadr] == signature); + int g1 = m->pair_geom1[pairadr]; + int g2 = m->pair_geom2[pairadr]; + mj_collideGeoms(m, d, pairadr, g1, g2); } // apply bitmask filtering at the bodyflex level @@ -520,10 +518,10 @@ void mj_collision(const mjModel* m, mjData* d) { } // finish merging predefined geom pairs - if (npair) { - while (pairadr < npair) { - mj_collideGeoms(m, d, pairadr++, -1); - } + for (; pairadr < npair; pairadr++) { + int g1 = m->pair_geom1[pairadr]; + int g2 = m->pair_geom2[pairadr]; + mj_collideGeoms(m, d, pairadr, g1, g2); } // flex self-collisions @@ -583,38 +581,27 @@ void mj_collision(const mjModel* m, mjData* d) { //------------------------------------ binary tree search ------------------------------------------ // collision tree node -struct mjCollisionTree_ { +typedef struct { int node1; int node2; -}; -typedef struct mjCollisionTree_ mjCollisionTree; +} mjCollisionTree; // checks if the proposed collision pair is already present in pair_geom and calls narrow phase void mj_collideGeomPair(const mjModel* m, mjData* d, int g1, int g2, int merged, int startadr, int pairadr) { - // merged: make sure geom pair is not repeated + // merged, find matching pair if (merged) { - // find matching pair - int found = 0; for (int k=startadr; k < pairadr; k++) { if ((m->pair_geom1[k] == g1 && m->pair_geom2[k] == g2) || (m->pair_geom1[k] == g2 && m->pair_geom2[k] == g1)) { - found = 1; - break; + return; } } - - // not found: test - if (!found) { - mj_collideGeoms(m, d, g1, g2); - } } - // not merged: always test - else { - mj_collideGeoms(m, d, g1, g2); - } + // not merged, always test + mj_collideGeoms(m, d, -1, g1, g2); } @@ -1558,19 +1545,13 @@ static void mj_makeCapsule(const mjModel* m, mjData* d, int f, const int vid[2], // test two geoms for collision, apply filters, add to contact list -void mj_collideGeoms(const mjModel* m, mjData* d, int g1, int g2) { +void mj_collideGeoms(const mjModel* m, mjData* d, int ipair, int g1, int g2) { int num, type1, type2, condim; mjtNum margin, gap, friction[5], solref[mjNREF], solimp[mjNIMP]; mjtNum solreffriction[mjNREF] = {0}; - int ipair = (g2 < 0 ? g1 : -1); - - // get explicit geom ids from pair + // sleep filtering for explicit pairs if (ipair >= 0) { - g1 = m->pair_geom1[ipair]; - g2 = m->pair_geom2[ipair]; - - // sleep filtering for explicit pairs if (mjENABLED(mjENBL_SLEEP)) { int b1 = m->geom_bodyid[g1]; int b2 = m->geom_bodyid[g2]; @@ -1648,77 +1629,6 @@ void mj_collideGeoms(const mjModel* m, mjData* d, int g1, int g2) { mjERROR("too many contacts returned by collision function"); } - // remove bad and repeated contacts in box-box - if (collisionFunc == mjc_BoxBox) { - // use dim field to mark: -1: bad, 0: good - for (int i=0; i < num; i++) { - con[i].dim = 0; - } - - // get box info - const mjtNum* pos1 = d->geom_xpos + 3 * g1; - const mjtNum* mat1 = d->geom_xmat + 9 * g1; - const mjtNum* size1 = m->geom_size + 3 * g1; - const mjtNum* pos2 = d->geom_xpos + 3 * g2; - const mjtNum* mat2 = d->geom_xmat + 9 * g2; - const mjtNum* size2 = m->geom_size + 3 * g2; - - // find bad: contacts outside one of the boxes - for (int i=0; i < num; i++) { - // box sizes with margin - mjtNum sz1[3] = {size1[0] + margin, size1[1] + margin, size1[2] + margin}; - mjtNum sz2[3] = {size2[0] + margin, size2[1] + margin, size2[2] + margin}; - - // relative distance from surface (1%) outside of which box-box contacts are removed - static mjtNum kRemoveRatio = 1.01; - - // is the contact outside: 1, inside: -1, within the removal width: 0 - int out1 = mju_outsideBox(con[i].pos, pos1, mat1, sz1, kRemoveRatio); - int out2 = mju_outsideBox(con[i].pos, pos2, mat2, sz2, kRemoveRatio); - - // mark as bad if outside one box and not inside the other box - if ((out1 == 1 && out2 != -1) || (out2 == 1 && out1 != -1)) { - con[i].dim = -1; - } - } - - // find duplicates - for (int i=0; i < num-1; i++) { - if (con[i].dim == -1) { - continue; // already marked bad: skip - } - for (int j=i+1; j < num; j++) { - if (con[j].dim == -1) { - continue; // already marked bad: skip - } - if (con[i].pos[0] == con[j].pos[0] && - con[i].pos[1] == con[j].pos[1] && - con[i].pos[2] == con[j].pos[2]) { - con[i].dim = -1; - break; - } - } - } - - // consolidate good - int i = 0; - for (int j=0; j < num; j++) { - // good: maybe copy - if (con[j].dim == 0) { - // different: copy - if (i < j) { - con[i] = con[j]; - } - - // advance either way - i++; - } - } - - // adjust size - num = i; - } - // set condim, gap, solref, solimp, friction: dynamic if (ipair < 0) { mj_contactParam(m, &condim, &gap, solref, solimp, friction, g1, g2, -1, -1); diff --git a/src/engine/engine_collision_driver.h b/src/engine/engine_collision_driver.h index 4667d91d..9583cf5f 100644 --- a/src/engine/engine_collision_driver.h +++ b/src/engine/engine_collision_driver.h @@ -50,7 +50,7 @@ void mj_collideTree(const mjModel* m, mjData* d, int bf1, int bf2, int mj_broadphase(const mjModel* m, mjData* d, int* bfpair, int maxpair); // test two geoms for collision, apply filters, add to contact list -void mj_collideGeoms(const mjModel* m, mjData* d, int g1, int g2); +void mj_collideGeoms(const mjModel* m, mjData* d, int ipair, 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); diff --git a/test/engine/engine_collision_box_test.cc b/test/engine/engine_collision_box_test.cc index ad23f936..e73f4ad9 100644 --- a/test/engine/engine_collision_box_test.cc +++ b/test/engine/engine_collision_box_test.cc @@ -95,7 +95,7 @@ TEST_F(MjCollisionBoxTest, BadContacts) { } // expect some contacts to have been removed - EXPECT_LT(nmatched, num) << local_path; + EXPECT_EQ(nmatched, num) << local_path; // get box info const mjtNum* pos1 = data->geom_xpos + 3 * g1; @@ -197,7 +197,7 @@ TEST_F(MjCollisionBoxTest, DuplicateContacts) { } // expect some contacts to have been removed - EXPECT_LT(nmatched, num); + EXPECT_EQ(nmatched, num); // loop over raw contacts, find removed for (int i = 0; i < num; i++) {