use crate::dense::{small_svd, svd_flip, tsqr};
use crate::error::{Result, SvdLibError};
use crate::matrix::SparseMatDense;
use crate::types::{Algorithm, Detail, Diagnostics, SvdFloat, SvdRec};
use ndarray::{s, Array1, Array2, Axis};
use rand::rngs::StdRng;
use rand::{rng, Rng, SeedableRng};
use rand_distr::{Distribution, Normal};
pub const DEFAULT_OVERSAMPLES: usize = 10;
pub const DEFAULT_POWER_ITERATIONS: usize = 2;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Normalizer {
#[default]
Tsqr,
ColumnNorm,
None,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Sketch {
PowerIteration { iterations: usize },
BlockKrylov { blocks: usize },
}
impl Default for Sketch {
fn default() -> Self {
Sketch::PowerIteration {
iterations: DEFAULT_POWER_ITERATIONS,
}
}
}
#[derive(Debug, Clone)]
pub struct RandomizedConfig {
pub rank: usize,
pub oversamples: usize,
pub sketch: Sketch,
pub normalizer: Normalizer,
pub mean_center: bool,
pub seed: Option<u64>,
}
impl RandomizedConfig {
pub fn new(rank: usize) -> Self {
Self {
rank,
oversamples: DEFAULT_OVERSAMPLES,
sketch: Sketch::default(),
normalizer: Normalizer::default(),
mean_center: false,
seed: None,
}
}
pub fn oversamples(mut self, n: usize) -> Self {
self.oversamples = n;
self
}
pub fn power_iterations(mut self, n: usize) -> Self {
self.sketch = Sketch::PowerIteration { iterations: n };
self
}
pub fn block_krylov(mut self, blocks: usize) -> Self {
self.sketch = Sketch::BlockKrylov { blocks };
self
}
pub fn normalizer(mut self, n: Normalizer) -> Self {
self.normalizer = n;
self
}
pub fn mean_center(mut self, yes: bool) -> Self {
self.mean_center = yes;
self
}
pub fn seed(mut self, seed: u64) -> Self {
self.seed = Some(seed);
self
}
}
pub type Progress<'a> = &'a (dyn Fn(&str) + Sync);
pub fn svd<T: SvdFloat, M: SparseMatDense<T>>(a: &M, rank: usize) -> Result<SvdRec<T>> {
svd_with(a, &RandomizedConfig::new(rank), None)
}
pub fn svd_seed<T: SvdFloat, M: SparseMatDense<T>>(
a: &M,
rank: usize,
seed: u64,
) -> Result<SvdRec<T>> {
svd_with(a, &RandomizedConfig::new(rank).seed(seed), None)
}
pub fn svd_block_krylov<T: SvdFloat, M: SparseMatDense<T>>(
a: &M,
rank: usize,
blocks: usize,
seed: Option<u64>,
) -> Result<SvdRec<T>> {
let mut cfg = RandomizedConfig::new(rank).block_krylov(blocks);
cfg.seed = seed;
svd_with(a, &cfg, None)
}
pub fn svd_centered<T: SvdFloat, M: SparseMatDense<T>>(
a: &M,
rank: usize,
seed: Option<u64>,
) -> Result<SvdRec<T>> {
let mut cfg = RandomizedConfig::new(rank).mean_center(true);
cfg.seed = seed;
svd_with(a, &cfg, None)
}
pub fn svd_with<T: SvdFloat, M: SparseMatDense<T>>(
a: &M,
cfg: &RandomizedConfig,
progress: Option<Progress<'_>>,
) -> Result<SvdRec<T>> {
let note = |msg: &str| {
if let Some(p) = progress {
p(msg);
}
};
let (rows, cols) = (a.rows(), a.cols());
let min_dim = rows.min(cols);
if cfg.rank == 0 {
return Err(SvdLibError::invalid("randomized: rank must be at least 1"));
}
if cfg.rank > min_dim {
return Err(SvdLibError::invalid(format!(
"randomized: rank {} exceeds min(rows, cols) = {min_dim}",
cfg.rank
)));
}
if let Sketch::BlockKrylov { blocks } = cfg.sketch {
if blocks == 0 {
return Err(SvdLibError::invalid(
"randomized: block_krylov needs at least one block",
));
}
}
let rank = cfg.rank;
let l = (rank + cfg.oversamples).min(min_dim);
let seed = cfg.seed.unwrap_or_else(|| rng().next_u64());
let mut rng_state = StdRng::seed_from_u64(seed);
let means: Option<Array1<T>> = if cfg.mean_center {
note("computing column means");
Some(a.col_means())
} else {
None
};
let mut matvecs = 0usize;
let mul = |rhs: &Array2<T>, out: &mut Array2<T>, trans: bool| match &means {
Some(m) => a.mul_dense_centered(rhs.view(), out.view_mut(), trans, m.view()),
None => a.mul_dense(rhs.view(), out.view_mut(), trans),
};
note("drawing the random sketch");
let omega = gaussian(cols, l, &mut rng_state);
let mut basis = match cfg.sketch {
Sketch::PowerIteration { iterations } => {
note("projecting");
let mut y = Array2::<T>::zeros((rows, l));
mul(&omega, &mut y, false);
matvecs += l;
normalize(&mut y, cfg.normalizer)?;
let mut z = Array2::<T>::zeros((cols, l));
for i in 0..iterations {
note(&format!("power iteration {}/{}", i + 1, iterations));
mul(&y, &mut z, true);
matvecs += l;
normalize(&mut z, cfg.normalizer)?;
mul(&z, &mut y, false);
matvecs += l;
normalize(&mut y, cfg.normalizer)?;
}
y
}
Sketch::BlockKrylov { blocks } => {
note("building the Krylov block basis");
let width = (blocks * l).min(min_dim);
let mut k = Array2::<T>::zeros((rows, width));
let mut y = Array2::<T>::zeros((rows, l));
let mut z = Array2::<T>::zeros((cols, l));
mul(&omega, &mut y, false);
matvecs += l;
normalize(&mut y, cfg.normalizer)?;
let mut filled = 0usize;
for b in 0..blocks {
if filled >= width {
break;
}
let take = l.min(width - filled);
k.slice_mut(s![.., filled..filled + take])
.assign(&y.slice(s![.., ..take]));
filled += take;
if b + 1 == blocks {
break;
}
note(&format!("krylov block {}/{}", b + 2, blocks));
mul(&y, &mut z, true);
matvecs += l;
normalize(&mut z, cfg.normalizer)?;
mul(&z, &mut y, false);
matvecs += l;
normalize(&mut y, cfg.normalizer)?;
}
if filled < width {
k = k.slice(s![.., ..filled]).to_owned();
}
k
}
};
note("orthonormalising the basis");
tsqr(&mut basis)?;
let width = basis.ncols();
note("projecting onto the basis");
let mut bt = Array2::<T>::zeros((cols, width));
mul(&basis, &mut bt, true);
matvecs += width;
note("reducing");
let r_c = tsqr(&mut bt)?; let small = small_svd(r_c.t())?;
let keep = rank.min(small.s.len());
let u_hat = small.u.slice(s![.., ..keep]);
let v_hat = small.vt.slice(s![..keep, ..]).t().to_owned();
let mut u = basis.dot(&u_hat);
let mut vt = bt
.dot(&v_hat)
.reversed_axes()
.as_standard_layout()
.to_owned();
let s = small.s.slice(s![..keep]).to_owned();
svd_flip(&mut u, &mut vt);
let (oversamples, power_iterations, block_size) = match cfg.sketch {
Sketch::PowerIteration { iterations } => (cfg.oversamples, iterations, l),
Sketch::BlockKrylov { blocks } => (cfg.oversamples, blocks, l),
};
Ok(SvdRec {
d: keep,
u,
s,
vt,
total_squared_norm: T::from_f64_val(crate::matrix::total_squared_norm(
a,
means.as_ref().map(|m| m.view()),
)),
diagnostics: Diagnostics {
algorithm: match cfg.sketch {
Sketch::PowerIteration { .. } => Algorithm::Randomized,
Sketch::BlockKrylov { .. } => Algorithm::BlockKrylov,
},
non_zero: a.nnz(),
dimensions: rank,
significant_values: keep,
transposed: false,
random_seed: seed,
matvecs,
detail: Detail::Randomized {
oversamples,
power_iterations,
block_size,
},
},
})
}
fn gaussian<T: SvdFloat>(rows: usize, cols: usize, rng: &mut StdRng) -> Array2<T> {
let normal = Normal::new(0.0, 1.0).expect("N(0,1) is well-formed");
Array2::from_shape_fn((rows, cols), |_| T::from_f64_val(normal.sample(rng)))
}
fn normalize<T: SvdFloat>(m: &mut Array2<T>, how: Normalizer) -> Result<()> {
match how {
Normalizer::Tsqr => {
tsqr(m)?;
Ok(())
}
Normalizer::ColumnNorm => {
let floor = T::from_f64_val(1e-10);
for mut col in m.axis_iter_mut(Axis(1)) {
let n = col.iter().map(|&x| x * x).sum::<T>().sqrt();
if n > floor {
let inv = T::one() / n;
col.map_inplace(|x| *x *= inv);
}
}
Ok(())
}
Normalizer::None => Ok(()),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::matrix::SvdMat;
use crate::testing::{dense_of, gen_lowrank, gen_sparse, reference_singular_values, Lcg};
use sprs::TriMatI;
fn diagonal(n: usize) -> SvdMat<f64> {
let mut t = TriMatI::<f64, u32>::new((n, n));
for i in 0..n {
t.add_triplet(i, i, (n - i) as f64);
}
t.to_csr::<u64>()
}
fn max_rel_error(a: &SvdMat<f64>, got: &SvdRec<f64>) -> f64 {
let want = reference_singular_values(&dense_of(a));
got.s
.iter()
.enumerate()
.map(|(i, &g)| (g - want[i]).abs() / want[i].abs().max(1e-30))
.fold(0.0f64, f64::max)
}
#[test]
fn accurate_on_decaying_spectrum() {
let a = gen_lowrank(400, 120, 10, 5);
let got = svd_seed(&a, 10, 42).unwrap();
let err = max_rel_error(&a, &got);
assert!(err < 1e-6, "max relative error {err:.3e}");
}
#[test]
fn power_iteration_converges_slowly_on_linear_decay() {
let a = diagonal(60);
let loose = svd_with(
&a,
&RandomizedConfig::new(10).seed(42).power_iterations(0),
None,
)
.unwrap();
let tight = svd_with(
&a,
&RandomizedConfig::new(10).seed(42).power_iterations(7),
None,
)
.unwrap();
let e0 = max_rel_error(&a, &loose);
let e7 = max_rel_error(&a, &tight);
assert!(
e0 > 1e-2,
"q=0 should be visibly inaccurate here, got {e0:.3e}"
);
assert!(
e7 < e0 / 50.0,
"7 power iterations should improve substantially: {e0:.3e} -> {e7:.3e}"
);
}
#[test]
fn block_krylov_is_near_exact_on_linear_decay() {
let a = diagonal(60);
let got = svd_with(
&a,
&RandomizedConfig::new(10).seed(42).block_krylov(4),
None,
)
.unwrap();
let err = max_rel_error(&a, &got);
assert!(err < 1e-10, "block krylov max relative error {err:.3e}");
}
#[test]
fn power_iterations_improve_accuracy() {
let a = gen_sparse(600, 200, 0.05, 13);
let mut prev = f64::INFINITY;
for q in [0usize, 1, 2, 4, 7] {
let cfg = RandomizedConfig::new(15).seed(42).power_iterations(q);
let got = svd_with(&a, &cfg, None).unwrap();
let err = max_rel_error(&a, &got);
assert!(
err <= prev * 1.5 + 1e-9,
"q={q} error {err:.3e} is worse than q's predecessor {prev:.3e}"
);
prev = err;
}
let plain = svd_with(
&a,
&RandomizedConfig::new(15).seed(42).power_iterations(0),
None,
)
.unwrap();
let e0 = max_rel_error(&a, &plain);
assert!(
prev < e0 / 10.0,
"7 power iterations ({prev:.3e}) should be well under the q=0 error ({e0:.3e})"
);
}
#[test]
fn block_krylov_beats_power_iteration_on_flat_spectrum() {
let a = gen_sparse(800, 200, 0.04, 29);
let rank = 20;
let power = svd_with(
&a,
&RandomizedConfig::new(rank).seed(42).power_iterations(3),
None,
)
.unwrap();
let krylov = svd_with(
&a,
&RandomizedConfig::new(rank).seed(42).block_krylov(4),
None,
)
.unwrap();
let e_power = max_rel_error(&a, &power);
let e_krylov = max_rel_error(&a, &krylov);
assert!(
e_krylov <= e_power,
"block krylov {e_krylov:.3e} did not improve on power iteration {e_power:.3e}"
);
assert_eq!(krylov.diagnostics.algorithm, Algorithm::BlockKrylov);
}
#[test]
fn orientation_is_correct_for_wide_and_tall() {
for (r, c) in [(400usize, 80usize), (80, 400)] {
let a = gen_sparse(r, c, 0.08, 11);
let got = svd_seed(&a, 10, 42).unwrap();
assert_eq!(got.u.dim(), (r, 10), "u shape for {r}x{c}");
assert_eq!(got.vt.dim(), (10, c), "vt shape for {r}x{c}");
}
}
#[test]
fn singular_vectors_are_orthonormal() {
let a = gen_lowrank(300, 100, 12, 71);
let got = svd_seed(&a, 12, 42).unwrap();
let ou = crate::dense::orthogonality_error(&got.u.view());
assert!(ou < 1e-8, "||UᵀU - I|| = {ou:.3e}");
let vt_t = got.vt.t().to_owned();
let ov = crate::dense::orthogonality_error(&vt_t.view());
assert!(ov < 1e-8, "||VᵀV - I|| = {ov:.3e}");
}
#[test]
fn works_on_csr_and_csc_without_panicking() {
let a = gen_sparse(300, 120, 0.05, 3);
let csc = a.to_other_storage();
let x = svd_seed(&a, 10, 42).unwrap();
let y = svd_seed(&csc, 10, 42).unwrap();
for (p, q) in x.s.iter().zip(y.s.iter()) {
approx::assert_relative_eq!(p, q, max_relative = 1e-9);
}
}
#[test]
fn works_on_masked_matrices() {
let a = gen_sparse(300, 60, 0.1, 19);
let cols: Vec<usize> = (0..60).filter(|c| c % 2 == 0).collect();
let masked = crate::matrix::MaskedCsMat::with_columns(&a, &cols);
let got = svd_seed(&masked, 8, 42).unwrap();
assert_eq!(got.u.nrows(), 300);
assert_eq!(got.vt.ncols(), 30);
for w in got.s.to_vec().windows(2) {
assert!(w[0] >= w[1]);
}
}
#[test]
fn mean_centering_matches_dense_pca() {
let a = gen_lowrank(300, 60, 8, 37);
let dense = dense_of(&a);
let means = dense.mean_axis(Axis(0)).unwrap();
let centered = &dense - &means.view().insert_axis(Axis(0));
let want = reference_singular_values(¢ered);
let cfg = RandomizedConfig::new(8)
.seed(42)
.mean_center(true)
.power_iterations(5);
let got = svd_with(&a, &cfg, None).unwrap();
for (i, &g) in got.s.iter().enumerate() {
let rel = (g - want[i]).abs() / want[i].abs().max(1e-30);
assert!(
rel < 1e-5,
"centered singular value {i}: {g:.9e} vs {:.9e} (rel {rel:.3e})",
want[i]
);
}
}
#[test]
fn unseeded_runs_differ() {
let a = gen_sparse(300, 100, 0.05, 47);
let cfg = RandomizedConfig::new(6).power_iterations(0);
let x = svd_with(&a, &cfg, None).unwrap();
let y = svd_with(&a, &cfg, None).unwrap();
assert_ne!(
x.diagnostics.random_seed, y.diagnostics.random_seed,
"an unseeded config produced the same seed twice"
);
assert_ne!(x.u, y.u, "unseeded runs produced identical bases");
}
#[test]
fn seeded_runs_are_reproducible() {
let a = gen_sparse(300, 100, 0.05, 51);
let x = svd_seed(&a, 8, 999).unwrap();
let y = svd_seed(&a, 8, 999).unwrap();
assert_eq!(x.s, y.s);
assert_eq!(x.u, y.u);
assert_eq!(x.vt, y.vt);
}
#[test]
fn agrees_with_irlba() {
let a = gen_lowrank(400, 150, 12, 61);
let rand = svd_with(
&a,
&RandomizedConfig::new(12).seed(42).power_iterations(6),
None,
)
.unwrap();
let exact = crate::irlba::svd_seed(&a, 12, 42).unwrap();
for i in 0..12 {
let rel = (rand.s[i] - exact.s[i]).abs() / exact.s[i];
assert!(rel < 1e-6, "triplet {i}: randomized vs irlba rel {rel:.3e}");
}
}
#[test]
fn normalizers_all_produce_usable_results() {
let a = gen_lowrank(400, 100, 10, 67);
for n in [Normalizer::Tsqr, Normalizer::ColumnNorm, Normalizer::None] {
let cfg = RandomizedConfig::new(10)
.seed(42)
.power_iterations(1)
.normalizer(n);
let got = svd_with(&a, &cfg, None).unwrap();
let err = max_rel_error(&a, &got);
assert!(err < 1e-2, "{n:?} gave max relative error {err:.3e}");
}
}
#[test]
fn progress_callback_is_invoked() {
let a = gen_sparse(200, 80, 0.1, 73);
let seen = std::sync::Mutex::new(Vec::<String>::new());
let sink = |msg: &str| seen.lock().unwrap().push(msg.to_string());
let cfg = RandomizedConfig::new(6).seed(42).power_iterations(2);
svd_with(&a, &cfg, Some(&sink)).unwrap();
let msgs = seen.into_inner().unwrap();
assert!(!msgs.is_empty(), "no progress reported");
assert!(
msgs.iter().any(|m| m.contains("power iteration")),
"power iterations were not reported: {msgs:?}"
);
}
#[test]
fn block_krylov_clamps_basis_to_matrix_rank() {
let a = gen_sparse(500, 60, 0.08, 7);
let got = svd_block_krylov(&a, 12, 4, Some(42)).expect("should clamp, not fail");
assert_eq!(got.d, 12);
assert_eq!(got.u.dim(), (500, 12));
assert_eq!(got.vt.dim(), (12, 60));
let b = gen_sparse(60, 500, 0.08, 11);
let got = svd_block_krylov(&b, 12, 4, Some(42)).expect("should clamp, not fail");
assert_eq!(got.u.dim(), (60, 12));
assert_eq!(got.vt.dim(), (12, 500));
}
#[test]
fn rejects_bad_configuration() {
let a = gen_sparse(50, 30, 0.2, 1);
assert!(matches!(svd(&a, 0), Err(SvdLibError::InvalidArgument(_))));
assert!(matches!(svd(&a, 31), Err(SvdLibError::InvalidArgument(_))));
let cfg = RandomizedConfig::new(5).block_krylov(0);
assert!(matches!(
svd_with(&a, &cfg, None),
Err(SvdLibError::InvalidArgument(_))
));
}
#[test]
fn f32_works() {
let a64 = gen_lowrank(300, 80, 8, 79);
let want = reference_singular_values(&dense_of(&a64));
let mut t = TriMatI::<f32, u32>::new((300, 80));
for (v, (i, j)) in a64.iter() {
t.add_triplet(i as usize, j as usize, *v as f32);
}
let a32: SvdMat<f32> = t.to_csr::<u64>();
let cfg = RandomizedConfig::new(8).seed(42).power_iterations(4);
let got = svd_with(&a32, &cfg, None).unwrap();
for (i, &g) in got.s.iter().enumerate() {
let rel = ((g as f64) - want[i]).abs() / want[i].abs().max(1e-30);
assert!(rel < 1e-3, "f32 singular value {i}: rel {rel:.3e}");
}
}
#[test]
fn oversampling_clamps_to_matrix_rank() {
let a = gen_sparse(40, 20, 0.3, 83);
let cfg = RandomizedConfig::new(5).seed(42).oversamples(1000);
let got = svd_with(&a, &cfg, None).unwrap();
assert_eq!(got.d, 5);
let mut rng = Lcg::new(1);
let _ = rng.next_u64();
}
#[test]
fn diagnostics_report_matvecs_and_algorithm() {
let a = gen_sparse(200, 80, 0.1, 89);
let got = svd_seed(&a, 6, 42).unwrap();
assert_eq!(got.diagnostics.algorithm, Algorithm::Randomized);
assert!(got.diagnostics.matvecs > 0);
match got.diagnostics.detail {
Detail::Randomized {
power_iterations, ..
} => {
assert_eq!(power_iterations, DEFAULT_POWER_ITERATIONS);
}
ref other => panic!("wrong detail variant: {other:?}"),
}
}
#[test]
#[ignore = "diagnostic, run explicitly"]
fn report_convergence_rates() {
let cases: Vec<(&str, SvdMat<f64>, usize)> = vec![
("diag_60_linear", diagonal(60), 10),
("lowrank_400x120_r10", gen_lowrank(400, 120, 10, 5), 10),
("sparse_600x200_flat", gen_sparse(600, 200, 0.05, 13), 15),
];
for (name, a, rank) in cases {
let want = reference_singular_values(&dense_of(&a));
print!(
"{name:<22} sigma_ratio={:.3} ",
want[rank] / want[rank - 1]
);
for q in [0usize, 1, 2, 4, 7] {
let cfg = RandomizedConfig::new(rank).seed(42).power_iterations(q);
let got = svd_with(&a, &cfg, None).unwrap();
print!("q{q}={:.2e} ", max_rel_error(&a, &got));
}
for b in [2usize, 4] {
let cfg = RandomizedConfig::new(rank).seed(42).block_krylov(b);
let got = svd_with(&a, &cfg, None).unwrap();
print!("bk{b}={:.2e} ", max_rel_error(&a, &got));
}
println!();
}
}
}