From f6c85287bf625b7991fdd9a3a7ade530e4b47074 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Wed, 27 May 2026 10:04:37 -0700 Subject: [PATCH] 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 --- src/engine/engine_core_constraint.c | 26 ++++++++--- src/engine/engine_util_misc.c | 70 ++++++++++++++++++++++++++--- src/engine/engine_util_misc.h | 3 ++ 3 files changed, 88 insertions(+), 11 deletions(-) diff --git a/src/engine/engine_core_constraint.c b/src/engine/engine_core_constraint.c index 2b7a1dbf..e7a521a0 100644 --- a/src/engine/engine_core_constraint.c +++ b/src/engine/engine_core_constraint.c @@ -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; diff --git a/src/engine/engine_util_misc.c b/src/engine/engine_util_misc.c index eada6949..65b0d6c2 100644 --- a/src/engine/engine_util_misc.c +++ b/src/engine/engine_util_misc.c @@ -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]); } } diff --git a/src/engine/engine_util_misc.h b/src/engine/engine_util_misc.h index f97bf4bb..42aa3704 100644 --- a/src/engine/engine_util_misc.h +++ b/src/engine/engine_util_misc.h @@ -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);