From 3aa973894ce646e2795c634aace1f81d1f7ea08f Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Wed, 19 Nov 2025 03:39:15 -0800 Subject: [PATCH] Fast flex: Interpolate contact points directly. For each contact point, we used to first compute the vertex weights on a triangle and then add a contact point per vertex using the flex interpolation (trilinear or quadratic), obtaining the total weight by multiplying the vertex weight and the basis function value at that vertex. After this change, each contact point is added directly by evaluating the basis function directly a the point. PiperOrigin-RevId: 834218909 Change-Id: Ic0fdb03fc0ef5478798c293df99204fd33116fd1 --- src/engine/engine_core_constraint.c | 163 ++++++++++++++++------------ 1 file changed, 95 insertions(+), 68 deletions(-) diff --git a/src/engine/engine_core_constraint.c b/src/engine/engine_core_constraint.c index 8cc5116f..092b0ecd 100644 --- a/src/engine/engine_core_constraint.c +++ b/src/engine/engine_core_constraint.c @@ -183,9 +183,16 @@ static int mj_elemBodyWeight(const mjModel* m, const mjData* d, int f, int e, in // compute body weights for a given contact vertex, return #bodies -static int mj_vertBodyWeight(const mjModel* m, const mjData* d, int f, int v, - int* body, mjtNum* weight, mjtNum bw) { - mjtNum* coord = m->flex_vert0 + 3*v; +static int mj_vertBodyWeight(const mjModel* m, const mjData* d, int f, int* v, + int* body, mjtNum* bweight, const mjtNum* vweight, int nw) { + if (nw == 0) { + return 0; + } + + mjtNum coord[3] = {0, 0, 0}; + for (int i = 0; i < nw; i++) { + mju_addToScl3(coord, m->flex_vert0 + 3*v[i], vweight[i]); + } int nstart = m->flex_nodeadr[f]; int nend = m->flex_nodeadr[f] + m->flex_nodenum[f]; int nb = 0; @@ -195,7 +202,7 @@ static int mj_vertBodyWeight(const mjModel* m, const mjData* d, int f, int v, if (w < 1e-5) { continue; } - if (weight) weight[nb] = w * bw; + if (bweight) bweight[nb] = w; body[nb++] = m->flex_nodebodyid[i]; } @@ -899,10 +906,6 @@ int mj_contactJacobian(const mjModel* m, mjData* d, const mjContact* con, int di int bid[729]; // 729 = 27*27 mjtNum bweight[729]; for (int side=0; side < 2; side++) { - int nw = 0; - int vid[4]; - mjtNum bw[4]; - // geom if (con->geom[side] >= 0) { bid[nb] = m->geom_bodyid[con->geom[side]]; @@ -910,32 +913,39 @@ int mj_contactJacobian(const mjModel* m, mjData* d, const mjContact* con, int di nb++; } - // flex vert - else if (con->vert[side] >= 0) { - vid[0] = m->flex_vertadr[con->flex[side]] + con->vert[side]; - bw[0] = side ? +1 : -1; - nw = 1; - } - - // flex elem + // flex else { - nw = mj_elemBodyWeight(m, d, con->flex[side], con->elem[side], - con->vert[1-side], con->pos, vid, bw); + int nw = 0; + int vid[4]; + mjtNum vweight[4]; - // negative sign for first side of contact - if (side == 0) { - mju_scl(bw, bw, -1, nw); + // vert + if (con->vert[side] >= 0) { + vid[0] = m->flex_vertadr[con->flex[side]] + con->vert[side]; + vweight[0] = side ? +1 : -1; + nw = 1; } - } - // get body or node ids and weights - for (int k=0; k < nw; k++) { + // elem + else { + nw = mj_elemBodyWeight(m, d, con->flex[side], con->elem[side], + con->vert[1-side], con->pos, vid, vweight); + + // negative sign for first side of contact + if (side == 0) { + mju_scl(vweight, vweight, -1, nw); + } + } + + // get body or node ids and weights if (m->flex_interp[con->flex[side]] == 0) { - bid[nb] = m->flex_vertbodyid[vid[k]]; - bweight[nb] = bw[k]; - nb++; + for (int k=0; k < nw; k++) { + bid[nb] = m->flex_vertbodyid[vid[k]]; + bweight[nb] = vweight[k]; + nb++; + } } else { - nb += mj_vertBodyWeight(m, d, con->flex[side], vid[k], bid+nb, bweight+nb, bw[k]); + nb += mj_vertBodyWeight(m, d, con->flex[side], vid, bid+nb, bweight+nb, vweight, nw); } } } @@ -1154,8 +1164,8 @@ void mj_diagApprox(const mjModel* m, mjData* d) { tran = rot = 0; for (int side=0; side < 2; side++) { // get bodies and weights - int nb = 0, bid[729], vid[4], nw = 0; - mjtNum bweight[729], bw[4]; + int nb = 0, bid[729]; + mjtNum bweight[729]; // geom if (con->geom[side] >= 0) { @@ -1164,27 +1174,34 @@ void mj_diagApprox(const mjModel* m, mjData* d) { nb = 1; } - // flex vert - else if (con->vert[side] >= 0) { - vid[0] = m->flex_vertadr[con->flex[side]] + con->vert[side]; - bw[0] = 1; - nw = 1; - } - - // flex elem + // flex else { - nw = mj_elemBodyWeight(m, d, con->flex[side], con->elem[side], - con->vert[1-side], con->pos, vid, bw); - } + int nw = 0; + int vid[4]; + mjtNum vweight[4]; - // get body or node ids and weights - for (int k=0; k < nw; k++) { + // vert + if (con->vert[side] >= 0) { + vid[0] = m->flex_vertadr[con->flex[side]] + con->vert[side]; + vweight[0] = 1; + nw = 1; + } + + // elem + else { + nw = mj_elemBodyWeight(m, d, con->flex[side], con->elem[side], + con->vert[1-side], con->pos, vid, vweight); + } + + // convert verted ids and weights to body ids and weights if (m->flex_interp[con->flex[side]] == 0) { - bid[k] = m->flex_vertbodyid[vid[k]]; - bweight[k] = bw[k]; - nb++; + for (int k=0; k < nw; k++) { + bid[k] = m->flex_vertbodyid[vid[k]]; + bweight[k] = vweight[k]; + nb++; + } } else { - nb = mj_vertBodyWeight(m, d, con->flex[side], vid[k], bid, bweight, bw[k]); + nb += mj_vertBodyWeight(m, d, con->flex[side], vid, bid, bweight, vweight, nw); } } @@ -1884,36 +1901,46 @@ static int mj_nc(const mjModel* m, mjData* d, int* nnz) { // get bodies int nb = 0, bid[729]; for (int side=0; side < 2; side++) { - int nw = 0; - int vid[4]; - // geom if (con->geom[side] >= 0) { bid[nb++] = m->geom_bodyid[con->geom[side]]; } - // flex vert - else if (con->vert[side] >= 0) { - vid[nw++] = m->flex_vertadr[con->flex[side]] + con->vert[side]; - } - - // flex elem + // flex else { - int f = con->flex[side]; - int fdim = m->flex_dim[f]; - const int* edata = m->flex_elem + m->flex_elemdataadr[f] + con->elem[side]*(fdim+1); - for (int k=0; k <= fdim; k++) { - vid[nw++] = m->flex_vertadr[f] + edata[k]; - } - } + int nw = 0; + int vid[4]; + mjtNum vweight[4]; - // get body or node ids and weights - for (int k=0; k < nw; k++) { + // flex vert + if (con->vert[side] >= 0) { + vid[nw++] = m->flex_vertadr[con->flex[side]] + con->vert[side]; + vweight[0] = 1; + } + + // flex elem + else { + int f = con->flex[side]; + int fdim = m->flex_dim[f]; + const int* edata = m->flex_elem + m->flex_elemdataadr[f] + con->elem[side]*(fdim+1); + for (int k=0; k <= fdim; k++) { + vid[nw++] = m->flex_vertadr[f] + edata[k]; + } + + if (m->flex_interp[f]) { + nw = mj_elemBodyWeight(m, d, con->flex[side], con->elem[side], + con->vert[1-side], con->pos, vid, vweight); + } + } + + // get body or node ids and weights if (m->flex_interp[con->flex[side]] == 0) { - bid[nb] = m->flex_vertbodyid[vid[k]]; - nb++; + for (int k=0; k < nw; k++) { + bid[nb] = m->flex_vertbodyid[vid[k]]; + nb++; + } } else { - nb += mj_vertBodyWeight(m, d, con->flex[side], vid[k], bid + nb, NULL, 0); + nb += mj_vertBodyWeight(m, d, con->flex[side], vid, bid+nb, NULL, vweight, nw); } } }