use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
use crate::cholesky::{Cholesky, CholeskyError};
use crate::utils::{norm_1, norm_inf};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ExpertCholeskySolveError {
NotSquare,
DimensionMismatch,
NotPositiveDefinite,
SingularToWorkingPrecision,
}
impl core::fmt::Display for ExpertCholeskySolveError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::NotSquare => write!(f, "Matrix must be square"),
Self::DimensionMismatch => write!(f, "Matrix and vector dimensions do not match"),
Self::NotPositiveDefinite => write!(f, "Matrix is not positive definite"),
Self::SingularToWorkingPrecision => {
write!(f, "Matrix is singular to working precision")
}
}
}
}
impl std::error::Error for ExpertCholeskySolveError {}
impl From<CholeskyError> for ExpertCholeskySolveError {
fn from(e: CholeskyError) -> Self {
match e {
CholeskyError::NotPositiveDefinite { .. } => Self::NotPositiveDefinite,
CholeskyError::NotSquare { .. } => Self::NotSquare,
CholeskyError::DimensionMismatch { .. } => Self::DimensionMismatch,
}
}
}
#[derive(Debug, Clone)]
pub struct ExpertCholeskySolveResult<T: Scalar> {
pub solution: Mat<T>,
pub rcond: T,
pub forward_error: Vec<T>,
pub backward_error: Vec<T>,
pub scale: Option<Vec<T>>,
pub equilibrated: bool,
}
pub fn solve_cholesky_expert<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
b: MatRef<'_, T>,
equilibrate: bool,
) -> Result<ExpertCholeskySolveResult<T>, ExpertCholeskySolveError> {
let n = a.nrows();
let nrhs = b.ncols();
if a.ncols() != n {
return Err(ExpertCholeskySolveError::NotSquare);
}
if b.nrows() != n {
return Err(ExpertCholeskySolveError::DimensionMismatch);
}
if n == 0 {
return Ok(ExpertCholeskySolveResult {
solution: Mat::zeros(0, nrhs),
rcond: T::one(),
forward_error: vec![],
backward_error: vec![],
scale: None,
equilibrated: false,
});
}
let (a_scaled, b_scaled, scale, did_equilibrate) = if equilibrate {
apply_symmetric_equilibration(a, b)
} else {
let mut a_copy = Mat::zeros(n, n);
let mut b_copy = Mat::zeros(n, nrhs);
for i in 0..n {
for j in 0..n {
a_copy[(i, j)] = a[(i, j)];
}
for j in 0..nrhs {
b_copy[(i, j)] = b[(i, j)];
}
}
(a_copy, b_copy, None, false)
};
let anorm = norm_1(a_scaled.as_ref());
let chol = Cholesky::compute(a_scaled.as_ref())?;
let x_scaled = chol.solve(b_scaled.as_ref())?;
let rcond = estimate_rcond_cholesky(&chol, anorm, n);
let eps = <T as Scalar>::epsilon();
if rcond < eps {
return Err(ExpertCholeskySolveError::SingularToWorkingPrecision);
}
let (forward_error, backward_error) =
compute_error_bounds_cholesky(&a_scaled, &b_scaled, &x_scaled, &chol, n, nrhs);
let solution = if let Some(ref s) = scale {
unscale_solution_symmetric(&x_scaled, s)
} else {
x_scaled
};
Ok(ExpertCholeskySolveResult {
solution,
rcond,
forward_error,
backward_error,
scale,
equilibrated: did_equilibrate,
})
}
fn apply_symmetric_equilibration<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
b: MatRef<'_, T>,
) -> (Mat<T>, Mat<T>, Option<Vec<T>>, bool) {
let n = a.nrows();
let nrhs = b.ncols();
let mut scale = vec![T::one(); n];
let mut needs_scaling = false;
for i in 0..n {
let diag = a[(i, i)];
if diag > T::zero() {
let s = T::one() / Real::sqrt(diag);
if Scalar::abs(s - T::one()) > T::from_f64(0.1).unwrap_or(T::zero()) {
scale[i] = s;
needs_scaling = true;
}
}
}
if !needs_scaling {
let mut a_copy = Mat::zeros(n, n);
let mut b_copy = Mat::zeros(n, nrhs);
for i in 0..n {
for j in 0..n {
a_copy[(i, j)] = a[(i, j)];
}
for j in 0..nrhs {
b_copy[(i, j)] = b[(i, j)];
}
}
return (a_copy, b_copy, None, false);
}
let mut a_scaled = Mat::zeros(n, n);
let mut b_scaled = Mat::zeros(n, nrhs);
for i in 0..n {
for j in 0..n {
a_scaled[(i, j)] = scale[i] * a[(i, j)] * scale[j];
}
for j in 0..nrhs {
b_scaled[(i, j)] = scale[i] * b[(i, j)];
}
}
(a_scaled, b_scaled, Some(scale), true)
}
fn estimate_rcond_cholesky<T: Field + Real + bytemuck::Zeroable>(
chol: &Cholesky<T>,
anorm: T,
n: usize,
) -> T {
if anorm <= T::zero() || n == 0 {
return T::zero();
}
let ainv_norm_est = match hager_higham_inv_1norm_spd(chol, n) {
Some(v) => v,
None => return T::zero(),
};
let kappa_est = anorm * ainv_norm_est;
if kappa_est <= T::zero() {
T::zero()
} else {
T::one() / kappa_est
}
}
const HAGER_HIGHAM_ITMAX: usize = 5;
fn hager_higham_inv_1norm_spd<T: Field + Real + bytemuck::Zeroable>(
chol: &Cholesky<T>,
n: usize,
) -> Option<T> {
if n == 0 {
return Some(T::zero());
}
let apply = |v: &Mat<T>| -> Option<Mat<T>> { chol.solve(v.as_ref()).ok() };
let one_norm = |v: &Mat<T>| -> T {
let mut s = T::zero();
for i in 0..n {
s = s + Scalar::abs(v[(i, 0)]);
}
s
};
let sign_of = |val: T| -> T {
if val >= T::zero() {
T::one()
} else {
-T::one()
}
};
let argmax_abs = |v: &Mat<T>| -> usize {
let mut j = 0usize;
let mut best = Scalar::abs(v[(0, 0)]);
for i in 1..n {
let a = Scalar::abs(v[(i, 0)]);
if a > best {
best = a;
j = i;
}
}
j
};
if n == 1 {
let mut e = Mat::zeros(1, 1);
e[(0, 0)] = T::one();
let y = apply(&e)?;
return Some(Scalar::abs(y[(0, 0)]));
}
let mut x = Mat::zeros(n, 1);
let scale = T::one() / T::from_usize(n)?;
for i in 0..n {
x[(i, 0)] = scale;
}
let y = apply(&x)?;
let mut est = one_norm(&y);
let mut isgn = vec![T::zero(); n];
for i in 0..n {
isgn[i] = sign_of(y[(i, 0)]);
x[(i, 0)] = isgn[i];
}
let z = apply(&x)?;
let mut j = argmax_abs(&z);
let mut iter = 2usize;
loop {
x = Mat::zeros(n, 1);
x[(j, 0)] = T::one();
let v = apply(&x)?;
let est_old = est;
est = one_norm(&v);
let sign_matches = (0..n).all(|i| sign_of(v[(i, 0)]) == isgn[i]);
if sign_matches {
break;
}
if est <= est_old {
break;
}
for i in 0..n {
isgn[i] = sign_of(v[(i, 0)]);
x[(i, 0)] = isgn[i];
}
let z2 = apply(&x)?;
let j_last = j;
j = argmax_abs(&z2);
let converged = z2[(j_last, 0)] == Scalar::abs(z2[(j, 0)]);
if converged || iter >= HAGER_HIGHAM_ITMAX {
break;
}
iter += 1;
}
let mut x_alt = Mat::zeros(n, 1);
let mut altsgn = T::one();
let denom = T::from_usize(n - 1)?;
for i in 0..n {
let weight = T::one() + T::from_usize(i)? / denom;
x_alt[(i, 0)] = altsgn * weight;
altsgn = -altsgn;
}
let y_alt = apply(&x_alt)?;
let temp = T::from_f64(2.0)? * (one_norm(&y_alt) / T::from_usize(3 * n)?);
if temp > est {
est = temp;
}
Some(est)
}
fn compute_error_bounds_cholesky<T: Field + Real + bytemuck::Zeroable>(
a: &Mat<T>,
b: &Mat<T>,
x: &Mat<T>,
chol: &Cholesky<T>,
n: usize,
nrhs: usize,
) -> (Vec<T>, Vec<T>) {
let eps = <T as Scalar>::epsilon();
let mut forward_error = vec![T::zero(); nrhs];
let mut backward_error = vec![T::zero(); nrhs];
for col in 0..nrhs {
let mut r = Mat::zeros(n, 1);
for i in 0..n {
let mut ax_i = T::zero();
for j in 0..n {
ax_i = ax_i + a[(i, j)] * x[(j, col)];
}
r[(i, 0)] = b[(i, col)] - ax_i;
}
let mut r_inf = T::zero();
let mut x_inf = T::zero();
let mut b_inf = T::zero();
for i in 0..n {
let abs_r = Scalar::abs(r[(i, 0)]);
let abs_x = Scalar::abs(x[(i, col)]);
let abs_b = Scalar::abs(b[(i, col)]);
if abs_r > r_inf {
r_inf = abs_r;
}
if abs_x > x_inf {
x_inf = abs_x;
}
if abs_b > b_inf {
b_inf = abs_b;
}
}
let a_inf = norm_inf(a.as_ref());
let denom = a_inf * x_inf + b_inf;
if denom > T::zero() {
backward_error[col] = r_inf / denom;
} else {
backward_error[col] = T::zero();
}
let e = match chol.solve(r.as_ref()) {
Ok(e) => e,
Err(_) => {
forward_error[col] = T::one();
continue;
}
};
let mut e_inf = T::zero();
for i in 0..n {
let abs_e = Scalar::abs(e[(i, 0)]);
if abs_e > e_inf {
e_inf = abs_e;
}
}
if x_inf > T::zero() {
forward_error[col] = e_inf / x_inf;
} else {
forward_error[col] = e_inf;
}
if forward_error[col] < eps {
forward_error[col] = eps;
}
if backward_error[col] < eps {
backward_error[col] = eps;
}
}
(forward_error, backward_error)
}
fn unscale_solution_symmetric<T: Field + Real + bytemuck::Zeroable>(
x: &Mat<T>,
scale: &[T],
) -> Mat<T> {
let n = x.nrows();
let nrhs = x.ncols();
let mut x_unscaled = Mat::zeros(n, nrhs);
for i in 0..n {
for j in 0..nrhs {
x_unscaled[(i, j)] = scale[i] * x[(i, j)];
}
}
x_unscaled
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_solve_cholesky_expert_simple() {
let a = Mat::from_rows(&[&[4.0f64, 2.0], &[2.0, 5.0]]);
let b = Mat::from_rows(&[&[8.0], &[11.0]]);
let result = solve_cholesky_expert(a.as_ref(), b.as_ref(), false).unwrap();
let x = &result.solution;
let ax0 = a[(0, 0)] * x[(0, 0)] + a[(0, 1)] * x[(1, 0)];
let ax1 = a[(1, 0)] * x[(0, 0)] + a[(1, 1)] * x[(1, 0)];
assert!(approx_eq(ax0, b[(0, 0)], 1e-10));
assert!(approx_eq(ax1, b[(1, 0)], 1e-10));
assert!(result.rcond > 0.0);
assert!(result.rcond <= 1.0);
}
#[test]
fn test_solve_cholesky_expert_identity() {
let a = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[0.0, 1.0, 0.0], &[0.0, 0.0, 1.0]]);
let b = Mat::from_rows(&[&[1.0], &[2.0], &[3.0]]);
let result = solve_cholesky_expert(a.as_ref(), b.as_ref(), false).unwrap();
for i in 0..3 {
assert!(approx_eq(result.solution[(i, 0)], b[(i, 0)], 1e-10));
}
assert!(approx_eq(result.rcond, 1.0, 0.1));
}
#[test]
fn test_solve_cholesky_expert_with_equilibration() {
let a = Mat::from_rows(&[&[1000.0f64, 1.0], &[1.0, 0.01]]);
let b = Mat::from_rows(&[&[1001.0], &[1.01]]);
let result = solve_cholesky_expert(a.as_ref(), b.as_ref(), true).unwrap();
let x = &result.solution;
let ax0 = a[(0, 0)] * x[(0, 0)] + a[(0, 1)] * x[(1, 0)];
let ax1 = a[(1, 0)] * x[(0, 0)] + a[(1, 1)] * x[(1, 0)];
assert!(approx_eq(ax0, b[(0, 0)], 1e-4));
assert!(approx_eq(ax1, b[(1, 0)], 1e-4));
assert!(result.equilibrated);
assert!(result.scale.is_some());
}
#[test]
fn test_solve_cholesky_expert_multiple_rhs() {
let a = Mat::from_rows(&[&[4.0f64, 2.0], &[2.0, 5.0]]);
let b = Mat::from_rows(&[&[8.0, 6.0], &[11.0, 9.0]]);
let result = solve_cholesky_expert(a.as_ref(), b.as_ref(), false).unwrap();
let x = &result.solution;
for col in 0..2 {
let ax0 = a[(0, 0)] * x[(0, col)] + a[(0, 1)] * x[(1, col)];
let ax1 = a[(1, 0)] * x[(0, col)] + a[(1, 1)] * x[(1, col)];
assert!(approx_eq(ax0, b[(0, col)], 1e-10));
assert!(approx_eq(ax1, b[(1, col)], 1e-10));
}
assert_eq!(result.forward_error.len(), 2);
assert_eq!(result.backward_error.len(), 2);
}
#[test]
fn test_solve_cholesky_expert_not_positive_definite() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[2.0, 1.0]]);
let b = Mat::from_rows(&[&[1.0], &[2.0]]);
let result = solve_cholesky_expert(a.as_ref(), b.as_ref(), false);
assert!(result.is_err());
assert_eq!(
result.unwrap_err(),
ExpertCholeskySolveError::NotPositiveDefinite
);
}
#[test]
fn test_solve_cholesky_expert_rcond_beyond_first_five_columns() {
let n = 7;
let mut a: Mat<f64> = Mat::zeros(n, n);
for i in 0..n - 1 {
a[(i, i)] = 1.0;
}
let tiny = 1e-8;
a[(n - 1, n - 1)] = tiny;
let mut b: Mat<f64> = Mat::zeros(n, 1);
for i in 0..n {
b[(i, 0)] = 1.0;
}
let result = solve_cholesky_expert(a.as_ref(), b.as_ref(), false).unwrap();
assert!(
result.rcond < 1e-6,
"rcond {} should reflect the true ill-conditioning (expected ~1e-8)",
result.rcond
);
}
#[test]
fn test_solve_cholesky_expert_f32() {
let a = Mat::from_rows(&[&[4.0f32, 2.0], &[2.0, 5.0]]);
let b = Mat::from_rows(&[&[8.0f32], &[11.0]]);
let result = solve_cholesky_expert(a.as_ref(), b.as_ref(), false).unwrap();
let x = &result.solution;
let ax0 = a[(0, 0)] * x[(0, 0)] + a[(0, 1)] * x[(1, 0)];
let ax1 = a[(1, 0)] * x[(0, 0)] + a[(1, 1)] * x[(1, 0)];
assert!((ax0 - b[(0, 0)]).abs() < 1e-5);
assert!((ax1 - b[(1, 0)]).abs() < 1e-5);
}
}