Allow CSR back-substitution to handle multiple vectors.

PiperOrigin-RevId: 713229673
Change-Id: I7a5b43fe966cf9e482bd41e2eea6c30dd3ffa1d4
This commit is contained in:
Yuval Tassa
2025-01-08 03:30:41 -08:00
committed by Copybara-Service
parent 4d82ab5762
commit 8a5f092081
5 changed files with 116 additions and 29 deletions
+39 -1
View File
@@ -495,7 +495,7 @@ TEST_F(CoreSmoothTest, SolveLDs) {
for (int i=0; i < nv; i+=2) vec[i] = vec2[i] = 0;
mj_solveLD(m, vec.data(), 1, d->qLD, d->qLDiagInv);
mj_solveLDs(vec2.data(), LDs.data(), d->qLDiagInv, nv,
mj_solveLDs(vec2.data(), LDs.data(), d->qLDiagInv, nv, 1,
d->C_rownnz, d->C_rowadr, m->dof_simplenum, d->C_colind);
// expect vectors to match up to floating point precision
@@ -507,6 +507,44 @@ TEST_F(CoreSmoothTest, SolveLDs) {
mj_deleteModel(m);
}
TEST_F(CoreSmoothTest, SolveLDmultipleVectors) {
const std::string xml_path = GetTestDataFilePath(kInertiaPath);
char error[1024];
mjModel* m = mj_loadXML(xml_path.c_str(), nullptr, error, sizeof(error));
ASSERT_THAT(m, NotNull()) << "Failed to load model: " << error;
mjData* d = mj_makeData(m);
mj_forward(m, d);
int nv = m->nv;
int nC = m->nC;
// copy LD into LDs: CSR format
vector<mjtNum> LDs(nC);
for (int i=0; i < nC; i++) {
LDs[i] = d->qLD[d->mapM2C[i]];
}
// compare n LD and LDs vector solve
int n = 3;
vector<mjtNum> vec(nv*n);
vector<mjtNum> vec2(nv*n);
for (int i=0; i < nv*n; i++) vec[i] = vec2[i] = 2 + 3*i;
for (int i=0; i < nv*n; i+=3) vec[i] = vec2[i] = 0;
mj_solveLD(m, vec.data(), n, d->qLD, d->qLDiagInv);
mj_solveLDs(vec2.data(), LDs.data(), d->qLDiagInv, nv, n,
d->C_rownnz, d->C_rowadr, m->dof_simplenum, d->C_colind);
// expect vectors to match up to floating point precision
for (int i=0; i < nv*n; i++) {
EXPECT_FLOAT_EQ(vec[i], vec2[i]);
}
mj_deleteData(d);
mj_deleteModel(m);
}
TEST_F(CoreSmoothTest, FactorIs) {
const std::string xml_path = GetTestDataFilePath(kInertiaPath);
char error[1024];