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
+16 -10
View File
@@ -2965,10 +2965,11 @@ void mj_projectConstraint(const mjModel* m, mjData* d) {
return;
}
// pre-count A nonzeros (compute AR_rownnz, AR_rowadr)
d->nA = mju_sqrMatTDSparseCount(d->efc_AR_rownnz, d->efc_AR_rowadr, nefc,
BT_rownnz, BT_rowadr, BT_colind,
B_rownnz, B_rowadr, B_colind, B_rowsuper, d, /*flg_upper=*/1);
int* diagind = mjSTACKALLOC(d, nefc, int);
d->nA = mju_sqrMatTDSparseSymbolic(
d->efc_AR_rownnz, d->efc_AR_rowadr, NULL, diagind,
nv, nefc, BT_rownnz, BT_rowadr, BT_colind,
B_rownnz, B_rowadr, B_colind, B_rowsuper, d);
// allocate A values and column indices on arena
d->efc_AR = mj_arenaAllocByte(d, sizeof(mjtNum) * d->nA, _Alignof(mjtNum));
@@ -2981,12 +2982,17 @@ void mj_projectConstraint(const mjModel* m, mjData* d) {
return;
}
// A = B * B'
int* diagind = mjSTACKALLOC(d, nefc, int);
mju_sqrMatTDSparse(d->efc_AR, BT, B, NULL, nv, nefc,
d->efc_AR_rownnz, d->efc_AR_rowadr, d->efc_AR_colind,
BT_rownnz, BT_rowadr, BT_colind, NULL,
B_rownnz, B_rowadr, B_colind, B_rowsuper, d, diagind);
// A = B * B': symbolic phase
mju_sqrMatTDSparseSymbolic(
d->efc_AR_rownnz, d->efc_AR_rowadr, d->efc_AR_colind, diagind,
nv, nefc, BT_rownnz, BT_rowadr, BT_colind,
B_rownnz, B_rowadr, B_colind, B_rowsuper, d);
// A = B * B': numeric phase
mju_sqrMatTDSparseNumeric(
d->efc_AR, nefc, d->efc_AR_rownnz, d->efc_AR_rowadr,
d->efc_AR_colind, diagind, BT, BT_rownnz, BT_rowadr,
BT_colind, B, B_rownnz, B_rowadr, B_colind, B_rowsuper, NULL, d);
// AR = A + diag(R)
for (int i=0; i < nefc; i++) {