From 1e1eb4826163426362b20bc319253ff272713465 Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Mon, 9 Mar 2026 08:19:40 -0700 Subject: [PATCH] No-op refactor of flex interpolation factorization and solve in mj_implicitSkip. PiperOrigin-RevId: 880867833 Change-Id: I2ab8ce5a027742af07ef0256bc4cdee62289282e --- src/engine/engine_forward.c | 377 ++++++++++++++++++++---------------- 1 file changed, 206 insertions(+), 171 deletions(-) diff --git a/src/engine/engine_forward.c b/src/engine/engine_forward.c index f294f69e..adc083af 100644 --- a/src/engine/engine_forward.c +++ b/src/engine/engine_forward.c @@ -1121,6 +1121,206 @@ void mj_RungeKutta(const mjModel* m, mjData* d, int N) { } +// context for flex interp reduced dense factorization/solve +typedef struct { + mjtNum* H; // dense Cholesky-factored matrix (ndof x ndof) + int* dof_indices; // global DOF index for each local flex DOF + int ndof; // number of flex DOFs + int ncoupling; // number of off-diagonal coupling terms + mjtNum* coupling_val; // coupling coefficient values + int* coupling_row; // local flex row index for each coupling term + int* coupling_col; // global DOF column index for each coupling term +} FlexInterpContext; + + +// collect flex DOFs for one flex, marking seen_dof and incrementing count +static void flexInterp_collect(const mjModel* m, int f, + int* chain_dofs, int* seen_dof, int* count) { + int nodenum = m->flex_nodenum[f]; + int nodeadr = m->flex_nodeadr[f]; + for (int n=0; n < nodenum; n++) { + int b = m->flex_nodebodyid[nodeadr+n]; + int chain_nnz; + if (m->body_dofnum[b] == 0) { + // pinned node: use bodyChain to get parent DOFs + chain_nnz = mj_bodyChain(m, b, chain_dofs); + } else { + // regular flex node: use body's own DOFs only + chain_nnz = m->body_dofnum[b]; + for (int j=0; j < chain_nnz; j++) { + chain_dofs[j] = m->body_dofadr[b] + j; + } + } + for (int i=0; i < chain_nnz; i++) { + int dof = chain_dofs[i]; + if (!seen_dof[dof]) { + seen_dof[dof] = 1; + (*count)++; + } + } + } +} + + +// build and factor the reduced dense matrix for flex interp DOFs +// mark/free stack handled by caller +static FlexInterpContext flexInterp_factor(const mjModel* m, mjData* d, int nv) { + FlexInterpContext ctx = {0}; + + int* chain_dofs = mjSTACKALLOC(d, nv, int); + int* seen_dof = mjSTACKALLOC(d, nv, int); + mju_fillInt(seen_dof, 0, nv); + + // count flex DOFs + int ndof = 0; + for (int f=0; f < m->nflex; f++) { + if (m->flex_interp[f]) { + flexInterp_collect(m, f, chain_dofs, seen_dof, &ndof); + } + } + if (ndof == 0) { + return ctx; + } + + // allocate and build global-to-local mapping + int* dof_indices = mjSTACKALLOC(d, ndof, int); + int* global2local = mjSTACKALLOC(d, nv, int); + mju_fillInt(global2local, -1, nv); + + // collect unique DOFs in order + int cnt = 0; + mju_fillInt(seen_dof, 0, nv); + for (int f=0; f < m->nflex; f++) { + if (m->flex_interp[f]) { + int nodenum = m->flex_nodenum[f]; + int nodeadr = m->flex_nodeadr[f]; + for (int n=0; n < nodenum; n++) { + int b = m->flex_nodebodyid[nodeadr+n]; + int chain_nnz; + if (m->body_dofnum[b] == 0) { + // pinned node: use bodyChain to get parent DOFs + chain_nnz = mj_bodyChain(m, b, chain_dofs); + } else { + // regular flex node: use body's own DOFs only + chain_nnz = m->body_dofnum[b]; + for (int j=0; j < chain_nnz; j++) { + chain_dofs[j] = m->body_dofadr[b] + j; + } + } + for (int i=0; i < chain_nnz; i++) { + int dof = chain_dofs[i]; + if (!seen_dof[dof]) { + seen_dof[dof] = 1; + dof_indices[cnt] = dof; + global2local[dof] = cnt; + cnt++; + } + } + } + } + } + + // select sparse matrix format based on integrator + int implicit = (m->opt.integrator == mjINT_IMPLICIT); + const int* rownnz = implicit ? m->D_rownnz : m->M_rownnz; + const int* rowadr = implicit ? m->D_rowadr : m->M_rowadr; + const int* colind = implicit ? m->D_colind : m->M_colind; + const mjtNum* source = implicit ? d->qLU : d->qH; + + // count coupling terms (off-diagonal: flex row, non-flex col) + int ncoupling = 0; + for (int i=0; i < ndof; i++) { + int row = dof_indices[i]; + int start = rowadr[row]; + int end = start + rownnz[row]; + for (int k=start; k < end; k++) { + if (global2local[colind[k]] < 0) { + ncoupling++; + } + } + } + + // allocate coupling storage + mjtNum* coupling_val = NULL; + int* coupling_row = NULL; + int* coupling_col = NULL; + if (ncoupling > 0) { + coupling_val = mjSTACKALLOC(d, ncoupling, mjtNum); + coupling_row = mjSTACKALLOC(d, ncoupling, int); + coupling_col = mjSTACKALLOC(d, ncoupling, int); + } + + // build H_flex (dense) from qLU (implicit) or qH (implicitfast) + mjtNum* H = mjSTACKALLOC(d, ndof*ndof, mjtNum); + mju_zero(H, ndof*ndof); + + int coup_cnt = 0; + for (int i=0; i < ndof; i++) { + int row = dof_indices[i]; + int start = rowadr[row]; + int end = start + rownnz[row]; + for (int k=start; k < end; k++) { + int col = colind[k]; + int local_j = global2local[col]; + if (local_j >= 0) { + H[i*ndof+local_j] = source[k]; + } else if (coup_cnt < ncoupling) { + coupling_val[coup_cnt] = source[k]; + coupling_row[coup_cnt] = i; + coupling_col[coup_cnt] = col; + coup_cnt++; + } + } + } + + // add flex stiffness and factorize + mjd_flexInterp_addH(m, d, H, dof_indices, ndof, m->opt.timestep); + mju_cholFactor(H, ndof, mjMINVAL); + + // store results in context + ctx.H = H; + ctx.dof_indices = dof_indices; + ctx.ndof = ndof; + ctx.ncoupling = ncoupling; + ctx.coupling_val = coupling_val; + ctx.coupling_row = coupling_row; + ctx.coupling_col = coupling_col; + return ctx; +} + + +// solve the reduced dense system for flex interp DOFs, overwrite qacc +static void flexInterp_solve(const mjModel* m, mjData* d, const FlexInterpContext* ctx, + mjtNum* qacc, const mjtNum* qfrc, int nv) { + int ndof = ctx->ndof; + mjtNum* qfrc_flex = mjSTACKALLOC(d, ndof, mjtNum); + mjtNum* res = mjSTACKALLOC(d, nv, mjtNum); + + mjtNum h = m->opt.timestep; + mjtNum damp = (m->nflex > 0 && m->flex_damping) ? m->flex_damping[0] : 0; + mjtNum scl = h*h + h*damp; + mjtNum factor = (scl > mjMINVAL) ? (h/scl) : 0; + + // velocity correction: -h * K * v + mju_zero(res, nv); + mjd_flexInterp_mulKD(m, d, res, d->qvel, h); + + for (int i=0; i < ndof; i++) { + int global_dof = ctx->dof_indices[i]; + qfrc_flex[i] = qfrc[global_dof] + res[global_dof] * factor; + } + + // coupling correction: qfrc_flex -= H_coupling * qacc_parent + for (int k=0; k < ctx->ncoupling; k++) { + qfrc_flex[ctx->coupling_row[k]] -= ctx->coupling_val[k] * qacc[ctx->coupling_col[k]]; + } + + // solve and scatter back + mju_cholSolve(qfrc_flex, ctx->H, qfrc_flex, ndof); + mju_scatter(qacc, qfrc_flex, ctx->dof_indices, ndof); +} + + // fully implicit in velocity, possibly skipping factorization void mj_implicitSkip(const mjModel* m, mjData* d, int skipfactor) { TM_START; @@ -1144,21 +1344,15 @@ void mj_implicitSkip(const mjModel* m, mjData* d, int skipfactor) { // check for flex_interp int has_flex_interp = 0; - for (int f = 0; f < m->nflex; f++) { + for (int f=0; f < m->nflex; f++) { if (m->flex_interp[f]) { has_flex_interp = 1; break; } } - // flex: data structures for reduced dense factorization - mjtNum* H_flex = NULL; - int* flex_dof_indices = NULL; - int nflexdofs = 0; - int ncoupling = 0; - mjtNum* coupling_val = NULL; - int* coupling_row = NULL; - int* coupling_col = NULL; + // flex interp context (populated during factorization) + FlexInterpContext flex = {0}; // factorization if (!skipfactor) { @@ -1190,135 +1384,7 @@ void mj_implicitSkip(const mjModel* m, mjData* d, int skipfactor) { // flex: reduced dense factorization if (has_flex_interp && !sleep_filter) { - // temporary allocations for body chain - int* chain_dofs = mjSTACKALLOC(d, nv, int); - int* seen_dof = mjSTACKALLOC(d, nv, int); - mju_fillInt(seen_dof, 0, nv); - - // identify flex DOFs - // For pinned nodes (body_dofnum==0): use bodyChain to include parent DOFs - // For regular flex nodes: use body_dofadr for one-way coupling - for (int f=0; f < m->nflex; f++) { - if (m->flex_interp[f]) { - int nodenum = m->flex_nodenum[f]; - int nodeadr = m->flex_nodeadr[f]; - for (int n=0; n < nodenum; n++) { - int b = m->flex_nodebodyid[nodeadr + n]; - int chain_nnz; - if (m->body_dofnum[b] == 0) { - // Pinned node: use bodyChain to get parent DOFs - chain_nnz = mj_bodyChain(m, b, chain_dofs); - } else { - // Regular flex node: use body's own DOFs only - chain_nnz = m->body_dofnum[b]; - for (int j = 0; j < chain_nnz; j++) { - chain_dofs[j] = m->body_dofadr[b] + j; - } - } - for (int i=0; i < chain_nnz; i++) { - int dof = chain_dofs[i]; - if (!seen_dof[dof]) { - seen_dof[dof] = 1; - nflexdofs++; - } - } - } - } - } - - // allocations - if (nflexdofs > 0) { - flex_dof_indices = mjSTACKALLOC(d, nflexdofs, int); - int* global2local = mjSTACKALLOC(d, nv, int); - mju_fillInt(global2local, -1, nv); - - // collect unique DOFs in order - int cnt = 0; - mju_fillInt(seen_dof, 0, nv); - for (int f=0; f < m->nflex; f++) { - if (m->flex_interp[f]) { - int nodenum = m->flex_nodenum[f]; - int nodeadr = m->flex_nodeadr[f]; - for (int n=0; n < nodenum; n++) { - int b = m->flex_nodebodyid[nodeadr + n]; - int chain_nnz; - if (m->body_dofnum[b] == 0) { - // Pinned node: use bodyChain to get parent DOFs - chain_nnz = mj_bodyChain(m, b, chain_dofs); - } else { - // Regular flex node: use body's own DOFs only - chain_nnz = m->body_dofnum[b]; - for (int j = 0; j < chain_nnz; j++) { - chain_dofs[j] = m->body_dofadr[b] + j; - } - } - for (int i=0; i < chain_nnz; i++) { - int dof = chain_dofs[i]; - if (!seen_dof[dof]) { - seen_dof[dof] = 1; - flex_dof_indices[cnt] = dof; - global2local[dof] = cnt; - cnt++; - } - } - } - } - } - - const int* rownnz = (m->opt.integrator == mjINT_IMPLICIT) ? m->D_rownnz : m->M_rownnz; - const int* rowadr = (m->opt.integrator == mjINT_IMPLICIT) ? m->D_rowadr : m->M_rowadr; - const int* colind = (m->opt.integrator == mjINT_IMPLICIT) ? m->D_colind : m->M_colind; - const mjtNum* source = (m->opt.integrator == mjINT_IMPLICIT) ? d->qLU : d->qH; - - // count coupling terms (off-diagonal: flex row, non-flex col) - for (int i=0; i < nflexdofs; i++) { - int row = flex_dof_indices[i]; - int start = rowadr[row]; - int end = start + rownnz[row]; - for (int k=start; k < end; k++) { - if (global2local[colind[k]] < 0) { - ncoupling++; - } - } - } - - // allocate coupling storage - if (ncoupling > 0) { - coupling_val = mjSTACKALLOC(d, ncoupling, mjtNum); - coupling_row = mjSTACKALLOC(d, ncoupling, int); - coupling_col = mjSTACKALLOC(d, ncoupling, int); - } - - // build H_flex (dense) from qLU (implicit) or qH (implicitfast) - H_flex = mjSTACKALLOC(d, nflexdofs*nflexdofs, mjtNum); - mju_zero(H_flex, nflexdofs*nflexdofs); - - int coup_cnt = 0; - for (int i=0; i < nflexdofs; i++) { - int row = flex_dof_indices[i]; - int start = rowadr[row]; - int end = start + rownnz[row]; - for (int k=start; k < end; k++) { - int col = colind[k]; - int local_j = global2local[col]; - if (local_j >= 0) { - H_flex[i*nflexdofs + local_j] = source[k]; - } else if (coup_cnt < ncoupling) { - coupling_val[coup_cnt] = source[k]; - coupling_row[coup_cnt] = i; // local flex index - coupling_col[coup_cnt] = col; // global parent index - coup_cnt++; - } - } - } - - // add stiffness to H_flex - mjtNum h = m->opt.timestep; - mjd_flexInterp_addH(m, d, H_flex, flex_dof_indices, nflexdofs, h); - - // factor H_flex - mju_cholFactor(H_flex, nflexdofs, mjMINVAL); - } + flex = flexInterp_factor(m, d, nv); } // standard factorization (implicit / implicitfast) @@ -1330,7 +1396,6 @@ void mj_implicitSkip(const mjModel* m, mjData* d, int skipfactor) { } } - // solve // standard sparse solve if (m->opt.integrator == mjINT_IMPLICIT) { mju_solveLUSparse(qacc, d->qLU, qfrc, nv, m->D_rownnz, m->D_rowadr, m->D_diag, m->D_colind, @@ -1346,38 +1411,8 @@ void mj_implicitSkip(const mjModel* m, mjData* d, int skipfactor) { } // flex: reduced dense solve - if (H_flex) { - // compute qfrc_flex - mjtNum* qfrc_flex = mjSTACKALLOC(d, nflexdofs, mjtNum); - mjtNum* res = mjSTACKALLOC(d, nv, mjtNum); - - mjtNum h = m->opt.timestep; - mjtNum damp = (m->nflex > 0 && m->flex_damping) ? m->flex_damping[0] : 0; - mjtNum scl = h * h + h * damp; - mjtNum factor = (scl > mjMINVAL) ? (h/scl) : 0; - - // velocity correction: -h * K * v - mju_zero(res, nv); - mjd_flexInterp_mulKD(m, d, res, d->qvel, h); // returns -scl * K * v - - for (int i=0; i < nflexdofs; i++) { - int global_dof = flex_dof_indices[i]; - qfrc_flex[i] = qfrc[global_dof] + res[global_dof] * factor; - } - - // apply coupling correction: qfrc_flex -= H_coupling * qacc_parent - if (ncoupling > 0) { - for (int k=0; k < ncoupling; k++) { - qfrc_flex[coupling_row[k]] -= coupling_val[k] * qacc[coupling_col[k]]; - } - } - - // solve H_flex * qacc_flex = qfrc_flex - // reuse qfrc_flex as result buffer (qacc_flex) - mju_cholSolve(qfrc_flex, H_flex, qfrc_flex, nflexdofs); - - // overwrite flex DOFs with reduced dense solution - mju_scatter(qacc, qfrc_flex, flex_dof_indices, nflexdofs); + if (flex.H) { + flexInterp_solve(m, d, &flex, qacc, qfrc, nv); } // advance state and time