oxicuda_cs/sparse_pca/
sparse_pca_witten.rs1use crate::error::{CsError, CsResult};
13use crate::linalg::{mat_t_vec, mat_vec};
14use crate::thresholding::iht::soft_threshold;
15
16#[derive(Debug, Clone)]
18pub struct SparsePcaResult {
19 pub u: Vec<f64>,
21 pub v: Vec<f64>,
23 pub sigmas: Vec<f64>,
24 pub iterations: Vec<usize>,
25}
26
27fn l2_normalise(x: &mut [f64]) {
28 let nrm: f64 = x.iter().map(|v| v * v).sum::<f64>().sqrt().max(1.0e-300);
29 for v in x.iter_mut() {
30 *v /= nrm;
31 }
32}
33
34fn l1_norm(x: &[f64]) -> f64 {
35 x.iter().map(|v| v.abs()).sum()
36}
37
38pub fn sparse_pca_witten(
40 x: &[f64],
41 n: usize,
42 p: usize,
43 k: usize,
44 c1: f64,
45 max_iter: usize,
46 tol: f64,
47) -> CsResult<SparsePcaResult> {
48 if x.len() != n * p {
49 return Err(CsError::ShapeMismatch {
50 expected: vec![n, p],
51 got: vec![x.len()],
52 });
53 }
54 if k == 0 || k > n.min(p) {
55 return Err(CsError::InvalidRank(k));
56 }
57 if c1 <= 0.0 {
58 return Err(CsError::InvalidParameter("c1 must be > 0".into()));
59 }
60 let mut residual = x.to_vec();
61 let mut u_cols = vec![0.0_f64; n * k];
62 let mut v_cols = vec![0.0_f64; p * k];
63 let mut sigmas = vec![0.0_f64; k];
64 let mut iters = vec![0_usize; k];
65 for comp in 0..k {
66 let mut v = vec![0.0_f64; p];
68 v[..p].copy_from_slice(&residual[..p]);
69 l2_normalise(&mut v);
70 let mut u = vec![0.0_f64; n];
71 let mut sigma = 0.0_f64;
72 for it in 0..max_iter {
73 u = mat_vec(&residual, n, p, &v)?;
75 l2_normalise(&mut u);
76 let xtu = mat_t_vec(&residual, n, p, &u)?;
78 let max_abs = xtu.iter().fold(0.0_f64, |a, &v| a.max(v.abs()));
79 let mut lo = 0.0_f64;
80 let mut hi = max_abs + 1.0;
81 for _ in 0..30 {
82 let mid = 0.5 * (lo + hi);
83 let cand = soft_threshold(&xtu, mid);
84 let nrm = cand.iter().map(|v| v * v).sum::<f64>().sqrt().max(1.0e-300);
85 let v_cand: Vec<f64> = cand.iter().map(|v| v / nrm).collect();
86 let l1 = l1_norm(&v_cand);
87 if l1 > c1 {
88 lo = mid;
89 } else {
90 hi = mid;
91 }
92 }
93 let cand = soft_threshold(&xtu, hi);
94 let nrm = cand.iter().map(|v| v * v).sum::<f64>().sqrt().max(1.0e-300);
95 let v_new: Vec<f64> = cand.iter().map(|v| v / nrm).collect();
96 sigma = mat_vec(&residual, n, p, &v_new)?
98 .iter()
99 .zip(u.iter())
100 .map(|(a, b)| a * b)
101 .sum::<f64>();
102 let mut delta = 0.0_f64;
104 for j in 0..p {
105 let d = v_new[j] - v[j];
106 delta += d * d;
107 }
108 v = v_new;
109 iters[comp] = it + 1;
110 if delta.sqrt() < tol {
111 break;
112 }
113 }
114 for i in 0..n {
116 u_cols[i * k + comp] = u[i];
117 }
118 for j in 0..p {
119 v_cols[j * k + comp] = v[j];
120 }
121 sigmas[comp] = sigma;
122 for i in 0..n {
124 for j in 0..p {
125 residual[i * p + j] -= sigma * u[i] * v[j];
126 }
127 }
128 }
129 Ok(SparsePcaResult {
130 u: u_cols,
131 v: v_cols,
132 sigmas,
133 iterations: iters,
134 })
135}
136
137#[cfg(test)]
138mod tests {
139 use super::*;
140
141 #[test]
142 fn sparse_pca_rank1_runs() {
143 let n = 8;
145 let p = 6;
146 let mut data = vec![0.0_f64; n * p];
147 for i in 0..n {
148 data[i * p] = (i as f64) - 3.5; }
150 let r = sparse_pca_witten(&data, n, p, 1, 1.5, 50, 1.0e-7).expect("ok");
151 assert!(r.v[0].abs() > 0.5);
153 }
154}