Add dense LU factorization and solve functions.

PiperOrigin-RevId: 909290647
Change-Id: I77ae2352b20abf96ef7b48ae32b03a1f4098604e
This commit is contained in:
Yuval Tassa
2026-05-02 13:26:14 -07:00
committed by Copybara-Service
parent dbd451138c
commit 25751a7b98
3 changed files with 236 additions and 1 deletions
+84 -1
View File
@@ -751,7 +751,90 @@ void mju_bandMulMatVec(mjtNum* res, const mjtNum* mat, const mjtNum* vec,
}
//------------------------------ LU factorization --------------------------------------------------
//------------------------------ dense LU factorization --------------------------------------------
// dense LU factorization with partial pivoting
// factorizes n x n row-major matrix A in-place into L and U
// L has unit diagonal (not stored), U has explicit diagonal
// pivot stores row permutation: row i of original = row pivot[i] of result
// return 1 if successful, 0 if singular (diagonal element < mjMINVAL)
int mju_factorLU(mjtNum* restrict A, int n, int* pivot) {
for (int k=0; k < n; k++) {
// initialize pivot
pivot[k] = k;
// find pivot: max absolute value in column k, rows k..n-1
mjtNum maxval = mju_abs(A[k*n+k]);
int maxrow = k;
for (int i=k+1; i < n; i++) {
mjtNum val = mju_abs(A[i*n+k]);
if (val > maxval) {
maxval = val;
maxrow = i;
}
}
// check singularity
if (maxval < mjMINVAL) {
return 0;
}
// swap rows k and maxrow
if (maxrow != k) {
pivot[k] = maxrow;
for (int j=0; j < n; j++) {
mjtNum tmp = A[k*n+j];
A[k*n+j] = A[maxrow*n+j];
A[maxrow*n+j] = tmp;
}
}
// compute multipliers and update trailing submatrix
mjtNum diaginv = 1.0 / A[k*n+k];
for (int i=k+1; i < n; i++) {
A[i*n+k] *= diaginv;
mjtNum Aik = A[i*n+k];
for (int j=k+1; j < n; j++) {
A[i*n+j] -= Aik * A[k*n+j];
}
}
}
return 1;
}
// solve A*x = b given LU factorization of A, LU and pivot are output of mju_factorLU
void mju_solveLU(mjtNum* restrict x, const mjtNum* LU, const mjtNum* b, const int* pivot, int n) {
// copy b into x
mju_copy(x, b, n);
// apply row permutation and forward substitution: solve L*y = P*b
for (int i=0; i < n; i++) {
// apply pivot swap
if (pivot[i] != i) {
mjtNum tmp = x[i];
x[i] = x[pivot[i]];
x[pivot[i]] = tmp;
}
// subtract known terms
for (int j=0; j < i; j++) {
x[i] -= LU[i*n+j] * x[j];
}
}
// back substitution: solve U*x = y
for (int i=n-1; i >= 0; i--) {
for (int j=i+1; j < n; j++) {
x[i] -= LU[i*n+j] * x[j];
}
x[i] /= LU[i*n+i];
}
}
//------------------------------ sparse LU factorization -------------------------------------------
// sparse reverse-order LU factorization, no fill-in (assuming tree topology)
// result: LU = L + U; original = (U+I) * L; scratch size is n