use num_traits::{FromPrimitive, One};
use oxiblas_core::scalar::{Field, Real, Scalar};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BandCholeskyError {
NotPositiveDefinite {
index: usize,
},
InvalidDimensions {
n: usize,
kd: usize,
},
InvalidStorageLength {
expected: usize,
actual: usize,
},
DimensionMismatch {
expected: usize,
actual: usize,
},
}
impl core::fmt::Display for BandCholeskyError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
BandCholeskyError::NotPositiveDefinite { index } => {
write!(
f,
"Band matrix is not positive definite (detected at index {index})"
)
}
BandCholeskyError::InvalidDimensions { n, kd } => {
write!(f, "Invalid band dimensions: n={n}, kd={kd}")
}
BandCholeskyError::InvalidStorageLength { expected, actual } => {
write!(
f,
"Invalid band storage length: expected {expected}, got {actual}"
)
}
BandCholeskyError::DimensionMismatch { expected, actual } => {
write!(f, "Dimension mismatch: expected {expected}, got {actual}")
}
}
}
}
impl std::error::Error for BandCholeskyError {}
#[derive(Clone, Debug)]
pub struct BandCholesky<T: Scalar> {
ab: Vec<T>,
n: usize,
kd: usize,
ldab: usize,
}
impl<T: Field + Real> BandCholesky<T> {
pub fn compute(n: usize, kd: usize, ab: &[T]) -> Result<Self, BandCholeskyError> {
if n == 0 {
return Ok(BandCholesky {
ab: Vec::new(),
n: 0,
kd,
ldab: kd + 1,
});
}
if kd >= n {
return Err(BandCholeskyError::InvalidDimensions { n, kd });
}
let ldab = kd + 1;
let expected_len = ldab * n;
if ab.len() != expected_len {
return Err(BandCholeskyError::InvalidStorageLength {
expected: expected_len,
actual: ab.len(),
});
}
let mut ab_work = ab.to_vec();
for j in 0..n {
let mut sum = T::zero();
let k_start = j.saturating_sub(kd);
for k in k_start..j {
let l_jk = ab_work[(j - k) + k * ldab];
sum = sum + l_jk * l_jk;
}
let diag = ab_work[j * ldab] - sum;
let tol = <T as Scalar>::epsilon()
* <T as FromPrimitive>::from_usize(n).unwrap_or(<T as One>::one());
if diag <= tol {
return Err(BandCholeskyError::NotPositiveDefinite { index: j });
}
ab_work[j * ldab] = Real::sqrt(diag);
let l_jj = ab_work[j * ldab];
let i_end = (j + kd).min(n - 1);
for i in (j + 1)..=i_end {
let mut sum_off = T::zero();
for k in k_start..j {
let i_row = i - k;
let j_row = j - k;
if i_row <= kd && j_row <= kd {
let l_ik = ab_work[i_row + k * ldab];
let l_jk = ab_work[j_row + k * ldab];
sum_off = sum_off + l_ik * l_jk;
}
}
let a_ij = ab_work[(i - j) + j * ldab];
ab_work[(i - j) + j * ldab] = (a_ij - sum_off) / l_jj;
}
}
Ok(BandCholesky {
ab: ab_work,
n,
kd,
ldab,
})
}
#[inline]
pub fn size(&self) -> usize {
self.n
}
#[inline]
pub fn kd(&self) -> usize {
self.kd
}
pub fn ab(&self) -> &[T] {
&self.ab
}
pub fn determinant(&self) -> T {
if self.n == 0 {
return T::one();
}
let mut det_l = T::one();
for j in 0..self.n {
det_l = det_l * self.ab[j * self.ldab]; }
det_l * det_l
}
pub fn log_determinant(&self) -> T {
if self.n == 0 {
return T::zero();
}
let mut log_det = T::zero();
let two = T::one() + T::one();
for j in 0..self.n {
log_det = log_det + Real::ln(self.ab[j * self.ldab]);
}
two * log_det
}
pub fn solve(&self, b: &[T]) -> Result<Vec<T>, BandCholeskyError> {
if b.len() != self.n {
return Err(BandCholeskyError::DimensionMismatch {
expected: self.n,
actual: b.len(),
});
}
if self.n == 0 {
return Ok(Vec::new());
}
let mut x = b.to_vec();
for j in 0..self.n {
x[j] = x[j] / self.ab[j * self.ldab];
let i_end = (j + self.kd).min(self.n - 1);
for i in (j + 1)..=i_end {
let l_ij = self.ab[(i - j) + j * self.ldab];
x[i] = x[i] - l_ij * x[j];
}
}
for j in (0..self.n).rev() {
x[j] = x[j] / self.ab[j * self.ldab];
let i_start = j.saturating_sub(self.kd);
for i in i_start..j {
let l_ji = self.ab[(j - i) + i * self.ldab];
x[i] = x[i] - l_ji * x[j];
}
}
Ok(x)
}
pub fn solve_multiple(&self, b: &[T], nrhs: usize) -> Result<Vec<T>, BandCholeskyError> {
if b.len() != self.n * nrhs {
return Err(BandCholeskyError::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;
for j in 0..self.n {
let l_jj = self.ab[j * self.ldab];
for col in 0..nrhs {
x[j * ldb + col] = x[j * ldb + col] / l_jj;
}
let i_end = (j + self.kd).min(self.n - 1);
for i in (j + 1)..=i_end {
let l_ij = self.ab[(i - j) + j * self.ldab];
for col in 0..nrhs {
x[i * ldb + col] = x[i * ldb + col] - l_ij * x[j * ldb + col];
}
}
}
for j in (0..self.n).rev() {
let l_jj = self.ab[j * self.ldab];
for col in 0..nrhs {
x[j * ldb + col] = x[j * ldb + col] / l_jj;
}
let i_start = j.saturating_sub(self.kd);
for i in i_start..j {
let l_ji = self.ab[(j - i) + i * self.ldab];
for col in 0..nrhs {
x[i * ldb + col] = x[i * ldb + col] - l_ji * x[j * ldb + col];
}
}
}
Ok(x)
}
}
pub fn dense_to_band_lower<T: Field + Real>(a: &[T], n: usize, kd: usize) -> Vec<T> {
let ldab = kd + 1;
let mut ab = vec![T::zero(); ldab * n];
for j in 0..n {
let i_end = (j + kd).min(n - 1);
for i in j..=i_end {
let row_in_band = i - j;
ab[row_in_band + j * ldab] = a[i * n + j];
}
}
ab
}
pub fn band_lower_to_dense<T: Field + Real>(ab: &[T], n: usize, kd: usize) -> Vec<T> {
let ldab = kd + 1;
let mut a = vec![T::zero(); n * n];
for j in 0..n {
let i_end = (j + kd).min(n - 1);
for i in j..=i_end {
let row_in_band = i - j;
let val = ab[row_in_band + j * ldab];
a[i * n + j] = val;
if i != j {
a[j * n + i] = val;
}
}
}
a
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_dense_to_band_lower_tridiagonal() {
let n = 4;
let kd = 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_lower(&a, n, kd);
let ldab = kd + 1;
assert_eq!(ab.len(), ldab * n);
assert!((ab[0] - 4.0).abs() < 1e-10);
assert!((ab[ldab] - 4.0).abs() < 1e-10);
assert!((ab[2 * ldab] - 4.0).abs() < 1e-10);
assert!((ab[3 * ldab] - 4.0).abs() < 1e-10);
assert!((ab[1] - (-1.0)).abs() < 1e-10);
assert!((ab[1 + ldab] - (-1.0)).abs() < 1e-10);
assert!((ab[1 + 2 * ldab] - (-1.0)).abs() < 1e-10);
}
#[test]
fn test_band_lower_to_dense() {
let n = 4;
let kd = 1;
#[rustfmt::skip]
let a_orig: 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_lower(&a_orig, n, kd);
let a_back = band_lower_to_dense(&ab, n, kd);
for i in 0..n * n {
assert!(
(a_orig[i] - a_back[i]).abs() < 1e-10,
"Mismatch at index {i}: {} vs {}",
a_orig[i],
a_back[i]
);
}
}
#[test]
fn test_band_cholesky_tridiagonal() {
let n = 4;
let kd = 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_lower(&a, n, kd);
let chol = BandCholesky::compute(n, kd, &ab).expect("Should be SPD");
let b = vec![3.0, 2.0, 2.0, 3.0];
let x = chol.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_cholesky_determinant() {
let n = 2;
let kd = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
4.0, 2.0,
2.0, 5.0,
];
let ab = dense_to_band_lower(&a, n, kd);
let chol = BandCholesky::compute(n, kd, &ab).expect("Should be SPD");
let det = chol.determinant();
assert!((det - 16.0).abs() < 1e-10, "det = {det}, expected 16");
}
#[test]
fn test_band_cholesky_pentadiagonal() {
let n = 5;
let kd = 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_lower(&a, n, kd);
let chol = BandCholesky::compute(n, kd, &ab).expect("Should be SPD");
let b = vec![7.0, 6.0, 4.0, 6.0, 7.0];
let x = chol.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_cholesky_not_spd() {
let n = 3;
let kd = 1;
#[rustfmt::skip]
let a: Vec<f64> = vec![
1.0, -2.0, 0.0,
-2.0, 1.0, -2.0,
0.0, -2.0, 1.0,
];
let ab = dense_to_band_lower(&a, n, kd);
let result = BandCholesky::<f64>::compute(n, kd, &ab);
assert!(result.is_err());
match result {
Err(BandCholeskyError::NotPositiveDefinite { index: _ }) => {}
_ => panic!("Expected NotPositiveDefinite error"),
}
}
#[test]
fn test_band_cholesky_solve_multiple() {
let n = 4;
let kd = 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_lower(&a, n, kd);
let chol = BandCholesky::compute(n, kd, &ab).expect("Should be SPD");
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 = chol.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_cholesky_empty() {
let result = BandCholesky::<f64>::compute(0, 0, &[]);
assert!(result.is_ok());
let chol = result.unwrap();
assert_eq!(chol.size(), 0);
}
#[test]
fn test_band_cholesky_f32() {
let n = 3;
let kd = 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_lower(&a, n, kd);
let chol = BandCholesky::compute(n, kd, &ab).expect("Should be SPD");
let b = vec![3.0f32, 2.0, 3.0];
let x = chol.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]
);
}
}
}