From ee8abdf854b46179ecb912f723ca9f371555354a Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Fri, 16 May 2025 06:41:08 -0700 Subject: [PATCH] Speed up mju_sqrMatTDSparse PiperOrigin-RevId: 759593417 Change-Id: I0168e1f96333769d09d61e835560910d43aac608 --- src/engine/engine_util_sparse.c | 207 +++++++++++++++++- src/engine/engine_util_sparse.h | 10 + .../engine_util_sparse_benchmark_test.cc | 15 +- test/engine/engine_util_sparse_test.cc | 124 +++++++++-- 4 files changed, 335 insertions(+), 21 deletions(-) diff --git a/src/engine/engine_util_sparse.c b/src/engine/engine_util_sparse.c index 17e960a3..44dddb92 100644 --- a/src/engine/engine_util_sparse.c +++ b/src/engine/engine_util_sparse.c @@ -18,6 +18,7 @@ #include #include +#include #include // IWYU pragma: keep #include #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); diff --git a/src/engine/engine_util_sparse.h b/src/engine/engine_util_sparse.h index ab7ea0f3..33f65513 100644 --- a/src/engine/engine_util_sparse.h +++ b/src/engine/engine_util_sparse.h @@ -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, diff --git a/test/benchmark/engine_util_sparse_benchmark_test.cc b/test/benchmark/engine_util_sparse_benchmark_test.cc index 6be10ded..6a1b3835 100644 --- a/test/benchmark/engine_util_sparse_benchmark_test.cc +++ b/test/benchmark/engine_util_sparse_benchmark_test.cc @@ -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 diff --git a/test/engine/engine_util_sparse_test.cc b/test/engine/engine_util_sparse_test.cc index 7fee2127..edbd2558 100644 --- a/test/engine/engine_util_sparse_test.cc +++ b/test/engine/engine_util_sparse_test.cc @@ -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(""); + 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(""); @@ -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(""); + 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(""); 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);