use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
use super::selective::{SelectiveSvd, SelectiveSvdError, SingularValueSelector};
use super::{SvdDc, SvdDcError};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TruncatedSvdError {
EmptyMatrix,
ZeroRank,
RankTooLarge {
requested: usize,
max_rank: usize,
},
NotConverged,
InternalError,
}
impl core::fmt::Display for TruncatedSvdError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::ZeroRank => write!(f, "Requested rank is zero"),
Self::RankTooLarge {
requested,
max_rank,
} => {
write!(
f,
"Requested rank {} exceeds maximum {}",
requested, max_rank
)
}
Self::NotConverged => write!(f, "Algorithm did not converge"),
Self::InternalError => write!(f, "Internal computation failed"),
}
}
}
impl std::error::Error for TruncatedSvdError {}
impl From<SelectiveSvdError> for TruncatedSvdError {
fn from(e: SelectiveSvdError) -> Self {
match e {
SelectiveSvdError::EmptyMatrix => Self::EmptyMatrix,
SelectiveSvdError::InvalidRange => Self::InternalError,
SelectiveSvdError::NoSingularValuesInRange => Self::InternalError,
SelectiveSvdError::NotConverged => Self::NotConverged,
SelectiveSvdError::InvalidIndexRange => Self::RankTooLarge {
requested: 0,
max_rank: 0,
},
}
}
}
impl From<SvdDcError> for TruncatedSvdError {
fn from(e: SvdDcError) -> Self {
match e {
SvdDcError::EmptyMatrix => Self::EmptyMatrix,
SvdDcError::NotConverged => Self::NotConverged,
SvdDcError::SecularEquationFailed => Self::NotConverged,
}
}
}
#[derive(Debug, Clone)]
pub struct TruncatedSvd<T: Scalar> {
u: Mat<T>,
sigma: Vec<T>,
vt: Mat<T>,
m: usize,
n: usize,
k: usize,
}
impl<T: Field + Real + bytemuck::Zeroable> TruncatedSvd<T> {
pub fn compute(a: MatRef<'_, T>, k: usize) -> Result<Self, TruncatedSvdError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(TruncatedSvdError::EmptyMatrix);
}
if k == 0 {
return Err(TruncatedSvdError::ZeroRank);
}
let max_rank = m.min(n);
if k > max_rank {
return Err(TruncatedSvdError::RankTooLarge {
requested: k,
max_rank,
});
}
let use_selective = k <= max_rank / 2;
if use_selective {
let selector = SingularValueSelector::IndexRange {
low: 0,
high: k - 1,
};
let svd = SelectiveSvd::compute(a, selector)?;
let u = svd.u().map_or_else(
|| Mat::zeros(m, k),
|u_ref| {
let mut u = Mat::zeros(m, k);
for i in 0..m {
for j in 0..k.min(u_ref.ncols()) {
u[(i, j)] = u_ref[(i, j)];
}
}
u
},
);
let vt = svd.vt().map_or_else(
|| Mat::zeros(k, n),
|vt_ref| {
let mut vt = Mat::zeros(k, n);
for i in 0..k.min(vt_ref.nrows()) {
for j in 0..n {
vt[(i, j)] = vt_ref[(i, j)];
}
}
vt
},
);
let sigma = svd.singular_values().to_vec();
Ok(Self {
u,
sigma,
vt,
m,
n,
k,
})
} else {
let svd = SvdDc::compute(a)?;
let u_full = svd.u();
let vt_full = svd.vt();
let sigma_full = svd.singular_values();
let mut u = Mat::zeros(m, k);
for i in 0..m {
for j in 0..k {
u[(i, j)] = u_full[(i, j)];
}
}
let mut vt = Mat::zeros(k, n);
for i in 0..k {
for j in 0..n {
vt[(i, j)] = vt_full[(i, j)];
}
}
let sigma = sigma_full[..k].to_vec();
Ok(Self {
u,
sigma,
vt,
m,
n,
k,
})
}
}
pub fn singular_values_only(a: MatRef<'_, T>, k: usize) -> Result<Vec<T>, TruncatedSvdError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(TruncatedSvdError::EmptyMatrix);
}
if k == 0 {
return Err(TruncatedSvdError::ZeroRank);
}
let max_rank = m.min(n);
if k > max_rank {
return Err(TruncatedSvdError::RankTooLarge {
requested: k,
max_rank,
});
}
let selector = SingularValueSelector::IndexRange {
low: 0,
high: k - 1,
};
let svd = SelectiveSvd::singular_values_only(a, selector)?;
Ok(svd.singular_values().to_vec())
}
pub fn u(&self) -> MatRef<'_, T> {
self.u.as_ref()
}
pub fn u_matrix(&self) -> Mat<T> {
self.u.clone()
}
pub fn singular_values(&self) -> &[T] {
&self.sigma
}
pub fn vt(&self) -> MatRef<'_, T> {
self.vt.as_ref()
}
pub fn vt_matrix(&self) -> Mat<T> {
self.vt.clone()
}
pub fn v(&self) -> Mat<T> {
let mut v = Mat::zeros(self.n, self.k);
for i in 0..self.n {
for j in 0..self.k {
v[(i, j)] = self.vt[(j, i)];
}
}
v
}
pub fn rank(&self) -> usize {
self.k
}
pub fn dimensions(&self) -> (usize, usize) {
(self.m, self.n)
}
pub fn reconstruct(&self) -> Mat<T> {
let mut result = Mat::zeros(self.m, self.n);
for i in 0..self.m {
for j in 0..self.n {
let mut sum = T::zero();
for l in 0..self.k {
sum = sum + self.u[(i, l)] * self.sigma[l] * self.vt[(l, j)];
}
result[(i, j)] = sum;
}
}
result
}
pub fn reconstruction_error(&self, a: MatRef<'_, T>) -> T {
let approx = self.reconstruct();
let mut error_sq = T::zero();
for i in 0..self.m {
for j in 0..self.n {
let diff = a[(i, j)] - approx[(i, j)];
error_sq = error_sq + diff * diff;
}
}
Real::sqrt(error_sq)
}
pub fn relative_error(&self, a: MatRef<'_, T>) -> T {
let error = self.reconstruction_error(a);
let mut a_norm_sq = T::zero();
for i in 0..self.m {
for j in 0..self.n {
a_norm_sq = a_norm_sq + a[(i, j)] * a[(i, j)];
}
}
let a_norm = Real::sqrt(a_norm_sq);
if a_norm > T::zero() {
error / a_norm
} else {
T::zero()
}
}
pub fn nuclear_norm(&self) -> T {
self.sigma.iter().copied().fold(T::zero(), |acc, s| acc + s)
}
pub fn energy_ratio(&self, a: MatRef<'_, T>) -> T {
let mut a_frob_sq = T::zero();
for i in 0..self.m {
for j in 0..self.n {
a_frob_sq = a_frob_sq + a[(i, j)] * a[(i, j)];
}
}
if a_frob_sq <= T::zero() {
return T::one();
}
let truncated_energy: T = self
.sigma
.iter()
.map(|&s| s * s)
.fold(T::zero(), |a, b| a + b);
truncated_energy / a_frob_sq
}
pub fn explained_variance_ratio(&self, a: MatRef<'_, T>) -> Vec<T> {
let mut a_frob_sq = T::zero();
for i in 0..self.m {
for j in 0..self.n {
a_frob_sq = a_frob_sq + a[(i, j)] * a[(i, j)];
}
}
if a_frob_sq <= T::zero() {
return vec![T::zero(); self.k];
}
self.sigma.iter().map(|&s| (s * s) / a_frob_sq).collect()
}
pub fn project(&self, x: &[T]) -> Vec<T> {
assert_eq!(x.len(), self.n);
let mut result = vec![T::zero(); self.k];
for i in 0..self.k {
for j in 0..self.n {
result[i] = result[i] + self.vt[(i, j)] * x[j];
}
}
result
}
pub fn inverse_project(&self, y: &[T]) -> Vec<T> {
assert_eq!(y.len(), self.k);
let mut result = vec![T::zero(); self.n];
for i in 0..self.n {
for j in 0..self.k {
result[i] = result[i] + self.vt[(j, i)] * y[j];
}
}
result
}
}
pub fn thin_svd<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
) -> Result<TruncatedSvd<T>, TruncatedSvdError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(TruncatedSvdError::EmptyMatrix);
}
let k = m.min(n);
let svd = SvdDc::compute(a)?;
let u_full = svd.u();
let vt_full = svd.vt();
let sigma_full = svd.singular_values();
let mut u = Mat::zeros(m, k);
for i in 0..m {
for j in 0..k {
u[(i, j)] = u_full[(i, j)];
}
}
let mut vt = Mat::zeros(k, n);
for i in 0..k {
for j in 0..n {
vt[(i, j)] = vt_full[(i, j)];
}
}
let sigma = sigma_full[..k].to_vec();
Ok(TruncatedSvd {
u,
sigma,
vt,
m,
n,
k,
})
}
pub fn rank_k_approximation<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
k: usize,
) -> Result<Mat<T>, TruncatedSvdError> {
let tsvd = TruncatedSvd::compute(a, k)?;
Ok(tsvd.reconstruct())
}
pub fn optimal_rank_for_energy<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
target_energy: T,
) -> Result<usize, TruncatedSvdError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(TruncatedSvdError::EmptyMatrix);
}
let mut total_energy = T::zero();
for i in 0..m {
for j in 0..n {
total_energy = total_energy + a[(i, j)] * a[(i, j)];
}
}
if total_energy <= T::zero() {
return Ok(0);
}
let thin = thin_svd(a)?;
let sigma = thin.singular_values();
let mut accumulated = T::zero();
let threshold = target_energy * total_energy;
for (k, &s) in sigma.iter().enumerate() {
accumulated = accumulated + s * s;
if accumulated >= threshold {
return Ok(k + 1);
}
}
Ok(sigma.len())
}
pub fn numerical_rank<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
tol: Option<T>,
) -> Result<usize, TruncatedSvdError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(TruncatedSvdError::EmptyMatrix);
}
let thin = thin_svd(a)?;
let sigma = thin.singular_values();
if sigma.is_empty() || sigma[0] <= T::zero() {
return Ok(0);
}
let eps = <T as Scalar>::epsilon();
let default_tol = eps * T::from_usize(m.max(n)).unwrap_or(T::one());
let threshold = tol.unwrap_or(default_tol) * sigma[0];
let rank = sigma.iter().take_while(|&&s| s > threshold).count();
Ok(rank)
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_truncated_svd_basic() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
assert_eq!(tsvd.rank(), 2);
assert_eq!(tsvd.singular_values().len(), 2);
assert!(tsvd.singular_values()[0] >= tsvd.singular_values()[1]);
}
#[test]
fn test_truncated_svd_rank1() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 1).unwrap();
assert_eq!(tsvd.rank(), 1);
assert_eq!(tsvd.u().nrows(), 2);
assert_eq!(tsvd.u().ncols(), 1);
assert_eq!(tsvd.vt().nrows(), 1);
assert_eq!(tsvd.vt().ncols(), 3);
}
#[test]
fn test_truncated_svd_dimensions() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0, 4.0, 5.0],
&[6.0, 7.0, 8.0, 9.0, 10.0],
&[11.0, 12.0, 13.0, 14.0, 15.0],
]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
assert_eq!(tsvd.dimensions(), (3, 5));
assert_eq!(tsvd.u().nrows(), 3);
assert_eq!(tsvd.u().ncols(), 2);
assert_eq!(tsvd.vt().nrows(), 2);
assert_eq!(tsvd.vt().ncols(), 5);
}
#[test]
fn test_truncated_svd_reconstruction() {
let a = Mat::from_rows(&[&[4.0f64, 2.0], &[1.0, 3.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
let approx = tsvd.reconstruct();
for i in 0..2 {
for j in 0..2 {
assert!(
approx_eq(approx[(i, j)], a[(i, j)], 0.1),
"approx[{},{}] = {}, expected {}",
i,
j,
approx[(i, j)],
a[(i, j)]
);
}
}
}
#[test]
fn test_thin_svd() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0, 4.0, 5.0],
&[6.0, 7.0, 8.0, 9.0, 10.0],
&[11.0, 12.0, 13.0, 14.0, 15.0],
]);
let thin = thin_svd(a.as_ref()).unwrap();
assert_eq!(thin.rank(), 3); assert_eq!(thin.singular_values().len(), 3);
}
#[test]
fn test_thin_svd_tall() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0],
&[3.0, 4.0],
&[5.0, 6.0],
&[7.0, 8.0],
&[9.0, 10.0],
]);
let thin = thin_svd(a.as_ref()).unwrap();
assert_eq!(thin.rank(), 2); }
#[test]
fn test_rank_k_approximation() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let approx = rank_k_approximation(a.as_ref(), 1).unwrap();
assert_eq!(approx.nrows(), 2);
assert_eq!(approx.ncols(), 3);
}
#[test]
fn test_singular_values_only() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let sigma = TruncatedSvd::singular_values_only(a.as_ref(), 2).unwrap();
assert_eq!(sigma.len(), 2);
assert!(sigma[0] >= sigma[1]);
}
#[test]
fn test_truncated_svd_error_empty() {
let a: Mat<f64> = Mat::zeros(0, 3);
let result = TruncatedSvd::compute(a.as_ref(), 1);
assert!(matches!(result, Err(TruncatedSvdError::EmptyMatrix)));
}
#[test]
fn test_truncated_svd_error_zero_rank() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let result = TruncatedSvd::compute(a.as_ref(), 0);
assert!(matches!(result, Err(TruncatedSvdError::ZeroRank)));
}
#[test]
fn test_truncated_svd_error_rank_too_large() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let result = TruncatedSvd::compute(a.as_ref(), 5);
assert!(matches!(
result,
Err(TruncatedSvdError::RankTooLarge { .. })
));
}
#[test]
fn test_truncated_svd_orthogonality() {
let a = Mat::from_rows(&[
&[1.0f64, 2.0, 3.0, 4.0],
&[5.0, 6.0, 7.0, 8.0],
&[9.0, 10.0, 11.0, 12.0],
]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
let u = tsvd.u();
let vt = tsvd.vt();
for i in 0..2 {
for j in 0..2 {
let mut sum = 0.0;
for k in 0..3 {
sum += u[(k, i)] * u[(k, j)];
}
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
approx_eq(sum, expected, 0.1),
"U^T*U[{},{}] = {}, expected {}",
i,
j,
sum,
expected
);
}
}
for i in 0..2 {
for j in 0..2 {
let mut sum = 0.0;
for k in 0..4 {
sum += vt[(i, k)] * vt[(j, k)];
}
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
approx_eq(sum, expected, 0.1),
"V*V^T[{},{}] = {}, expected {}",
i,
j,
sum,
expected
);
}
}
}
#[test]
fn test_energy_ratio() {
let a = Mat::from_rows(&[&[3.0f64, 0.0, 0.0], &[0.0, 4.0, 0.0], &[0.0, 0.0, 5.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 1).unwrap();
let energy = tsvd.energy_ratio(a.as_ref());
assert!(
approx_eq(energy, 0.5, 0.1),
"Energy ratio {} should be ~0.5",
energy
);
}
#[test]
fn test_explained_variance_ratio() {
let a = Mat::from_rows(&[&[3.0f64, 0.0], &[0.0, 4.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
let ratios = tsvd.explained_variance_ratio(a.as_ref());
assert!(
approx_eq(ratios[0], 0.64, 0.1),
"First ratio {} should be ~0.64",
ratios[0]
);
assert!(
approx_eq(ratios[1], 0.36, 0.1),
"Second ratio {} should be ~0.36",
ratios[1]
);
}
#[test]
fn test_project_and_inverse_project() {
let a = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[0.0, 2.0, 0.0], &[0.0, 0.0, 3.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
let x = vec![1.0, 1.0, 1.0];
let projected = tsvd.project(&x);
assert_eq!(projected.len(), 2);
let reconstructed = tsvd.inverse_project(&projected);
assert_eq!(reconstructed.len(), 3);
}
#[test]
fn test_optimal_rank_for_energy() {
let a = Mat::from_rows(&[
&[10.0f64, 0.0, 0.0, 0.0],
&[0.0, 5.0, 0.0, 0.0],
&[0.0, 0.0, 2.0, 0.0],
&[0.0, 0.0, 0.0, 1.0],
]);
let k = optimal_rank_for_energy(a.as_ref(), 0.9).unwrap();
assert!(k <= 3, "Should need at most 3 components for 90% energy");
let k = optimal_rank_for_energy(a.as_ref(), 0.99).unwrap();
assert!(k >= 2, "Should need at least 2 components for 99% energy");
}
#[test]
fn test_numerical_rank() {
let a = Mat::from_rows(&[&[1.0f64, 0.0, 0.0], &[0.0, 2.0, 0.0], &[0.0, 0.0, 3.0]]);
let rank = numerical_rank(a.as_ref(), None).unwrap();
assert_eq!(rank, 3);
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[2.0, 4.0, 6.0], &[3.0, 6.0, 9.0]]);
let rank = numerical_rank(a.as_ref(), Some(1e-8)).unwrap();
assert!(rank <= 2, "Rank should be at most 2 for rank-1 matrix");
}
#[test]
fn test_v_matrix() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
let v = tsvd.v();
assert_eq!(v.nrows(), 3);
assert_eq!(v.ncols(), 2);
let vt = tsvd.vt();
for i in 0..3 {
for j in 0..2 {
assert!(
approx_eq(v[(i, j)], vt[(j, i)], 1e-10),
"V[{},{}] = {} should equal Vt[{},{}] = {}",
i,
j,
v[(i, j)],
j,
i,
vt[(j, i)]
);
}
}
}
#[test]
fn test_truncated_svd_f32() {
let a = Mat::from_rows(&[&[1.0f32, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
assert_eq!(tsvd.rank(), 2);
}
#[test]
fn test_truncated_svd_relative_error() {
let a = Mat::from_rows(&[&[4.0f64, 1.0], &[1.0, 3.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
let error = tsvd.relative_error(a.as_ref());
assert!(
error < 0.1,
"Full reconstruction error {} should be small",
error
);
}
#[test]
fn test_nuclear_norm() {
let a = Mat::from_rows(&[&[2.0f64, 0.0], &[0.0, 3.0]]);
let tsvd = TruncatedSvd::compute(a.as_ref(), 2).unwrap();
let nuclear = tsvd.nuclear_norm();
assert!(
approx_eq(nuclear, 5.0, 0.5),
"Nuclear norm {} should be ~5",
nuclear
);
}
}