diff --git a/src/engine/engine_util_sparse.c b/src/engine/engine_util_sparse.c index d72a9fca..23c038ec 100644 --- a/src/engine/engine_util_sparse.c +++ b/src/engine/engine_util_sparse.c @@ -795,3 +795,103 @@ void mju_sqrMatTDSparse(mjtNum* res, const mjtNum* mat, const mjtNum* matT, mj_freeStack(d); } + + + +// block-diagonalize a dense matrix +// res output matrix +// mat input matrix +// nc_mat number of columns in mat +// nc_res number of columns in res +// nb number of blocks +// perm_r reverse permutation of rows (res -> mat) +// perm_c reverse permutation of columns (res -> mat) +// block_nr number of rows in each block +// block_nc number of columns in each block +// block_r first row of each block +// block_c first column of each block +void mju_blockDiag(mjtNum* restrict res, const mjtNum* restrict mat, + int nc_mat, int nc_res, int nb, + const int* restrict perm_r, const int* restrict perm_c, + const int* restrict block_nr, const int* restrict block_nc, + const int* restrict block_r, const int* restrict block_c) { + for (int b=0; b < nb; b++) { + int bnr = block_nr[b]; + int bnc = block_nc[b]; + const int* adr_r = perm_r + block_r[b]; + const int* adr_c = perm_c + block_c[b]; + int adr = nc_res * block_r[b]; + for (int r = 0; r < bnr; r++) { + for (int c = 0; c < bnc; c++) { + res[adr++] = mat[nc_mat * adr_r[r] + adr_c[c]]; + } + } + } +} + + +// block-diagonalize a sparse matrix +// res values of the target matrix res +// res_rownnz number of non-zeros in each row of res +// res_rowadr row address of each non-zero in res +// res_colind column index of each non-zero in res +// mat values of the source matrix mat +// mat_rownnz number of non-zeros in each row of mat +// mat_rowadr row address of each non-zero in mat +// mat_colind column index of each non-zero in mat +// nr number of rows in mat/res +// nb number of blocks +// perm_r reverse permutation of rows (res -> mat) +// perm_c forward permutation of columns (mat -> res) +// block_r first row of each block in res +// block_c first column of each block in res +// mat2 optional additional source matrix (same structure as mat) +// res2 optional additional target matrix (same structure as res) +void mju_blockDiagSparse(mjtNum* restrict res, int* restrict res_rownnz, + int* restrict res_rowadr, int* restrict res_colind, + const mjtNum* restrict mat, const int* restrict rownnz, + const int* restrict rowadr, const int* restrict colind, + int nr, int nb, + const int* restrict perm_r, const int* restrict perm_c, + const int* restrict block_r, const int* restrict block_c, + mjtNum* restrict res2, const mjtNum* restrict mat2) { + int block = 0; + int col_offset = block_c[block]; + int row_next = block + 1 < nb ? block_r[block + 1] : nr; + for (int r=0; r < nr; r++) { + // row k in mat goes to row r in res + int k = perm_r[r]; + + // rownnz + int nnz = rownnz[k]; + res_rownnz[r] = nnz; + + // rowadr + int res_adr = (r == 0) ? 0 : res_rowadr[r-1] + res_rownnz[r-1]; + res_rowadr[r] = res_adr; + + // colind + int* res_colind_r = res_colind + res_adr; + mjtNum* res_r = res + res_adr; + int mat_adr = rowadr[k]; + const int* colind_k = colind + mat_adr; + const mjtNum* mat_k = mat + mat_adr; + for (int j=0; j < nnz; j++) { + res_colind_r[j] = perm_c[colind_k[j]] - col_offset; + } + + // values (dense copy: partial order within block is guaranteed) + mju_copy(res_r, mat_k, nnz); + if (mat2 && res2) { + mju_copy(res2 + res_adr, mat2 + mat_adr, nnz); + } + + // end of block reached: update block counter, column offset, next row + if (r + 1 >= row_next && block + 1 < nb ) { + block++; + col_offset = block_c[block]; + row_next = block + 1 < nb ? block_r[block + 1] : nr; + } + } +} + diff --git a/src/engine/engine_util_sparse.h b/src/engine/engine_util_sparse.h index 08a663b2..1f75739b 100644 --- a/src/engine/engine_util_sparse.h +++ b/src/engine/engine_util_sparse.h @@ -103,6 +103,21 @@ MJAPI void mju_sqrMatTDSparseCount(int* res_rownnz, int* res_rowadr, int nr, // precompute res_rowadr for mju_sqrMatTDSparse using uncompressed memory MJAPI void mju_sqrMatTDUncompressedInit(int* res_rowadr, int nc); +// block-diagonalize a dense matrix +MJAPI void mju_blockDiag(mjtNum* res, const mjtNum* mat, + int nc_mat, int nc_res, int nb, + const int* perm_r, const int* perm_c, + const int* block_nr, const int* block_nc, + const int* blockadr_r, const int* blockadr_c); + +// block-diagonalize a sparse matrix +MJAPI void mju_blockDiagSparse( + mjtNum* res, int* res_rownnz, int* res_rowadr, int* res_colind, + const mjtNum* mat, const int* rownnz, const int* rowadr, const int* colind, + int nr, int nb, + const int* perm_r, const int* perm_c, + const int* block_r, const int* block_c, + mjtNum* res2, const mjtNum* mat2); // ------------------------------ inlined functions ------------------------------------------------ diff --git a/test/engine/engine_util_sparse_test.cc b/test/engine/engine_util_sparse_test.cc index 94da0d5d..41231c2c 100644 --- a/test/engine/engine_util_sparse_test.cc +++ b/test/engine/engine_util_sparse_test.cc @@ -1088,5 +1088,227 @@ TEST_F(EngineUtilSparseTest, MergeSorted) { EXPECT_THAT(merged_b, ElementsAre(1, 2, 3, 4, 5, 6, 7, 8)); } +TEST_F(EngineUtilSparseTest, BlockDiag) { + // 4x5 matrix with 3 blocks + constexpr int nr = 4; + constexpr int nc = 5; + const mjtNum mat[nr*nc] = { + 1, 2, 0, 0, 0, + 0, 0, 3, 4, 0, + 0, 0, 5, 6, 0, + 0, 0, 0, 0, 7 + }; + + // block structure + constexpr int nb = 3; + const int block_nr[nb] = {1, 2, 1}; + const int block_nc[nb] = {2, 2, 1}; + const int block_r[nb] = {0, 1, 3}; + const int block_c[nb] = {0, 2, 4}; + + // test with identity permutations + const int perm_r[nr] = {0, 1, 2, 3}; + const int perm_c[nc] = {0, 1, 2, 3, 4}; + mjtNum res[nr*nc] = {0}; + mju_blockDiag(res, mat, nc, nc, nb, + perm_r, perm_c, + block_nr, block_nc, + block_r, block_c); + EXPECT_THAT(res, ElementsAre(1, 2, 0, 0, 0, + 3, 4, 5, 6, 0, + 0, 0, 0, 0, 0, + 7, 0, 0, 0, 0)); +} + +void PermuteMat(mjtNum* res, const mjtNum* mat, int nr, int nc, + const int* perm_r, const int* perm_c, + bool scatter_r, bool scatter_c); + + +TEST_F(EngineUtilSparseTest, BlockDiagPerm) { + // 4x5 matrix with 3 blocks + constexpr int nr = 4; + constexpr int nc = 5; + const mjtNum mat[nr*nc] = { + 1, 2, 0, 0, 0, + 0, 0, 3, 4, 0, + 0, 0, 5, 6, 0, + 0, 0, 0, 0, 7 + }; + + // block structure + constexpr int nb = 3; + const int block_nr[nb] = {1, 2, 1}; + const int block_nc[nb] = {2, 2, 1}; + const int block_r[nb] = {0, 1, 3}; + const int block_c[nb] = {0, 2, 4}; + + // scatter mat into mat_p + const int perm_r[nr] = {1, 3, 2, 0}; + const int perm_c[nc] = {2, 0, 4, 3, 1}; + mjtNum mat_p[nr*nc]; + PermuteMat(mat_p, mat, nr, nc, perm_r, perm_c, true, true); + + // test with permutation + mjtNum res[nr*nc] = {0}; + mju_blockDiag(res, mat_p, nc, nc, nb, + perm_r, perm_c, + block_nr, block_nc, + block_r, block_c); + EXPECT_THAT(res, ElementsAre(1, 2, 0, 0, 0, + 3, 4, 5, 6, 0, + 0, 0, 0, 0, 0, + 7, 0, 0, 0, 0)); +} + +TEST_F(EngineUtilSparseTest, BlockDiagLessCols) { + // 4x5 matrix with 3 blocks + constexpr int nr = 4; + constexpr int nc = 5; + const mjtNum mat[nr*nc] = { + 1, 2, 0, 0, 0, + 0, 0, 3, 4, 0, + 0, 0, 5, 6, 0, + 0, 0, 0, 0, 7 + }; + + // block structure (ignore middle block) + constexpr int nb = 2; + const int block_nr[nb] = {1, 1}; + const int block_nc[nb] = {2, 1}; + const int block_r[nb] = {0, 3}; + const int block_c[nb] = {0, 4}; + + // scatter mat into mat_p + const int perm_r[nr] = {1, 3, 2, 0}; + const int perm_c[nc] = {2, 0, 4, 3, 1}; + mjtNum mat_p[nr*nc]; + PermuteMat(mat_p, mat, nr, nc, perm_r, perm_c, true, true); + + // test with permutation and less columns (ignore middle block) + constexpr int nc_res = 3; + mjtNum res2[nr*nc_res] = {0}; + mju_blockDiag(res2, mat_p, nc, nc_res, nb, + perm_r, perm_c, + block_nr, block_nc, + block_r, block_c); + EXPECT_THAT(res2, ElementsAre(1, 2, 0, + 0, 0, 0, + 0, 0, 0, + 7, 0, 0)); +} + +TEST_F(EngineUtilSparseTest, BlockDiagSparse) { + // 4x5 matrix with 3 blocks + constexpr int nr = 4; + constexpr int nc = 5; + const mjtNum mat[nr*nc] = { + 1, 2, 0, 0, 0, + 0, 0, 3, 4, 0, + 0, 0, 5, 6, 0, + 0, 0, 0, 0, 7 + }; + constexpr int nnz = 7; + + // block structure + constexpr int nb = 3; + const int block_r[nb] = {0, 1, 3}; + const int block_c[nb] = {0, 2, 4}; + + // convert to sparse + int rownnz[nr]; + int rowadr[nr]; + int colind[nnz]; + mjtNum mat_sparse[nnz]; + mju_dense2sparse(mat_sparse, mat, nr, nc, rownnz, rowadr, colind, nnz); + + // test with identity permutations + const int perm_r[nr] = {0, 1, 2, 3}; + const int perm_c[nc] = {0, 1, 2, 3, 4}; + int res_rownnz[nr]; + int res_rowadr[nr]; + int res_colind[nnz]; + mjtNum res[nnz]; + mju_blockDiagSparse(res, res_rownnz, res_rowadr, res_colind, + mat_sparse, rownnz, rowadr, colind, nr, nb, + perm_r, perm_c, + block_r, block_c, nullptr, nullptr); + mjtNum dense_res[nr*nc]; + mju_sparse2dense(dense_res, res, nr, nc, res_rownnz, res_rowadr, res_colind); + EXPECT_THAT(dense_res, ElementsAre(1, 2, 0, 0, 0, + 3, 4, 0, 0, 0, + 5, 6, 0, 0, 0, + 7, 0, 0, 0, 0)); + + // permute mat into mat_p (scatter rows, gather columns) + const int perm_r2[nr] = {3, 1, 0, 2}; + const int perm_c2[nc] = {4, 0, 2, 1, 3}; + mjtNum mat_p[nr*nc]; + PermuteMat(mat_p, mat, nr, nc, perm_r2, perm_c2, true, false); + mju_dense2sparse(mat_sparse, mat_p, nr, nc, rownnz, rowadr, colind, nnz); + + // test with permutation + mju_blockDiagSparse(res, res_rownnz, res_rowadr, res_colind, + mat_sparse, rownnz, rowadr, colind, nr, nb, + perm_r2, perm_c2, + block_r, block_c, nullptr, nullptr); + mju_sparse2dense(dense_res, res, nr, nc, res_rownnz, res_rowadr, res_colind); + EXPECT_THAT(dense_res, ElementsAre(1, 2, 0, 0, 0, + 3, 4, 0, 0, 0, + 5, 6, 0, 0, 0, + 7, 0, 0, 0, 0)); +} + +TEST_F(EngineUtilSparseTest, PermuteMat) { + const mjtNum mat[] = {1, 2, 0, 0, + 0, 0, 3, 4, + 0, 0, 5, 6}; + const int perm_r[] = {2, 0, 1}; + const int perm_c[] = {3, 2, 0, 1}; + mjtNum gather[3*4]; + PermuteMat(gather, mat, 3, 4, perm_r, perm_c, false, false); + EXPECT_THAT(gather, ElementsAre(6, 5, 0, 0, + 0, 0, 1, 2, + 4, 3, 0, 0)); + mjtNum scatter[3*4]; + PermuteMat(scatter, gather, 3, 4, perm_r, perm_c, true, true); + EXPECT_THAT(scatter, ElementsAre(1, 2, 0, 0, + 0, 0, 3, 4, + 0, 0, 5, 6)); + mjtNum mixed[3*4]; + PermuteMat(mixed, mat, 3, 4, perm_r, perm_c, true, false); + EXPECT_THAT(mixed, ElementsAre(4, 3, 0, 0, + 6, 5, 0, 0, + 0, 0, 1, 2)); + mjtNum mixed_back[3*4]; + PermuteMat(mixed_back, mixed, 3, 4, perm_r, perm_c, false, true); + EXPECT_THAT(mixed_back, ElementsAre(1, 2, 0, 0, + 0, 0, 3, 4, + 0, 0, 5, 6)); +} + +// local function for permuting the rows and columns of a dense matrix +void PermuteMat(mjtNum* res, const mjtNum* mat, int nr, int nc, + const int* perm_r, const int* perm_c, + bool scatter_r, bool scatter_c) { + for (int r = 0; r < nr; r++) { + for (int c = 0; c < nc; c++) { + if (scatter_r && scatter_c) { + // scatter both + res[perm_r[r] * nc + perm_c[c]] = mat[r * nc + c]; + } else if (scatter_r && !scatter_c) { + // scatter rows, gather columns + res[perm_r[r] * nc + c] = mat[r * nc + perm_c[c]]; + } else if (!scatter_r && scatter_c) { + // gather rows, scatter columns + res[r * nc + perm_c[c]] = mat[perm_r[r] * nc + c]; + } else { + // gather both + res[r * nc + c] = mat[perm_r[r] * nc + perm_c[c]]; + } + } + } +} + } // namespace } // namespace mujoco