Skip to main content

oxicuda_cs/sparse_pca/
sparse_pca_witten.rs

1//! Penalised matrix decomposition for Sparse PCA (Witten, Tibshirani & Hastie 2009).
2//!
3//! For a centred data matrix `X` (`n × p`), find rank-1 components `u`, `v` minimising
4//! `½ ||X − σ u v^T||²` subject to `||u||₂ ≤ 1`, `||v||₂ ≤ 1`, `||v||₁ ≤ c₁`.
5//!
6//! Solved by alternation:
7//!   u = X v / ||X v||
8//!   v = S_λ(X^T u) / ||S_λ(X^T u)||  (with λ chosen by binary search to enforce ||v||₁ ≤ c₁)
9//!
10//! For multiple components, deflate `X ← X − σ u v^T` and repeat.
11
12use crate::error::{CsError, CsResult};
13use crate::linalg::{mat_t_vec, mat_vec};
14use crate::thresholding::iht::soft_threshold;
15
16/// Sparse PCA result containing components and singular values.
17#[derive(Debug, Clone)]
18pub struct SparsePcaResult {
19    /// `u` columns stacked as `n × k` row-major (one column per component).
20    pub u: Vec<f64>,
21    /// `v` columns stacked as `p × k` row-major.
22    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
38/// Compute `k` sparse principal components with L1 budget `c1` on each `v`.
39pub 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        // Initialise v as the first row of `residual` (normalised) — could randomise instead.
67        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 update: u = X v / ||X v||
74            u = mat_vec(&residual, n, p, &v)?;
75            l2_normalise(&mut u);
76            // v update: pick lambda s.t. ||S_lambda(X^T u)||_1 ≤ c1; binary-search.
77            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            // Compute sigma.
97            sigma = mat_vec(&residual, n, p, &v_new)?
98                .iter()
99                .zip(u.iter())
100                .map(|(a, b)| a * b)
101                .sum::<f64>();
102            // Convergence check.
103            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        // Store and deflate.
115        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        // residual = residual - sigma * u v^T
123        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        // Rank-1 data: x = u v^T with sparse v.
144        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; // column 0 only
149        }
150        let r = sparse_pca_witten(&data, n, p, 1, 1.5, 50, 1.0e-7).expect("ok");
151        // v should put most mass on j=0.
152        assert!(r.v[0].abs() > 0.5);
153    }
154}