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
+28 -16
View File
@@ -1530,10 +1530,10 @@ static void MakeHessian(mjData* d, mjPrimalContext* ctx) {
// sparse
if (ctx->is_sparse) {
// initialize Hessian rowadr, rownnz; get total nonzeros
ctx->nH = mju_sqrMatTDSparseCount(ctx->H_rownnz, ctx->H_rowadr, nv,
ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind,
ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind,
ctx->JT_rowsuper, d, /*flg_upper=*/0);
ctx->nH = mju_sqrMatTDSparseSymbolic(
ctx->H_rownnz, ctx->H_rowadr, NULL, NULL,
nefc, nv, ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind,
ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind, ctx->JT_rowsuper, d);
// add M nonzeros to Hessian total (unavoidable overcounting since H_colind is still unknown)
ctx->nH += ctx->M_rowadr[nv - 1] + ctx->M_rownnz[nv - 1];
@@ -1549,12 +1549,18 @@ static void MakeHessian(mjData* d, mjPrimalContext* ctx) {
ctx->H_colind = mjSTACKALLOC(d, ctx->nH, int);
ctx->H = mjSTACKALLOC(d, ctx->nH, mjtNum);
// compute H = J'*D*J
mju_sqrMatTDSparse(ctx->H, ctx->J, ctx->JT, ctx->D, nefc, nv,
ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind,
ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind, NULL,
ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind, ctx->JT_rowsuper,
d, /*diagind=*/NULL);
// compute H = J'*D*J: symbolic phase
mju_sqrMatTDSparseSymbolic(
ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind, NULL,
nefc, nv, ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind,
ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind, ctx->JT_rowsuper, d);
// compute H = J'*D*J: numeric phase
mju_sqrMatTDSparseNumeric(
ctx->H, nv, ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind,
NULL, ctx->J, ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind,
ctx->JT, ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind,
ctx->JT_rowsuper, ctx->D, d);
// add mass matrix: H = J'*D*J + C
mju_addToMatSparse(ctx->H, ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind, nv,
@@ -1626,12 +1632,18 @@ static void FactorizeHessian(mjData* d, mjPrimalContext* ctx, int flg_recompute)
if (ctx->is_sparse) {
// maybe compute H = M + J'*D*J
if (flg_recompute) {
// compute H = J'*D*J
mju_sqrMatTDSparse(ctx->H, ctx->J, ctx->JT, ctx->D, nefc, nv,
ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind,
ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind, NULL,
ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind, ctx->JT_rowsuper,
d, /*diagind=*/NULL);
// compute H = J'*D*J: symbolic phase
mju_sqrMatTDSparseSymbolic(
ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind, NULL,
nefc, nv, ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind,
ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind, ctx->JT_rowsuper, d);
// compute H = J'*D*J: numeric phase
mju_sqrMatTDSparseNumeric(
ctx->H, nv, ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind,
NULL, ctx->J, ctx->J_rownnz, ctx->J_rowadr, ctx->J_colind,
ctx->JT, ctx->JT_rownnz, ctx->JT_rowadr, ctx->JT_colind,
ctx->JT_rowsuper, ctx->D, d);
// add mass matrix: H = J'*D*J + C
mju_addToMatSparse(ctx->H, ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind, nv,