use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SelectiveSvdError {
EmptyMatrix,
InvalidRange,
NoSingularValuesInRange,
NotConverged,
InvalidIndexRange,
}
impl core::fmt::Display for SelectiveSvdError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::InvalidRange => write!(f, "Invalid value range specified"),
Self::NoSingularValuesInRange => {
write!(f, "No singular values found in the specified range")
}
Self::NotConverged => write!(f, "Algorithm did not converge"),
Self::InvalidIndexRange => write!(f, "Invalid index range"),
}
}
}
impl std::error::Error for SelectiveSvdError {}
#[derive(Debug, Clone)]
pub enum SingularValueSelector<T> {
All,
ValueRange {
low: T,
high: T,
},
IndexRange {
low: usize,
high: usize,
},
}
#[derive(Debug, Clone)]
pub struct SelectiveSvd<T: Scalar> {
sigma: Vec<T>,
u: Option<Mat<T>>,
vt: Option<Mat<T>>,
m: usize,
n: usize,
index_offset: usize,
}
const MAX_BISECTION_ITER: usize = 1000;
const MAX_INVERSE_ITER: usize = 100;
impl<T: Field + Real + bytemuck::Zeroable> SelectiveSvd<T> {
pub fn compute(
a: MatRef<'_, T>,
selector: SingularValueSelector<T>,
) -> Result<Self, SelectiveSvdError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(SelectiveSvdError::EmptyMatrix);
}
if m < n {
let mut at = Mat::zeros(n, m);
for i in 0..m {
for j in 0..n {
at[(j, i)] = a[(i, j)];
}
}
let svd_t = Self::compute_tall(at.as_ref(), selector)?;
let u = svd_t.vt.map(|vt| {
let mut u = Mat::zeros(m, vt.ncols());
for i in 0..vt.nrows().min(m) {
for j in 0..vt.ncols() {
u[(i, j)] = vt[(i, j)];
}
}
u
});
let vt = svd_t.u.map(|u| {
let mut vt = Mat::zeros(u.ncols(), n);
for i in 0..u.ncols() {
for j in 0..u.nrows().min(n) {
vt[(i, j)] = u[(j, i)];
}
}
vt
});
return Ok(Self {
sigma: svd_t.sigma,
u,
vt,
m,
n,
index_offset: svd_t.index_offset,
});
}
Self::compute_tall(a, selector)
}
pub fn singular_values_only(
a: MatRef<'_, T>,
selector: SingularValueSelector<T>,
) -> Result<Self, SelectiveSvdError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(SelectiveSvdError::EmptyMatrix);
}
if m < n {
let mut at = Mat::zeros(n, m);
for i in 0..m {
for j in 0..n {
at[(j, i)] = a[(i, j)];
}
}
let mut svd_t = Self::compute_tall_values_only(at.as_ref(), selector)?;
svd_t.m = m;
svd_t.n = n;
return Ok(svd_t);
}
Self::compute_tall_values_only(a, selector)
}
fn compute_tall(
a: MatRef<'_, T>,
selector: SingularValueSelector<T>,
) -> Result<Self, SelectiveSvdError> {
let m = a.nrows();
let n = a.ncols();
let k = m.min(n);
let (u_b, d, e, vt_b) = bidiagonalize(a)?;
let (t_diag, t_offdiag) = form_btb_tridiagonal(&d, &e);
let (index_low, _index_high, sigma) = match selector {
SingularValueSelector::All => {
let sigma = compute_all_singular_values(&t_diag, &t_offdiag)?;
(0, sigma.len().saturating_sub(1), sigma)
}
SingularValueSelector::ValueRange { low, high } => {
if low < T::zero() || high < low {
return Err(SelectiveSvdError::InvalidRange);
}
let eig_low = low * low;
let eig_high = high * high;
let (sigma, idx_low) =
compute_singular_values_in_range(&t_diag, &t_offdiag, eig_low, eig_high)?;
if sigma.is_empty() {
return Err(SelectiveSvdError::NoSingularValuesInRange);
}
(idx_low, idx_low + sigma.len() - 1, sigma)
}
SingularValueSelector::IndexRange { low, high } => {
if low > high || high >= k {
return Err(SelectiveSvdError::InvalidIndexRange);
}
let sigma = compute_singular_values_by_index(&t_diag, &t_offdiag, low, high)?;
(low, high, sigma)
}
};
let num_sv = sigma.len();
let v_bidiag = compute_right_singular_vectors(&t_diag, &t_offdiag, &sigma)?;
let u_bidiag = compute_left_singular_vectors(&d, &e, &sigma, &v_bidiag);
let mut u = Mat::zeros(m, num_sv);
for j in 0..num_sv {
for i in 0..m {
let mut sum = T::zero();
for l in 0..k {
sum = sum + u_b[(i, l)] * u_bidiag[(l, j)];
}
u[(i, j)] = sum;
}
}
let mut vt = Mat::zeros(num_sv, n);
for i in 0..num_sv {
for j in 0..n {
let mut sum = T::zero();
for l in 0..k {
sum = sum + v_bidiag[(l, i)] * vt_b[(l, j)];
}
vt[(i, j)] = sum;
}
}
Ok(Self {
sigma,
u: Some(u),
vt: Some(vt),
m,
n,
index_offset: index_low,
})
}
fn compute_tall_values_only(
a: MatRef<'_, T>,
selector: SingularValueSelector<T>,
) -> Result<Self, SelectiveSvdError> {
let m = a.nrows();
let n = a.ncols();
let k = m.min(n);
let (_u_b, d, e, _vt_b) = bidiagonalize(a)?;
let (t_diag, t_offdiag) = form_btb_tridiagonal(&d, &e);
let (index_offset, sigma) = match selector {
SingularValueSelector::All => {
let sigma = compute_all_singular_values(&t_diag, &t_offdiag)?;
(0, sigma)
}
SingularValueSelector::ValueRange { low, high } => {
if low < T::zero() || high < low {
return Err(SelectiveSvdError::InvalidRange);
}
let eig_low = low * low;
let eig_high = high * high;
let (sigma, idx) =
compute_singular_values_in_range(&t_diag, &t_offdiag, eig_low, eig_high)?;
if sigma.is_empty() {
return Err(SelectiveSvdError::NoSingularValuesInRange);
}
(idx, sigma)
}
SingularValueSelector::IndexRange { low, high } => {
if low > high || high >= k {
return Err(SelectiveSvdError::InvalidIndexRange);
}
let sigma = compute_singular_values_by_index(&t_diag, &t_offdiag, low, high)?;
(low, sigma)
}
};
Ok(Self {
sigma,
u: None,
vt: None,
m,
n,
index_offset,
})
}
pub fn singular_values(&self) -> &[T] {
&self.sigma
}
pub fn u(&self) -> Option<MatRef<'_, T>> {
self.u.as_ref().map(|u| u.as_ref())
}
pub fn vt(&self) -> Option<MatRef<'_, T>> {
self.vt.as_ref().map(|vt| vt.as_ref())
}
pub fn index_offset(&self) -> usize {
self.index_offset
}
pub fn count(&self) -> usize {
self.sigma.len()
}
pub fn reconstruct(&self) -> Option<Mat<T>> {
let u = self.u.as_ref()?;
let vt = self.vt.as_ref()?;
let mut a = 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.sigma.len() {
sum = sum + u[(i, l)] * self.sigma[l] * vt[(l, j)];
}
a[(i, j)] = sum;
}
}
Some(a)
}
}
fn bidiagonalize<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
) -> Result<(Mat<T>, Vec<T>, Vec<T>, Mat<T>), SelectiveSvdError> {
let m = a.nrows();
let n = a.ncols();
let k = m.min(n);
let mut work = Mat::zeros(m, n);
for i in 0..m {
for j in 0..n {
work[(i, j)] = a[(i, j)];
}
}
let mut tau_left = vec![T::zero(); k];
let num_right = k.saturating_sub(1);
let mut tau_right = vec![T::zero(); num_right];
let mut d = vec![T::zero(); k];
let mut e = vec![T::zero(); num_right];
for j in 0..k {
let (tau, beta) = householder_left(&mut work, j, m);
d[j] = beta;
tau_left[j] = tau;
apply_householder_left(&mut work, j, m, n, tau);
if j < n - 1 {
let (tau, beta) = householder_right(&mut work, j, n);
if j < e.len() {
e[j] = beta;
tau_right[j] = tau;
}
apply_householder_right(&mut work, j, m, n, tau);
}
}
let mut u = Mat::zeros(m, m);
for i in 0..m {
u[(i, i)] = T::one();
}
for j in 0..k {
let tau = tau_left[j];
if tau != T::zero() {
for r in 0..m {
let mut w = u[(r, j)];
for i in (j + 1)..m {
w = w + u[(r, i)] * work[(i, j)];
}
let tw = tau * w;
u[(r, j)] = u[(r, j)] - tw;
for i in (j + 1)..m {
u[(r, i)] = u[(r, i)] - tw * work[(i, j)];
}
}
}
}
let mut v = Mat::zeros(n, n);
for i in 0..n {
v[(i, i)] = T::one();
}
for j in 0..tau_right.len() {
let tau = tau_right[j];
if tau != T::zero() {
let start = j + 1;
for r in 0..n {
let mut w = v[(r, start)];
for i in (start + 1)..n {
w = w + v[(r, i)] * work[(j, i)];
}
let tw = tau * w;
v[(r, start)] = v[(r, start)] - tw;
for i in (start + 1)..n {
v[(r, i)] = v[(r, i)] - tw * work[(j, i)];
}
}
}
}
let mut vt = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
vt[(i, j)] = v[(j, i)];
}
}
Ok((u, d, e, vt))
}
fn householder_left<T: Field + Real>(work: &mut Mat<T>, j: usize, m: usize) -> (T, T) {
let mut norm_sq = T::zero();
for i in j..m {
norm_sq = norm_sq + work[(i, j)] * work[(i, j)];
}
let norm = Real::sqrt(norm_sq);
if norm == T::zero() {
return (T::zero(), T::zero());
}
let x_j = work[(j, j)];
let beta = if x_j >= T::zero() { -norm } else { norm };
let tau = (beta - x_j) / beta;
let scale = T::one() / (x_j - beta);
for i in (j + 1)..m {
work[(i, j)] = work[(i, j)] * scale;
}
(tau, beta)
}
fn householder_right<T: Field + Real>(work: &mut Mat<T>, j: usize, n: usize) -> (T, T) {
let start = j + 1;
let mut norm_sq = T::zero();
for i in start..n {
norm_sq = norm_sq + work[(j, i)] * work[(j, i)];
}
let norm = Real::sqrt(norm_sq);
if norm == T::zero() {
return (T::zero(), T::zero());
}
let x_j = work[(j, start)];
let beta = if x_j >= T::zero() { -norm } else { norm };
let tau = (beta - x_j) / beta;
let scale = T::one() / (x_j - beta);
for i in (start + 1)..n {
work[(j, i)] = work[(j, i)] * scale;
}
(tau, beta)
}
fn apply_householder_left<T: Field + Real>(
work: &mut Mat<T>,
j: usize,
m: usize,
n: usize,
tau: T,
) {
if tau == T::zero() {
return;
}
for col in (j + 1)..n {
let mut w = work[(j, col)];
for i in (j + 1)..m {
w = w + work[(i, j)] * work[(i, col)];
}
let tw = tau * w;
work[(j, col)] = work[(j, col)] - tw;
for i in (j + 1)..m {
work[(i, col)] = work[(i, col)] - tw * work[(i, j)];
}
}
}
fn apply_householder_right<T: Field + Real>(
work: &mut Mat<T>,
j: usize,
m: usize,
n: usize,
tau: T,
) {
if tau == T::zero() {
return;
}
let start = j + 1;
for row in (j + 1)..m {
let mut w = work[(row, start)];
for i in (start + 1)..n {
w = w + work[(j, i)] * work[(row, i)];
}
let tw = tau * w;
work[(row, start)] = work[(row, start)] - tw;
for i in (start + 1)..n {
work[(row, i)] = work[(row, i)] - tw * work[(j, i)];
}
}
}
fn form_btb_tridiagonal<T: Field + Real>(d: &[T], e: &[T]) -> (Vec<T>, Vec<T>) {
let n = d.len();
let mut t_diag = vec![T::zero(); n];
let mut t_offdiag = vec![T::zero(); n.saturating_sub(1)];
for i in 0..n {
t_diag[i] = d[i] * d[i];
if i > 0 && i - 1 < e.len() {
t_diag[i] = t_diag[i] + e[i - 1] * e[i - 1];
}
}
for i in 0..t_offdiag.len() {
if i < e.len() {
t_offdiag[i] = d[i] * e[i];
}
}
(t_diag, t_offdiag)
}
fn compute_all_singular_values<T: Field + Real>(
t_diag: &[T],
t_offdiag: &[T],
) -> Result<Vec<T>, SelectiveSvdError> {
let n = t_diag.len();
if n == 0 {
return Ok(Vec::new());
}
let (low, high) = gershgorin_bounds(t_diag, t_offdiag);
let eigenvalues = bisection_eigenvalues(t_diag, t_offdiag, 0, n - 1, low, high)?;
let mut sigma: Vec<T> = eigenvalues
.iter()
.map(|&eig| {
if eig > T::zero() {
Real::sqrt(eig)
} else {
T::zero()
}
})
.collect();
sigma.sort_by(|a, b| {
if *b > *a {
core::cmp::Ordering::Greater
} else if *b < *a {
core::cmp::Ordering::Less
} else {
core::cmp::Ordering::Equal
}
});
Ok(sigma)
}
fn compute_singular_values_in_range<T: Field + Real>(
t_diag: &[T],
t_offdiag: &[T],
eig_low: T,
eig_high: T,
) -> Result<(Vec<T>, usize), SelectiveSvdError> {
let n = t_diag.len();
if n == 0 {
return Ok((Vec::new(), 0));
}
let count_below_low = sturm_count(t_diag, t_offdiag, eig_low);
let count_below_high = sturm_count(t_diag, t_offdiag, eig_high);
if count_below_high <= count_below_low {
return Ok((Vec::new(), 0));
}
let index_low = count_below_low;
let index_high = count_below_high - 1;
let (global_low, global_high) = gershgorin_bounds(t_diag, t_offdiag);
let search_low = if eig_low < global_low {
global_low
} else {
eig_low
};
let search_high = if eig_high > global_high {
global_high
} else {
eig_high
};
let eigenvalues = bisection_eigenvalues(
t_diag,
t_offdiag,
index_low,
index_high,
search_low,
search_high,
)?;
let mut sigma: Vec<T> = eigenvalues
.iter()
.map(|&eig| {
if eig > T::zero() {
Real::sqrt(eig)
} else {
T::zero()
}
})
.collect();
sigma.sort_by(|a, b| {
if *b > *a {
core::cmp::Ordering::Greater
} else if *b < *a {
core::cmp::Ordering::Less
} else {
core::cmp::Ordering::Equal
}
});
let all_sigma = compute_all_singular_values(t_diag, t_offdiag)?;
let mut idx_offset = 0;
if !sigma.is_empty() && !all_sigma.is_empty() {
let max_computed = sigma[0];
for s in &all_sigma {
if *s > max_computed + <T as Scalar>::epsilon() * T::from_f64(100.0).unwrap_or(T::one())
{
idx_offset += 1;
} else {
break;
}
}
}
Ok((sigma, idx_offset))
}
fn compute_singular_values_by_index<T: Field + Real>(
t_diag: &[T],
t_offdiag: &[T],
index_low: usize,
index_high: usize,
) -> Result<Vec<T>, SelectiveSvdError> {
let n = t_diag.len();
if n == 0 || index_high >= n {
return Err(SelectiveSvdError::InvalidIndexRange);
}
let all_sigma = compute_all_singular_values(t_diag, t_offdiag)?;
let sigma: Vec<T> = all_sigma
.iter()
.skip(index_low)
.take(index_high - index_low + 1)
.copied()
.collect();
Ok(sigma)
}
fn gershgorin_bounds<T: Field + Real>(diag: &[T], offdiag: &[T]) -> (T, T) {
let n = diag.len();
if n == 0 {
return (T::zero(), T::zero());
}
let mut low = diag[0];
let mut high = diag[0];
if !offdiag.is_empty() {
let r = Scalar::abs(offdiag[0]);
low = if diag[0] - r < low { diag[0] - r } else { low };
high = if diag[0] + r > high {
diag[0] + r
} else {
high
};
}
for i in 1..n - 1 {
let r = Scalar::abs(offdiag[i - 1]) + Scalar::abs(offdiag[i]);
let center = diag[i];
if center - r < low {
low = center - r;
}
if center + r > high {
high = center + r;
}
}
if n > 1 {
let r = Scalar::abs(offdiag[n - 2]);
let center = diag[n - 1];
if center - r < low {
low = center - r;
}
if center + r > high {
high = center + r;
}
}
if low < T::zero() {
low = T::zero();
}
let margin = (high - low) * T::from_f64(0.01).unwrap_or(T::one());
(low - margin, high + margin)
}
fn sturm_count<T: Field + Real>(diag: &[T], offdiag: &[T], x: T) -> usize {
let n = diag.len();
if n == 0 {
return 0;
}
let eps = <T as Scalar>::epsilon() * T::from_f64(100.0).unwrap_or(T::one());
let mut count = 0;
let mut d_prev = diag[0] - x;
if d_prev <= T::zero() {
count += 1;
}
for i in 1..n {
let e_sq = offdiag[i - 1] * offdiag[i - 1];
let d_curr = if Scalar::abs(d_prev) < eps {
diag[i] - x - e_sq / (if d_prev >= T::zero() { eps } else { -eps })
} else {
diag[i] - x - e_sq / d_prev
};
if d_curr <= T::zero() {
count += 1;
}
d_prev = d_curr;
}
count
}
fn bisection_eigenvalues<T: Field + Real>(
diag: &[T],
offdiag: &[T],
index_low: usize,
index_high: usize,
value_low: T,
value_high: T,
) -> Result<Vec<T>, SelectiveSvdError> {
let num_eigenvalues = index_high - index_low + 1;
let mut eigenvalues = vec![T::zero(); num_eigenvalues];
let eps = <T as Scalar>::epsilon();
let tol = eps
* (Scalar::abs(value_low) + Scalar::abs(value_high))
* T::from_f64(100.0).unwrap_or(T::one());
for k in 0..num_eigenvalues {
let target_index = index_low + k;
let mut lo = value_low;
let mut hi = value_high;
for _ in 0..MAX_BISECTION_ITER {
let mid = (lo + hi) / T::from_f64(2.0).unwrap_or_else(T::zero);
let count = sturm_count(diag, offdiag, mid);
if count <= target_index {
lo = mid;
} else {
hi = mid;
}
if hi - lo < tol {
break;
}
}
eigenvalues[k] = (lo + hi) / T::from_f64(2.0).unwrap_or_else(T::zero);
}
Ok(eigenvalues)
}
fn compute_right_singular_vectors<T: Field + Real + bytemuck::Zeroable>(
t_diag: &[T],
t_offdiag: &[T],
sigma: &[T],
) -> Result<Mat<T>, SelectiveSvdError> {
let n = t_diag.len();
let k = sigma.len();
if n == 0 || k == 0 {
return Ok(Mat::zeros(0, 0));
}
let mut v = Mat::zeros(n, k);
let eps = <T as Scalar>::epsilon();
let tol = eps * T::from_f64(100.0).unwrap_or(T::one());
for (col, &s) in sigma.iter().enumerate() {
let lambda = s * s;
let mut x = vec![T::one(); n];
let norm_init = Real::sqrt(T::from_usize(n).unwrap_or(T::one()));
for xi in &mut x {
*xi = *xi / norm_init;
}
for _ in 0..MAX_INVERSE_ITER {
match solve_tridiagonal_shifted(t_diag, t_offdiag, &x, lambda) {
Ok(x_new) => {
let mut norm_sq = T::zero();
for &xi in &x_new {
norm_sq = norm_sq + xi * xi;
}
let norm = Real::sqrt(norm_sq);
if norm < tol {
break;
}
let mut diff = T::zero();
for i in 0..n {
let new_val = x_new[i] / norm;
diff = diff + (new_val - x[i]) * (new_val - x[i]);
x[i] = new_val;
}
if Real::sqrt(diff) < tol {
break;
}
}
Err(_) => break,
}
}
for prev in 0..col {
let mut dot = T::zero();
for i in 0..n {
dot = dot + x[i] * v[(i, prev)];
}
for i in 0..n {
x[i] = x[i] - dot * v[(i, prev)];
}
}
let mut norm_sq = T::zero();
for &xi in &x {
norm_sq = norm_sq + xi * xi;
}
let norm = Real::sqrt(norm_sq);
if norm > tol {
for i in 0..n {
v[(i, col)] = x[i] / norm;
}
} else {
v[(col.min(n - 1), col)] = T::one();
}
}
Ok(v)
}
fn compute_left_singular_vectors<T: Field + Real + bytemuck::Zeroable>(
d: &[T],
e: &[T],
sigma: &[T],
v: &Mat<T>,
) -> Mat<T> {
let n = d.len();
let k = sigma.len();
let eps = <T as Scalar>::epsilon() * T::from_f64(100.0).unwrap_or(T::one());
let mut u = Mat::zeros(n, k);
for col in 0..k {
let s = sigma[col];
if s < eps {
u[(col.min(n - 1), col)] = T::one();
continue;
}
for i in 0..n {
let mut sum = d[i] * v[(i, col)];
if i < e.len() {
sum = sum + e[i] * v[(i + 1, col)];
}
u[(i, col)] = sum / s;
}
let mut norm_sq = T::zero();
for i in 0..n {
norm_sq = norm_sq + u[(i, col)] * u[(i, col)];
}
let norm = Real::sqrt(norm_sq);
if norm > eps {
for i in 0..n {
u[(i, col)] = u[(i, col)] / norm;
}
}
}
u
}
fn solve_tridiagonal_shifted<T: Field + Real>(
diag: &[T],
offdiag: &[T],
b: &[T],
lambda: T,
) -> Result<Vec<T>, SelectiveSvdError> {
let n = diag.len();
if n == 0 {
return Ok(Vec::new());
}
let eps = <T as Scalar>::epsilon() * T::from_f64(1000.0).unwrap_or(T::one());
let mut c_prime = vec![T::zero(); n];
let mut d_prime = vec![T::zero(); n];
let diag_shifted = diag[0] - lambda;
if Scalar::abs(diag_shifted) < eps {
let reg = if diag_shifted >= T::zero() { eps } else { -eps };
c_prime[0] = if !offdiag.is_empty() {
offdiag[0] / reg
} else {
T::zero()
};
d_prime[0] = b[0] / reg;
} else {
c_prime[0] = if !offdiag.is_empty() {
offdiag[0] / diag_shifted
} else {
T::zero()
};
d_prime[0] = b[0] / diag_shifted;
}
for i in 1..n {
let a_i = offdiag[i - 1];
let diag_shifted = diag[i] - lambda;
let denom = diag_shifted - a_i * c_prime[i - 1];
if Scalar::abs(denom) < eps {
let reg = if denom >= T::zero() { eps } else { -eps };
c_prime[i] = if i < offdiag.len() {
offdiag[i] / reg
} else {
T::zero()
};
d_prime[i] = (b[i] - a_i * d_prime[i - 1]) / reg;
} else {
c_prime[i] = if i < offdiag.len() {
offdiag[i] / denom
} else {
T::zero()
};
d_prime[i] = (b[i] - a_i * d_prime[i - 1]) / denom;
}
}
let mut x = vec![T::zero(); n];
x[n - 1] = d_prime[n - 1];
for i in (0..n - 1).rev() {
x[i] = d_prime[i] - c_prime[i] * x[i + 1];
}
Ok(x)
}
pub fn count_singular_values_above<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
threshold: T,
) -> Result<usize, SelectiveSvdError> {
if a.nrows() == 0 || a.ncols() == 0 {
return Err(SelectiveSvdError::EmptyMatrix);
}
if threshold < T::zero() {
return Err(SelectiveSvdError::InvalidRange);
}
let (_, d, e, _) = bidiagonalize(a)?;
let (t_diag, t_offdiag) = form_btb_tridiagonal(&d, &e);
let threshold_sq = threshold * threshold;
let (_, _high) = gershgorin_bounds(&t_diag, &t_offdiag);
let count_below = sturm_count(&t_diag, &t_offdiag, threshold_sq);
let total = t_diag.len();
Ok(total - count_below)
}
pub fn singular_value_bounds<T: Field + Real + bytemuck::Zeroable>(
a: MatRef<'_, T>,
) -> Result<(T, T), SelectiveSvdError> {
if a.nrows() == 0 || a.ncols() == 0 {
return Err(SelectiveSvdError::EmptyMatrix);
}
let (_, d, e, _) = bidiagonalize(a)?;
let (t_diag, t_offdiag) = form_btb_tridiagonal(&d, &e);
let (low, high) = gershgorin_bounds(&t_diag, &t_offdiag);
let sv_low = if low > T::zero() {
Real::sqrt(low)
} else {
T::zero()
};
let sv_high = if high > T::zero() {
Real::sqrt(high)
} else {
T::zero()
};
Ok((sv_low, sv_high))
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_selective_svd_all() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
assert_eq!(svd.singular_values().len(), 2);
let sigma = svd.singular_values();
assert!(sigma[0] >= sigma[1]);
}
#[test]
fn test_selective_svd_index_range() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let svd = SelectiveSvd::compute(
a.as_ref(),
SingularValueSelector::IndexRange { low: 0, high: 0 },
)
.unwrap();
assert_eq!(svd.count(), 1);
assert_eq!(svd.index_offset(), 0);
}
#[test]
fn test_selective_svd_index_range_middle() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
let svd = SelectiveSvd::compute(
a.as_ref(),
SingularValueSelector::IndexRange { low: 1, high: 1 },
)
.unwrap();
assert_eq!(svd.count(), 1);
}
#[test]
fn test_selective_svd_values_only() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0], &[5.0, 6.0]]);
let svd =
SelectiveSvd::singular_values_only(a.as_ref(), SingularValueSelector::All).unwrap();
assert_eq!(svd.count(), 2);
assert!(svd.u().is_none());
assert!(svd.vt().is_none());
}
#[test]
fn test_selective_svd_diagonal() {
let a = Mat::from_rows(&[&[3.0f64, 0.0], &[0.0, 5.0]]);
let svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
let sigma = svd.singular_values();
assert!(approx_eq(sigma[0], 5.0, 1e-8));
assert!(approx_eq(sigma[1], 3.0, 1e-8));
}
#[test]
fn test_selective_svd_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 svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
for &s in svd.singular_values() {
assert!(approx_eq(s, 1.0, 1e-8));
}
}
#[test]
fn test_selective_svd_tall() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0], &[5.0, 6.0], &[7.0, 8.0]]);
let svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
assert_eq!(svd.count(), 2);
let u = svd.u().unwrap();
let vt = svd.vt().unwrap();
assert_eq!(u.nrows(), 4);
assert_eq!(u.ncols(), 2);
assert_eq!(vt.nrows(), 2);
assert_eq!(vt.ncols(), 2);
}
#[test]
fn test_selective_svd_wide() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0, 4.0], &[5.0, 6.0, 7.0, 8.0]]);
let svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
assert_eq!(svd.count(), 2);
}
#[test]
fn test_selective_svd_reconstruction() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
let reconstructed = svd.reconstruct().unwrap();
for i in 0..2 {
for j in 0..2 {
assert!(
approx_eq(reconstructed[(i, j)], a[(i, j)], 1e-6),
"reconstructed[{},{}] = {}, a = {}",
i,
j,
reconstructed[(i, j)],
a[(i, j)]
);
}
}
}
#[test]
fn test_selective_svd_1x1() {
let a = Mat::from_rows(&[&[5.0f64]]);
let svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
assert_eq!(svd.count(), 1);
assert!(approx_eq(svd.singular_values()[0], 5.0, 1e-8));
}
#[test]
fn test_count_singular_values() {
let a = Mat::from_rows(&[&[3.0f64, 0.0], &[0.0, 5.0]]);
let count = count_singular_values_above(a.as_ref(), 4.0).unwrap();
assert_eq!(count, 1);
let count = count_singular_values_above(a.as_ref(), 2.0).unwrap();
assert_eq!(count, 2);
let count = count_singular_values_above(a.as_ref(), 6.0).unwrap();
assert_eq!(count, 0); }
#[test]
fn test_singular_value_bounds() {
let a = Mat::from_rows(&[&[3.0f64, 0.0], &[0.0, 5.0]]);
let (low, high) = singular_value_bounds(a.as_ref()).unwrap();
assert!(low <= 3.0 + 0.5); assert!(high >= 5.0 - 0.5);
}
#[test]
fn test_selective_svd_invalid_index() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let result = SelectiveSvd::compute(
a.as_ref(),
SingularValueSelector::IndexRange { low: 0, high: 5 },
);
assert!(matches!(result, Err(SelectiveSvdError::InvalidIndexRange)));
}
#[test]
fn test_selective_svd_orthogonal_vectors() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 10.0]]);
let svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
let u = svd.u().unwrap();
let vt = svd.vt().unwrap();
for i in 0..u.ncols() {
for j in 0..u.ncols() {
let mut dot = 0.0;
for k in 0..u.nrows() {
dot += u[(k, i)] * u[(k, j)];
}
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
approx_eq(dot, expected, 1e-6),
"U^T*U[{},{}] = {}, expected {}",
i,
j,
dot,
expected
);
}
}
for i in 0..vt.nrows() {
for j in 0..vt.nrows() {
let mut dot = 0.0;
for k in 0..vt.ncols() {
dot += vt[(i, k)] * vt[(j, k)];
}
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
approx_eq(dot, expected, 1e-6),
"V*V^T[{},{}] = {}, expected {}",
i,
j,
dot,
expected
);
}
}
}
#[test]
fn test_form_btb_tridiagonal() {
let d = vec![2.0f64, 3.0, 1.0];
let e = vec![1.0f64, 2.0];
let (t_diag, t_offdiag) = form_btb_tridiagonal(&d, &e);
assert!(approx_eq(t_diag[0], 4.0, 1e-10)); assert!(approx_eq(t_diag[1], 10.0, 1e-10)); assert!(approx_eq(t_diag[2], 5.0, 1e-10));
assert!(approx_eq(t_offdiag[0], 2.0, 1e-10)); assert!(approx_eq(t_offdiag[1], 6.0, 1e-10)); }
#[test]
fn test_selective_svd_f32() {
let a = Mat::from_rows(&[&[1.0f32, 2.0], &[3.0, 4.0]]);
let svd = SelectiveSvd::compute(a.as_ref(), SingularValueSelector::All).unwrap();
assert_eq!(svd.count(), 2);
}
}