Speed up mju_sqrMatTDSparse
PiperOrigin-RevId: 759593417 Change-Id: I0168e1f96333769d09d61e835560910d43aac608
This commit is contained in:
committed by
Copybara-Service
parent
81442e06a0
commit
ee8abdf854
@@ -18,6 +18,7 @@
|
||||
#include <string.h>
|
||||
|
||||
#include <mujoco/mjdata.h>
|
||||
#include <mujoco/mjmacro.h>
|
||||
#include <mujoco/mjsan.h> // IWYU pragma: keep
|
||||
#include <mujoco/mjtnum.h>
|
||||
#include "engine/engine_io.h"
|
||||
@@ -711,8 +712,10 @@ void mju_sqrMatTDUncompressedInit(int* res_rowadr, int nc) {
|
||||
|
||||
|
||||
|
||||
// compute sparse M'*diag*M (diag=NULL: compute M'*M), res has uncompressed layout
|
||||
// res_rowadr is required to be precomputed
|
||||
// max number of supernodes handled
|
||||
#define mjMAXSUPER 8
|
||||
|
||||
// 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,
|
||||
int* res_rownnz, const int* res_rowadr, int* res_colind,
|
||||
@@ -721,6 +724,206 @@ void mju_sqrMatTDSparse(mjtNum* res, const mjtNum* mat, const mjtNum* matT,
|
||||
const int* rownnzT, const int* rowadrT,
|
||||
const int* colindT, const int* rowsuperT,
|
||||
mjData* d, int* diagind) {
|
||||
mj_markStack(d);
|
||||
|
||||
// reinterpret transposed matrices as compressed sparse column
|
||||
const mjtNum* mat_csc = matT;
|
||||
const int* colnnz = rownnzT;
|
||||
const int* coladr = rowadrT;
|
||||
const int* rowind = colindT;
|
||||
const int* colsuper = rowsuperT;
|
||||
const mjtNum* matT_csc = mat;
|
||||
const int* colnnzT = rownnz;
|
||||
const int* coladrT = rowadr;
|
||||
const int* rowindT = colind;
|
||||
// rowsuper is unused
|
||||
|
||||
// marker[i] = 1 if row i is set in current column
|
||||
int* marker = mjSTACKALLOC(d, nc, int);
|
||||
mju_zeroInt(marker, nc);
|
||||
|
||||
// dense buffer (considered column-major) containing up to mjMAXSUPER columns
|
||||
mjtNum* buffer = mjSTACKALLOC(d, nc*mjMAXSUPER, mjtNum);
|
||||
|
||||
// dense index vector of the current column (unsorted)
|
||||
int* buffer_idx = mjSTACKALLOC(d, nc, int);
|
||||
|
||||
// rowstart[i]: address of first row in column mat'[:, i] with index > current column
|
||||
int* rowstart = mjSTACKALLOC(d, nr, int);
|
||||
mju_zeroInt(rowstart, nr);
|
||||
|
||||
// clear res_rownnz
|
||||
mju_zeroInt(res_rownnz, nc);
|
||||
|
||||
// construct res[lower+diagonal], by column
|
||||
for (int c=0; c < nc; c++) {
|
||||
int buffer_nnz = 0;
|
||||
|
||||
// prepare column c of mat
|
||||
int nnz = colnnz[c];
|
||||
int adr = coladr[c];
|
||||
const int* ind = rowind + adr;
|
||||
|
||||
// val: array of ns > 0 column pointers with identical pattern to c
|
||||
const mjtNum* val[mjMAXSUPER];
|
||||
|
||||
// first column is c
|
||||
int ns = 1;
|
||||
val[0] = mat_csc + adr;
|
||||
|
||||
// add c's supernodes, if any
|
||||
int cs;
|
||||
if (colsuper && (cs = colsuper[c])) {
|
||||
ns += mjMIN(cs, mjMAXSUPER - 1);
|
||||
for (int s=1; s < ns; s++) {
|
||||
val[s] = mat_csc + coladr[c + s];
|
||||
}
|
||||
}
|
||||
|
||||
// diagonal special-case: dense dot product of column c, with/out diag
|
||||
mjtNum diag_c[mjMAXSUPER];
|
||||
if (diag) {
|
||||
for (int s=0; s < ns; s++) {
|
||||
mjtNum ds = 0;
|
||||
for (int k=0; k < nnz; k++) {
|
||||
ds += (val[s][k] * val[s][k]) * diag[ind[k]];
|
||||
}
|
||||
diag_c[s] = ds;
|
||||
}
|
||||
} else {
|
||||
for (int s=0; s < ns; s++) {
|
||||
diag_c[s] = mju_dot(val[s], val[s], nnz);
|
||||
}
|
||||
}
|
||||
|
||||
// in the strict lower triangle, compute
|
||||
// res[:, c] = mat' * mat[:, c] = sum_r(diag[r] * mat'[:, r] * mat[:, c])
|
||||
for (int i=0; i < nnz; i++) {
|
||||
// prepare column r of mat'
|
||||
int r = ind[i];
|
||||
int adrT = coladrT[r];
|
||||
int nnzT = colnnzT[r];
|
||||
const int* indT = rowindT + adrT;
|
||||
const mjtNum* valT = matT_csc + adrT;
|
||||
|
||||
// get v[s] = diag[r] * mat[r, c + s] for s in [0, ns)
|
||||
mjtNum v[mjMAXSUPER];
|
||||
if (diag) {
|
||||
mjtNum diag_r = diag[r];
|
||||
for (int s=0; s < ns; s++) {
|
||||
v[s] = diag_r * val[s][i];
|
||||
}
|
||||
} else {
|
||||
for (int s=0; s < ns; s++) {
|
||||
v[s] = val[s][i];
|
||||
}
|
||||
}
|
||||
|
||||
// gather to dense buffer columns: buffer[:, s] += mat'[:, r] * v[s]
|
||||
for (int k=rowstart[r]; k < nnzT; k++) {
|
||||
int j = indT[k];
|
||||
|
||||
// if j is not in the strict lower triangle, increment rowstart and continue
|
||||
if (j <= c) {
|
||||
rowstart[r]++;
|
||||
continue;
|
||||
}
|
||||
|
||||
// first nonzero in row j: mark and set value
|
||||
if (!marker[j]) {
|
||||
// mark j and save it
|
||||
marker[j] = 1;
|
||||
buffer_idx[buffer_nnz++] = j;
|
||||
|
||||
// set value
|
||||
mjtNum vk = valT[k];
|
||||
for (int s=0; s < ns; s++) {
|
||||
buffer[s*nc + j] = vk * v[s];
|
||||
}
|
||||
}
|
||||
|
||||
// otherwise existing nonzero in row j: add to value
|
||||
else {
|
||||
mjtNum vk = valT[k];
|
||||
for (int s=0; s < ns; s++) {
|
||||
buffer[s*nc + j] += vk * v[s];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// scatter to res from dense buffer: res[:, c + s] = buffer[:, s] for s in [0, ns)
|
||||
|
||||
// write values under diagonal
|
||||
for (int i=0; i < buffer_nnz; i++) {
|
||||
int j = buffer_idx[i];
|
||||
marker[j] = 0;
|
||||
int adr_j = res_rowadr[j] + res_rownnz[j];
|
||||
|
||||
// truncate row to strict lower triangle
|
||||
int lower = j - c;
|
||||
int nm = mjMIN(ns, lower);
|
||||
|
||||
// increment nonzeros
|
||||
res_rownnz[j] += nm;
|
||||
|
||||
// write value
|
||||
for (int s=0; s < nm; s++) {
|
||||
res[adr_j + s] = buffer[s*nc + j];
|
||||
}
|
||||
|
||||
// write index
|
||||
for (int s=0; s < nm; s++) {
|
||||
res_colind[adr_j + s] = c + s;
|
||||
}
|
||||
}
|
||||
|
||||
// write diagonal value
|
||||
for (int s=0; s < ns; s++) {
|
||||
int adr_s = res_rowadr[c + s] + res_rownnz[c + s]++;
|
||||
res_colind[adr_s] = c + s;
|
||||
res[adr_s] = diag_c[s];
|
||||
}
|
||||
|
||||
// supernode: skip ahead if ns > 1
|
||||
c += ns - 1;
|
||||
}
|
||||
|
||||
// upper triangle requested: save diagonal indices and fill
|
||||
if (diagind) {
|
||||
// save diagonal indices
|
||||
for (int i=0; i < nc; i++) {
|
||||
diagind[i] = res_rowadr[i] + res_rownnz[i] - 1;
|
||||
}
|
||||
|
||||
// fill upper triangle
|
||||
for (int i=0; i < nc; i++) {
|
||||
int start = res_rowadr[i];
|
||||
int end = start + res_rownnz[i] - 1;
|
||||
for (int j=start; j < end; j++) {
|
||||
int adr = res_rowadr[res_colind[j]] + res_rownnz[res_colind[j]]++;
|
||||
res[adr] = res[j];
|
||||
res_colind[adr] = i;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mj_freeStack(d);
|
||||
}
|
||||
|
||||
#undef mjMAXSUPER
|
||||
|
||||
|
||||
|
||||
// legacy row-based implementation (reference)
|
||||
void mju_sqrMatTDSparse_row(mjtNum* res, const mjtNum* mat, const mjtNum* matT,
|
||||
const mjtNum* diag, int nr, int nc,
|
||||
int* res_rownnz, const 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* diagind) {
|
||||
// allocate space for accumulation buffer and matT
|
||||
mj_markStack(d);
|
||||
|
||||
|
||||
Reference in New Issue
Block a user