use num_traits::{FromPrimitive, One};
use oxiblas_core::scalar::{Field, Real, Scalar};
#[inline]
fn scalar_abs<T: Scalar>(x: T) -> T::Real {
Scalar::abs(x)
}
#[inline]
fn scalar_epsilon<T: Scalar>() -> T::Real {
<T as Scalar>::epsilon()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BandLuError {
Singular {
index: usize,
},
InvalidDimensions {
n: usize,
kl: usize,
ku: usize,
},
InvalidStorageLength {
expected: usize,
actual: usize,
},
DimensionMismatch {
expected: usize,
actual: usize,
},
}
impl core::fmt::Display for BandLuError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
BandLuError::Singular { index } => {
write!(f, "Band matrix is singular at index {index}")
}
BandLuError::InvalidDimensions { n, kl, ku } => {
write!(f, "Invalid band dimensions: n={n}, kl={kl}, ku={ku}")
}
BandLuError::InvalidStorageLength { expected, actual } => {
write!(
f,
"Invalid band storage length: expected {expected}, got {actual}"
)
}
BandLuError::DimensionMismatch { expected, actual } => {
write!(f, "Dimension mismatch: expected {expected}, got {actual}")
}
}
}
}
impl std::error::Error for BandLuError {}
#[derive(Clone, Debug)]
pub struct BandLu<T: Scalar> {
ab: Vec<T>,
n: usize,
kl: usize,
ku: usize,
ldab: usize,
pivot: Vec<usize>,
}
impl<T: Field + Real> BandLu<T> {
pub fn compute(n: usize, kl: usize, ku: usize, ab: &[T]) -> Result<Self, BandLuError> {
if n == 0 {
return Ok(BandLu {
ab: Vec::new(),
n: 0,
kl,
ku,
ldab: 2 * kl + ku + 1,
pivot: Vec::new(),
});
}
if kl >= n || ku >= n {
return Err(BandLuError::InvalidDimensions { n, kl, ku });
}
let ldab = 2 * kl + ku + 1;
let expected_len = ldab * n;
if ab.len() != expected_len {
return Err(BandLuError::InvalidStorageLength {
expected: expected_len,
actual: ab.len(),
});
}
let mut ab_work = ab.to_vec();
for j in 0..n {
for r in 0..kl {
ab_work[r + j * ldab] = T::zero();
}
}
let mut pivot = vec![0usize; n];
let mut ju = 0usize;
for j in 0..n {
let km = kl.min(n - 1 - j);
let mut jp_rel = 0usize;
let mut pivot_val = scalar_abs(ab_work[band_idx(ldab, kl, ku, j, j)]);
for i in 1..=km {
let val = scalar_abs(ab_work[band_idx(ldab, kl, ku, j + i, j)]);
if val > pivot_val {
pivot_val = val;
jp_rel = i;
}
}
pivot[j] = j + jp_rel;
let tol = scalar_epsilon::<T>()
* <T::Real as FromPrimitive>::from_usize(n).unwrap_or(<T::Real as One>::one());
if pivot_val <= tol {
return Err(BandLuError::Singular { index: j });
}
ju = ju.max((j + ku + jp_rel).min(n - 1));
if jp_rel != 0 {
for c in j..=ju {
let idx_j = band_idx(ldab, kl, ku, j, c);
let idx_p = band_idx(ldab, kl, ku, j + jp_rel, c);
ab_work.swap(idx_j, idx_p);
}
}
if km > 0 {
let pivot_inv = T::one() / ab_work[band_idx(ldab, kl, ku, j, j)];
for i in 1..=km {
let idx = band_idx(ldab, kl, ku, j + i, j);
ab_work[idx] = ab_work[idx] * pivot_inv;
}
for c in (j + 1)..=ju {
let u_jc = ab_work[band_idx(ldab, kl, ku, j, c)];
if u_jc != T::zero() {
for i in 1..=km {
let l_ij = ab_work[band_idx(ldab, kl, ku, j + i, j)];
let idx = band_idx(ldab, kl, ku, j + i, c);
ab_work[idx] = ab_work[idx] - l_ij * u_jc;
}
}
}
}
}
Ok(BandLu {
ab: ab_work,
n,
kl,
ku,
ldab,
pivot,
})
}
#[inline]
pub fn size(&self) -> usize {
self.n
}
#[inline]
pub fn kl(&self) -> usize {
self.kl
}
#[inline]
pub fn ku(&self) -> usize {
self.ku
}
pub fn pivot(&self) -> &[usize] {
&self.pivot
}
pub fn ab(&self) -> &[T] {
&self.ab
}
pub fn solve(&self, b: &[T]) -> Result<Vec<T>, BandLuError> {
if b.len() != self.n {
return Err(BandLuError::DimensionMismatch {
expected: self.n,
actual: b.len(),
});
}
if self.n == 0 {
return Ok(Vec::new());
}
let mut x = b.to_vec();
if self.kl > 0 {
for j in 0..self.n {
let lm = self.kl.min(self.n - 1 - j);
let p = self.pivot[j];
if p != j {
x.swap(j, p);
}
let xj = x[j];
for i in 1..=lm {
let l_elem = self.ab[band_idx(self.ldab, self.kl, self.ku, j + i, j)];
x[j + i] = x[j + i] - l_elem * xj;
}
}
}
let kmax = self.ku + self.kl;
for j in (0..self.n).rev() {
let diag = self.ab[band_idx(self.ldab, self.kl, self.ku, j, j)];
x[j] = x[j] / diag;
let xj = x[j];
for i in j.saturating_sub(kmax)..j {
let u_elem = self.ab[band_idx(self.ldab, self.kl, self.ku, i, j)];
x[i] = x[i] - u_elem * xj;
}
}
Ok(x)
}
pub fn solve_multiple(&self, b: &[T], nrhs: usize) -> Result<Vec<T>, BandLuError> {
if b.len() != self.n * nrhs {
return Err(BandLuError::DimensionMismatch {
expected: self.n * nrhs,
actual: b.len(),
});
}
if self.n == 0 || nrhs == 0 {
return Ok(Vec::new());
}
let mut x = b.to_vec();
let ldb = nrhs;
if self.kl > 0 {
for j in 0..self.n {
let lm = self.kl.min(self.n - 1 - j);
let p = self.pivot[j];
if p != j {
for col in 0..nrhs {
x.swap(j * ldb + col, p * ldb + col);
}
}
for i in 1..=lm {
let l_elem = self.ab[band_idx(self.ldab, self.kl, self.ku, j + i, j)];
for col in 0..nrhs {
x[(j + i) * ldb + col] = x[(j + i) * ldb + col] - l_elem * x[j * ldb + col];
}
}
}
}
let kmax = self.ku + self.kl;
for j in (0..self.n).rev() {
let diag = self.ab[band_idx(self.ldab, self.kl, self.ku, j, j)];
for col in 0..nrhs {
x[j * ldb + col] = x[j * ldb + col] / diag;
}
for i in j.saturating_sub(kmax)..j {
let u_elem = self.ab[band_idx(self.ldab, self.kl, self.ku, i, j)];
for col in 0..nrhs {
x[i * ldb + col] = x[i * ldb + col] - u_elem * x[j * ldb + col];
}
}
}
Ok(x)
}
pub fn solve_transpose(&self, b: &[T]) -> Result<Vec<T>, BandLuError> {
if b.len() != self.n {
return Err(BandLuError::DimensionMismatch {
expected: self.n,
actual: b.len(),
});
}
if self.n == 0 {
return Ok(Vec::new());
}
let mut x = b.to_vec();
let kmax = self.ku + self.kl;
for j in 0..self.n {
let diag = self.ab[band_idx(self.ldab, self.kl, self.ku, j, j)];
x[j] = x[j] / diag;
let xj = x[j];
let i_end = (j + kmax).min(self.n - 1);
for i in (j + 1)..=i_end {
let u_elem = self.ab[band_idx(self.ldab, self.kl, self.ku, j, i)];
x[i] = x[i] - u_elem * xj;
}
}
if self.kl > 0 {
for j in (0..self.n).rev() {
let lm = self.kl.min(self.n - 1 - j);
let mut acc = x[j];
for i in 1..=lm {
let l_elem = self.ab[band_idx(self.ldab, self.kl, self.ku, j + i, j)];
acc = acc - l_elem * x[j + i];
}
x[j] = acc;
let p = self.pivot[j];
if p != j {
x.swap(j, p);
}
}
}
Ok(x)
}
pub fn rcond(&self, anorm_1: T) -> T {
if self.n == 0 || anorm_1 == T::zero() {
return T::zero();
}
let ainv_norm_1 = self.estimate_inv_norm_1();
if ainv_norm_1 == T::zero() {
return T::one();
}
T::one() / (anorm_1 * ainv_norm_1)
}
fn estimate_inv_norm_1(&self) -> T {
let n = self.n;
if n == 0 {
return T::zero();
}
let one_over_n = T::one() / T::from_usize(n).unwrap_or(T::one());
let mut x = vec![one_over_n; n];
const MAX_ITER: usize = 5;
let mut gamma = T::zero();
for _iter in 0..MAX_ITER {
let w = match self.solve(&x) {
Ok(w) => w,
Err(_) => return T::from_f64(1e30).unwrap_or(T::one() / <T as Scalar>::epsilon()),
};
let mut gamma_new = T::zero();
for &wi in &w {
gamma_new = gamma_new + Scalar::abs(wi);
}
if gamma_new <= gamma {
return gamma;
}
gamma = gamma_new;
for i in 0..n {
x[i] = if w[i] >= T::zero() {
T::one()
} else {
-T::one()
};
}
let z = match self.solve_transpose(&x) {
Ok(z) => z,
Err(_) => return gamma,
};
let mut j_max = 0;
let mut z_max = Scalar::abs(z[0]);
for j in 1..n {
let z_abs = Scalar::abs(z[j]);
if z_abs > z_max {
z_max = z_abs;
j_max = j;
}
}
let mut z_dot_xi = T::zero();
for i in 0..n {
z_dot_xi = z_dot_xi + z[i] * x[i];
}
if z_max <= z_dot_xi {
return gamma;
}
for i in 0..n {
x[i] = T::zero();
}
x[j_max] = T::one();
}
gamma
}
pub fn condition_number_estimate(&self) -> T {
if self.n == 0 {
return T::one();
}
let mut max_diag = T::zero();
let mut min_diag = T::from_f64(1e30).unwrap_or(T::one() / <T as Scalar>::epsilon());
for j in 0..self.n {
let diag = Scalar::abs(self.ab[band_idx(self.ldab, self.kl, self.ku, j, j)]);
if diag > max_diag {
max_diag = diag;
}
if diag < min_diag && diag > T::zero() {
min_diag = diag;
}
}
if min_diag > T::zero() {
max_diag / min_diag
} else {
T::from_f64(1e30).unwrap_or(T::one() / <T as Scalar>::epsilon())
}
}
}
pub fn band_norm_1<T: Field + Real>(ab: &[T], n: usize, kl: usize, ku: usize) -> T {
let ldab = 2 * kl + ku + 1;
let mut max_col_sum = T::zero();
for j in 0..n {
let mut col_sum = T::zero();
let i_start = j.saturating_sub(ku);
let i_end = (j + kl).min(n - 1);
for i in i_start..=i_end {
let row_in_band = kl + ku + i - j;
col_sum = col_sum + Scalar::abs(ab[row_in_band + j * ldab]);
}
if col_sum > max_col_sum {
max_col_sum = col_sum;
}
}
max_col_sum
}
pub fn band_norm_inf<T: Field + Real>(ab: &[T], n: usize, kl: usize, ku: usize) -> T {
let ldab = 2 * kl + ku + 1;
let mut row_sums = vec![T::zero(); n];
for j in 0..n {
let i_start = j.saturating_sub(ku);
let i_end = (j + kl).min(n - 1);
for i in i_start..=i_end {
let row_in_band = kl + ku + i - j;
row_sums[i] = row_sums[i] + Scalar::abs(ab[row_in_band + j * ldab]);
}
}
let mut max_row_sum = T::zero();
for &sum in &row_sums {
if sum > max_row_sum {
max_row_sum = sum;
}
}
max_row_sum
}
#[inline]
fn band_idx(ldab: usize, kl: usize, ku: usize, i: usize, j: usize) -> usize {
let row_in_band = kl + ku + i - j;
row_in_band + j * ldab
}
pub fn dense_to_band<T: Field + Real>(a: &[T], n: usize, kl: usize, ku: usize) -> Vec<T> {
let ldab = 2 * kl + ku + 1;
let mut ab = vec![T::zero(); ldab * n];
for j in 0..n {
let i_start = j.saturating_sub(ku);
let i_end = (j + kl).min(n - 1);
for i in i_start..=i_end {
let row_in_band = kl + ku + i - j;
ab[row_in_band + j * ldab] = a[i * n + j];
}
}
ab
}
pub fn band_to_dense<T: Field + Real>(ab: &[T], n: usize, kl: usize, ku: usize) -> Vec<T> {
let ldab = 2 * kl + ku + 1;
let mut a = vec![T::zero(); n * n];
for j in 0..n {
let i_start = j.saturating_sub(ku);
let i_end = (j + kl).min(n - 1);
for i in i_start..=i_end {
let row_in_band = kl + ku + i - j;
a[i * n + j] = ab[row_in_band + j * ldab];
}
}
a
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_dense_to_band_tridiagonal() {
let n = 4;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
2.0, -1.0, 0.0, 0.0,
-1.0, 2.0, -1.0, 0.0,
0.0, -1.0, 2.0, -1.0,
0.0, 0.0, -1.0, 2.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let ldab = 4;
assert_eq!(ab.len(), ldab * n);
assert!((ab[2] - 2.0).abs() < 1e-10); assert!((ab[2 + ldab] - 2.0).abs() < 1e-10); assert!((ab[2 + 2 * ldab] - 2.0).abs() < 1e-10); assert!((ab[2 + 3 * ldab] - 2.0).abs() < 1e-10);
assert!((ab[1 + ldab] - (-1.0)).abs() < 1e-10); assert!((ab[1 + 2 * ldab] - (-1.0)).abs() < 1e-10); assert!((ab[1 + 3 * ldab] - (-1.0)).abs() < 1e-10);
assert!((ab[3] - (-1.0)).abs() < 1e-10); assert!((ab[3 + ldab] - (-1.0)).abs() < 1e-10); assert!((ab[3 + 2 * ldab] - (-1.0)).abs() < 1e-10); }
#[test]
fn test_band_to_dense() {
let n = 4;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a_orig: Vec<f64> = vec![
2.0, -1.0, 0.0, 0.0,
-1.0, 2.0, -1.0, 0.0,
0.0, -1.0, 2.0, -1.0,
0.0, 0.0, -1.0, 2.0,
];
let ab = dense_to_band(&a_orig, n, kl, ku);
let a_back = band_to_dense(&ab, n, kl, ku);
for i in 0..n * n {
assert!(
(a_orig[i] - a_back[i]).abs() < 1e-10,
"Mismatch at index {i}"
);
}
}
#[test]
fn test_band_lu_tridiagonal() {
let n = 4;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
4.0, -1.0, 0.0, 0.0,
-1.0, 4.0, -1.0, 0.0,
0.0, -1.0, 4.0, -1.0,
0.0, 0.0, -1.0, 4.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("Should not be singular");
let b = vec![3.0, 2.0, 2.0, 3.0];
let x = lu.solve(&b).expect("Should solve");
for i in 0..n {
assert!(
(x[i] - 1.0).abs() < 1e-10,
"x[{i}] = {}, expected 1.0",
x[i]
);
}
}
#[test]
fn test_band_lu_pentadiagonal() {
let n = 5;
let kl = 2;
let ku = 2;
#[rustfmt::skip]
let a: Vec<f64> = vec![
10.0, -1.0, -2.0, 0.0, 0.0,
-1.0, 10.0, -1.0, -2.0, 0.0,
-2.0, -1.0, 10.0, -1.0, -2.0,
0.0, -2.0, -1.0, 10.0, -1.0,
0.0, 0.0, -2.0, -1.0, 10.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("Should not be singular");
let b = vec![7.0, 6.0, 4.0, 6.0, 7.0];
let x = lu.solve(&b).expect("Should solve");
for i in 0..n {
let mut ax_i = 0.0;
for j in 0..n {
ax_i += a[i * n + j] * x[j];
}
assert!(
(ax_i - b[i]).abs() < 1e-9,
"Ax[{i}] = {ax_i}, expected {}",
b[i]
);
}
}
#[test]
fn test_band_lu_solve_multiple() {
let n = 4;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
4.0, -1.0, 0.0, 0.0,
-1.0, 4.0, -1.0, 0.0,
0.0, -1.0, 4.0, -1.0,
0.0, 0.0, -1.0, 4.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("Should not be singular");
let nrhs = 2;
#[rustfmt::skip]
let b = vec![
3.0, 1.0, 2.0, 2.0, 2.0, 2.0, 3.0, 1.0, ];
let x = lu.solve_multiple(&b, nrhs).expect("Should solve");
for rhs in 0..nrhs {
for i in 0..n {
let mut ax_i = 0.0;
for j in 0..n {
ax_i += a[i * n + j] * x[j * nrhs + rhs];
}
let b_i = b[i * nrhs + rhs];
assert!(
(ax_i - b_i).abs() < 1e-9,
"RHS {rhs}: Ax[{i}] = {ax_i}, expected {b_i}"
);
}
}
}
#[test]
fn test_band_lu_singular() {
let n = 3;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
1.0, -1.0, 0.0,
-1.0, 1.0, 0.0, 0.0, 0.0, 1.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let result = BandLu::<f64>::compute(n, kl, ku, &ab);
assert!(result.is_err());
match result {
Err(BandLuError::Singular { index: _ }) => {}
_ => panic!("Expected Singular error"),
}
}
#[test]
fn test_band_lu_asymmetric_bandwidth() {
let n = 4;
let kl = 1;
let ku = 2;
#[rustfmt::skip]
let a: Vec<f64> = vec![
10.0, -1.0, -2.0, 0.0,
-1.0, 10.0, -1.0, -2.0,
0.0, -1.0, 10.0, -1.0,
0.0, 0.0, -1.0, 10.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("Should not be singular");
let b = vec![7.0, 6.0, 8.0, 9.0];
let x = lu.solve(&b).expect("Should solve");
for i in 0..n {
let mut ax_i = 0.0;
for j in 0..n {
ax_i += a[i * n + j] * x[j];
}
assert!(
(ax_i - b[i]).abs() < 1e-9,
"Ax[{i}] = {ax_i}, expected {}",
b[i]
);
}
}
#[test]
fn test_band_lu_empty() {
let result = BandLu::<f64>::compute(0, 0, 0, &[]);
assert!(result.is_ok());
let lu = result.unwrap();
assert_eq!(lu.size(), 0);
}
#[test]
fn test_band_lu_f32() {
let n = 3;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f32> = vec![
4.0, -1.0, 0.0,
-1.0, 4.0, -1.0,
0.0, -1.0, 4.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("Should not be singular");
let b = vec![3.0f32, 2.0, 3.0];
let x = lu.solve(&b).expect("Should solve");
for i in 0..n {
let mut ax_i = 0.0f32;
for j in 0..n {
ax_i += a[i * n + j] * x[j];
}
assert!(
(ax_i - b[i]).abs() < 1e-5,
"Ax[{i}] = {ax_i}, expected {}",
b[i]
);
}
}
#[test]
fn test_band_norm_1() {
let n = 3;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
4.0, -1.0, 0.0,
-1.0, 4.0, -1.0,
0.0, -1.0, 4.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let norm = band_norm_1(&ab, n, kl, ku);
assert!(
(norm - 6.0).abs() < 1e-10,
"norm_1 = {}, expected 6.0",
norm
);
}
#[test]
fn test_band_norm_inf() {
let n = 3;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
4.0, -1.0, 0.0,
-1.0, 4.0, -1.0,
0.0, -1.0, 4.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let norm = band_norm_inf(&ab, n, kl, ku);
assert!(
(norm - 6.0).abs() < 1e-10,
"norm_inf = {}, expected 6.0",
norm
);
}
#[test]
fn test_band_rcond() {
let n = 4;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
4.0, -1.0, 0.0, 0.0,
-1.0, 4.0, -1.0, 0.0,
0.0, -1.0, 4.0, -1.0,
0.0, 0.0, -1.0, 4.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let anorm = band_norm_1(&ab, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("Should not be singular");
let rcond = lu.rcond(anorm);
assert!(rcond > 0.0, "rcond = {}, should be > 0", rcond);
assert!(rcond <= 1.0, "rcond = {}, should be <= 1", rcond);
assert!(
rcond > 0.1,
"rcond = {}, matrix seems ill-conditioned",
rcond
);
}
#[test]
fn test_band_condition_number_estimate() {
let n = 3;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
4.0, -1.0, 0.0,
-1.0, 4.0, -1.0,
0.0, -1.0, 4.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("Should not be singular");
let cond = lu.condition_number_estimate();
assert!(cond >= 1.0, "cond = {}, should be >= 1", cond);
}
#[test]
fn test_band_solve_transpose() {
let n = 3;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
4.0, -1.0, 0.0,
-1.0, 4.0, -1.0,
0.0, -1.0, 4.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("Should not be singular");
let b = vec![1.0, 2.0, 3.0];
let x = lu.solve_transpose(&b).expect("Should solve transpose");
for i in 0..n {
let mut atx_i = 0.0;
for j in 0..n {
atx_i += a[j * n + i] * x[j]; }
assert!(
(atx_i - b[i]).abs() < 1e-10,
"A^T*x[{i}] = {atx_i}, expected {}",
b[i]
);
}
}
fn reconstruct_plu(lu: &BandLu<f64>) -> Vec<f64> {
let n = lu.size();
let kl = lu.kl();
let ku = lu.ku();
let ldab = 2 * kl + ku + 1;
let ab = lu.ab();
let pivot = lu.pivot();
let mut x = vec![0.0f64; n * n];
for j in 0..n {
let i_start = j.saturating_sub(kl + ku);
for i in i_start..=j {
x[i * n + j] = ab[band_idx(ldab, kl, ku, i, j)];
}
}
for j in (0..n.saturating_sub(1)).rev() {
let km = kl.min(n - 1 - j);
for i in 1..=km {
let l = ab[band_idx(ldab, kl, ku, j + i, j)];
for c in 0..n {
x[(j + i) * n + c] += l * x[j * n + c];
}
}
let p = pivot[j];
if p != j {
for c in 0..n {
x.swap(j * n + c, p * n + c);
}
}
}
x
}
#[test]
fn test_band_lu_pivot_swap_tridiagonal() {
let n = 5;
let kl = 1;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
1.0, 2.0, 0.0, 0.0, 0.0,
3.0, 1.0, 2.0, 0.0, 0.0,
0.0, 4.0, 1.0, 2.0, 0.0,
0.0, 0.0, 5.0, 1.0, 2.0,
0.0, 0.0, 0.0, 6.0, 1.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("nonsingular");
assert!(
lu.pivot().iter().enumerate().any(|(j, &p)| p != j),
"expected a real pivot swap, got pivot = {:?}",
lu.pivot()
);
let a_rec = reconstruct_plu(&lu);
for i in 0..n * n {
assert!(
(a_rec[i] - a[i]).abs() < 1e-12,
"reconstruction mismatch at {i}: got {}, expected {}",
a_rec[i],
a[i]
);
}
let x_true = [1.0, 2.0, 3.0, 4.0, 5.0];
let mut b = vec![0.0f64; n];
for i in 0..n {
for j in 0..n {
b[i] += a[i * n + j] * x_true[j];
}
}
let x = lu.solve(&b).expect("solve");
for i in 0..n {
assert!(
(x[i] - x_true[i]).abs() < 1e-10,
"solve x[{i}] = {}, expected {}",
x[i],
x_true[i]
);
}
let mut bt = vec![0.0f64; n];
for i in 0..n {
for j in 0..n {
bt[i] += a[j * n + i] * x_true[j]; }
}
let xt = lu.solve_transpose(&bt).expect("solve_transpose");
for i in 0..n {
assert!(
(xt[i] - x_true[i]).abs() < 1e-10,
"solve_transpose x[{i}] = {}, expected {}",
xt[i],
x_true[i]
);
}
let nrhs = 2;
let mut bm = vec![0.0f64; n * nrhs];
for i in 0..n {
bm[i * nrhs] = b[i];
bm[i * nrhs + 1] = 2.0 * b[i];
}
let xm = lu.solve_multiple(&bm, nrhs).expect("solve_multiple");
for i in 0..n {
assert!(
(xm[i * nrhs] - x_true[i]).abs() < 1e-10,
"solve_multiple rhs0 x[{i}] = {}, expected {}",
xm[i * nrhs],
x_true[i]
);
assert!(
(xm[i * nrhs + 1] - 2.0 * x_true[i]).abs() < 1e-10,
"solve_multiple rhs1 x[{i}] = {}, expected {}",
xm[i * nrhs + 1],
2.0 * x_true[i]
);
}
}
#[test]
fn test_band_lu_pivot_swap_wide_band() {
let n = 5;
let kl = 2;
let ku = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
1.0, 2.0, 0.0, 0.0, 0.0,
3.0, 1.0, 2.0, 0.0, 0.0,
6.0, 4.0, 1.0, 2.0, 0.0,
0.0, 7.0, 5.0, 1.0, 2.0,
0.0, 0.0, 8.0, 6.0, 1.0,
];
let ab = dense_to_band(&a, n, kl, ku);
let lu = BandLu::compute(n, kl, ku, &ab).expect("nonsingular");
assert_eq!(lu.pivot()[0], 2, "column 0 must pivot two rows down");
assert!(
lu.pivot().iter().enumerate().any(|(j, &p)| p != j),
"expected a real pivot swap, got pivot = {:?}",
lu.pivot()
);
let a_rec = reconstruct_plu(&lu);
for i in 0..n * n {
assert!(
(a_rec[i] - a[i]).abs() < 1e-12,
"reconstruction mismatch at {i}: got {}, expected {}",
a_rec[i],
a[i]
);
}
let x_true = [2.0, -1.0, 3.0, 0.5, -4.0];
let mut b = vec![0.0f64; n];
for i in 0..n {
for j in 0..n {
b[i] += a[i * n + j] * x_true[j];
}
}
let x = lu.solve(&b).expect("solve");
for i in 0..n {
assert!(
(x[i] - x_true[i]).abs() < 1e-10,
"solve x[{i}] = {}, expected {}",
x[i],
x_true[i]
);
}
let mut bt = vec![0.0f64; n];
for i in 0..n {
for j in 0..n {
bt[i] += a[j * n + i] * x_true[j];
}
}
let xt = lu.solve_transpose(&bt).expect("solve_transpose");
for i in 0..n {
assert!(
(xt[i] - x_true[i]).abs() < 1e-10,
"solve_transpose x[{i}] = {}, expected {}",
xt[i],
x_true[i]
);
}
}
}