CSR implementation of mj_solveLD.

PiperOrigin-RevId: 704316200
Change-Id: Ibaff0284e40b3ebbe43bb489b6211ce739270e27
This commit is contained in:
Yuval Tassa
2024-12-09 09:33:20 -08:00
committed by Copybara-Service
parent 537fc2ff45
commit 1c4c7b012c
4 changed files with 152 additions and 87 deletions
+32
View File
@@ -1575,6 +1575,38 @@ void mj_solveLD(const mjModel* m, mjtNum* restrict x, int n,
}
// in-place sparse backsubstitution: x = inv(L'*D*L)*x
// like mj_solveLD, but using the CSR representation of L
void mj_solveLDs(mjtNum* x, const mjtNum* qLDs, const mjtNum* qLDiagInv, int nv,
const int* rownnz, const int* rowadr, const int* diag, const int* colind) {
// x <- L^-T x
for (int i=nv-2; i >= 0; i--) {
int d1 = diag[i] + 1;
int nnz = rownnz[i] - d1;
if (nnz > 0) {
int adr = rowadr[i] + d1;
x[i] -= mju_dotSparse(qLDs+adr, x, nnz, colind+adr, /*flg_unc1=*/0);
}
}
// x(i) /= D(i,i)
for (int i=0; i < nv; i++) {
x[i] *= qLDiagInv[i];
}
// x <- L^-1 x
for (int i=1; i < nv; i++) {
int d = diag[i];
if (d > 0) {
int adr = rowadr[i];
x[i] -= mju_dotSparse(qLDs+adr, x, d, colind+adr, /*flg_unc1=*/0);
}
}
}
// sparse backsubstitution: x = inv(L'*D*L)*y
// use factorization in d
void mj_solveM(const mjModel* m, mjData* d, mjtNum* x, const mjtNum* y, int n) {
+6 -1
View File
@@ -55,10 +55,15 @@ MJAPI void mj_factorI(const mjModel* m, mjData* d, const mjtNum* M, mjtNum* qLD,
// sparse L'*D*L factorizaton of the inertia matrix M, assumed spd
MJAPI void mj_factorM(const mjModel* m, mjData* d);
// sparse backsubstitution: x = inv(L'*D*L)*y
// sparse backsubstitution: x = inv(L'*D*L)*x
MJAPI void mj_solveLD(const mjModel* m, mjtNum* x, int n,
const mjtNum* qLD, const mjtNum* qLDiagInv);
// in-place sparse backsubstitution: x = inv(L'*D*L)*x
// like mj_solveLD, but using the CSR representation of L
MJAPI void mj_solveLDs(mjtNum* x, const mjtNum* qLDs, const mjtNum* qLDiagInv, int nv,
const int* rownnz, const int* rowadr, const int* diag, const int* colind);
// sparse backsubstitution: x = inv(L'*D*L)*y, use factorization in d
MJAPI void mj_solveM(const mjModel* m, mjData* d, mjtNum* x, const mjtNum* y, int n);