use crate::decomposition::jacobi;
use crate::decomposition::pca::matmul_flat;
use crate::model_selection::rng::Rng;
fn gaussian_matrix(rows: usize, cols: usize, seed: u64) -> Vec<f64> {
let mut rng = Rng::new(seed);
let mut out = vec![0.0; rows * cols];
for v in out.iter_mut() {
*v = rng.next_normal();
}
out
}
#[allow(dead_code)]
pub(crate) struct RandomizedSvd {
pub singular_values: Vec<f64>,
pub u: Vec<f64>,
pub vt: Vec<f64>,
}
pub(crate) fn randomized_svd(
x: &[f64],
n: usize,
p: usize,
k: usize,
oversample: usize,
n_iter: usize,
seed: u64,
) -> Option<RandomizedSvd> {
if n == 0 || p == 0 || k == 0 {
return None;
}
let k = k.min(n).min(p);
let l = (k + oversample).min(n).min(p);
let omega = gaussian_matrix(p, l, seed);
let mut y = vec![0.0; n * l];
matmul_flat(&mut y, x, &omega, n, p, l);
for _ in 0..n_iter {
let mut z = vec![0.0; p * l];
matmul_transpose_at(&mut z, x, &y, n, p, l);
let mut y2 = vec![0.0; n * l];
matmul_flat(&mut y2, x, &z, n, p, l);
y = y2;
}
let q = modified_gram_schmidt(&y, n, l);
let mut b = vec![0.0; l * p];
matmul_qt_x(&mut b, &q, x, n, p, l);
let mut c = vec![0.0; l * l];
for i in 0..l {
for j in 0..l {
let mut s = 0.0;
for t in 0..p {
s += b[i * p + t] * b[j * p + t];
}
c[i * l + j] = s;
}
}
let (eigvals, eigvecs) = jacobi::eigh_flat(&mut c, l)?;
let mut svals: Vec<f64> = eigvals.iter().map(|&v| v.max(0.0).sqrt()).collect();
let mut u = vec![0.0; n * k];
let mut vt = vec![0.0; k * p];
for j in 0..k {
let sigma = svals[j];
for r in 0..n {
let mut s = 0.0;
for t in 0..l {
s += q[r * l + t] * eigvecs[j * l + t];
}
u[r * k + j] = s;
}
if sigma > 1e-300 {
let inv_sigma = 1.0 / sigma;
for c_idx in 0..p {
let mut s = 0.0;
for t in 0..l {
s += b[t * p + c_idx] * eigvecs[j * l + t];
}
vt[j * p + c_idx] = s * inv_sigma;
}
} else {
}
}
svals.truncate(k);
Some(RandomizedSvd {
singular_values: svals,
u,
vt,
})
}
fn modified_gram_schmidt(a: &[f64], m: usize, n: usize) -> Vec<f64> {
let mut q = a.to_vec();
for j in 0..n {
for i in 0..j {
let mut dot = 0.0;
for r in 0..m {
dot += q[r * n + i] * q[r * n + j];
}
for r in 0..m {
q[r * n + j] -= dot * q[r * n + i];
}
}
let mut norm2 = 0.0;
for r in 0..m {
norm2 += q[r * n + j] * q[r * n + j];
}
let norm = norm2.sqrt();
if norm > 1e-300 {
let inv = 1.0 / norm;
for r in 0..m {
q[r * n + j] *= inv;
}
}
}
q
}
fn matmul_transpose_at(c: &mut [f64], a: &[f64], b: &[f64], n: usize, p: usize, l: usize) {
#[cfg(feature = "matrixmultiply")]
{
unsafe {
matrixmultiply::dgemm(
p,
n,
l,
1.0,
a.as_ptr(),
1,
p as isize, b.as_ptr(),
l as isize,
1,
0.0,
c.as_mut_ptr(),
l as isize,
1,
);
}
}
#[cfg(not(feature = "matrixmultiply"))]
{
for i in 0..p {
for j in 0..l {
let mut s = 0.0;
for t in 0..n {
s += a[t * p + i] * b[t * l + j];
}
c[i * l + j] = s;
}
}
}
}
fn matmul_qt_x(c: &mut [f64], q: &[f64], x: &[f64], n: usize, p: usize, l: usize) {
for i in 0..l {
for j in 0..p {
let mut s = 0.0;
for t in 0..n {
s += q[t * l + i] * x[t * p + j];
}
c[i * p + j] = s;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn randomized_svd_reconstructs_low_rank() {
let n = 6;
let p = 4;
let sqrt3 = 3f64.sqrt();
let sqrt2 = 2f64.sqrt();
let u1 = [1.0 / sqrt3, 0.0, 0.0, 1.0 / sqrt3, 0.0, 1.0 / sqrt3];
let v1 = [1.0 / sqrt2, 1.0 / sqrt2, 0.0, 0.0];
let u2 = [0.0, 1.0 / sqrt3, 1.0 / sqrt3, 0.0, 1.0 / sqrt3, 0.0];
let v2 = [0.0, 0.0, 1.0 / sqrt2, 1.0 / sqrt2];
let mut x = vec![0.0; n * p];
for i in 0..n {
for j in 0..p {
x[i * p + j] = 10.0 * u1[i] * v1[j] + 3.0 * u2[i] * v2[j];
}
}
let svd = randomized_svd(&x, n, p, 2, 0, 7, 12345).unwrap();
assert_eq!(svd.singular_values.len(), 2);
assert!(
svd.singular_values[0] > svd.singular_values[1],
"sigma0={} should exceed sigma1={}",
svd.singular_values[0],
svd.singular_values[1]
);
assert!(
(svd.singular_values[0] - 10.0).abs() < 1e-6,
"sigma0 = {}",
svd.singular_values[0]
);
assert!(
(svd.singular_values[1] - 3.0).abs() < 1e-6,
"sigma1 = {}",
svd.singular_values[1]
);
let mut recon = vec![0.0; n * p];
let mut sv = vec![0.0; 2 * p];
for j in 0..2 {
for c in 0..p {
sv[j * p + c] = svd.vt[j * p + c] * svd.singular_values[j];
}
}
matmul_flat(&mut recon, &svd.u, &sv, n, 2, p);
let mut max_err: f64 = 0.0;
for i in 0..n {
for j in 0..p {
max_err = max_err.max((recon[i * p + j] - x[i * p + j]).abs());
}
}
assert!(max_err < 1e-6, "max reconstruction error = {max_err}");
}
#[test]
fn modified_gram_schmidt_orthonormal() {
let m = 5;
let n = 3;
let a = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 0.0, 1.0, 1.0, 0.0, 2.0, 1.0, 0.0, 1.0, 1.0, 0.0, ];
let mut rm = vec![0.0; m * n];
for i in 0..m {
rm[i * n] = a[i];
rm[i * n + 1] = a[m + i];
rm[i * n + 2] = a[2 * m + i];
}
let q = modified_gram_schmidt(&rm, m, n);
for i in 0..n {
for j in 0..n {
let mut dot = 0.0;
for r in 0..m {
dot += q[r * n + i] * q[r * n + j];
}
let expect = if i == j { 1.0 } else { 0.0 };
assert!(
(dot - expect).abs() < 1e-9,
"QᵀQ[{i},{j}] = {dot} ≠ {expect}"
);
}
}
}
#[test]
fn randomized_svd_singular_values_nonneg() {
let x = vec![0.0; 4 * 3];
let svd = randomized_svd(&x, 4, 3, 2, 2, 3, 1).unwrap();
for &s in &svd.singular_values {
assert!(approx(s, 0.0, 1e-9), "expected zero sigma, got {s}");
}
}
#[test]
fn randomized_svd_empty_returns_none() {
assert!(randomized_svd(&[], 0, 4, 2, 2, 3, 1).is_none());
assert!(randomized_svd(&[], 4, 0, 2, 2, 3, 1).is_none());
assert!(randomized_svd(&[0.0; 4], 4, 1, 0, 2, 3, 1).is_none());
}
}