Speed up mju_sqrMatTDSparse

PiperOrigin-RevId: 759593417
Change-Id: I0168e1f96333769d09d61e835560910d43aac608
This commit is contained in:
Yuval Tassa
2025-05-16 06:41:08 -07:00
committed by Copybara-Service
parent 81442e06a0
commit ee8abdf854
4 changed files with 335 additions and 21 deletions
+205 -2
View File
@@ -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);
+10
View File
@@ -110,6 +110,16 @@ MJAPI void mju_sqrMatTDSparse(mjtNum* res, const mjtNum* mat, const mjtNum* matT
const int* colindT, const int* rowsuperT,
mjData* d, int* diagind);
// LEGACY: row-based implementation
MJAPI 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);
// precount res_rownnz and precompute res_rowadr for mju_sqrMatTDSparse
MJAPI void mju_sqrMatTDSparseCount(int* res_rownnz, int* res_rowadr, int nr,
const int* rownnz, const int* rowadr, const int* colind,
@@ -613,18 +613,25 @@ static void BM_sqrMatTDSparse(benchmark::State& state, SqrMatTDFuncPtr func) {
}
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_sqrMatTDSparse_new(benchmark::State& state) {
BM_sqrMatTDSparse_col(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_sqrMatTDSparse(state, &mju_sqrMatTDSparse);
}
BENCHMARK(BM_sqrMatTDSparse_new);
BENCHMARK(BM_sqrMatTDSparse_col);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_sqrMatTDSparse_old(benchmark::State& state) {
BM_sqrMatTDSparse_row(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_sqrMatTDSparse(state, &mju_sqrMatTDSparse_row);
}
BENCHMARK(BM_sqrMatTDSparse_row);
void ABSL_ATTRIBUTE_NO_TAIL_CALL
BM_sqrMatTDSparse_uncompressed(benchmark::State& state) {
MujocoErrorTestGuard guard;
BM_sqrMatTDSparse(state, nullptr);
}
BENCHMARK(BM_sqrMatTDSparse_old);
BENCHMARK(BM_sqrMatTDSparse_uncompressed);
} // namespace
} // namespace mujoco
+109 -15
View File
@@ -377,9 +377,53 @@ TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse1) {
mj_deleteModel(model);
}
TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparseLower) {
// 2 -1 1
// M = 2 -1 2
// 2 2 3
mjModel* model = LoadModelFromString("<mujoco/>");
mjData* data = mj_makeData(model);
mjtNum mat[] = {2, -1, 1, 2, -1, 2, 2, 2, 3};
int colind[] = {0, 1, 2, 0, 1, 2, 0, 1, 2};
int rownnz[] = {3, 3, 3};
int rowadr[] = {0, 3, 6};
mjtNum matT[] = {2, 2, 2, -1, -1, 2, 1, 2, 3};
int colindT[] = {0, 1, 2, 0, 1, 2, 0, 1, 2};
int rownnzT[] = {3, 3, 3};
int rowadrT[] = {0, 3, 6};
mjtNum matH[] = {0, 0, 0, 0, 0, 0, 0, 0, 0};
int colindH[] = {0, 0, 0, 0, 0, 0, 0, 0, 0};
int rownnzH[] = {0, 0, 0};
int rowadrH[] = {0, 0, 0};
// test precount
mju_sqrMatTDSparseCount(rownnzH, rowadrH, 3, rownnz, rowadr, colind,
rownnzT, rowadrT, colindT, nullptr, data, 0);
EXPECT_THAT(rownnzH, ElementsAre(1, 2, 3));
EXPECT_THAT(rowadrH, ElementsAre(0, 1, 3));
// test computation
mju_sqrMatTDUncompressedInit(rowadrH, 3);
mju_sqrMatTDSparse(matH, mat, matT, nullptr, 3, 3, rownnzH, rowadrH, colindH,
rownnz, rowadr, colind, nullptr, rownnzT, rowadrT, colindT,
nullptr, data, nullptr);
EXPECT_THAT(matH, ElementsAre(12, 0, 0, 0, 6, 0, 12, 3, 14));
EXPECT_THAT(colindH, ElementsAre(0, 0, 0, 0, 1, 0, 0, 1, 2));
EXPECT_THAT(rownnzH, ElementsAre(1, 2, 3));
EXPECT_THAT(rowadrH, ElementsAre(0, 3, 6));
mj_deleteData(data);
mj_deleteModel(model);
}
TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse2) {
// 2 -1 1
// M = 1 2 -1
// M = 2 -1 2
// 2 2 3
mjModel* model = LoadModelFromString("<mujoco/>");
@@ -419,6 +463,7 @@ TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse2) {
EXPECT_THAT(colindH, ElementsAre(0, 1, 2, 0, 1, 2, 0, 1, 2));
EXPECT_THAT(rownnzH, ElementsAre(3, 3, 3));
EXPECT_THAT(rowadrH, ElementsAre(0, 3, 6));
EXPECT_THAT(diagindH, ElementsAre(0, 4, 8));
mj_deleteData(data);
mj_deleteModel(model);
@@ -464,14 +509,63 @@ TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse3) {
nullptr, data, diagindH);
EXPECT_THAT(matH, ElementsAre(66, 4, 0, 4, 35, 0, 0, 0, 0));
EXPECT_THAT(colindH, ElementsAre(0, 1, 0, 0, 1, 0, 0, 0, 0));
EXPECT_THAT(rownnzH, ElementsAre(2, 2, 0));
EXPECT_THAT(colindH, ElementsAre(0, 1, 0, 0, 1, 0, 2, 0, 0));
EXPECT_THAT(rownnzH, ElementsAre(2, 2, 1));
EXPECT_THAT(rowadrH, ElementsAre(0, 3, 6));
mj_deleteData(data);
mj_deleteModel(model);
}
TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse3b) {
// 1 2 0
// M = 0 3 4
// 5 0 0
mjModel* model = LoadModelFromString("<mujoco/>");
mjData* data = mj_makeData(model);
mjtNum mat[] = {1, 2, 3, 4, 5};
int colind[] = {0, 1, 1, 2, 0};
int rownnz[] = {2, 2, 1};
int rowadr[] = {0, 2, 4};
mjtNum matT[] = {1, 5, 2, 3, 4};
int colindT[] = {0, 2, 0, 1, 1};
int rownnzT[] = {2, 2, 1};
int rowadrT[] = {0, 2, 4};
mjtNum matH[] = {0, 0, 0, 0, 0, 0, 0, 0, 0};
int colindH[] = {0, 0, 0, 0, 0, 0, 0, 0, 0};
int rownnzH[] = {0, 0, 0};
int rowadrH[] = {0, 0, 0};
int diagindH[] = {0, 0, 0};
mjtNum diag[] = {1, 1, 1};
// test precount
mju_sqrMatTDSparseCount(rownnzH, rowadrH, 3, rownnz, rowadr, colind,
rownnzT, rowadrT, colindT, nullptr, data, 1);
EXPECT_THAT(rownnzH, ElementsAre(2, 3, 2));
EXPECT_THAT(rowadrH, ElementsAre(0, 2, 5));
// test computation
mju_sqrMatTDUncompressedInit(rowadrH, 3);
mju_sqrMatTDSparse(matH, mat, matT, diag, 3, 3, rownnzH, rowadrH, colindH,
rownnz, rowadr, colind, nullptr, rownnzT, rowadrT, colindT,
nullptr, data, diagindH);
EXPECT_THAT(matH, ElementsAre(26, 2, 0, 2, 13, 12, 12, 16, 0));
EXPECT_THAT(colindH, ElementsAre(0, 1, 0, 0, 1, 2, 1, 2, 0));
EXPECT_THAT(rownnzH, ElementsAre(2, 3, 2));
EXPECT_THAT(rowadrH, ElementsAre(0, 3, 6));
EXPECT_THAT(diagindH, ElementsAre(0, 4, 7));
mj_deleteData(data);
mj_deleteModel(model);
}
TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse4) {
// 1 0 2
// M = 0 0 3
@@ -513,8 +607,8 @@ TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse4) {
nullptr, data, diagindH);
EXPECT_THAT(matH, ElementsAre(66, 4, 0, 0, 0, 0, 4, 35, 0));
EXPECT_THAT(colindH, ElementsAre(0, 2, 0, 0, 0, 0, 0, 2, 0));
EXPECT_THAT(rownnzH, ElementsAre(2, 0, 2));
EXPECT_THAT(colindH, ElementsAre(0, 2, 0, 1, 0, 0, 0, 2, 0));
EXPECT_THAT(rownnzH, ElementsAre(2, 1, 2));
EXPECT_THAT(rowadrH, ElementsAre(0, 3, 6));
mj_deleteData(data);
@@ -759,19 +853,19 @@ TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse9) {
}
TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse10) {
// 1 1 1
// M = 2 2 2
// 3 3 3
// 1 2 3
// M = 2 3 2
// 3 1 1
mjModel* model = LoadModelFromString("<mujoco/>");
mjData* data = mj_makeData(model);
mjtNum mat[] = {1, 1, 1, 2, 2, 2, 3, 3, 3};
mjtNum mat[] = {1, 2, 3, 2, 3, 2, 3, 1, 1};
int colind[] = {0, 1, 2, 0, 1, 2, 0, 1, 2};
int rownnz[] = {3, 3, 3};
int rowadr[] = {0, 3, 6};
mjtNum matT[] = {1, 2, 3, 1, 2, 3, 1, 2, 3};
mjtNum matT[] = {1, 2, 3, 2, 3, 1, 3, 2, 1};
int colindT[] = {0, 1, 2, 0, 1, 2, 0, 1, 2};
int rownnzT[] = {3, 3, 3};
int rowadrT[] = {0, 3, 6};
@@ -783,7 +877,7 @@ TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse10) {
int rowadrH[] = {0, 0, 0};
int diagindH[] = {0, 0, 0};
mjtNum diag[] = {1, 1, 1};
mjtNum diag[] = {1, 2, 1};
// test precount
mju_sqrMatTDSparseCount(rownnzH, rowadrH, 3, rownnz, rowadr, colind,
@@ -798,7 +892,7 @@ TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse10) {
rownnz, rowadr, colind, nullptr, rownnzT, rowadrT, colindT,
rowsuperT, data, diagindH);
EXPECT_THAT(matH, ElementsAre(14, 14, 14, 14, 14, 14, 14, 14, 14));
EXPECT_THAT(matH, ElementsAre(18, 17, 14, 17, 23, 19, 14, 19, 18));
EXPECT_THAT(colindH, ElementsAre(0, 1, 2, 0, 1, 2, 0, 1, 2));
EXPECT_THAT(rownnzH, ElementsAre(3, 3, 3));
EXPECT_THAT(rowadrH, ElementsAre(0, 3, 6));
@@ -951,9 +1045,9 @@ TEST_F(EngineUtilSparseTest, MjuSqrMatTDSparse13) {
EXPECT_THAT(matH, ElementsAre(3, 3, 0, 0, 0, 3, 3, 0, 0, 0, 0, 0, 0, 0, 0, 0,
0, 0, 0, 0, 0, 0, 0, 0, 0));
EXPECT_THAT(colindH, ElementsAre(0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0,
0, 0, 0, 0, 0, 0, 0, 0, 0, 0));
EXPECT_THAT(rownnzH, ElementsAre(2, 2, 0, 0, 0));
EXPECT_THAT(colindH, ElementsAre(0, 1, 0, 0, 0, 0, 1, 0, 0, 0, 2, 0, 0, 0, 0,
3, 0, 0, 0, 0, 4, 0, 0, 0, 0));
EXPECT_THAT(rownnzH, ElementsAre(2, 2, 1, 1, 1));
EXPECT_THAT(rowadrH, ElementsAre(0, 5, 10, 15, 20));
mj_deleteData(data);