use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::MatRef;
use crate::svd::{Svd, SvdError};
pub fn norm_1<T: Real>(a: MatRef<'_, T>) -> T {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return T::zero();
}
let mut max_col_sum = T::zero();
for j in 0..n {
let mut col_sum = T::zero();
for i in 0..m {
col_sum = col_sum + Scalar::abs(a[(i, j)]);
}
if col_sum > max_col_sum {
max_col_sum = col_sum;
}
}
max_col_sum
}
pub fn norm_inf<T: Real>(a: MatRef<'_, T>) -> T {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return T::zero();
}
let mut max_row_sum = T::zero();
for i in 0..m {
let mut row_sum = T::zero();
for j in 0..n {
row_sum = row_sum + Scalar::abs(a[(i, j)]);
}
if row_sum > max_row_sum {
max_row_sum = row_sum;
}
}
max_row_sum
}
pub fn norm_frobenius<T: Real>(a: MatRef<'_, T>) -> T {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return T::zero();
}
let mut sum_sq = T::zero();
for i in 0..m {
for j in 0..n {
let val = a[(i, j)];
sum_sq = sum_sq + val * val;
}
}
Real::sqrt(sum_sq)
}
pub fn norm_2<T: Field + Real + bytemuck::Zeroable>(a: MatRef<'_, T>) -> Result<T, SvdError> {
let svd = Svd::compute(a)?;
Ok(svd.norm2())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TraceError {
NotSquare {
nrows: usize,
ncols: usize,
},
}
impl core::fmt::Display for TraceError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::NotSquare { nrows, ncols } => {
write!(f, "Matrix is not square: {nrows}×{ncols}")
}
}
}
}
impl std::error::Error for TraceError {}
pub fn trace<T: Field>(a: MatRef<'_, T>) -> Result<T, TraceError> {
let n = a.nrows();
if n != a.ncols() {
return Err(TraceError::NotSquare {
nrows: n,
ncols: a.ncols(),
});
}
if n == 0 {
return Ok(T::zero());
}
let mut tr = T::zero();
for i in 0..n {
tr = tr + a[(i, i)];
}
Ok(tr)
}
pub fn norm_nuclear<T: Field + Real + bytemuck::Zeroable>(a: MatRef<'_, T>) -> Result<T, SvdError> {
let svd = Svd::compute(a)?;
let sigma = svd.singular_values();
let mut sum = T::zero();
for &s in sigma {
sum = sum + s;
}
Ok(sum)
}
pub fn norm_max<T: Real>(a: MatRef<'_, T>) -> T {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return T::zero();
}
let mut max_val = T::zero();
for i in 0..m {
for j in 0..n {
let abs_val = Scalar::abs(a[(i, j)]);
if abs_val > max_val {
max_val = abs_val;
}
}
}
max_val
}
#[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_norm_1() {
let a = Mat::from_rows(&[&[1.0f64, -2.0, 3.0], &[4.0, 5.0, -6.0]]);
let n1 = norm_1(a.as_ref());
assert!(approx_eq(n1, 9.0, 1e-10));
}
#[test]
fn test_norm_inf() {
let a = Mat::from_rows(&[&[1.0f64, -2.0, 3.0], &[4.0, 5.0, -6.0]]);
let ni = norm_inf(a.as_ref());
assert!(approx_eq(ni, 15.0, 1e-10));
}
#[test]
fn test_norm_frobenius() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let nf = norm_frobenius(a.as_ref());
assert!(approx_eq(nf, 30.0f64.sqrt(), 1e-10));
}
#[test]
fn test_norm_2() {
let a = Mat::from_rows(&[&[3.0f64, 0.0], &[0.0, 4.0]]);
let n2 = norm_2(a.as_ref()).unwrap();
assert!(approx_eq(n2, 4.0, 1e-10));
}
#[test]
fn test_norm_2_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 n2 = norm_2(eye.as_ref()).unwrap();
assert!(approx_eq(n2, 1.0, 1e-10));
}
#[test]
fn test_trace() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let t = trace(a.as_ref()).unwrap();
assert!(approx_eq(t, 15.0, 1e-10)); }
#[test]
fn test_trace_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 t = trace(eye.as_ref()).unwrap();
assert!(approx_eq(t, 3.0, 1e-10));
}
#[test]
fn test_trace_not_square() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let result = trace(a.as_ref());
assert!(matches!(
result,
Err(TraceError::NotSquare { nrows: 2, ncols: 3 })
));
}
#[test]
fn test_trace_empty() {
let a: Mat<f64> = Mat::zeros(0, 0);
let t = trace(a.as_ref()).unwrap();
assert!(approx_eq(t, 0.0, 1e-10));
}
#[test]
fn test_norm_nuclear() {
let a = Mat::from_rows(&[&[3.0f64, 0.0], &[0.0, 4.0]]);
let nn = norm_nuclear(a.as_ref()).unwrap();
assert!(approx_eq(nn, 7.0, 1e-10));
}
#[test]
fn test_norm_max() {
let a = Mat::from_rows(&[&[1.0f64, -5.0], &[3.0, 2.0]]);
let nm = norm_max(a.as_ref());
assert!(approx_eq(nm, 5.0, 1e-10));
}
#[test]
fn test_norm_empty() {
let a: Mat<f64> = Mat::zeros(0, 0);
assert!(approx_eq(norm_1(a.as_ref()), 0.0, 1e-10));
assert!(approx_eq(norm_inf(a.as_ref()), 0.0, 1e-10));
assert!(approx_eq(norm_frobenius(a.as_ref()), 0.0, 1e-10));
assert!(approx_eq(norm_max(a.as_ref()), 0.0, 1e-10));
}
#[test]
fn test_norm_relations() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let n2 = norm_2(a.as_ref()).unwrap();
let nf = norm_frobenius(a.as_ref());
assert!(n2 <= nf + 1e-10);
assert!(nf <= 2.0f64.sqrt() * n2 + 1e-10);
}
#[test]
fn test_norm_1_inf_relation() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let mut at = Mat::zeros(3, 2);
for i in 0..2 {
for j in 0..3 {
at[(j, i)] = a[(i, j)];
}
}
let n1_a = norm_1(a.as_ref());
let ninf_at = norm_inf(at.as_ref());
assert!(approx_eq(n1_a, ninf_at, 1e-10));
}
#[test]
fn test_norms_f32() {
let a = Mat::from_rows(&[&[1.0f32, 2.0], &[3.0, 4.0]]);
let n1 = norm_1(a.as_ref());
let ninf = norm_inf(a.as_ref());
let nf = norm_frobenius(a.as_ref());
assert!((n1 - 6.0).abs() < 1e-5); assert!((ninf - 7.0).abs() < 1e-5); assert!((nf - 30.0f32.sqrt()).abs() < 1e-5);
}
}