use faer::prelude::Solve;
use nalgebra::DVector;
use crate::context::error::OxiflowError;
pub trait SparseLinearSolver: Send + Sync {
fn solve(
&self,
a: &faer::sparse::SparseColMat<usize, f64>,
b: &DVector<f64>,
) -> Result<DVector<f64>, OxiflowError>;
}
#[derive(Debug, Default, Clone, Copy)]
pub struct FaerSparseSolver;
impl SparseLinearSolver for FaerSparseSolver {
fn solve(
&self,
a: &faer::sparse::SparseColMat<usize, f64>,
b: &DVector<f64>,
) -> Result<DVector<f64>, OxiflowError> {
let n = b.len();
if a.nrows() != n || a.ncols() != n {
return Err(OxiflowError::PreconditionFailed {
context: "FaerSparseSolver",
message: format!(
"dimension mismatch: a is {}x{}, b has length {n}",
a.nrows(),
a.ncols()
),
});
}
let b_faer = faer::Col::from_fn(n, |i| b[i]);
let lu = a.sp_lu().map_err(|e| OxiflowError::PreconditionFailed {
context: "FaerSparseSolver",
message: format!(
"sparse LU factorization failed (singular or near-singular system): {e:?}"
),
})?;
let x_faer = lu.solve(&b_faer);
Ok(DVector::from_fn(n, |i, _| x_faer[i]))
}
}
#[cfg(test)]
mod tests {
use super::*;
use faer::sparse::{SparseColMat, Triplet};
fn tridiagonal(n: usize) -> SparseColMat<usize, f64> {
let mut triplets = Vec::with_capacity(3 * n);
for i in 0..n {
triplets.push(Triplet::new(i, i, 2.0));
if i > 0 {
triplets.push(Triplet::new(i, i - 1, -1.0));
}
if i + 1 < n {
triplets.push(Triplet::new(i, i + 1, -1.0));
}
}
SparseColMat::try_new_from_triplets(n, n, &triplets)
.expect("valid triplets for a tridiagonal system")
}
#[test]
fn solves_small_tridiagonal_system() {
let n = 5;
let a = tridiagonal(n);
let b = DVector::from_element(n, 1.0);
let x = FaerSparseSolver.solve(&a, &b).unwrap();
for i in 0..n {
let mut residual = 2.0 * x[i];
if i > 0 {
residual -= x[i - 1];
}
if i + 1 < n {
residual -= x[i + 1];
}
assert!(
(residual - b[i]).abs() < 1e-9,
"residual too large at row {i}: {residual} vs {}",
b[i]
);
}
}
#[test]
fn solves_500x500_tridiagonal_system() {
let n = 500;
let a = tridiagonal(n);
let b = DVector::from_element(n, 1.0);
let x = FaerSparseSolver.solve(&a, &b).unwrap();
for i in 0..n {
let mut residual = 2.0 * x[i];
if i > 0 {
residual -= x[i - 1];
}
if i + 1 < n {
residual -= x[i + 1];
}
assert!((residual - b[i]).abs() < 1e-6);
}
}
#[test]
fn dimension_mismatch_returns_error() {
let a = tridiagonal(5);
let b = DVector::from_element(3, 1.0);
let err = FaerSparseSolver.solve(&a, &b).unwrap_err();
assert!(matches!(err, OxiflowError::PreconditionFailed { .. }));
}
}