use super::*;
use crate::matrix::traits::{RandomizedAlgs, RsvdArgs, SampleOps};
fn spiked(n: usize, p: usize, strength: f32, seed: u64) -> (DMatrix<f32>, DVector<f32>) {
let noise = DMatrix::<f32>::rnorm_seeded(n, p, seed);
let mut u = DVector::<f32>::from_fn(n, |i, _| if i % 2 == 0 { 1.0 } else { -1.0 });
u /= u.norm();
let mut v = DVector::<f32>::from_fn(p, |j, _| if j < p / 3 { 1.0 } else { 0.0 });
v /= v.norm();
(noise + strength * (&u * v.transpose()), v)
}
#[test]
fn rsvd_recovers_a_planted_spike() {
let (x, v) = spiked(120, 200, 40.0, 1);
let full = x.clone().svd(false, true);
let exact = full.singular_values[0];
let exact_v = full.v_t.as_ref().unwrap().row(0).transpose();
let (_, s, vv) = x.rsvd(2).unwrap();
let s1 = s[0];
assert!(
(s1 - exact).abs() <= 0.02 * exact,
"leading singular value {s1} vs exact {exact}"
);
let cos_exact = vv.column(0).dot(&exact_v).abs();
assert!(
cos_exact >= 0.99,
"cosine with the exact leading vector {cos_exact}"
);
let cos_plant = exact_v.dot(&v).abs();
assert!(cos_plant >= 0.85, "exact vs planted cosine {cos_plant}");
}
#[test]
fn rsvd_leading_value_of_noise_stays_at_the_edge() {
let x = DMatrix::<f32>::rnorm_seeded(120, 200, 2);
let exact = x.clone().svd(false, false).singular_values[0];
let (_, s, _) = x.rsvd(2).unwrap();
assert!(
s[0] <= exact * 1.001,
"randomised value cannot exceed the exact one"
);
assert!(
s[0] >= 0.9 * exact,
"randomised value {} far below exact {exact}",
s[0]
);
}
fn clustered(n: usize, seed: u64) -> (DMatrix<f64>, Vec<f64>) {
let q = DMatrix::<f64>::rnorm_seeded(n, n, seed).qr().q();
let lambda: Vec<f64> = (0..n).map(|i| 0.99f64.powi(i as i32)).collect();
(
&q * DMatrix::from_diagonal(&DVector::from_vec(lambda.clone())) * q.transpose(),
lambda,
)
}
#[test]
fn rsvd_with_more_iterations_and_oversampling_separates_a_clustered_spectrum() {
let (x, lambda) = clustered(300, 3);
let worst = |s: &DVector<f64>| {
(0..15)
.map(|i| (s[i] - lambda[i]).abs())
.fold(0.0f64, f64::max)
};
let (_, s_default, _) = x.rsvd(15).unwrap();
let args = RsvdArgs {
power_iters: 20,
oversample: 10,
};
let (_, s_more, _) = x.rsvd_with(15, &args).unwrap();
assert!(
worst(&s_more) < 1e-3,
"20 iterations, 10 oversample: worst error {}",
worst(&s_more)
);
assert!(
worst(&s_more) < worst(&s_default) / 5.0,
"more iterations and oversampling should be clearly better: {} vs default {}",
worst(&s_more),
worst(&s_default)
);
}
#[test]
fn rsvd_defaults_are_the_long_standing_settings() {
assert_eq!(
RsvdArgs::default(),
RsvdArgs {
power_iters: 5,
oversample: 5
}
);
}
#[test]
fn rsvd_with_caps_an_oversized_oversample() {
let x = DMatrix::<f64>::rnorm_seeded(40, 30, 5);
let args = RsvdArgs {
power_iters: 2,
oversample: usize::MAX,
};
let (_, s, _) = x.rsvd_with(10, &args).unwrap();
assert_eq!(s.len(), 10);
}
#[test]
fn rsvd_of_an_empty_matrix_is_an_error() {
assert!(DMatrix::<f64>::zeros(0, 5).rsvd(3).is_err());
}