Add private functions mju_blockDiag and mju_blockDiagSparse
PiperOrigin-RevId: 752815446 Change-Id: Ia7f73160315ce57b78a2c34bc781f00e1b0f05c4
This commit is contained in:
committed by
Copybara-Service
parent
d657e628f0
commit
8078a45727
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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 ------------------------------------------------
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user