Optimize mj_flex performance.
Two performance optimizations for interpolated flex objects: 1. Hoist loop-invariant stride calculations in `mju_cellLookup` out of nested loops. This reduces multiplications from 16 to 6 (linear) and 54 to 12 (quadratic) per vertex. 2. Extract and consolidate optimized 3D interpolation logic into a new reusable utility `mju_evalBasisArray` in `engine_util_misc.c`. This function uses nested loops and precomputed 1D shape functions to avoid expensive dynamic `phi` calls and branching, and leverages stack-buffered outputs to eliminate compiler pointer-aliasing barriers. We propagate this optimization to both kinematics (`mju_interpolate3D`) and constraint setup (`engine_core_constraint.c`). Together these changes yield a ~50% overall speedup in `mj_fwdKinematics` for interpolated flexes with ~10k vertices in the collision meshes and ~10 nodes in the deformation grid. PiperOrigin-RevId: 922198093 Change-Id: I9c804e65cd532e2f9ac6e9a71d01c30cd64186fa
This commit is contained in:
committed by
Copybara-Service
parent
af4be63cd8
commit
f6c85287bf
@@ -290,13 +290,27 @@ static int mj_vertBodyWeight(const mjModel* m, const mjData* d, int f, int* v,
|
||||
int nstart = m->flex_nodeadr[f];
|
||||
int nb = 0;
|
||||
|
||||
for (int j = 0; j < npc; j++) {
|
||||
mjtNum w = mju_evalBasis(local, j, order);
|
||||
if (w < 1e-5) {
|
||||
continue;
|
||||
if (npc > 27) {
|
||||
for (int j = 0; j < npc; j++) {
|
||||
mjtNum w = mju_evalBasis(local, j, order);
|
||||
if (w < 1e-5) {
|
||||
continue;
|
||||
}
|
||||
if (bweight) bweight[nb] = sign * w;
|
||||
body[nb++] = m->flex_nodebodyid[nstart + nodeindices[j]];
|
||||
}
|
||||
} else {
|
||||
mjtNum basis[27];
|
||||
mju_evalBasisArray(basis, local, order);
|
||||
|
||||
for (int j = 0; j < npc; j++) {
|
||||
mjtNum w = basis[j];
|
||||
if (w < 1e-5) {
|
||||
continue;
|
||||
}
|
||||
if (bweight) bweight[nb] = sign * w;
|
||||
body[nb++] = m->flex_nodebodyid[nstart + nodeindices[j]];
|
||||
}
|
||||
if (bweight) bweight[nb] = sign * w;
|
||||
body[nb++] = m->flex_nodebodyid[nstart + nodeindices[j]];
|
||||
}
|
||||
|
||||
return nb;
|
||||
|
||||
@@ -574,6 +574,49 @@ mjtNum mju_evalBasis(const mjtNum x[3], int i, int order) {
|
||||
}
|
||||
}
|
||||
|
||||
// evaluate the basis functions at x for all nodes in the cell
|
||||
void mju_evalBasisArray(mjtNum* basis, const mjtNum x[3], int order) {
|
||||
if (order == 1) {
|
||||
mjtNum p[3][2] = {
|
||||
{1 - x[0], x[0]},
|
||||
{1 - x[1], x[1]},
|
||||
{1 - x[2], x[2]}
|
||||
};
|
||||
int j = 0;
|
||||
for (int i0=0; i0<2; i0++) {
|
||||
mjtNum w0 = p[0][i0];
|
||||
for (int i1=0; i1<2; i1++) {
|
||||
mjtNum w01 = w0 * p[1][i1];
|
||||
for (int i2=0; i2<2; i2++) {
|
||||
basis[j++] = w01 * p[2][i2];
|
||||
}
|
||||
}
|
||||
}
|
||||
} else if (order == 2) {
|
||||
mjtNum p[3][3];
|
||||
for (int d=0; d<3; d++) {
|
||||
for (int i=0; i<3; i++) {
|
||||
p[d][i] = phi(x[d], i, 2);
|
||||
}
|
||||
}
|
||||
int j = 0;
|
||||
for (int i0=0; i0<3; i0++) {
|
||||
mjtNum w0 = p[0][i0];
|
||||
for (int i1=0; i1<3; i1++) {
|
||||
mjtNum w01 = w0 * p[1][i1];
|
||||
for (int i2=0; i2<3; i2++) {
|
||||
basis[j++] = w01 * p[2][i2];
|
||||
}
|
||||
}
|
||||
}
|
||||
} else {
|
||||
int npoint = (order + 1) * (order + 1) * (order + 1);
|
||||
for (int j=0; j < npoint; j++) {
|
||||
basis[j] = mju_evalBasis(x, j, order);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// map global parametric coord to cell-local coord and build node indices
|
||||
// coord: [0,1]^3 parametric coordinates
|
||||
// cellnum: cell counts (cx, cy, cz)
|
||||
@@ -600,16 +643,21 @@ int mju_cellLookup(const mjtNum coord[3], const int cellnum[3], int order, mjtNu
|
||||
|
||||
// build node indices for this cell
|
||||
if (nodeindices) {
|
||||
int gi_base = ci * order;
|
||||
int gj_base = cj * order;
|
||||
int gk_base = ck * order;
|
||||
int ny_g = cy * order + 1;
|
||||
int nz_g = cz * order + 1;
|
||||
int ni = 0;
|
||||
for (int li = 0; li <= order; li++) {
|
||||
int gi = gi_base + li;
|
||||
int gi_stride = gi * ny_g * nz_g;
|
||||
for (int lj = 0; lj <= order; lj++) {
|
||||
int gj = gj_base + lj;
|
||||
int gj_stride = gi_stride + gj * nz_g;
|
||||
for (int lk = 0; lk <= order; lk++) {
|
||||
int gi = ci*order + li;
|
||||
int gj = cj*order + lj;
|
||||
int gk = ck*order + lk;
|
||||
nodeindices[ni++] = gi*ny_g*nz_g + gj*nz_g + gk;
|
||||
int gk = gk_base + lk;
|
||||
nodeindices[ni++] = gj_stride + gk;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -624,9 +672,21 @@ int mju_cellLookup(const mjtNum coord[3], const int cellnum[3], int order, mjtNu
|
||||
void mju_interpolate3D(mjtNum res[3], const mjtNum x[3], const mjtNum* coeff, int order,
|
||||
const int* nodeindices) {
|
||||
int npoint = (order + 1) * (order + 1) * (order + 1);
|
||||
|
||||
if (npoint > 27) {
|
||||
for (int j=0; j < npoint; j++) {
|
||||
int idx = nodeindices ? nodeindices[j] : j;
|
||||
mju_addToScl3(res, coeff+3*idx, mju_evalBasis(x, j, order));
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
mjtNum basis[27];
|
||||
mju_evalBasisArray(basis, x, order);
|
||||
|
||||
for (int j=0; j < npoint; j++) {
|
||||
int idx = nodeindices ? nodeindices[j] : j;
|
||||
mju_addToScl3(res, coeff+3*idx, mju_evalBasis(x, j, order));
|
||||
mju_addToScl3(res, coeff+3*idx, basis[j]);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -89,6 +89,9 @@ MJAPI void mju_defGradient(mjtNum res[9], const mjtNum p[3], const mjtNum* dof,
|
||||
// evaluate the basis function at x for the i-th node
|
||||
MJAPI mjtNum mju_evalBasis(const mjtNum x[3], int i, int order);
|
||||
|
||||
// evaluate the basis functions at x for all nodes in the cell
|
||||
MJAPI void mju_evalBasisArray(mjtNum* basis, const mjtNum x[3], int order);
|
||||
|
||||
// map global parametric coord to cell-local coord and build node indices
|
||||
MJAPI int mju_cellLookup(const mjtNum coord[3], const int cellnum[3], int order, mjtNum local[3],
|
||||
int* nodeindices);
|
||||
|
||||
Reference in New Issue
Block a user