use crate::linear_algebra::{Matrix, Vector};
use crate::scalar::Numeric;
use crate::utils::error_codes::CalcError;
#[derive(Debug, Clone, Copy)]
#[must_use]
pub struct Lu<const N: usize, T = f64> {
pub(crate) lu: Matrix<N, N, T>,
pub(crate) perm: [usize; N],
pub(crate) sign: T,
}
impl<const N: usize, T: Numeric> Matrix<N, N, T> {
pub fn lu(self) -> Result<Lu<N, T>, CalcError> {
let mut a = self;
let mut perm: [usize; N] = core::array::from_fn(|i| i);
let mut sign = T::ONE;
for k in 0..N {
let mut p = k;
let mut best = a[(k, k)].abs();
for i in (k + 1)..N {
let magnitude = a[(i, k)].abs();
if magnitude > best {
best = magnitude;
p = i;
}
}
if a[(p, k)] == T::ZERO {
return Err(CalcError::SingularMatrix);
}
if p != k {
for c in 0..N {
let tmp = a[(k, c)];
a[(k, c)] = a[(p, c)];
a[(p, c)] = tmp;
}
perm.swap(k, p);
sign = -sign;
}
for i in (k + 1)..N {
let factor = a[(i, k)] / a[(k, k)];
a[(i, k)] = factor;
for j in (k + 1)..N {
let term = factor * a[(k, j)];
a[(i, j)] -= term;
}
}
}
Ok(Lu { lu: a, perm, sign })
}
pub fn solve(self, b: Vector<N, T>) -> Result<Vector<N, T>, CalcError> {
Ok(self.lu()?.solve(b))
}
}
impl<const N: usize, T: Numeric> Lu<N, T> {
pub fn l(&self) -> Matrix<N, N, T> {
Matrix::from_fn(|r, c| {
if r == c {
T::ONE
} else if c < r {
self.lu[(r, c)]
} else {
T::ZERO
}
})
}
pub fn u(&self) -> Matrix<N, N, T> {
Matrix::from_fn(|r, c| if c >= r { self.lu[(r, c)] } else { T::ZERO })
}
#[inline]
#[must_use]
pub fn permutation(&self) -> [usize; N] {
self.perm
}
#[inline]
#[must_use]
pub fn determinant(&self) -> T {
let mut det = self.sign;
for i in 0..N {
det *= self.lu[(i, i)];
}
det
}
pub fn solve(&self, b: Vector<N, T>) -> Vector<N, T> {
let mut x: [T; N] = core::array::from_fn(|i| b[self.perm[i]]);
for i in 0..N {
let mut sum = x[i];
for (j, &xj) in x.iter().enumerate().take(i) {
sum -= self.lu[(i, j)] * xj;
}
x[i] = sum;
}
for i in (0..N).rev() {
let mut sum = x[i];
for (j, &xj) in x.iter().enumerate().skip(i + 1) {
sum -= self.lu[(i, j)] * xj;
}
x[i] = sum / self.lu[(i, i)];
}
Vector::new(x)
}
pub fn solve_matrix<const K: usize>(&self, b: Matrix<N, K, T>) -> Matrix<N, K, T> {
let mut result = Matrix::zeros();
for c in 0..K {
let x = self.solve(b.column(c));
for r in 0..N {
result[(r, c)] = x[r];
}
}
result
}
pub fn inverse(&self) -> Matrix<N, N, T> {
self.solve_matrix(Matrix::identity())
}
}