Move local function to top of test file.
PiperOrigin-RevId: 757722664 Change-Id: Ief7d8ca9fda990bf582ca0feaa4b645f0a4bb6e7
This commit is contained in:
committed by
Copybara-Service
parent
af088adaee
commit
9ac5fff41d
@@ -26,6 +26,29 @@
|
||||
namespace mujoco {
|
||||
namespace {
|
||||
|
||||
// permute the rows and columns of a dense matrix
|
||||
inline 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]];
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
using ::testing::ElementsAre;
|
||||
using EngineUtilSparseTest = MujocoTest;
|
||||
|
||||
@@ -1128,12 +1151,10 @@ 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
|
||||
};
|
||||
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;
|
||||
@@ -1145,7 +1166,7 @@ TEST_F(EngineUtilSparseTest, BlockDiag) {
|
||||
// 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};
|
||||
mjtNum res[nr * nc] = {0};
|
||||
mju_blockDiag(res, mat, nc, nc, nb,
|
||||
perm_r, perm_c,
|
||||
block_nr, block_nc,
|
||||
@@ -1156,16 +1177,11 @@ TEST_F(EngineUtilSparseTest, BlockDiag) {
|
||||
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] = {
|
||||
const mjtNum mat[nr * nc] = {
|
||||
1, 2, 0, 0, 0,
|
||||
0, 0, 3, 4, 0,
|
||||
0, 0, 5, 6, 0,
|
||||
@@ -1182,11 +1198,11 @@ TEST_F(EngineUtilSparseTest, BlockDiagPerm) {
|
||||
// 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];
|
||||
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};
|
||||
mjtNum res[nr * nc] = {0};
|
||||
mju_blockDiag(res, mat_p, nc, nc, nb,
|
||||
perm_r, perm_c,
|
||||
block_nr, block_nc,
|
||||
@@ -1201,7 +1217,7 @@ TEST_F(EngineUtilSparseTest, BlockDiagLessCols) {
|
||||
// 4x5 matrix with 3 blocks
|
||||
constexpr int nr = 4;
|
||||
constexpr int nc = 5;
|
||||
const mjtNum mat[nr*nc] = {
|
||||
const mjtNum mat[nr * nc] = {
|
||||
1, 2, 0, 0, 0,
|
||||
0, 0, 3, 4, 0,
|
||||
0, 0, 5, 6, 0,
|
||||
@@ -1218,12 +1234,12 @@ TEST_F(EngineUtilSparseTest, BlockDiagLessCols) {
|
||||
// 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];
|
||||
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};
|
||||
mjtNum res2[nr * nc_res] = {0};
|
||||
mju_blockDiag(res2, mat_p, nc, nc_res, nb,
|
||||
perm_r, perm_c,
|
||||
block_nr, block_nc,
|
||||
@@ -1238,7 +1254,7 @@ TEST_F(EngineUtilSparseTest, BlockDiagSparse) {
|
||||
// 4x5 matrix with 3 blocks
|
||||
constexpr int nr = 4;
|
||||
constexpr int nc = 5;
|
||||
const mjtNum mat[nr*nc] = {
|
||||
const mjtNum mat[nr * nc] = {
|
||||
1, 2, 0, 0, 0,
|
||||
0, 0, 3, 4, 0,
|
||||
0, 0, 5, 6, 0,
|
||||
@@ -1269,7 +1285,7 @@ TEST_F(EngineUtilSparseTest, BlockDiagSparse) {
|
||||
mat_sparse, rownnz, rowadr, colind, nr, nb,
|
||||
perm_r, perm_c,
|
||||
block_r, block_c, nullptr, nullptr);
|
||||
mjtNum dense_res[nr*nc];
|
||||
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,
|
||||
@@ -1301,50 +1317,27 @@ TEST_F(EngineUtilSparseTest, PermuteMat) {
|
||||
0, 0, 5, 6};
|
||||
const int perm_r[] = {2, 0, 1};
|
||||
const int perm_c[] = {3, 2, 0, 1};
|
||||
mjtNum gather[3*4];
|
||||
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];
|
||||
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];
|
||||
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];
|
||||
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