diff --git a/src/engine/engine_core_constraint.c b/src/engine/engine_core_constraint.c index 42b4ebf8..c96a67a7 100644 --- a/src/engine/engine_core_constraint.c +++ b/src/engine/engine_core_constraint.c @@ -420,7 +420,7 @@ void mj_mulJacVec_island(const mjModel* m, const mjData* d, int Jrowadr = d->efc_J_rowadr[row]; int* Jind = d->efc_J_colind + Jrowadr; mjtNum* J = d->efc_J + Jrowadr; - res[i] = mju_dotSparse2(vec, J, vecnnz, vecind, Jnnz, Jind); + res[i] = mju_dotSparse2(vec, J, vecnnz, vecind, Jnnz, Jind, /*flg_unc2=*/0); } } @@ -428,7 +428,7 @@ void mj_mulJacVec_island(const mjModel* m, const mjData* d, else { int nv = m->nv; for (int i=0; i < resnnz; i++) { - res[i] = mju_dotSparse(vec, d->efc_J + nv*resind[i], vecnnz, vecind); + res[i] = mju_dotSparse(vec, d->efc_J + nv*resind[i], vecnnz, vecind, /*flg_unc1=*/0); } } } @@ -481,7 +481,7 @@ void mj_mulJacTVec_island(const mjModel* m, const mjData* d, int JTrowadr = d->efc_JT_rowadr[row]; int* JTind = d->efc_JT_colind + JTrowadr; mjtNum* JT = d->efc_JT + JTrowadr; - res[i] = mju_dotSparse2(vec, JT, vecnnz, vecind, JTnnz, JTind); + res[i] = mju_dotSparse2(vec, JT, vecnnz, vecind, JTnnz, JTind, /*flg_unc2=*/0); } } @@ -489,7 +489,7 @@ void mj_mulJacTVec_island(const mjModel* m, const mjData* d, else { int nefc = d->nefc; for (int i=0; i < resnnz; i++) { - res[i] = mju_dotSparse(vec, d->efc_JT + nefc*resind[i], vecnnz, vecind); + res[i] = mju_dotSparse(vec, d->efc_JT + nefc*resind[i], vecnnz, vecind, /*flg_unc1=*/0); } } } diff --git a/src/engine/engine_solver.c b/src/engine/engine_solver.c index 2884616d..dd058a04 100644 --- a/src/engine/engine_solver.c +++ b/src/engine/engine_solver.c @@ -168,7 +168,8 @@ static void residual(const mjModel* m, mjData* d, mjtNum* res, int i, int dim, i for (int j=0; j < dim; j++) { res[j] = d->efc_b[i+j] + mju_dotSparse(d->efc_AR + d->efc_AR_rowadr[i+j], d->efc_force, d->efc_AR_rownnz[i+j], - d->efc_AR_colind + d->efc_AR_rowadr[i+j]); + d->efc_AR_colind + d->efc_AR_rowadr[i+j], + /*flg_unc1=*/0); } } diff --git a/src/engine/engine_util_solve.c b/src/engine/engine_util_solve.c index 4fdc82a9..933b2d90 100644 --- a/src/engine/engine_util_solve.c +++ b/src/engine/engine_util_solve.c @@ -235,7 +235,7 @@ void mju_cholSolveSparse(mjtNum* res, const mjtNum* mat, const mjtNum* vec, int // x(i) -= sum_j L(i,j)*x(j), j=0:i-1 if (nnz > 1) { - res[i] -= mju_dotSparse(mat+adr, res, nnz-1, colind+adr); + res[i] -= mju_dotSparse(mat+adr, res, nnz-1, colind+adr, /*flg_unc1=*/0); // modulo AVX, the above line does // for (int j=0; j0) { - res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]); + res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r], /*flg_unc2=*/0); r++; rs--; @@ -213,7 +243,7 @@ void mju_mulMatVecSparse_avx(mjtNum* res, const mjtNum* mat, const mjtNum* vec, } else { - res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r]); + res[r] = mju_dotSparse_avx(mat+rowadr[r], vec, rownnz[r], colind+rowadr[r], /*flg_unc2=*/0); } } } diff --git a/test/engine/engine_util_sparse_test.cc b/test/engine/engine_util_sparse_test.cc index 668fda6f..495a45c9 100644 --- a/test/engine/engine_util_sparse_test.cc +++ b/test/engine/engine_util_sparse_test.cc @@ -30,19 +30,58 @@ using ::testing::ElementsAre; using EngineUtilSparseTest = MujocoTest; TEST_F(EngineUtilSparseTest, MjuDot) { - mjtNum a[] = {1, 2, 3, 4, 5, 6, 7}; - mjtNum b[] = {7, 0, 6, 0, 0, 5, 0, 0, 0, 4, 0, 0, 0, 3, 0, 0, 2, 0, 1}; + mjtNum a[] = {2, 3, 4, 5, 6, 7, 8}; + mjtNum u[] = {2, 1, 3, 1, 1, 4, 1, 1, 1, 5, 1, 1, 1, 6, 1, 1, 7, 1, 8}; + mjtNum b[] = {8, 1, 7, 1, 1, 6, 1, 1, 1, 5, 1, 1, 1, 4, 1, 1, 3, 1, 2}; int i[] = {0, 2, 5, 9, 13, 16, 18}; // test various vector lengths as mju_dotSparse adds numbers in groups of four - EXPECT_EQ(mju_dotSparse(a, b, 0, i), 0); - EXPECT_EQ(mju_dotSparse(a, b, 1, i), 7); - EXPECT_EQ(mju_dotSparse(a, b, 2, i), 7 + 2*6); - EXPECT_EQ(mju_dotSparse(a, b, 3, i), 7 + 2*6 + 3*5); - EXPECT_EQ(mju_dotSparse(a, b, 4, i), 7 + 2*6 + 3*5 + 4*4); - EXPECT_EQ(mju_dotSparse(a, b, 5, i), 7 + 2*6 + 3*5 + 4*4 + 5*3); - EXPECT_EQ(mju_dotSparse(a, b, 6, i), 7 + 2*6 + 3*5 + 4*4 + 5*3 + 6*2); - EXPECT_EQ(mju_dotSparse(a, b, 7, i), 7 + 2*6 + 3*5 + 4*4 + 5*3 + 6*2 + 7); + + // a is compressed + int flg_unc1 = 0; + EXPECT_EQ(mju_dotSparse(a, b, 0, i, flg_unc1), 0); + EXPECT_EQ(mju_dotSparse(a, b, 1, i, flg_unc1), 2*8); + EXPECT_EQ(mju_dotSparse(a, b, 2, i, flg_unc1), 2*8 + 3*7); + EXPECT_EQ(mju_dotSparse(a, b, 3, i, flg_unc1), 2*8 + 3*7 + 4*6); + EXPECT_EQ(mju_dotSparse(a, b, 4, i, flg_unc1), 2*8 + 3*7 + 4*6 + 5*5); + EXPECT_EQ(mju_dotSparse(a, b, 5, i, flg_unc1), 2*8 + 3*7 + 4*6 + 5*5 + 6*4); + EXPECT_EQ(mju_dotSparse(a, b, 6, i, flg_unc1), + 2*8 + 3*7 + 4*6 + 5*5 + 6*4 + 7*3); + EXPECT_EQ(mju_dotSparse(a, b, 7, i, flg_unc1), + 2*8 + 3*7 + 4*6 + 5*5 + 6*4 + 7*3 + 8*2); + + // u is compressed + flg_unc1 = 1; + EXPECT_EQ(mju_dotSparse(u, b, 0, i, flg_unc1), 0); + EXPECT_EQ(mju_dotSparse(u, b, 1, i, flg_unc1), 2*8); + EXPECT_EQ(mju_dotSparse(u, b, 2, i, flg_unc1), 2*8 + 3*7); + EXPECT_EQ(mju_dotSparse(u, b, 3, i, flg_unc1), 2*8 + 3*7 + 4*6); + EXPECT_EQ(mju_dotSparse(u, b, 4, i, flg_unc1), 2*8 + 3*7 + 4*6 + 5*5); + EXPECT_EQ(mju_dotSparse(u, b, 5, i, flg_unc1), 2*8 + 3*7 + 4*6 + 5*5 + 6*4); + EXPECT_EQ(mju_dotSparse(u, b, 6, i, flg_unc1), + 2*8 + 3*7 + 4*6 + 5*5 + 6*4 + 7*3); + EXPECT_EQ(mju_dotSparse(u, b, 7, i, flg_unc1), + 2*8 + 3*7 + 4*6 + 5*5 + 6*4 + 7*3 + 8*2); +} + +TEST_F(EngineUtilSparseTest, MjuDot2) { + constexpr int annz = 6; + constexpr int bnnz = 5; + int ia[annz] = {0, 2, 5, 6, 7}; + mjtNum a[annz] = {2, 3, 4, 5, 6}; + int ib[bnnz] = { 1, 2, 3, 5, 7}; + mjtNum b[bnnz] = { 8, 7, 6, 5, 4}; + mjtNum u[] = {1, 8, 7, 6, 1, 5, 1, 4}; + + // test various vector lengths as mju_dotSparse adds numbers in groups of four + + // a is compressed + int flg_unc2 = 0; + EXPECT_EQ(mju_dotSparse2(a, b, annz, ia, bnnz, ib, flg_unc2), 3*7+4*5+6*4); + + // u is uncompressed + flg_unc2 = 1; + EXPECT_EQ(mju_dotSparse2(a, u, annz, ia, bnnz, ib, flg_unc2), 3*7+4*5+6*4); } TEST_F(EngineUtilSparseTest, CombineSparseCount) {