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:
committed by
Copybara-Service
parent
62cffb1536
commit
a2d0e33c0f
@@ -13,8 +13,6 @@
|
||||
// limitations under the License.
|
||||
|
||||
#include "engine/engine_util_sparse.h"
|
||||
#include "engine/engine_util_sparse_avx.h" // IWYU pragma: keep
|
||||
|
||||
|
||||
#include <mujoco/mjdata.h>
|
||||
#include <mujoco/mjmacro.h>
|
||||
@@ -23,7 +21,7 @@
|
||||
#include "engine/engine_memory.h"
|
||||
#include "engine/engine_util_blas.h"
|
||||
#include "engine/engine_util_misc.h"
|
||||
|
||||
#include "engine/engine_util_sparse_avx.h" // IWYU pragma: keep
|
||||
|
||||
//------------------------------ sparse operations -------------------------------------------------
|
||||
|
||||
@@ -723,9 +721,317 @@ void mju_sqrMatTDUncompressedInit(int* res_rowadr, int nc) {
|
||||
}
|
||||
|
||||
|
||||
// max number of supernodes handled
|
||||
// max number of supernodes handled by column-based matrix squaring functions
|
||||
#define mjMAXSUPER 8
|
||||
|
||||
// column-based symbolic phase for sparse matrix squaring: compute sparsity pattern of M'*M
|
||||
// if res_colind is NULL: count mode, fill res_rownnz/res_rowadr, return nnz
|
||||
// if res_colind is not NULL: fill mode, write sorted column indices
|
||||
// if res_diagind is not NULL: also fill upper triangle and output diagonal indices
|
||||
int mju_sqrMatTDSparseSymbolic(
|
||||
int* restrict res_rownnz, int* restrict res_rowadr,
|
||||
int* restrict res_colind, int* restrict res_diagind, int nr, int nc,
|
||||
const int* rownnz, const int* rowadr, const int* colind,
|
||||
const int* rownnzT, const int* rowadrT, const int* colindT, const int* rowsuperT, mjData* d) {
|
||||
mj_markStack(d);
|
||||
|
||||
// reinterpret M^T as CSC
|
||||
const int* colnnz = rownnzT;
|
||||
const int* coladr = rowadrT;
|
||||
const int* rowind = colindT;
|
||||
const int* colsuper = rowsuperT;
|
||||
|
||||
// reinterpret M as CSC
|
||||
const int* colnnzT = rownnz;
|
||||
const int* coladrT = rowadr;
|
||||
const int* rowindT = colind;
|
||||
|
||||
// marker[j] = 1 if row j has been visited in current column batch
|
||||
int* marker = mjSTACKALLOC(d, nc, int);
|
||||
mju_zeroInt(marker, nc);
|
||||
|
||||
// buffer_idx: list of row indices with nonzeros in current column batch
|
||||
int* buffer_idx = mjSTACKALLOC(d, nc, int);
|
||||
|
||||
// rowstart[r]: first index in row r of M where column > current result column
|
||||
int* rowstart = mjSTACKALLOC(d, nr, int);
|
||||
mju_zeroInt(rowstart, nr);
|
||||
|
||||
// clear res_rownnz (used for both counting and filling)
|
||||
mju_zeroInt(res_rownnz, nc);
|
||||
|
||||
// process result columns c = 0, 1, ..., nc-1
|
||||
int ns; // set in the loop
|
||||
for (int c = 0; c < nc; c += ns) {
|
||||
int buffer_nnz = 0;
|
||||
|
||||
// column c of M^T
|
||||
int nnz_c = colnnz[c];
|
||||
int adr_c = coladr[c];
|
||||
const int* ind_c = rowind + adr_c;
|
||||
|
||||
// supernode size: how many consecutive columns share the same sparsity pattern
|
||||
ns = 1;
|
||||
int cs;
|
||||
if (colsuper && (cs = colsuper[c])) {
|
||||
ns += mjMIN(cs, mjMAXSUPER - 1);
|
||||
}
|
||||
|
||||
// for each row r where M^T[r, c] != 0, look at row r of M
|
||||
for (int i = 0; i < nnz_c; i++) {
|
||||
int r = ind_c[i];
|
||||
int adrT = coladrT[r];
|
||||
int nnzT = colnnzT[r];
|
||||
const int* indT = rowindT + adrT;
|
||||
|
||||
// scan row r of M, starting from rowstart[r]
|
||||
for (int k = rowstart[r]; k < nnzT; k++) {
|
||||
int j = indT[k];
|
||||
|
||||
// skip if j <= c: only fill the strict lower triangle
|
||||
if (j <= c) {
|
||||
rowstart[r]++;
|
||||
continue;
|
||||
}
|
||||
|
||||
// new nonzero in row j of result
|
||||
if (!marker[j]) {
|
||||
marker[j] = 1;
|
||||
buffer_idx[buffer_nnz++] = j;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// scatter: update result rows j > c that have nonzeros in this column batch
|
||||
|
||||
// fill mode: write column indices, clear markers
|
||||
if (res_colind) {
|
||||
for (int i = 0; i < buffer_nnz; i++) {
|
||||
int j = buffer_idx[i];
|
||||
marker[j] = 0;
|
||||
int nm = mjMIN(ns, j - c);
|
||||
int adr_j = res_rowadr[j] + res_rownnz[j];
|
||||
for (int s = 0; s < nm; s++) {
|
||||
res_colind[adr_j + s] = c + s;
|
||||
}
|
||||
res_rownnz[j] += nm;
|
||||
}
|
||||
|
||||
// write diagonal entries
|
||||
for (int s = 0; s < ns; s++) {
|
||||
int col = c + s;
|
||||
if (colnnz[col]) {
|
||||
res_colind[res_rowadr[col] + res_rownnz[col]] = col;
|
||||
res_rownnz[col]++;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// count mode: just count and clear markers
|
||||
else {
|
||||
for (int i = 0; i < buffer_nnz; i++) {
|
||||
int j = buffer_idx[i];
|
||||
marker[j] = 0;
|
||||
int nm = mjMIN(ns, j - c);
|
||||
res_rownnz[j] += nm;
|
||||
if (res_diagind) {
|
||||
for (int s = 0; s < nm; s++) {
|
||||
res_rownnz[c + s]++;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// count diagonal entries
|
||||
for (int s = 0; s < ns; s++) {
|
||||
int col = c + s;
|
||||
if (colnnz[col]) {
|
||||
res_rownnz[col]++;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// count mode: compute res_rowadr from res_rownnz
|
||||
if (!res_colind) {
|
||||
res_rowadr[0] = 0;
|
||||
for (int r = 1; r < nc; r++) {
|
||||
res_rowadr[r] = res_rowadr[r - 1] + res_rownnz[r - 1];
|
||||
}
|
||||
}
|
||||
|
||||
// fill mode with upper triangle: record diagonal positions and mirror from lower
|
||||
if (res_colind && res_diagind) {
|
||||
// save current counts (lower + diagonal)
|
||||
int* lower_nnz = mjSTACKALLOC(d, nc, int);
|
||||
mju_copyInt(lower_nnz, res_rownnz, nc);
|
||||
|
||||
// save diagonal indices
|
||||
for (int r = 0; r < nc; r++) {
|
||||
res_diagind[r] = res_rowadr[r] + lower_nnz[r] - 1;
|
||||
}
|
||||
|
||||
// fill upper triangle: for each (r, c) with c < r, write to (c, r)
|
||||
for (int r = 0; r < nc; r++) {
|
||||
int adr = res_rowadr[r];
|
||||
int nnz = lower_nnz[r];
|
||||
for (int j = 0; j < nnz; j++) {
|
||||
int col = res_colind[adr + j];
|
||||
if (col < r) {
|
||||
res_colind[res_rowadr[col] + res_rownnz[col]++] = r;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mj_freeStack(d);
|
||||
|
||||
return res_rowadr[nc - 1] + res_rownnz[nc - 1];
|
||||
}
|
||||
|
||||
|
||||
// numeric phase for sparse matrix squaring: compute values given pre-computed sparsity
|
||||
// diagind can be NULL, otherwise fills upper triangle and saves diagonal indices
|
||||
void mju_sqrMatTDSparseNumeric(
|
||||
mjtNum* restrict res, int nc,
|
||||
const int* res_rownnz, const int* res_rowadr, const int* res_colind, const int* res_diagind,
|
||||
const mjtNum* mat, const int* rownnz, const int* rowadr, const int* colind,
|
||||
const mjtNum* matT, const int* rownnzT, const int* rowadrT, const int* colindT,
|
||||
const int* rowsuperT, const mjtNum* diag, mjData* d) {
|
||||
mj_markStack(d);
|
||||
|
||||
// dense accumulator for current result row (or batch of rows)
|
||||
mjtNum* restrict buffer = mjSTACKALLOC(d, nc * mjMAXSUPER, mjtNum);
|
||||
mju_zero(buffer, nc * mjMAXSUPER);
|
||||
|
||||
// process result rows
|
||||
int ns; // set in the loop
|
||||
for (int r = 0; r < nc; r += ns) {
|
||||
// determine supernode size
|
||||
ns = 1;
|
||||
if (rowsuperT) {
|
||||
ns = rowsuperT[r] + 1;
|
||||
if (ns > mjMAXSUPER) ns = mjMAXSUPER;
|
||||
}
|
||||
|
||||
// single row
|
||||
if (ns == 1) {
|
||||
int nnzT_r = rownnzT[r];
|
||||
int adr_r = rowadrT[r];
|
||||
|
||||
// accumulate: res[r, :] = sum over k in M'[r, :] of diag[k] * M'[r, k] * M[k, :]
|
||||
for (int i = 0; i < nnzT_r; i++) {
|
||||
int k = colindT[adr_r + i];
|
||||
mjtNum valT = matT[adr_r + i];
|
||||
mjtNum scale = diag ? diag[k] * valT : valT;
|
||||
if (scale == 0) continue;
|
||||
|
||||
int adr_k = rowadr[k];
|
||||
int nnz_k = rownnz[k];
|
||||
const int* ind_k = colind + adr_k;
|
||||
const mjtNum* val_k = mat + adr_k;
|
||||
|
||||
for (int j = 0; j < nnz_k; j++) {
|
||||
int c = ind_k[j];
|
||||
if (c > r) break;
|
||||
buffer[c] += scale * val_k[j];
|
||||
}
|
||||
}
|
||||
|
||||
// scatter from dense buffer to sparse result
|
||||
int res_adr = res_rowadr[r];
|
||||
int res_nnz = res_rownnz[r];
|
||||
const int* res_ind = res_colind + res_adr;
|
||||
mjtNum* res_val = res + res_adr;
|
||||
|
||||
for (int j = 0; j < res_nnz; j++) {
|
||||
int c = res_ind[j];
|
||||
res_val[j] = buffer[c];
|
||||
buffer[c] = 0;
|
||||
}
|
||||
}
|
||||
|
||||
// supernode: ns > 1 rows share the same sparsity pattern
|
||||
else {
|
||||
int nnzT_r = rownnzT[r];
|
||||
int adr_r = rowadrT[r];
|
||||
|
||||
// accumulate for ns rows
|
||||
for (int i = 0; i < nnzT_r; i++) {
|
||||
int k = colindT[adr_r + i];
|
||||
|
||||
// compute scale for all rows
|
||||
mjtNum scale[mjMAXSUPER];
|
||||
if (diag) {
|
||||
mjtNum dk = diag[k];
|
||||
if (dk == 0) continue;
|
||||
for (int s = 0; s < ns; s++) {
|
||||
scale[s] = dk * matT[rowadrT[r + s] + i];
|
||||
}
|
||||
} else {
|
||||
for (int s = 0; s < ns; s++) {
|
||||
scale[s] = matT[rowadrT[r + s] + i];
|
||||
}
|
||||
}
|
||||
|
||||
int adr_k = rowadr[k];
|
||||
int nnz_k = rownnz[k];
|
||||
const int* ind_k = colind + adr_k;
|
||||
const mjtNum* val_k = mat + adr_k;
|
||||
|
||||
for (int j = 0; j < nnz_k; j++) {
|
||||
int c = ind_k[j];
|
||||
if (c > r + ns - 1) break; // skip if beyond block
|
||||
mjtNum v = val_k[j];
|
||||
|
||||
for (int s = 0; s < ns; s++) {
|
||||
if (c <= r + s) {
|
||||
buffer[s * nc + c] += scale[s] * v;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// scatter
|
||||
for (int s = 0; s < ns; s++) {
|
||||
int row = r + s;
|
||||
int res_adr = res_rowadr[row];
|
||||
int res_nnz = res_rownnz[row];
|
||||
const int* res_ind = res_colind + res_adr;
|
||||
mjtNum* res_val = res + res_adr;
|
||||
for (int j = 0; j < res_nnz; j++) {
|
||||
int c = res_ind[j];
|
||||
res_val[j] = buffer[s*nc + c];
|
||||
buffer[s*nc + c] = 0;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// fill upper triangle: mirror values from lower triangle
|
||||
if (res_diagind) {
|
||||
// initialize write positions after diagonal
|
||||
int* upper_pos = mjSTACKALLOC(d, nc, int);
|
||||
for (int r = 0; r < nc; r++) {
|
||||
upper_pos[r] = res_diagind[r] + 1;
|
||||
}
|
||||
|
||||
// for each (r, c) with c < r, write r to row c
|
||||
for (int r = 0; r < nc; r++) {
|
||||
int adr = res_rowadr[r];
|
||||
int lower_nnz = res_diagind[r] - adr + 1;
|
||||
for (int j = 0; j < lower_nnz; j++) {
|
||||
int c = res_colind[adr + j];
|
||||
if (c < r) {
|
||||
res[upper_pos[c]++] = res[adr + j];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mj_freeStack(d);
|
||||
}
|
||||
|
||||
|
||||
// compute sparse M'*diag*M (diag=NULL: compute M'*M), res_rowadr must be precomputed
|
||||
void mju_sqrMatTDSparse(mjtNum* res, const mjtNum* mat, const mjtNum* matT,
|
||||
const mjtNum* diag, int nr, int nc,
|
||||
@@ -1160,7 +1466,7 @@ void mju_blockDiagSparse(mjtNum* restrict res, int* restrict res_rownnz,
|
||||
}
|
||||
|
||||
// end of block reached: update block counter, column offset, next row
|
||||
if (r + 1 >= row_next && block + 1 < nb ) {
|
||||
if (r + 1 >= row_next && block + 1 < nb) {
|
||||
block++;
|
||||
col_offset = block_c[block];
|
||||
row_next = block + 1 < nb ? block_r[block + 1] : nr;
|
||||
|
||||
Reference in New Issue
Block a user