use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
use crate::svd::{Svd, SvdError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RankError {
EmptyMatrix,
SvdFailed,
}
impl core::fmt::Display for RankError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::SvdFailed => write!(f, "SVD computation failed"),
}
}
}
impl std::error::Error for RankError {}
impl From<SvdError> for RankError {
fn from(e: SvdError) -> Self {
match e {
SvdError::EmptyMatrix => Self::EmptyMatrix,
SvdError::NotConverged => Self::SvdFailed,
}
}
}
pub fn rank<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
tol: Option<T>,
) -> Result<usize, RankError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Ok(0);
}
let svd = Svd::compute(a)?;
let sigma = svd.singular_values();
let tolerance = match tol {
Some(t) => t,
None => {
let eps = <T as Scalar>::epsilon();
let sigma_max = if sigma.is_empty() { T::one() } else { sigma[0] };
eps * T::from_f64(m.max(n) as f64).unwrap_or(T::one()) * sigma_max
}
};
Ok(svd.rank(tolerance))
}
pub fn nullity<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
tol: Option<T>,
) -> Result<usize, RankError> {
let n = a.ncols();
let r = rank(a, tol)?;
Ok(n - r)
}
pub fn null_space<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
tol: Option<T>,
) -> Result<Mat<T>, RankError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Ok(Mat::zeros(n, 0));
}
let svd = Svd::compute(a)?;
let sigma = svd.singular_values();
let vt = svd.vt();
let tolerance = match tol {
Some(t) => t,
None => {
let eps = <T as Scalar>::epsilon();
let sigma_max = if sigma.is_empty() { T::one() } else { sigma[0] };
eps * T::from_f64(m.max(n) as f64).unwrap_or(T::one()) * sigma_max
}
};
let r = svd.rank(tolerance);
let null_dim = n - r;
if null_dim == 0 {
return Ok(Mat::zeros(n, 0));
}
let mut null_basis = Mat::zeros(n, null_dim);
for i in 0..n {
for j in 0..null_dim {
null_basis[(i, j)] = vt[(r + j, i)];
}
}
Ok(null_basis)
}
pub fn col_space<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
tol: Option<T>,
) -> Result<Mat<T>, RankError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Ok(Mat::zeros(m, 0));
}
let svd = Svd::compute(a)?;
let sigma = svd.singular_values();
let u = svd.u();
let tolerance = match tol {
Some(t) => t,
None => {
let eps = <T as Scalar>::epsilon();
let sigma_max = if sigma.is_empty() { T::one() } else { sigma[0] };
eps * T::from_f64(m.max(n) as f64).unwrap_or(T::one()) * sigma_max
}
};
let r = svd.rank(tolerance);
if r == 0 {
return Ok(Mat::zeros(m, 0));
}
let mut col_basis = Mat::zeros(m, r);
for i in 0..m {
for j in 0..r {
col_basis[(i, j)] = u[(i, j)];
}
}
Ok(col_basis)
}
pub fn row_space<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
tol: Option<T>,
) -> Result<Mat<T>, RankError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Ok(Mat::zeros(n, 0));
}
let svd = Svd::compute(a)?;
let sigma = svd.singular_values();
let vt = svd.vt();
let tolerance = match tol {
Some(t) => t,
None => {
let eps = <T as Scalar>::epsilon();
let sigma_max = if sigma.is_empty() { T::one() } else { sigma[0] };
eps * T::from_f64(m.max(n) as f64).unwrap_or(T::one()) * sigma_max
}
};
let r = svd.rank(tolerance);
if r == 0 {
return Ok(Mat::zeros(n, 0));
}
let mut row_basis = Mat::zeros(n, r);
for i in 0..n {
for j in 0..r {
row_basis[(i, j)] = vt[(j, i)];
}
}
Ok(row_basis)
}
pub fn left_null_space<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
tol: Option<T>,
) -> Result<Mat<T>, RankError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Ok(Mat::zeros(m, 0));
}
let svd = Svd::compute(a)?;
let sigma = svd.singular_values();
let u = svd.u();
let tolerance = match tol {
Some(t) => t,
None => {
let eps = <T as Scalar>::epsilon();
let sigma_max = if sigma.is_empty() { T::one() } else { sigma[0] };
eps * T::from_f64(m.max(n) as f64).unwrap_or(T::one()) * sigma_max
}
};
let r = svd.rank(tolerance);
let left_null_dim = m - r;
if left_null_dim == 0 {
return Ok(Mat::zeros(m, 0));
}
let mut left_null_basis = Mat::zeros(m, left_null_dim);
for i in 0..m {
for j in 0..left_null_dim {
left_null_basis[(i, j)] = u[(i, r + j)];
}
}
Ok(left_null_basis)
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_rank_full() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
assert_eq!(rank(a.as_ref(), None).unwrap(), 2);
}
#[test]
fn test_rank_deficient() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[2.0, 4.0]]);
assert_eq!(rank(a.as_ref(), None).unwrap(), 1);
}
#[test]
fn test_rank_identity() {
let eye = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[0.0, 1.0, 0.0], &[0.0, 0.0, 1.0]]);
assert_eq!(rank(eye.as_ref(), None).unwrap(), 3);
}
#[test]
fn test_rank_tall() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0], &[5.0, 6.0]]);
assert_eq!(rank(a.as_ref(), None).unwrap(), 2);
}
#[test]
fn test_rank_wide() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
assert_eq!(rank(a.as_ref(), None).unwrap(), 2);
}
#[test]
fn test_nullity() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[2.0, 4.0]]);
assert_eq!(nullity(a.as_ref(), None).unwrap(), 1);
}
#[test]
fn test_nullity_full_rank() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
assert_eq!(nullity(a.as_ref(), None).unwrap(), 0);
}
#[test]
fn test_null_space_rank_1() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[2.0, 4.0]]);
let ns = null_space(a.as_ref(), None).unwrap();
assert_eq!(ns.ncols(), 1);
assert_eq!(ns.nrows(), 2);
let prod_0 = a[(0, 0)] * ns[(0, 0)] + a[(0, 1)] * ns[(1, 0)];
let prod_1 = a[(1, 0)] * ns[(0, 0)] + a[(1, 1)] * ns[(1, 0)];
assert!(approx_eq(prod_0, 0.0, 1e-10));
assert!(approx_eq(prod_1, 0.0, 1e-10));
}
#[test]
fn test_null_space_empty() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let ns = null_space(a.as_ref(), None).unwrap();
assert_eq!(ns.ncols(), 0);
}
#[test]
fn test_col_space_full_rank() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let cs = col_space(a.as_ref(), None).unwrap();
assert_eq!(cs.ncols(), 2);
assert_eq!(cs.nrows(), 2);
for i in 0..2 {
for j in 0..2 {
let mut dot = 0.0;
for k in 0..2 {
dot += cs[(k, i)] * cs[(k, j)];
}
let expected = if i == j { 1.0 } else { 0.0 };
assert!(approx_eq(dot, expected, 1e-10));
}
}
}
#[test]
fn test_col_space_rank_deficient() {
let a = Mat::from_rows(&[&[1.0f64, 0.0, 1.0], &[0.0, 1.0, 1.0], &[1.0, 1.0, 2.0]]);
let cs = col_space(a.as_ref(), None).unwrap();
assert_eq!(cs.ncols(), 2); assert_eq!(cs.nrows(), 3);
}
#[test]
fn test_row_space() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let rs = row_space(a.as_ref(), None).unwrap();
assert_eq!(rs.ncols(), 2);
assert_eq!(rs.nrows(), 2);
}
#[test]
fn test_left_null_space() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[2.0, 4.0]]);
let lns = left_null_space(a.as_ref(), None).unwrap();
assert_eq!(lns.ncols(), 1);
assert_eq!(lns.nrows(), 2);
let prod_0 = lns[(0, 0)] * a[(0, 0)] + lns[(1, 0)] * a[(1, 0)];
let prod_1 = lns[(0, 0)] * a[(0, 1)] + lns[(1, 0)] * a[(1, 1)];
assert!(approx_eq(prod_0, 0.0, 1e-10));
assert!(approx_eq(prod_1, 0.0, 1e-10));
}
#[test]
fn test_four_fundamental_subspaces() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let r = rank(a.as_ref(), None).unwrap();
let m = 3;
let n = 3;
let cs = col_space(a.as_ref(), None).unwrap();
let ns = null_space(a.as_ref(), None).unwrap();
let rs = row_space(a.as_ref(), None).unwrap();
let lns = left_null_space(a.as_ref(), None).unwrap();
assert_eq!(cs.ncols(), r);
assert_eq!(ns.ncols(), n - r);
assert_eq!(rs.ncols(), r);
assert_eq!(lns.ncols(), m - r);
}
#[test]
fn test_rank_f32() {
let a = Mat::from_rows(&[&[1.0f32, 2.0], &[3.0, 4.0]]);
assert_eq!(rank(a.as_ref(), None).unwrap(), 2);
}
}