use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
use super::norms::{norm_1, norm_inf};
use crate::lu::{Lu, LuError};
use crate::svd::{Svd, SvdError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CondError {
NotSquare,
EmptyMatrix,
SvdFailed,
LuFailed,
}
impl core::fmt::Display for CondError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::NotSquare => write!(f, "Matrix must be square"),
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::SvdFailed => write!(f, "SVD computation failed"),
Self::LuFailed => write!(f, "LU computation failed"),
}
}
}
impl std::error::Error for CondError {}
impl From<SvdError> for CondError {
fn from(e: SvdError) -> Self {
match e {
SvdError::EmptyMatrix => Self::EmptyMatrix,
SvdError::NotConverged => Self::SvdFailed,
}
}
}
impl From<LuError> for CondError {
fn from(e: LuError) -> Self {
match e {
LuError::NotSquare { .. } => Self::NotSquare,
_ => Self::LuFailed,
}
}
}
pub fn cond<T: Field + Real + bytemuck::Zeroable>(a: MatRef<'_, T>) -> Result<T, CondError> {
let svd = Svd::compute(a)?;
Ok(svd.condition_number())
}
pub fn cond_1<T: Field + Real + bytemuck::Zeroable>(a: MatRef<'_, T>) -> Result<T, CondError> {
let n = a.nrows();
if n != a.ncols() {
return Err(CondError::NotSquare);
}
if n == 0 {
return Err(CondError::EmptyMatrix);
}
let norm_a = norm_1(a);
let lu = match Lu::compute(a) {
Ok(lu) => lu,
Err(_) => return Ok(<T as Scalar>::max_value()), };
let a_inv = match lu.inverse() {
Ok(inv) => inv,
Err(_) => return Ok(<T as Scalar>::max_value()),
};
let norm_a_inv = norm_1(a_inv.as_ref());
Ok(norm_a * norm_a_inv)
}
pub fn cond_inf<T: Field + Real + bytemuck::Zeroable>(a: MatRef<'_, T>) -> Result<T, CondError> {
let n = a.nrows();
if n != a.ncols() {
return Err(CondError::NotSquare);
}
if n == 0 {
return Err(CondError::EmptyMatrix);
}
let norm_a = norm_inf(a);
let lu = match Lu::compute(a) {
Ok(lu) => lu,
Err(_) => return Ok(<T as Scalar>::max_value()),
};
let a_inv = match lu.inverse() {
Ok(inv) => inv,
Err(_) => return Ok(<T as Scalar>::max_value()),
};
let norm_a_inv = norm_inf(a_inv.as_ref());
Ok(norm_a * norm_a_inv)
}
pub fn rcond<T: Field + Real + bytemuck::Zeroable>(a: MatRef<'_, T>) -> Result<T, CondError> {
let kappa = cond_1(a)?;
if kappa >= <T as Scalar>::max_value() / T::from_f64(2.0).unwrap_or(T::one()) {
Ok(T::zero())
} else {
Ok(T::one() / kappa)
}
}
pub fn rcond_estimate<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
) -> Result<T, CondError> {
let n = a.nrows();
if n != a.ncols() {
return Err(CondError::NotSquare);
}
if n == 0 {
return Err(CondError::EmptyMatrix);
}
let norm_a = norm_1(a);
let lu = match Lu::compute(a) {
Ok(lu) => lu,
Err(_) => return Ok(T::zero()), };
let norm_a_inv_est = match hager_higham_inv_norm(&lu, n) {
Some(v) => v,
None => return Ok(T::zero()), };
let kappa_est = norm_a * norm_a_inv_est;
if kappa_est <= T::zero()
|| kappa_est >= <T as Scalar>::max_value() / T::from_f64(2.0).unwrap_or(T::one())
{
Ok(T::zero())
} else {
Ok(T::one() / kappa_est)
}
}
fn column_one_norm<T: Real>(x: &Mat<T>, n: usize) -> T {
let mut sum = T::zero();
for i in 0..n {
sum = sum + Scalar::abs(x[(i, 0)]);
}
sum
}
fn column_argmax_abs<T: Real>(x: &Mat<T>, n: usize) -> usize {
let mut best = 0usize;
let mut best_val = Scalar::abs(x[(0, 0)]);
for i in 1..n {
let val = Scalar::abs(x[(i, 0)]);
if val > best_val {
best_val = val;
best = i;
}
}
best
}
fn hager_higham_inv_norm<T: Field + Real + bytemuck::Zeroable>(lu: &Lu<T>, n: usize) -> Option<T> {
const ITMAX: usize = 5;
let sign_of = |val: T| -> T {
if val >= T::zero() {
T::one()
} else {
-T::one()
}
};
let mut x = Mat::<T>::zeros(n, 1);
let inv_n = T::one() / T::from_f64(n as f64)?;
for i in 0..n {
x[(i, 0)] = inv_n;
}
let mut v = lu.solve(x.as_ref()).ok()?;
if n == 1 {
return Some(Scalar::abs(v[(0, 0)]));
}
let mut est = column_one_norm(&v, n);
let mut isgn = vec![T::one(); n];
for i in 0..n {
let s = sign_of(v[(i, 0)]);
isgn[i] = s;
x[(i, 0)] = s;
}
let mut xt = lu.solve_transpose(x.as_ref()).ok()?;
let mut j = column_argmax_abs(&xt, n);
let mut iter = 2usize;
loop {
for i in 0..n {
x[(i, 0)] = T::zero();
}
x[(j, 0)] = T::one();
v = lu.solve(x.as_ref()).ok()?;
let est_old = est;
est = column_one_norm(&v, n);
let mut sign_changed = false;
for i in 0..n {
if sign_of(v[(i, 0)]) != isgn[i] {
sign_changed = true;
break;
}
}
if !sign_changed {
break;
}
if est <= est_old {
break;
}
for i in 0..n {
let s = sign_of(v[(i, 0)]);
isgn[i] = s;
x[(i, 0)] = s;
}
xt = lu.solve_transpose(x.as_ref()).ok()?;
let j_last = j;
j = column_argmax_abs(&xt, n);
if Scalar::abs(xt[(j_last, 0)]) >= Scalar::abs(xt[(j, 0)]) || iter >= ITMAX {
break;
}
iter += 1;
}
let mut alt = T::one();
let denom = T::from_f64((n - 1) as f64)?;
for i in 0..n {
let frac = T::from_f64(i as f64)? / denom;
x[(i, 0)] = alt * (T::one() + frac);
alt = -alt;
}
let vf = lu.solve(x.as_ref()).ok()?;
let three_n = T::from_f64((3 * n) as f64)?;
let temp = (T::from_f64(2.0)? * column_one_norm(&vf, n)) / three_n;
if temp > est {
est = temp;
}
Some(est)
}
#[cfg(test)]
mod tests {
use super::*;
use oxiblas_matrix::Mat;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_cond_diagonal() {
let a = Mat::from_rows(&[&[2.0f64, 0.0], &[0.0, 4.0]]);
let kappa = cond(a.as_ref()).unwrap();
assert!(approx_eq(kappa, 2.0, 1e-10));
}
#[test]
fn test_cond_identity() {
let eye = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[0.0, 1.0, 0.0], &[0.0, 0.0, 1.0]]);
let kappa = cond(eye.as_ref()).unwrap();
assert!(approx_eq(kappa, 1.0, 1e-10));
}
#[test]
fn test_cond_ill_conditioned() {
let a = Mat::from_rows(&[&[1.0f64, 1.0 / 2.0], &[1.0 / 2.0, 1.0 / 3.0]]);
let kappa = cond(a.as_ref()).unwrap();
assert!(kappa > 15.0);
assert!(kappa < 25.0);
}
#[test]
fn test_cond_1_identity() {
let eye = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.0]]);
let kappa = cond_1(eye.as_ref()).unwrap();
assert!(approx_eq(kappa, 1.0, 1e-10));
}
#[test]
fn test_cond_inf_identity() {
let eye = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.0]]);
let kappa = cond_inf(eye.as_ref()).unwrap();
assert!(approx_eq(kappa, 1.0, 1e-10));
}
#[test]
fn test_cond_1_diagonal() {
let a = Mat::from_rows(&[&[2.0f64, 0.0], &[0.0, 4.0]]);
let kappa = cond_1(a.as_ref()).unwrap();
assert!(approx_eq(kappa, 2.0, 1e-10));
}
#[test]
fn test_cond_singular() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[2.0, 4.0]]);
let kappa = cond(a.as_ref()).unwrap();
assert!(kappa > 1e10);
}
#[test]
fn test_rcond_identity() {
let eye = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.0]]);
let rc = rcond(eye.as_ref()).unwrap();
assert!(approx_eq(rc, 1.0, 1e-10));
}
#[test]
fn test_rcond_singular() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[2.0, 4.0]]);
let rc = cond_1(a.as_ref()).unwrap();
assert!(rc > 1e10);
let rc2 = rcond(a.as_ref()).unwrap();
assert!(rc2 < 1e-10);
}
#[test]
fn test_rcond_estimate() {
let a = Mat::from_rows(&[&[2.0f64, 1.0], &[1.0, 3.0]]);
let rc_est = rcond_estimate(a.as_ref()).unwrap();
let rc_exact = rcond(a.as_ref()).unwrap();
assert!(rc_est > 0.0);
assert!(rc_est / rc_exact < 10.0);
assert!(rc_exact / rc_est < 10.0);
}
#[test]
fn test_rcond_estimate_nonsymmetric() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[0.0, 1e-2, 5.0], &[4.0, 0.0, 1.0]]);
let rc_est = rcond_estimate(a.as_ref()).unwrap();
let rc_exact = rcond(a.as_ref()).unwrap();
assert!(rc_est > 0.0);
assert!(
rc_est >= rc_exact * (1.0 - 1e-9),
"rc_est = {rc_est} should not be below rc_exact = {rc_exact}"
);
assert!(
rc_est <= rc_exact * 3.0 + 1e-12,
"rc_est = {rc_est} over-estimates rcond vs exact = {rc_exact}"
);
}
#[test]
fn test_rcond_estimate_ill_conditioned_nonsymmetric() {
let a = Mat::from_rows(&[&[1.0f64, 1.0, 1.0], &[0.0, 1e-8, 1.0], &[0.0, 0.0, 1e-8]]);
let rc_est = rcond_estimate(a.as_ref()).unwrap();
let rc_exact = rcond(a.as_ref()).unwrap();
assert!(rc_est > 0.0);
assert!(
rc_est < 1e-6,
"near-singular matrix should have tiny rcond, got {rc_est}"
);
assert!(rc_est >= rc_exact * (1.0 - 1e-9));
assert!(rc_est <= rc_exact * 5.0 + 1e-12);
}
#[test]
fn test_cond_not_square() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let result = cond_1(a.as_ref());
assert!(matches!(result, Err(CondError::NotSquare)));
let kappa = cond(a.as_ref()).unwrap();
assert!(kappa > 0.0);
}
#[test]
fn test_cond_f32() {
let a = Mat::from_rows(&[&[2.0f32, 0.0], &[0.0, 4.0]]);
let kappa = cond(a.as_ref()).unwrap();
assert!((kappa - 2.0).abs() < 1e-4);
}
#[test]
fn test_cond_relationship() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let k2 = cond(a.as_ref()).unwrap();
let k1 = cond_1(a.as_ref()).unwrap();
assert!(k1 > 0.0);
assert!(k2 > 0.0);
}
}