2-3x speedup of sparse matrix squaring.

Split symbolic and numeric phases for sparse `M'*diag*M` computation. Microseconds per call for the monolithic vs the split approach for the 100_humanoids and 2humanoid100 models:

```
+-------+------+----------+------------+---------+
| Model | Arch | Col (µs) | Split (µs) | Speedup |
+-------+------+----------+------------+---------+
| 2H100 | x86  | 238.3    | 74.5       | 3.2x    |
+-------+------+----------+------------+---------+
|       | ARM  | 111.6    | 53.2       | 2.1x    |
+-------+------+----------+------------+---------+
| 100H  | x86  | 1325.3   | 656.2      | 2.0x    |
+-------+------+----------+------------+---------+
|       | ARM  | 594.8    | 306.6      | 1.9x    |
+-------+------+----------+------------+---------+
```

PiperOrigin-RevId: 900154308
Change-Id: Ia6e9b8e196e2ed37b723a0faf60e9731303a9619
This commit is contained in:
Yuval Tassa
2026-04-15 07:19:03 -07:00
committed by Copybara-Service
parent 62cffb1536
commit a2d0e33c0f
8 changed files with 1301 additions and 626 deletions
@@ -38,116 +38,7 @@ static const int kNumWarmupSteps = 500;
// ----------------------------- old functions --------------------------------
void ABSL_ATTRIBUTE_NOINLINE mju_sqrMatTDSparse_baseline(
mjtNum* res, const mjtNum* mat, const mjtNum* matT, const mjtNum* diag,
int nr, int nc, int* res_rownnz, int* res_rowadr, int* res_colind,
const int* rownnz, const int* rowadr, const int* colind,
const int* rowsuper, const int* rownnzT, const int* rowadrT,
const int* colindT, const int* rowsuperT, mjData* d, int* unused) {
mj_markStack(d);
int* chain = mj_stackAllocInt(d, 2 * nc);
mjtNum* buffer = mj_stackAllocNum(d, nc);
for (int r = 0; r < nc; r++) {
res_rowadr[r] = r * nc;
}
for (int r = 0; r < nc; r++) {
if (rowsuperT && r > 0 && rowsuperT[r - 1] > 0) {
res_rownnz[r] = res_rownnz[r - 1];
memcpy(res_colind + res_rowadr[r], res_colind + res_rowadr[r - 1],
res_rownnz[r] * sizeof(int));
if (rownnzT[r]) {
res_colind[res_rowadr[r] + res_rownnz[r]] = r;
res_rownnz[r]++;
}
} else {
int nchain = 0;
int inew = 0, iold = nc;
int lastadded = -1;
for (int i = 0; i < rownnzT[r]; i++) {
int c = colindT[rowadrT[r] + i];
if (rowsuper && lastadded >= 0 &&
(c - lastadded) <= rowsuper[lastadded]) {
continue;
} else {
lastadded = c;
}
int adr = inew;
inew = iold;
iold = adr;
int nnewchain = 0;
adr = 0;
int end = rowadr[c] + rownnz[c];
for (int adr1 = rowadr[c]; adr1 < end; adr1++) {
int col_mat = colind[adr1];
while (adr < nchain && chain[iold + adr] < col_mat &&
chain[iold + adr] <= r) {
chain[inew + nnewchain++] = chain[iold + adr++];
}
if (col_mat > r) {
break;
}
if (adr < nchain && chain[iold + adr] == col_mat) {
adr++;
}
chain[inew + nnewchain++] = col_mat;
}
while (adr < nchain && chain[iold + adr] <= r) {
chain[inew + nnewchain++] = chain[iold + adr++];
}
nchain = nnewchain;
}
res_rownnz[r] = nchain;
if (nchain) {
memcpy(res_colind + res_rowadr[r], chain + inew, nchain * sizeof(int));
}
}
}
for (int r = 0; r < nc; r++) {
int adr = res_rowadr[r];
for (int i = 0; i < res_rownnz[r]; i++) {
buffer[res_colind[adr + i]] = 0;
}
for (int i = 0; i < rownnzT[r]; i++) {
int c = colindT[rowadrT[r] + i];
mjtNum matTrc = matT[rowadrT[r] + i];
if (diag) {
matTrc *= diag[c];
}
int end = rowadr[c] + rownnz[c];
for (int adr = rowadr[c]; adr < end; adr++) {
int adr1;
if ((adr1 = colind[adr]) > r) {
break;
}
buffer[adr1] += matTrc * mat[adr];
}
}
adr = res_rowadr[r];
for (int i = 0; i < res_rownnz[r]; i++) {
res[adr + i] = buffer[res_colind[adr + i]];
}
}
for (int r = 1; r < nc; r++) {
int end = res_rowadr[r] + res_rownnz[r] - 1;
for (int adr = res_rowadr[r]; adr < end; adr++) {
int adr1 = res_rowadr[res_colind[adr]] + res_rownnz[res_colind[adr]]++;
res[adr1] = res[adr];
res_colind[adr1] = r;
}
}
mj_freeStack(d);
}
// transpose sparse matrix (uncompressed)
void ABSL_ATTRIBUTE_NOINLINE transposeSparse_baseline(
@@ -506,15 +397,27 @@ void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_combineSparse_old(
}
BENCHMARK(BM_combineSparse_old);
enum class Size { H2_100, H100 };
template <Size S>
const char* ModelPath() {
if constexpr (S == Size::H2_100) {
return "../test/benchmark/testdata/2humanoid100_chol.xml";
} else {
return "../test/benchmark/testdata/100_humanoids_chol.xml";
}
}
enum class Supernode {
None,
PostProcess,
Inline
};
template <Size S>
static void BM_transposeSparse(benchmark::State& state, TransposeFuncPtr func,
Supernode super) {
static mjModel* m = LoadModelFromPath("humanoid/humanoid100.xml");
static mjModel* m = LoadModelFromPath(ModelPath<S>());
// force use of sparse matrices
m->opt.jacobian = mjJAC_SPARSE;
@@ -553,131 +456,67 @@ static void BM_transposeSparse(benchmark::State& state, TransposeFuncPtr func,
}
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_transposeSparse_old(benchmark::State& state) {
BM_transposeSparse_2H100_old(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_transposeSparse(state, &transposeSparse_baseline, Supernode::None);
BM_transposeSparse<Size::H2_100>(state, &transposeSparse_baseline,
Supernode::None);
}
BENCHMARK(BM_transposeSparse_old);
BENCHMARK(BM_transposeSparse_2H100_old);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_transposeSparse_new(benchmark::State& state) {
BM_transposeSparse_2H100_new(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_transposeSparse(state, &mju_transposeSparse, Supernode::None);
BM_transposeSparse<Size::H2_100>(state, &mju_transposeSparse,
Supernode::None);
}
BENCHMARK(BM_transposeSparse_new);
BENCHMARK(BM_transposeSparse_2H100_new);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_transposeSparse_superpost(benchmark::State& state) {
BM_transposeSparse_2H100_superpost(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_transposeSparse(state, &mju_transposeSparse, Supernode::PostProcess);
BM_transposeSparse<Size::H2_100>(state, &mju_transposeSparse,
Supernode::PostProcess);
}
BENCHMARK(BM_transposeSparse_superpost);
BENCHMARK(BM_transposeSparse_2H100_superpost);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_transposeSparse_superinline(benchmark::State& state) {
BM_transposeSparse_2H100_superinline(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_transposeSparse(state, &mju_transposeSparse, Supernode::Inline);
}
BENCHMARK(BM_transposeSparse_superinline);
static void BM_sqrMatTDSparse(benchmark::State& state, SqrMatTDFuncPtr func) {
static mjModel* m =
LoadModelFromPath("../test/benchmark/testdata/2humanoid100.xml");
// force use of sparse matrices, Newton solver, no islands
m->opt.jacobian = mjJAC_SPARSE;
m->opt.solver = mjSOL_NEWTON;
m->opt.disableflags |= mjDSBL_ISLAND;
mjData* d = mj_makeData(m);
// warm-up rollout to get a typical state
while (d->time < 2) {
mj_step(m, d);
}
// allocate
mj_markStack(d);
mjtNum* H = mj_stackAllocNum(d, m->nv * m->nv);
int* rownnz = mj_stackAllocInt(d, m->nv);
int* rowadr = mj_stackAllocInt(d, m->nv);
int* colind = mj_stackAllocInt(d, m->nv * m->nv);
int* diagind = mj_stackAllocInt(d, m->nv);
// compute D corresponding to quad states
mjtNum* D = mj_stackAllocNum(d, d->nefc);
for (int i = 0; i < d->nefc; i++) {
if (d->efc_state[i] == mjCNSTRSTATE_QUADRATIC) {
D[i] = d->efc_D[i];
} else {
D[i] = 0;
}
}
int* JT_rownnz = mj_stackAllocInt(d, m->nv);
int* JT_rowadr = mj_stackAllocInt(d, m->nv);
int* JT_rowsuper = mj_stackAllocInt(d, m->nv);
int* JT_colind = mj_stackAllocInt(d, d->nJ);
mjtNum* JT = mj_stackAllocNum(d, d->nJ);
mju_transposeSparse(JT, d->efc_J, d->nefc, m->nv,
JT_rownnz, JT_rowadr, JT_colind, JT_rowsuper,
d->efc_J_rownnz, d->efc_J_rowadr, d->efc_J_colind);
// time benchmark
if (func) {
mju_sqrMatTDSparseCount(rownnz, rowadr, m->nv,
d->efc_J_rownnz, d->efc_J_rowadr, d->efc_J_colind,
JT_rownnz, JT_rowadr,
JT_colind, nullptr, d, 1);
for (auto s : state) {
// compute H = J'*D*J, compressed layout
func(H, d->efc_J, JT, D, d->nefc, m->nv, rownnz, rowadr, colind,
d->efc_J_rownnz, d->efc_J_rowadr, d->efc_J_colind, NULL,
JT_rownnz, JT_rowadr, JT_colind,
JT_rowsuper, d, diagind);
}
} else {
for (auto s : state) {
// baseline depends on efc_J_rowsuper
mju_superSparse(d->nefc, d->efc_J_rowsuper,
d->efc_J_rownnz, d->efc_J_rowadr, d->efc_J_colind);
// compute H = J'*D*J, uncompressed layout
mju_sqrMatTDSparse_baseline(
H, d->efc_J, JT, D, d->nefc, m->nv, rownnz, rowadr, colind,
d->efc_J_rownnz, d->efc_J_rowadr, d->efc_J_colind, d->efc_J_rowsuper,
JT_rownnz, JT_rowadr, JT_colind,
JT_rowsuper, d, /*unused=*/nullptr);
}
}
// finalize
mj_freeStack(d);
mj_deleteData(d);
state.SetItemsProcessed(state.iterations());
BM_transposeSparse<Size::H2_100>(state, &mju_transposeSparse,
Supernode::Inline);
}
BENCHMARK(BM_transposeSparse_2H100_superinline);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_sqrMatTDSparse_col(benchmark::State& state) {
BM_transposeSparse_100H_old(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_sqrMatTDSparse(state, &mju_sqrMatTDSparse);
BM_transposeSparse<Size::H100>(state, &transposeSparse_baseline,
Supernode::None);
}
BENCHMARK(BM_sqrMatTDSparse_col);
BENCHMARK(BM_transposeSparse_100H_old);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_sqrMatTDSparse_row(benchmark::State& state) {
BM_transposeSparse_100H_new(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_sqrMatTDSparse(state, &mju_sqrMatTDSparse_row);
BM_transposeSparse<Size::H100>(state, &mju_transposeSparse, Supernode::None);
}
BENCHMARK(BM_sqrMatTDSparse_row);
BENCHMARK(BM_transposeSparse_100H_new);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_sqrMatTDSparse_uncompressed(benchmark::State& state) {
BM_transposeSparse_100H_superpost(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_sqrMatTDSparse(state, nullptr);
BM_transposeSparse<Size::H100>(state, &mju_transposeSparse,
Supernode::PostProcess);
}
BENCHMARK(BM_sqrMatTDSparse_uncompressed);
BENCHMARK(BM_transposeSparse_100H_superpost);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_transposeSparse_100H_superinline(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_transposeSparse<Size::H100>(state, &mju_transposeSparse,
Supernode::Inline);
}
BENCHMARK(BM_transposeSparse_100H_superinline);
} // namespace
} // namespace mujoco