Speed up sparse supernode detection by combining it with transposition.

PiperOrigin-RevId: 760673008
Change-Id: I6d66580e675fd86e5b974859383f93f87482ca16
This commit is contained in:
Yuval Tassa
2025-05-19 10:13:20 -07:00
committed by Copybara-Service
parent ced630181d
commit 45fc15b844
6 changed files with 115 additions and 36 deletions
+2 -6
View File
@@ -1996,7 +1996,7 @@ void mj_makeConstraint(const mjModel* m, mjData* d) {
if (mj_isSparse(m)) {
// transpose
mju_transposeSparse(d->efc_JT, d->efc_J, d->nefc, m->nv,
d->efc_JT_rownnz, d->efc_JT_rowadr, d->efc_JT_colind,
d->efc_JT_rownnz, d->efc_JT_rowadr, d->efc_JT_colind, d->efc_JT_rowsuper,
d->efc_J_rownnz, d->efc_J_rowadr, d->efc_J_colind);
@@ -2010,10 +2010,6 @@ void mj_makeConstraint(const mjModel* m, mjData* d) {
__msan_allocated_memory(d->efc_J_rowsuper, d->nefc);
#endif // MEMORY_SANITIZER
#endif // mjUSEAVX
// supernodes of JT
mju_superSparse(m->nv, d->efc_JT_rowsuper,
d->efc_JT_rownnz, d->efc_JT_rowadr, d->efc_JT_colind);
}
// compute diagApprox
@@ -2182,7 +2178,7 @@ void mj_projectConstraint(const mjModel* m, mjData* d) {
int* BT_colind = mjSTACKALLOC(d, nB, int);
mjtNum* BT = mjSTACKALLOC(d, nB, mjtNum);
mju_transposeSparse(BT, B, nefc, nv,
BT_rownnz, BT_rowadr, BT_colind,
BT_rownnz, BT_rowadr, BT_colind, NULL,
B_rownnz, B_rowadr, B_colind);
// allocate AR row nonzeros and addresses on arena
+1 -1
View File
@@ -1572,7 +1572,7 @@ static void MakeHessian(mjData* d, mjCGContext* ctx) {
int* HT_rowadr = mjSTACKALLOC(d, nv, int);
int* HT_colind = mjSTACKALLOC(d, ctx->nH, int);
mju_transposeSparse(NULL, NULL, nv, nv,
HT_rownnz, HT_rowadr, HT_colind,
HT_rownnz, HT_rowadr, HT_colind, NULL,
ctx->H_rownnz, ctx->H_rowadr, ctx->H_colind);
// count total and row non-zeros of reverse-Cholesky factor L
+33 -7
View File
@@ -531,9 +531,9 @@ int mju_compressSparse(mjtNum* mat, int nr, int nc, int* rownnz, int* rowadr, in
// transpose sparse matrix
// transpose sparse matrix, optionally compute row supernodes
void mju_transposeSparse(mjtNum* res, const mjtNum* mat, int nr, int nc,
int* res_rownnz, int* res_rowadr, int* res_colind,
int* res_rownnz, int* res_rowadr, int* res_colind, int* res_rowsuper,
const int* rownnz, const int* rowadr, const int* colind) {
// clear number of non-zeros for each row of transposed
mju_zeroInt(res_rownnz, nc);
@@ -547,22 +547,40 @@ void mju_transposeSparse(mjtNum* res, const mjtNum* mat, int nr, int nc,
}
}
// init res_rowsuper
if (res_rowsuper) {
for (int i = 0; i < nc - 1; i++) {
res_rowsuper[i] = (res_rownnz[i] == res_rownnz[i + 1]);
}
res_rowsuper[nc - 1] = 0;
}
// compute the row addresses for the transposed matrix
res_rowadr[0] = 0;
for (int i = 1; i < nc; i++) {
res_rowadr[i] = res_rowadr[i-1] + res_rownnz[i-1];
}
// iterate through each non-zero entry of mat
// iterate through each row (column) of mat (res)
for (int r = 0; r < nr; r++) {
int c_prev = -1;
int start = rowadr[r];
int end = start + rownnz[r];
for (int i = start; i < end; i++) {
// swap rows with columns and increment res_rowadr
int c = res_rowadr[colind[i]]++;
res_colind[c] = r;
int c = colind[i];
int adr = res_rowadr[c]++;
res_colind[adr] = r;
if (res) {
res[c] = mat[i];
res[adr] = mat[i];
}
// mark non-supernodes
if (res_rowsuper) {
if (c > 0 && c != c_prev + 1 && res_rowsuper[c - 1]) {
res_rowsuper[c - 1] = 0;
}
c_prev = c;
}
}
}
@@ -571,8 +589,16 @@ void mju_transposeSparse(mjtNum* res, const mjtNum* mat, int nr, int nc,
for (int i = nc-1; i > 0; i--) {
res_rowadr[i] = res_rowadr[i-1];
}
res_rowadr[0] = 0;
// accumulate supernodes
if (res_rowsuper) {
for (int i = nc - 2; i >= 0; i--) {
if (res_rowsuper[i]) {
res_rowsuper[i] += res_rowsuper[i + 1];
}
}
}
}
+2 -2
View File
@@ -91,9 +91,9 @@ int mju_addToSparseMat(mjtNum* dst, const mjtNum* src, int n, int nrow, mjtNum s
int mju_addChains(int* res, int n, int NV1, int NV2,
const int* chain1, const int* chain2);
// transpose sparse matrix
// transpose sparse matrix, optionally compute row supernodes
MJAPI void mju_transposeSparse(mjtNum* res, const mjtNum* mat, int nr, int nc,
int* res_rownnz, int* res_rowadr, int* res_colind,
int* res_rownnz, int* res_rowadr, int* res_colind, int* res_rowsuper,
const int* rownnz, const int* rowadr, const int* colind);
// construct row supernodes