No-op refactor of flex interpolation factorization and solve in mj_implicitSkip.

PiperOrigin-RevId: 880867833
Change-Id: I2ab8ce5a027742af07ef0256bc4cdee62289282e
This commit is contained in:
Yuval Tassa
2026-03-09 08:19:40 -07:00
committed by Copybara-Service
parent 79d14035f9
commit 1e1eb48261
+206 -171
View File
@@ -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