use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TridiagError {
Singular {
index: usize,
},
DimensionMismatch {
expected: usize,
actual: usize,
},
EmptySystem,
}
impl core::fmt::Display for TridiagError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::Singular { index } => {
write!(f, "Tridiagonal system is singular at index {index}")
}
Self::DimensionMismatch { expected, actual } => {
write!(f, "Dimension mismatch: expected {expected}, got {actual}")
}
Self::EmptySystem => write!(f, "Empty tridiagonal system"),
}
}
}
impl std::error::Error for TridiagError {}
pub fn tridiag_solve<T: Field + Real>(
dl: &[T],
d_diag: &[T],
du: &[T],
b: &[T],
) -> Result<Vec<T>, TridiagError> {
let n = d_diag.len();
if n == 0 {
return Err(TridiagError::EmptySystem);
}
if dl.len() != n - 1 {
return Err(TridiagError::DimensionMismatch {
expected: n - 1,
actual: dl.len(),
});
}
if du.len() != n - 1 {
return Err(TridiagError::DimensionMismatch {
expected: n - 1,
actual: du.len(),
});
}
if b.len() != n {
return Err(TridiagError::DimensionMismatch {
expected: n,
actual: b.len(),
});
}
if n == 1 {
let eps = <T as Scalar>::epsilon();
if Scalar::abs(d_diag[0]) <= eps {
return Err(TridiagError::Singular { index: 0 });
}
return Ok(vec![b[0] / d_diag[0]]);
}
let mut c_prime = vec![T::zero(); n - 1];
let mut d_prime = vec![T::zero(); n];
let eps = <T as Scalar>::epsilon();
if Scalar::abs(d_diag[0]) <= eps {
return Err(TridiagError::Singular { index: 0 });
}
c_prime[0] = du[0] / d_diag[0];
d_prime[0] = b[0] / d_diag[0];
for i in 1..n {
let denom = d_diag[i] - dl[i - 1] * c_prime[i - 1];
if Scalar::abs(denom) <= eps {
return Err(TridiagError::Singular { index: i });
}
if i < n - 1 {
c_prime[i] = du[i] / denom;
}
d_prime[i] = (b[i] - dl[i - 1] * d_prime[i - 1]) / denom;
}
let mut x = vec![T::zero(); n];
x[n - 1] = d_prime[n - 1];
for i in (0..n - 1).rev() {
x[i] = d_prime[i] - c_prime[i] * x[i + 1];
}
Ok(x)
}
pub fn tridiag_solve_multiple<T: Field + Real + bytemuck::Zeroable>(
dl: &[T],
d_diag: &[T],
du: &[T],
b: MatRef<'_, T>,
) -> Result<Mat<T>, TridiagError> {
let n = d_diag.len();
let nrhs = b.ncols();
if n == 0 {
return Err(TridiagError::EmptySystem);
}
if b.nrows() != n {
return Err(TridiagError::DimensionMismatch {
expected: n,
actual: b.nrows(),
});
}
let mut result = Mat::zeros(n, nrhs);
for j in 0..nrhs {
let col: Vec<T> = (0..n).map(|i| b[(i, j)]).collect();
let x = tridiag_solve(dl, d_diag, du, &col)?;
for i in 0..n {
result[(i, j)] = x[i];
}
}
Ok(result)
}
#[derive(Debug, Clone)]
pub struct TridiagFactors<T: Scalar> {
pub dl_modified: Vec<T>,
pub d_modified: Vec<T>,
pub du: Vec<T>,
pub n: usize,
}
pub fn tridiag_factor<T: Field + Real>(
dl: &[T],
d_diag: &[T],
du: &[T],
) -> Result<TridiagFactors<T>, TridiagError> {
let n = d_diag.len();
if n == 0 {
return Err(TridiagError::EmptySystem);
}
if dl.len() != n - 1 || du.len() != n - 1 {
return Err(TridiagError::DimensionMismatch {
expected: n - 1,
actual: dl.len().min(du.len()),
});
}
let eps = <T as Scalar>::epsilon();
let mut dl_modified = vec![T::zero(); n - 1];
let mut d_modified = vec![T::zero(); n];
d_modified[0] = d_diag[0];
if Scalar::abs(d_modified[0]) <= eps {
return Err(TridiagError::Singular { index: 0 });
}
for i in 1..n {
dl_modified[i - 1] = dl[i - 1] / d_modified[i - 1];
d_modified[i] = d_diag[i] - dl_modified[i - 1] * du[i - 1];
if Scalar::abs(d_modified[i]) <= eps {
return Err(TridiagError::Singular { index: i });
}
}
Ok(TridiagFactors {
dl_modified,
d_modified,
du: du.to_vec(),
n,
})
}
pub fn tridiag_solve_factored<T: Field + Real>(
factors: &TridiagFactors<T>,
b: &[T],
) -> Result<Vec<T>, TridiagError> {
let n = factors.n;
if b.len() != n {
return Err(TridiagError::DimensionMismatch {
expected: n,
actual: b.len(),
});
}
let mut y = vec![T::zero(); n];
y[0] = b[0];
for i in 1..n {
y[i] = b[i] - factors.dl_modified[i - 1] * y[i - 1];
}
let mut x = vec![T::zero(); n];
x[n - 1] = y[n - 1] / factors.d_modified[n - 1];
for i in (0..n - 1).rev() {
x[i] = (y[i] - factors.du[i] * x[i + 1]) / factors.d_modified[i];
}
Ok(x)
}
#[derive(Debug, Clone)]
pub struct TridiagSPDFactors<T: Scalar> {
pub d_factor: Vec<T>,
pub l_factor: Vec<T>,
pub n: usize,
}
pub fn tridiag_factor_spd<T: Field + Real>(
d_diag: &[T],
e: &[T],
) -> Result<TridiagSPDFactors<T>, TridiagError> {
let n = d_diag.len();
if n == 0 {
return Err(TridiagError::EmptySystem);
}
if e.len() != n.saturating_sub(1) && n > 1 {
return Err(TridiagError::DimensionMismatch {
expected: n - 1,
actual: e.len(),
});
}
let eps = <T as Scalar>::epsilon();
let mut d_factor = vec![T::zero(); n];
let mut l_factor = vec![T::zero(); n.saturating_sub(1)];
d_factor[0] = d_diag[0];
if d_factor[0] <= eps {
return Err(TridiagError::Singular { index: 0 });
}
for i in 1..n {
l_factor[i - 1] = e[i - 1] / d_factor[i - 1];
d_factor[i] = d_diag[i] - l_factor[i - 1] * l_factor[i - 1] * d_factor[i - 1];
if d_factor[i] <= eps {
return Err(TridiagError::Singular { index: i });
}
}
Ok(TridiagSPDFactors {
d_factor,
l_factor,
n,
})
}
pub fn tridiag_solve_factored_spd<T: Field + Real>(
factors: &TridiagSPDFactors<T>,
b: &[T],
) -> Result<Vec<T>, TridiagError> {
let n = factors.n;
if b.len() != n {
return Err(TridiagError::DimensionMismatch {
expected: n,
actual: b.len(),
});
}
if n == 0 {
return Ok(Vec::new());
}
let mut y = vec![T::zero(); n];
y[0] = b[0];
for i in 1..n {
y[i] = b[i] - factors.l_factor[i - 1] * y[i - 1];
}
let mut z = vec![T::zero(); n];
for i in 0..n {
z[i] = y[i] / factors.d_factor[i];
}
let mut x = vec![T::zero(); n];
x[n - 1] = z[n - 1];
for i in (0..n - 1).rev() {
x[i] = z[i] - factors.l_factor[i] * x[i + 1];
}
Ok(x)
}
pub fn tridiag_solve_spd<T: Field + Real>(
d_diag: &[T],
e: &[T],
b: &[T],
) -> Result<Vec<T>, TridiagError> {
let n = d_diag.len();
if n == 0 {
return Err(TridiagError::EmptySystem);
}
if e.len() != n - 1 {
return Err(TridiagError::DimensionMismatch {
expected: n - 1,
actual: e.len(),
});
}
if b.len() != n {
return Err(TridiagError::DimensionMismatch {
expected: n,
actual: b.len(),
});
}
let eps = <T as Scalar>::epsilon();
let mut d_factor = vec![T::zero(); n];
let mut l_factor = vec![T::zero(); n - 1];
d_factor[0] = d_diag[0];
if d_factor[0] <= eps {
return Err(TridiagError::Singular { index: 0 });
}
for i in 1..n {
l_factor[i - 1] = e[i - 1] / d_factor[i - 1];
d_factor[i] = d_diag[i] - l_factor[i - 1] * l_factor[i - 1] * d_factor[i - 1];
if d_factor[i] <= eps {
return Err(TridiagError::Singular { index: i });
}
}
let mut y = vec![T::zero(); n];
y[0] = b[0];
for i in 1..n {
y[i] = b[i] - l_factor[i - 1] * y[i - 1];
}
let mut z = vec![T::zero(); n];
for i in 0..n {
z[i] = y[i] / d_factor[i];
}
let mut x = vec![T::zero(); n];
x[n - 1] = z[n - 1];
for i in (0..n - 1).rev() {
x[i] = z[i] - l_factor[i] * x[i + 1];
}
Ok(x)
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_tridiag_solve_simple() {
let dl = [-1.0f64, -1.0];
let d = [2.0f64, 2.0, 2.0];
let du = [-1.0f64, -1.0];
let b = [1.0f64, 0.0, 1.0];
let x = tridiag_solve(&dl, &d, &du, &b).unwrap();
assert!(approx_eq(x[0], 1.0, 1e-10));
assert!(approx_eq(x[1], 1.0, 1e-10));
assert!(approx_eq(x[2], 1.0, 1e-10));
}
#[test]
fn test_tridiag_solve_2x2() {
let dl = [1.0f64];
let d = [2.0f64, 3.0];
let du = [1.0f64];
let b = [5.0f64, 7.0];
let x = tridiag_solve(&dl, &d, &du, &b).unwrap();
assert!(approx_eq(x[0], 1.6, 1e-10));
assert!(approx_eq(x[1], 1.8, 1e-10));
}
#[test]
fn test_tridiag_solve_1x1() {
let dl: [f64; 0] = [];
let d = [2.0f64];
let du: [f64; 0] = [];
let b = [4.0f64];
let x = tridiag_solve(&dl, &d, &du, &b).unwrap();
assert!(approx_eq(x[0], 2.0, 1e-10));
}
#[test]
fn test_tridiag_solve_diagonal() {
let dl = [0.0f64, 0.0];
let d = [2.0f64, 3.0, 4.0];
let du = [0.0f64, 0.0];
let b = [4.0f64, 9.0, 16.0];
let x = tridiag_solve(&dl, &d, &du, &b).unwrap();
assert!(approx_eq(x[0], 2.0, 1e-10));
assert!(approx_eq(x[1], 3.0, 1e-10));
assert!(approx_eq(x[2], 4.0, 1e-10));
}
#[test]
fn test_tridiag_solve_singular() {
let dl = [1.0f64];
let d = [0.0f64, 1.0]; let du = [1.0f64];
let b = [1.0f64, 1.0];
let result = tridiag_solve(&dl, &d, &du, &b);
assert!(matches!(result, Err(TridiagError::Singular { index: 0 })));
}
#[test]
fn test_tridiag_solve_verify() {
let dl = [1.0f64, 2.0, 1.5];
let d = [4.0f64, 5.0, 6.0, 7.0];
let du = [1.0f64, 1.0, 2.0];
let b = [10.0f64, 20.0, 30.0, 40.0];
let x = tridiag_solve(&dl, &d, &du, &b).unwrap();
let ax0 = d[0] * x[0] + du[0] * x[1];
let ax1 = dl[0] * x[0] + d[1] * x[1] + du[1] * x[2];
let ax2 = dl[1] * x[1] + d[2] * x[2] + du[2] * x[3];
let ax3 = dl[2] * x[2] + d[3] * x[3];
assert!(approx_eq(ax0, b[0], 1e-10));
assert!(approx_eq(ax1, b[1], 1e-10));
assert!(approx_eq(ax2, b[2], 1e-10));
assert!(approx_eq(ax3, b[3], 1e-10));
}
#[test]
fn test_tridiag_factor_solve() {
let dl = [-1.0f64, -1.0];
let d = [2.0f64, 2.0, 2.0];
let du = [-1.0f64, -1.0];
let b = [1.0f64, 0.0, 1.0];
let factors = tridiag_factor(&dl, &d, &du).unwrap();
let x = tridiag_solve_factored(&factors, &b).unwrap();
assert!(approx_eq(x[0], 1.0, 1e-10));
assert!(approx_eq(x[1], 1.0, 1e-10));
assert!(approx_eq(x[2], 1.0, 1e-10));
}
#[test]
fn test_tridiag_factor_multiple_rhs() {
let dl = [-1.0f64, -1.0];
let d = [2.0f64, 2.0, 2.0];
let du = [-1.0f64, -1.0];
let factors = tridiag_factor(&dl, &d, &du).unwrap();
let b1 = [1.0f64, 0.0, 1.0];
let x1 = tridiag_solve_factored(&factors, &b1).unwrap();
let b2 = [2.0f64, 0.0, 2.0];
let x2 = tridiag_solve_factored(&factors, &b2).unwrap();
for i in 0..3 {
assert!(approx_eq(x2[i], 2.0 * x1[i], 1e-10));
}
}
#[test]
fn test_tridiag_solve_spd() {
let d = [4.0f64, 4.0, 4.0];
let e = [-1.0f64, -1.0];
let b = [3.0f64, 2.0, 3.0];
let x = tridiag_solve_spd(&d, &e, &b).unwrap();
let ax0 = d[0] * x[0] + e[0] * x[1];
let ax1 = e[0] * x[0] + d[1] * x[1] + e[1] * x[2];
let ax2 = e[1] * x[1] + d[2] * x[2];
assert!(approx_eq(ax0, b[0], 1e-10));
assert!(approx_eq(ax1, b[1], 1e-10));
assert!(approx_eq(ax2, b[2], 1e-10));
}
#[test]
fn test_tridiag_solve_multiple() {
let dl = [-1.0f64, -1.0];
let d = [2.0f64, 2.0, 2.0];
let du = [-1.0f64, -1.0];
let b = Mat::from_rows(&[&[1.0f64, 2.0], &[0.0, 0.0], &[1.0, 2.0]]);
let x = tridiag_solve_multiple(&dl, &d, &du, b.as_ref()).unwrap();
assert!(approx_eq(x[(0, 0)], 1.0, 1e-10));
assert!(approx_eq(x[(1, 0)], 1.0, 1e-10));
assert!(approx_eq(x[(2, 0)], 1.0, 1e-10));
assert!(approx_eq(x[(0, 1)], 2.0, 1e-10));
assert!(approx_eq(x[(1, 1)], 2.0, 1e-10));
assert!(approx_eq(x[(2, 1)], 2.0, 1e-10));
}
#[test]
fn test_tridiag_solve_f32() {
let dl = [-1.0f32, -1.0];
let d = [2.0f32, 2.0, 2.0];
let du = [-1.0f32, -1.0];
let b = [1.0f32, 0.0, 1.0];
let x = tridiag_solve(&dl, &d, &du, &b).unwrap();
assert!((x[0] - 1.0).abs() < 1e-5);
assert!((x[1] - 1.0).abs() < 1e-5);
assert!((x[2] - 1.0).abs() < 1e-5);
}
#[test]
fn test_tridiag_dimension_mismatch() {
let dl = [-1.0f64]; let d = [2.0f64, 2.0, 2.0];
let du = [-1.0f64, -1.0];
let b = [1.0f64, 0.0, 1.0];
let result = tridiag_solve(&dl, &d, &du, &b);
assert!(matches!(
result,
Err(TridiagError::DimensionMismatch { .. })
));
}
#[test]
fn test_tridiag_empty() {
let dl: [f64; 0] = [];
let d: [f64; 0] = [];
let du: [f64; 0] = [];
let b: [f64; 0] = [];
let result = tridiag_solve(&dl, &d, &du, &b);
assert!(matches!(result, Err(TridiagError::EmptySystem)));
}
#[test]
fn test_tridiag_large_system() {
let n = 100;
let dl: Vec<f64> = vec![-1.0; n - 1];
let d: Vec<f64> = vec![2.0; n];
let du: Vec<f64> = vec![-1.0; n - 1];
let b: Vec<f64> = vec![1.0; n];
let x = tridiag_solve(&dl, &d, &du, &b).unwrap();
let ax0 = d[0] * x[0] + du[0] * x[1];
let ax_mid = dl[n / 2 - 1] * x[n / 2 - 1] + d[n / 2] * x[n / 2] + du[n / 2] * x[n / 2 + 1];
let ax_last = dl[n - 2] * x[n - 2] + d[n - 1] * x[n - 1];
assert!(approx_eq(ax0, b[0], 1e-10));
assert!(approx_eq(ax_mid, b[n / 2], 1e-10));
assert!(approx_eq(ax_last, b[n - 1], 1e-10));
}
#[test]
fn test_tridiag_factor_spd_basic() {
let d = [4.0f64, 4.0, 4.0];
let e = [-1.0f64, -1.0];
let factors = tridiag_factor_spd(&d, &e).unwrap();
assert_eq!(factors.n, 3);
assert_eq!(factors.d_factor.len(), 3);
assert_eq!(factors.l_factor.len(), 2);
for &df in &factors.d_factor {
assert!(df > 0.0);
}
}
#[test]
fn test_tridiag_factor_solve_spd() {
let d = [4.0f64, 4.0, 4.0];
let e = [-1.0f64, -1.0];
let b = [3.0f64, 2.0, 3.0];
let factors = tridiag_factor_spd(&d, &e).unwrap();
let x = tridiag_solve_factored_spd(&factors, &b).unwrap();
let ax0 = d[0] * x[0] + e[0] * x[1];
let ax1 = e[0] * x[0] + d[1] * x[1] + e[1] * x[2];
let ax2 = e[1] * x[1] + d[2] * x[2];
assert!(approx_eq(ax0, b[0], 1e-10));
assert!(approx_eq(ax1, b[1], 1e-10));
assert!(approx_eq(ax2, b[2], 1e-10));
}
#[test]
fn test_tridiag_factor_spd_multiple_rhs() {
let d = [4.0f64, 4.0, 4.0];
let e = [-1.0f64, -1.0];
let factors = tridiag_factor_spd(&d, &e).unwrap();
let b1 = [3.0f64, 2.0, 3.0];
let x1 = tridiag_solve_factored_spd(&factors, &b1).unwrap();
let b2 = [6.0f64, 4.0, 6.0];
let x2 = tridiag_solve_factored_spd(&factors, &b2).unwrap();
for i in 0..3 {
assert!(approx_eq(x2[i], 2.0 * x1[i], 1e-10));
}
}
#[test]
fn test_tridiag_factor_spd_1x1() {
let d = [4.0f64];
let e: [f64; 0] = [];
let factors = tridiag_factor_spd(&d, &e).unwrap();
assert_eq!(factors.n, 1);
assert!(approx_eq(factors.d_factor[0], 4.0, 1e-10));
let b = [8.0f64];
let x = tridiag_solve_factored_spd(&factors, &b).unwrap();
assert!(approx_eq(x[0], 2.0, 1e-10));
}
#[test]
fn test_tridiag_factor_spd_not_positive() {
let d = [-4.0f64, 4.0, 4.0];
let e = [-1.0f64, -1.0];
let result = tridiag_factor_spd(&d, &e);
assert!(matches!(result, Err(TridiagError::Singular { index: 0 })));
}
#[test]
fn test_tridiag_factor_spd_not_definite() {
let d = [1.0f64, 1.0, 1.0];
let e = [2.0f64, 2.0];
let result = tridiag_factor_spd(&d, &e);
assert!(matches!(result, Err(TridiagError::Singular { .. })));
}
#[test]
fn test_tridiag_factor_spd_large() {
let n = 100;
let d: Vec<f64> = vec![2.0; n];
let e: Vec<f64> = vec![-1.0; n - 1];
let b: Vec<f64> = vec![1.0; n];
let factors = tridiag_factor_spd(&d, &e).unwrap();
let x = tridiag_solve_factored_spd(&factors, &b).unwrap();
let ax0 = d[0] * x[0] + e[0] * x[1];
let ax_mid = e[n / 2 - 1] * x[n / 2 - 1] + d[n / 2] * x[n / 2] + e[n / 2] * x[n / 2 + 1];
let ax_last = e[n - 2] * x[n - 2] + d[n - 1] * x[n - 1];
assert!(approx_eq(ax0, b[0], 1e-10));
assert!(approx_eq(ax_mid, b[n / 2], 1e-10));
assert!(approx_eq(ax_last, b[n - 1], 1e-10));
}
#[test]
fn test_tridiag_factor_spd_f32() {
let d = [4.0f32, 4.0, 4.0];
let e = [-1.0f32, -1.0];
let b = [3.0f32, 2.0, 3.0];
let factors = tridiag_factor_spd(&d, &e).unwrap();
let x = tridiag_solve_factored_spd(&factors, &b).unwrap();
let ax0 = d[0] * x[0] + e[0] * x[1];
let ax1 = e[0] * x[0] + d[1] * x[1] + e[1] * x[2];
let ax2 = e[1] * x[1] + d[2] * x[2];
assert!((ax0 - b[0]).abs() < 1e-5);
assert!((ax1 - b[1]).abs() < 1e-5);
assert!((ax2 - b[2]).abs() < 1e-5);
}
#[test]
fn test_tridiag_factor_spd_consistency() {
let d = [4.0f64, 4.0, 4.0, 4.0];
let e = [-1.0f64, -1.0, -1.0];
let b = [3.0f64, 2.0, 2.0, 3.0];
let x_direct = tridiag_solve_spd(&d, &e, &b).unwrap();
let factors = tridiag_factor_spd(&d, &e).unwrap();
let x_factored = tridiag_solve_factored_spd(&factors, &b).unwrap();
for i in 0..4 {
assert!(approx_eq(x_direct[i], x_factored[i], 1e-14));
}
}
#[test]
fn test_tridiag_factor_spd_empty() {
let d: [f64; 0] = [];
let e: [f64; 0] = [];
let result = tridiag_factor_spd(&d, &e);
assert!(matches!(result, Err(TridiagError::EmptySystem)));
}
#[test]
fn test_tridiag_factor_spd_dimension_mismatch() {
let d = [4.0f64, 4.0, 4.0];
let e = [-1.0f64];
let result = tridiag_factor_spd(&d, &e);
assert!(matches!(
result,
Err(TridiagError::DimensionMismatch { .. })
));
}
}