use num_traits::FromPrimitive;
use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BunchKaufmanError {
EmptyMatrix,
NotSquare {
nrows: usize,
ncols: usize,
},
Singular {
index: usize,
},
DimensionMismatch {
expected: usize,
actual: usize,
},
}
impl core::fmt::Display for BunchKaufmanError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::NotSquare { nrows, ncols } => {
write!(f, "Matrix is not square: {nrows}×{ncols}")
}
Self::Singular { index } => {
write!(f, "Matrix is singular at index {index}")
}
Self::DimensionMismatch { expected, actual } => {
write!(f, "Dimension mismatch: expected {expected}, got {actual}")
}
}
}
}
impl std::error::Error for BunchKaufmanError {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Uplo {
Upper,
Lower,
}
#[derive(Clone, Debug)]
pub struct BunchKaufman<T: Scalar> {
factors: Mat<T>,
ipiv: Vec<i32>,
uplo: Uplo,
n: usize,
}
fn alpha<T: Field + Real>() -> T {
let one = T::one();
let seventeen = T::from_f64(17.0).unwrap_or(one);
let eight = T::from_f64(8.0).unwrap_or(one);
(one + Real::sqrt(seventeen)) / eight
}
impl<T: Field + Real + bytemuck::Zeroable + FromPrimitive> BunchKaufman<T> {
pub fn compute(a: MatRef<'_, T>) -> Result<Self, BunchKaufmanError> {
Self::compute_with_uplo(a, Uplo::Lower)
}
pub fn compute_with_uplo(a: MatRef<'_, T>, uplo: Uplo) -> Result<Self, BunchKaufmanError> {
let n = a.nrows();
if n == 0 {
return Err(BunchKaufmanError::EmptyMatrix);
}
if n != a.ncols() {
return Err(BunchKaufmanError::NotSquare {
nrows: n,
ncols: a.ncols(),
});
}
let mut factors = Mat::zeros(n, n);
match uplo {
Uplo::Lower => {
for i in 0..n {
for j in 0..=i {
factors[(i, j)] = a[(i, j)];
}
}
}
Uplo::Upper => {
for i in 0..n {
for j in i..n {
factors[(i, j)] = a[(i, j)];
}
}
}
}
let mut ipiv = vec![0i32; n];
match uplo {
Uplo::Lower => Self::factor_lower(&mut factors, &mut ipiv, n)?,
Uplo::Upper => Self::factor_upper(&mut factors, &mut ipiv, n)?,
}
Ok(Self {
factors,
ipiv,
uplo,
n,
})
}
fn factor_lower(a: &mut Mat<T>, ipiv: &mut [i32], n: usize) -> Result<(), BunchKaufmanError> {
let alpha_val = alpha::<T>();
let mut k = 0;
while k < n {
let kstep;
let absakk = Scalar::abs(a[(k, k)]);
let mut imax = k;
let mut colmax = T::zero();
if k + 1 < n {
for i in (k + 1)..n {
let absval = Scalar::abs(a[(i, k)]);
if absval > colmax {
colmax = absval;
imax = i;
}
}
}
if colmax == T::zero() && absakk == T::zero() {
ipiv[k] = (k + 1) as i32;
return Err(BunchKaufmanError::Singular { index: k });
}
if absakk >= alpha_val * colmax {
kstep = 1;
ipiv[k] = (k + 1) as i32;
} else {
let mut rowmax = T::zero();
for j in k..imax {
let absval = Scalar::abs(a[(imax, j)]);
if absval > rowmax {
rowmax = absval;
}
}
if imax + 1 < n {
for i in (imax + 1)..n {
let absval = Scalar::abs(a[(i, imax)]);
if absval > rowmax {
rowmax = absval;
}
}
}
let absaimax = Scalar::abs(a[(imax, imax)]);
if absakk >= alpha_val * colmax * (colmax / rowmax) {
kstep = 1;
ipiv[k] = (k + 1) as i32;
} else if absaimax >= alpha_val * rowmax {
kstep = 1;
Self::swap_rows_cols_lower(a, k, imax, n);
ipiv[k] = (imax + 1) as i32;
} else {
kstep = 2;
if imax != k + 1 {
Self::swap_rows_cols_lower(a, k + 1, imax, n);
}
ipiv[k] = -((imax + 1) as i32);
ipiv[k + 1] = -((imax + 1) as i32);
}
}
if kstep == 1 {
let akk = a[(k, k)];
if akk == T::zero() {
return Err(BunchKaufmanError::Singular { index: k });
}
let akk_inv = T::one() / akk;
for i in (k + 1)..n {
a[(i, k)] = a[(i, k)] * akk_inv;
}
for j in (k + 1)..n {
let ajk = a[(j, k)];
for i in j..n {
a[(i, j)] = a[(i, j)] - ajk * a[(i, k)] * akk;
}
}
} else {
let akk = a[(k, k)];
let akp1k = a[(k + 1, k)];
let akp1kp1 = a[(k + 1, k + 1)];
let det = akk * akp1kp1 - akp1k * akp1k;
if det == T::zero() {
return Err(BunchKaufmanError::Singular { index: k });
}
let det_inv = T::one() / det;
let d11 = akp1kp1 * det_inv;
let d22 = akk * det_inv;
let d21 = -akp1k * det_inv;
for i in (k + 2)..n {
let wk = a[(i, k)];
let wkp1 = a[(i, k + 1)];
a[(i, k)] = wk * d11 + wkp1 * d21;
a[(i, k + 1)] = wk * d21 + wkp1 * d22;
}
for j in (k + 2)..n {
let ljk = a[(j, k)];
let ljkp1 = a[(j, k + 1)];
let djk = akk * ljk + akp1k * ljkp1;
let djkp1 = akp1k * ljk + akp1kp1 * ljkp1;
for i in j..n {
let lik = a[(i, k)];
let likp1 = a[(i, k + 1)];
a[(i, j)] = a[(i, j)] - lik * djk - likp1 * djkp1;
}
}
}
k += kstep;
}
Ok(())
}
fn factor_upper(a: &mut Mat<T>, ipiv: &mut [i32], n: usize) -> Result<(), BunchKaufmanError> {
let alpha_val = alpha::<T>();
let mut k = n;
while k > 0 {
k -= 1;
let kstep;
let absakk = Scalar::abs(a[(k, k)]);
let mut imax = 0;
let mut colmax = T::zero();
if k > 0 {
for i in 0..k {
let absval = Scalar::abs(a[(i, k)]);
if absval > colmax {
colmax = absval;
imax = i;
}
}
}
if colmax == T::zero() && absakk == T::zero() {
ipiv[k] = (k + 1) as i32;
return Err(BunchKaufmanError::Singular { index: k });
}
if absakk >= alpha_val * colmax {
kstep = 1;
ipiv[k] = (k + 1) as i32;
} else {
let mut rowmax = T::zero();
for j in (imax + 1)..=k {
let absval = Scalar::abs(a[(imax, j)]);
if absval > rowmax {
rowmax = absval;
}
}
if imax > 0 {
for i in 0..imax {
let absval = Scalar::abs(a[(i, imax)]);
if absval > rowmax {
rowmax = absval;
}
}
}
let absaimax = Scalar::abs(a[(imax, imax)]);
if absakk >= alpha_val * colmax * (colmax / rowmax) {
kstep = 1;
ipiv[k] = (k + 1) as i32;
} else if absaimax >= alpha_val * rowmax {
kstep = 1;
Self::swap_rows_cols_upper(a, k, imax, n);
ipiv[k] = (imax + 1) as i32;
} else {
kstep = 2;
if imax != k - 1 {
Self::swap_rows_cols_upper(a, k - 1, imax, n);
}
ipiv[k] = -((imax + 1) as i32);
ipiv[k - 1] = -((imax + 1) as i32);
}
}
if kstep == 1 {
let akk = a[(k, k)];
if akk == T::zero() {
return Err(BunchKaufmanError::Singular { index: k });
}
let akk_inv = T::one() / akk;
for i in 0..k {
a[(i, k)] = a[(i, k)] * akk_inv;
}
for j in 0..k {
let ajk = a[(j, k)];
for i in 0..=j {
a[(i, j)] = a[(i, j)] - ajk * a[(i, k)] * akk;
}
}
} else {
let km1 = k - 1;
let akm1km1 = a[(km1, km1)];
let akkm1 = a[(km1, k)];
let akk = a[(k, k)];
let det = akm1km1 * akk - akkm1 * akkm1;
if det == T::zero() {
return Err(BunchKaufmanError::Singular { index: k });
}
let det_inv = T::one() / det;
let d11 = akk * det_inv;
let d22 = akm1km1 * det_inv;
let d12 = -akkm1 * det_inv;
for i in 0..km1 {
let wkm1 = a[(i, km1)];
let wk = a[(i, k)];
a[(i, km1)] = wkm1 * d11 + wk * d12;
a[(i, k)] = wkm1 * d12 + wk * d22;
}
for j in 0..km1 {
let ljkm1 = a[(j, km1)];
let ljk = a[(j, k)];
let djkm1 = akm1km1 * ljkm1 + akkm1 * ljk;
let djk = akkm1 * ljkm1 + akk * ljk;
for i in 0..=j {
let likm1 = a[(i, km1)];
let lik = a[(i, k)];
a[(i, j)] = a[(i, j)] - likm1 * djkm1 - lik * djk;
}
}
k -= 1; }
}
Ok(())
}
fn swap_rows_cols_lower(a: &mut Mat<T>, i: usize, j: usize, n: usize) {
if i == j {
return;
}
let (i, j) = if i < j { (i, j) } else { (j, i) };
let tmp = a[(i, i)];
a[(i, i)] = a[(j, j)];
a[(j, j)] = tmp;
for k in 0..i {
let tmp = a[(i, k)];
a[(i, k)] = a[(j, k)];
a[(j, k)] = tmp;
}
for k in (i + 1)..j {
let tmp = a[(k, i)];
a[(k, i)] = a[(j, k)];
a[(j, k)] = tmp;
}
for k in (j + 1)..n {
let tmp = a[(k, i)];
a[(k, i)] = a[(k, j)];
a[(k, j)] = tmp;
}
}
fn swap_rows_cols_upper(a: &mut Mat<T>, i: usize, j: usize, n: usize) {
if i == j {
return;
}
let (i, j) = if i < j { (i, j) } else { (j, i) };
let tmp = a[(i, i)];
a[(i, i)] = a[(j, j)];
a[(j, j)] = tmp;
for k in 0..i {
let tmp = a[(k, i)];
a[(k, i)] = a[(k, j)];
a[(k, j)] = tmp;
}
for k in (i + 1)..j {
let tmp = a[(i, k)];
a[(i, k)] = a[(k, j)];
a[(k, j)] = tmp;
}
for k in (j + 1)..n {
let tmp = a[(i, k)];
a[(i, k)] = a[(j, k)];
a[(j, k)] = tmp;
}
}
pub fn solve(&self, b: MatRef<'_, T>) -> Result<Mat<T>, BunchKaufmanError> {
if b.nrows() != self.n {
return Err(BunchKaufmanError::DimensionMismatch {
expected: self.n,
actual: b.nrows(),
});
}
let nrhs = b.ncols();
let mut x = Mat::zeros(self.n, nrhs);
for i in 0..self.n {
for j in 0..nrhs {
x[(i, j)] = b[(i, j)];
}
}
match self.uplo {
Uplo::Lower => self.solve_lower(&mut x, nrhs),
Uplo::Upper => self.solve_upper(&mut x, nrhs),
}
Ok(x)
}
fn solve_lower(&self, x: &mut Mat<T>, nrhs: usize) {
let n = self.n;
let mut k = 0;
while k < n {
if self.ipiv[k] > 0 {
let kp = (self.ipiv[k] - 1) as usize;
if kp != k {
for j in 0..nrhs {
let tmp = x[(k, j)];
x[(k, j)] = x[(kp, j)];
x[(kp, j)] = tmp;
}
}
for i in (k + 1)..n {
let lik = self.factors[(i, k)];
for j in 0..nrhs {
x[(i, j)] = x[(i, j)] - lik * x[(k, j)];
}
}
k += 1;
} else {
let kp = (-self.ipiv[k] - 1) as usize;
if kp != k + 1 {
for j in 0..nrhs {
let tmp = x[(k + 1, j)];
x[(k + 1, j)] = x[(kp, j)];
x[(kp, j)] = tmp;
}
}
for i in (k + 2)..n {
let lik = self.factors[(i, k)];
let likp1 = self.factors[(i, k + 1)];
for j in 0..nrhs {
x[(i, j)] = x[(i, j)] - lik * x[(k, j)] - likp1 * x[(k + 1, j)];
}
}
k += 2;
}
}
k = 0;
while k < n {
if self.ipiv[k] > 0 {
let dkk = self.factors[(k, k)];
for j in 0..nrhs {
x[(k, j)] = x[(k, j)] / dkk;
}
k += 1;
} else {
let dkk = self.factors[(k, k)];
let dkp1k = self.factors[(k + 1, k)];
let dkp1kp1 = self.factors[(k + 1, k + 1)];
let det = dkk * dkp1kp1 - dkp1k * dkp1k;
for j in 0..nrhs {
let xk = x[(k, j)];
let xkp1 = x[(k + 1, j)];
x[(k, j)] = (dkp1kp1 * xk - dkp1k * xkp1) / det;
x[(k + 1, j)] = (dkk * xkp1 - dkp1k * xk) / det;
}
k += 2;
}
}
k = n;
while k > 0 {
k -= 1;
if self.ipiv[k] > 0 {
for i in (k + 1)..n {
let lik = self.factors[(i, k)];
for j in 0..nrhs {
x[(k, j)] = x[(k, j)] - lik * x[(i, j)];
}
}
let kp = (self.ipiv[k] - 1) as usize;
if kp != k {
for j in 0..nrhs {
let tmp = x[(k, j)];
x[(k, j)] = x[(kp, j)];
x[(kp, j)] = tmp;
}
}
} else if k > 0 && self.ipiv[k - 1] < 0 {
k -= 1;
for i in (k + 2)..n {
let lik = self.factors[(i, k)];
let likp1 = self.factors[(i, k + 1)];
for j in 0..nrhs {
x[(k, j)] = x[(k, j)] - lik * x[(i, j)];
x[(k + 1, j)] = x[(k + 1, j)] - likp1 * x[(i, j)];
}
}
let kp = (-self.ipiv[k] - 1) as usize;
if kp != k + 1 {
for j in 0..nrhs {
let tmp = x[(k + 1, j)];
x[(k + 1, j)] = x[(kp, j)];
x[(kp, j)] = tmp;
}
}
}
}
}
fn solve_upper(&self, x: &mut Mat<T>, nrhs: usize) {
let n = self.n;
let mut k = n;
while k > 0 {
k -= 1;
if self.ipiv[k] > 0 {
let kp = (self.ipiv[k] - 1) as usize;
if kp != k {
for j in 0..nrhs {
let tmp = x[(k, j)];
x[(k, j)] = x[(kp, j)];
x[(kp, j)] = tmp;
}
}
for i in 0..k {
let uik = self.factors[(i, k)];
for j in 0..nrhs {
x[(i, j)] = x[(i, j)] - uik * x[(k, j)];
}
}
} else if k > 0 && self.ipiv[k - 1] < 0 {
k -= 1;
let kp = (-self.ipiv[k] - 1) as usize;
if kp != k {
for j in 0..nrhs {
let tmp = x[(k, j)];
x[(k, j)] = x[(kp, j)];
x[(kp, j)] = tmp;
}
}
for i in 0..k {
let uik = self.factors[(i, k)];
let uikp1 = self.factors[(i, k + 1)];
for j in 0..nrhs {
x[(i, j)] = x[(i, j)] - uik * x[(k, j)] - uikp1 * x[(k + 1, j)];
}
}
}
}
k = n;
while k > 0 {
k -= 1;
if self.ipiv[k] > 0 {
let dkk = self.factors[(k, k)];
for j in 0..nrhs {
x[(k, j)] = x[(k, j)] / dkk;
}
} else if k > 0 && self.ipiv[k - 1] < 0 {
k -= 1;
let dkk = self.factors[(k, k)];
let dkkp1 = self.factors[(k, k + 1)];
let dkp1kp1 = self.factors[(k + 1, k + 1)];
let det = dkk * dkp1kp1 - dkkp1 * dkkp1;
for j in 0..nrhs {
let xk = x[(k, j)];
let xkp1 = x[(k + 1, j)];
x[(k, j)] = (dkp1kp1 * xk - dkkp1 * xkp1) / det;
x[(k + 1, j)] = (dkk * xkp1 - dkkp1 * xk) / det;
}
}
}
k = 0;
while k < n {
if self.ipiv[k] > 0 {
for i in 0..k {
let uik = self.factors[(i, k)];
for j in 0..nrhs {
x[(k, j)] = x[(k, j)] - uik * x[(i, j)];
}
}
let kp = (self.ipiv[k] - 1) as usize;
if kp != k {
for j in 0..nrhs {
let tmp = x[(k, j)];
x[(k, j)] = x[(kp, j)];
x[(kp, j)] = tmp;
}
}
k += 1;
} else {
for i in 0..k {
let uik = self.factors[(i, k)];
let uikp1 = self.factors[(i, k + 1)];
for j in 0..nrhs {
x[(k, j)] = x[(k, j)] - uik * x[(i, j)];
x[(k + 1, j)] = x[(k + 1, j)] - uikp1 * x[(i, j)];
}
}
let kp = (-self.ipiv[k] - 1) as usize;
if kp != k + 1 {
for j in 0..nrhs {
let tmp = x[(k + 1, j)];
x[(k + 1, j)] = x[(kp, j)];
x[(kp, j)] = tmp;
}
}
k += 2;
}
}
}
#[must_use]
pub fn ipiv(&self) -> &[i32] {
&self.ipiv
}
#[must_use]
pub fn n(&self) -> usize {
self.n
}
#[must_use]
pub fn uplo(&self) -> Uplo {
self.uplo
}
#[must_use]
pub fn inertia(&self) -> (usize, usize, usize) {
let mut n_pos = 0;
let mut n_neg = 0;
let mut n_zero = 0;
let mut k = 0;
while k < self.n {
if self.ipiv[k] > 0 {
let d = self.factors[(k, k)];
if d.real() > T::Real::zero() {
n_pos += 1;
} else if d.real() < T::Real::zero() {
n_neg += 1;
} else {
n_zero += 1;
}
k += 1;
} else {
let (d11, d21, d22) = match self.uplo {
Uplo::Lower => {
let d11 = self.factors[(k, k)];
let d21 = self.factors[(k + 1, k)];
let d22 = self.factors[(k + 1, k + 1)];
(d11, d21, d22)
}
Uplo::Upper => {
let d11 = self.factors[(k, k)];
let d12 = self.factors[(k, k + 1)];
let d22 = self.factors[(k + 1, k + 1)];
(d11, d12, d22)
}
};
let trace = d11 + d22;
let det = d11 * d22 - d21 * d21;
let discriminant = trace * trace - T::from_f64(4.0).unwrap_or_else(T::zero) * det;
if discriminant.real() >= T::Real::zero() {
let sqrt_disc = Real::sqrt(discriminant);
let two = T::from_f64(2.0).unwrap_or_else(T::zero);
let lambda1 = (trace + sqrt_disc) / two;
let lambda2 = (trace - sqrt_disc) / two;
if lambda1.real() > T::Real::zero() {
n_pos += 1;
} else if lambda1.real() < T::Real::zero() {
n_neg += 1;
} else {
n_zero += 1;
}
if lambda2.real() > T::Real::zero() {
n_pos += 1;
} else if lambda2.real() < T::Real::zero() {
n_neg += 1;
} else {
n_zero += 1;
}
} else {
n_zero += 2;
}
k += 2;
}
}
(n_pos, n_neg, n_zero)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_bunch_kaufman_positive_definite() {
let a = Mat::from_rows(&[&[4.0f64, 2.0, 1.0], &[2.0, 5.0, 2.0], &[1.0, 2.0, 6.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let b = Mat::from_rows(&[&[1.0], &[2.0], &[3.0]]);
let x = bk.solve(b.as_ref()).unwrap();
for i in 0..3 {
let mut ax_i = 0.0;
for j in 0..3 {
ax_i += a[(i, j)] * x[(j, 0)];
}
assert!(
approx_eq(ax_i, b[(i, 0)], 1e-10),
"Ax[{}] = {}, b = {}",
i,
ax_i,
b[(i, 0)]
);
}
}
#[test]
fn test_bunch_kaufman_indefinite() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 0.0], &[2.0, 0.0, 3.0], &[0.0, 3.0, 4.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let b = Mat::from_rows(&[&[1.0], &[2.0], &[3.0]]);
let x = bk.solve(b.as_ref()).unwrap();
for i in 0..3 {
let mut ax_i = 0.0;
for j in 0..3 {
ax_i += a[(i, j)] * x[(j, 0)];
}
assert!(
approx_eq(ax_i, b[(i, 0)], 1e-10),
"Ax[{}] = {}, b = {}",
i,
ax_i,
b[(i, 0)]
);
}
}
#[test]
fn test_bunch_kaufman_2x2() {
let a = Mat::from_rows(&[&[0.0f64, 1.0], &[1.0, 0.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let b = Mat::from_rows(&[&[1.0], &[2.0]]);
let x = bk.solve(b.as_ref()).unwrap();
assert!(approx_eq(x[(0, 0)], 2.0, 1e-10));
assert!(approx_eq(x[(1, 0)], 1.0, 1e-10));
}
#[test]
fn test_bunch_kaufman_diagonal() {
let a = Mat::from_rows(&[&[2.0f64, 0.0, 0.0], &[0.0, -3.0, 0.0], &[0.0, 0.0, 4.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let b = Mat::from_rows(&[&[2.0], &[-6.0], &[12.0]]);
let x = bk.solve(b.as_ref()).unwrap();
assert!(approx_eq(x[(0, 0)], 1.0, 1e-10));
assert!(approx_eq(x[(1, 0)], 2.0, 1e-10));
assert!(approx_eq(x[(2, 0)], 3.0, 1e-10));
}
#[test]
fn test_bunch_kaufman_multiple_rhs() {
let a = Mat::from_rows(&[&[4.0f64, 2.0], &[2.0, 5.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let b = Mat::from_rows(&[&[1.0, 4.0], &[2.0, 5.0]]);
let x = bk.solve(b.as_ref()).unwrap();
for col in 0..2 {
for i in 0..2 {
let mut ax_i = 0.0;
for j in 0..2 {
ax_i += a[(i, j)] * x[(j, col)];
}
assert!(
approx_eq(ax_i, b[(i, col)], 1e-10),
"Ax[{},{}] = {}, b = {}",
i,
col,
ax_i,
b[(i, col)]
);
}
}
}
#[test]
fn test_bunch_kaufman_inertia_positive_definite() {
let a = Mat::from_rows(&[&[4.0f64, 2.0], &[2.0, 5.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let (pos, neg, zero) = bk.inertia();
assert_eq!(pos, 2);
assert_eq!(neg, 0);
assert_eq!(zero, 0);
}
#[test]
fn test_bunch_kaufman_inertia_indefinite() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[2.0, 1.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let (pos, neg, zero) = bk.inertia();
assert_eq!(pos, 1);
assert_eq!(neg, 1);
assert_eq!(zero, 0);
}
#[test]
fn test_bunch_kaufman_upper() {
let a = Mat::from_rows(&[&[4.0f64, 2.0, 1.0], &[2.0, 5.0, 2.0], &[1.0, 2.0, 6.0]]);
let bk = BunchKaufman::compute_with_uplo(a.as_ref(), Uplo::Upper).unwrap();
let b = Mat::from_rows(&[&[1.0], &[2.0], &[3.0]]);
let x = bk.solve(b.as_ref()).unwrap();
for i in 0..3 {
let mut ax_i = 0.0;
for j in 0..3 {
ax_i += a[(i, j)] * x[(j, 0)];
}
assert!(
approx_eq(ax_i, b[(i, 0)], 1e-10),
"Ax[{}] = {}, b = {}",
i,
ax_i,
b[(i, 0)]
);
}
}
#[test]
fn test_bunch_kaufman_f32() {
let a = Mat::from_rows(&[&[4.0f32, 2.0, 1.0], &[2.0, 5.0, 2.0], &[1.0, 2.0, 6.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let b = Mat::from_rows(&[&[1.0f32], &[2.0], &[3.0]]);
let x = bk.solve(b.as_ref()).unwrap();
for i in 0..3 {
let mut ax_i = 0.0f32;
for j in 0..3 {
ax_i += a[(i, j)] * x[(j, 0)];
}
assert!(
(ax_i - b[(i, 0)]).abs() < 1e-5,
"Ax[{}] = {}, b = {}",
i,
ax_i,
b[(i, 0)]
);
}
}
#[test]
fn test_bunch_kaufman_identity() {
let a: Mat<f64> = Mat::eye(4);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let b = Mat::from_rows(&[&[1.0], &[2.0], &[3.0], &[4.0]]);
let x = bk.solve(b.as_ref()).unwrap();
for i in 0..4 {
assert!(approx_eq(x[(i, 0)], b[(i, 0)], 1e-10));
}
}
#[test]
fn test_bunch_kaufman_large() {
let n = 8;
let mut a = Mat::zeros(n, n);
for i in 0..n {
a[(i, i)] = if i % 2 == 0 { 10.0 } else { -10.0 };
for j in (i + 1)..n {
let val = 0.1 / ((i + j + 1) as f64);
a[(i, j)] = val;
a[(j, i)] = val;
}
}
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let mut b = Mat::zeros(n, 1);
for i in 0..n {
b[(i, 0)] = (i as f64) + 1.0;
}
let x = bk.solve(b.as_ref()).unwrap();
for i in 0..n {
let mut ax_i = 0.0;
for j in 0..n {
ax_i += a[(i, j)] * x[(j, 0)];
}
assert!(
approx_eq(ax_i, b[(i, 0)], 1e-8),
"Ax[{}] = {}, b = {}",
i,
ax_i,
b[(i, 0)]
);
}
}
#[test]
fn test_bunch_kaufman_negative_definite() {
let a = Mat::from_rows(&[&[-4.0f64, -2.0], &[-2.0, -5.0]]);
let bk = BunchKaufman::compute(a.as_ref()).unwrap();
let (pos, neg, zero) = bk.inertia();
assert_eq!(pos, 0);
assert_eq!(neg, 2);
assert_eq!(zero, 0);
}
}