Refactor mju_combineSparse to eliminate temporary buffers.

Combine sparse vectors in-place by first counting total `nnz` and then working backwards from the end. This removes the need for temporary buffers in `mju_combineSparse` and its callers and speeds up the function by ~10%.

PiperOrigin-RevId: 902530210
Change-Id: I4f48c327103552ab968d3915399c6067367bec9f
This commit is contained in:
Yuval Tassa
2026-04-20 03:09:31 -07:00
committed by Copybara-Service
parent 3230cf99f9
commit e6d77650f7
8 changed files with 57 additions and 86 deletions
+1 -7
View File
@@ -22,7 +22,6 @@
#include <absl/base/attributes.h>
#include <mujoco/mjdata.h>
#include <mujoco/mujoco.h>
#include "src/engine/engine_memory.h"
#include "src/engine/engine_support.h"
#include "src/engine/engine_util_solve.h"
#include "src/engine/engine_util_sparse.h"
@@ -352,10 +351,6 @@ constexpr int kNumUpdateVectors = 25;
int ABSL_ATTRIBUTE_NOINLINE mju_cholUpdateSparse_old(
mjtNum* mat, mjtNum* x, int n, int flg_plus, const int* rownnz,
const int* rowadr, const int* colind, int x_nnz, int* x_ind, mjData* d) {
mj_markStack(d);
int* buf_ind = mjSTACKALLOC(d, n, int);
mjtNum* sparse_buf = mjSTACKALLOC(d, n, mjtNum);
int rank = n, i = x_nnz - 1;
while (i >= 0) {
int nnz = rownnz[x_ind[i]], adr = rowadr[x_ind[i]];
@@ -372,10 +367,9 @@ int ABSL_ATTRIBUTE_NOINLINE mju_cholUpdateSparse_old(
mju_combineSparseInc(mat + adr, x, n, 1 / c, (flg_plus ? s / c : -s / c),
nnz - 1, i, colind + adr, x_ind);
int new_x_nnz = mju_combineSparse(x, mat + adr, c, -s, i, nnz - 1, x_ind,
colind + adr, sparse_buf, buf_ind);
colind + adr);
i = i - 1 + (new_x_nnz - i);
}
mj_freeStack(d);
return rank;
}
@@ -99,8 +99,7 @@ int ABSL_ATTRIBUTE_NOINLINE combineSparse_baseline(mjtNum* dst,
mjtNum a, mjtNum b,
int dst_nnz, int src_nnz,
int* dst_ind,
const int* src_ind,
mjtNum* buf, int* buf_ind) {
const int* src_ind) {
// check for identical pattern
if (compare_baseline(dst_ind, src_ind, dst_nnz)) {
// combine mjtNum data directly
@@ -116,8 +115,7 @@ int ABSL_ATTRIBUTE_NOINLINE combineSparse_new(mjtNum* dst,
mjtNum a, mjtNum b,
int dst_nnz, int src_nnz,
int* dst_ind,
const int* src_ind,
mjtNum* buf, int* buf_ind) {
const int* src_ind) {
// check for identical pattern
if (compare_memcmp(dst_ind, src_ind, dst_nnz)) {
// combine mjtNum data directly
@@ -372,7 +370,7 @@ static void BM_combineSparse(benchmark::State& state, CombineFuncPtr func) {
// in order to trigger all if's in combineSparse
func(H+rowadr[c], H+rowadr[r], 1, -H[adr+i],
rownnz[c], rownnz[c],
colind+rowadr[c], colind+rowadr[c], NULL, NULL);
colind+rowadr[c], colind+rowadr[c]);
}
}
}