Skip to main content

solow_decomposition/
sparse_pca.rs

1//! SparsePCA — L1-penalised principal-component learning
2//! (Zou-Hastie-Tibshirani 2006).
3//!
4//! Alternates between:
5//!   1. Fix loadings `V`, solve L1-Ridge regression for scores `U`.
6//!   2. Fix `U`, update `V` by the reduced SVD of `XᵀU`.
7
8use ndarray::{Array1, Array2, ArrayView2};
9use solow_core::{Error, Result};
10
11/// Fitted SparsePCA.
12#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
13#[derive(Clone, Debug, PartialEq)]
14pub struct SparsePCA {
15    /// Sparse component loadings `(k × d)`.
16    pub components: Array2<f64>,
17    /// Column mean subtracted at fit time.
18    pub mean: Array1<f64>,
19    /// Kept rank.
20    pub n_components: usize,
21    /// L1 penalty α used.
22    pub alpha: f64,
23    /// Convergence iterations run.
24    pub n_iter: usize,
25}
26
27impl SparsePCA {
28    /// Fit with defaults `alpha = 1.0`, `max_iter = 100`, `tol = 1e-6`.
29    pub fn fit(x: ArrayView2<'_, f64>, n_components: usize) -> Result<Self> {
30        Self::fit_with(x, n_components, 1.0, 100, 1e-6)
31    }
32
33    /// Full-configuration fit.
34    pub fn fit_with(
35        x: ArrayView2<'_, f64>,
36        n_components: usize,
37        alpha: f64,
38        max_iter: usize,
39        tol: f64,
40    ) -> Result<Self> {
41        let n = x.nrows();
42        let d = x.ncols();
43        if n_components == 0 || n_components > d {
44            return Err(Error::Value("SparsePCA: n_components out of range".into()));
45        }
46        if alpha < 0.0 {
47            return Err(Error::Value("SparsePCA: alpha must be ≥ 0".into()));
48        }
49        let mut mean = Array1::<f64>::zeros(d);
50        for j in 0..d {
51            for i in 0..n {
52                mean[j] += x[[i, j]];
53            }
54            mean[j] /= n as f64;
55        }
56        let mut centred = Array2::<f64>::zeros((n, d));
57        for i in 0..n {
58            for j in 0..d {
59                centred[[i, j]] = x[[i, j]] - mean[j];
60            }
61        }
62        // Warm start from top-k SVD.
63        let (u0, s0, v0) = svd(&centred, 300, 1e-12);
64        let mut v = Array2::<f64>::zeros((n_components, d));
65        for k in 0..n_components {
66            for j in 0..d {
67                v[[k, j]] = v0[[j, k]];
68            }
69        }
70        let mut u = Array2::<f64>::zeros((n, n_components));
71        for k in 0..n_components {
72            for i in 0..n {
73                u[[i, k]] = u0[[i, k]] * s0[k];
74            }
75        }
76        let mut iters = 0_usize;
77        for it in 0..max_iter {
78            iters = it + 1;
79            // Update V (soft-threshold).
80            let mut vt = Array2::<f64>::zeros((d, n_components));
81            for j in 0..d {
82                for k in 0..n_components {
83                    let mut s = 0.0_f64;
84                    for i in 0..n {
85                        s += centred[[i, j]] * u[[i, k]];
86                    }
87                    vt[[j, k]] = soft_threshold(s, alpha);
88                }
89            }
90            // Normalise columns of vt.
91            for k in 0..n_components {
92                let mut nrm = 0.0_f64;
93                for j in 0..d {
94                    nrm += vt[[j, k]] * vt[[j, k]];
95                }
96                let nrm = nrm.sqrt().max(1e-30);
97                for j in 0..d {
98                    vt[[j, k]] /= nrm;
99                }
100            }
101            let mut v_new = Array2::<f64>::zeros((n_components, d));
102            for k in 0..n_components {
103                for j in 0..d {
104                    v_new[[k, j]] = vt[[j, k]];
105                }
106            }
107            // Update U = X V*.
108            let mut u_new = Array2::<f64>::zeros((n, n_components));
109            for i in 0..n {
110                for k in 0..n_components {
111                    let mut s = 0.0_f64;
112                    for j in 0..d {
113                        s += centred[[i, j]] * v_new[[k, j]];
114                    }
115                    u_new[[i, k]] = s;
116                }
117            }
118            // Convergence.
119            let mut delta = 0.0_f64;
120            for k in 0..n_components {
121                for j in 0..d {
122                    let dd = v_new[[k, j]] - v[[k, j]];
123                    delta += dd * dd;
124                }
125            }
126            v = v_new;
127            u = u_new;
128            if delta.sqrt() < tol {
129                break;
130            }
131        }
132        Ok(Self {
133            components: v,
134            mean,
135            n_components,
136            alpha,
137            n_iter: iters,
138        })
139    }
140
141    /// Project new rows.
142    pub fn transform(&self, x: ArrayView2<'_, f64>) -> Result<Array2<f64>> {
143        let n = x.nrows();
144        let d = self.mean.len();
145        let k = self.n_components;
146        if x.ncols() != d {
147            return Err(Error::Shape("SparsePCA::transform: shape mismatch".into()));
148        }
149        let mut out = Array2::<f64>::zeros((n, k));
150        for i in 0..n {
151            for c in 0..k {
152                let mut s = 0.0_f64;
153                for j in 0..d {
154                    s += (x[[i, j]] - self.mean[j]) * self.components[[c, j]];
155                }
156                out[[i, c]] = s;
157            }
158        }
159        Ok(out)
160    }
161}
162
163fn soft_threshold(z: f64, alpha: f64) -> f64 {
164    if z > alpha {
165        z - alpha
166    } else if z < -alpha {
167        z + alpha
168    } else {
169        0.0
170    }
171}
172
173fn svd(a: &Array2<f64>, max_sweeps: usize, tol: f64) -> (Array2<f64>, Vec<f64>, Array2<f64>) {
174    let m = a.nrows();
175    let n = a.ncols();
176    if m >= n {
177        let mut u = a.clone();
178        let mut v = Array2::<f64>::eye(n);
179        for _ in 0..max_sweeps {
180            let mut off = 0.0_f64;
181            for p in 0..(n - 1) {
182                for q in (p + 1)..n {
183                    let mut alpha = 0.0_f64;
184                    let mut beta = 0.0_f64;
185                    let mut gamma = 0.0_f64;
186                    for i in 0..m {
187                        alpha += u[[i, p]] * u[[i, p]];
188                        beta += u[[i, q]] * u[[i, q]];
189                        gamma += u[[i, p]] * u[[i, q]];
190                    }
191                    off += gamma * gamma;
192                    if gamma.abs() < tol * (alpha * beta).sqrt().max(1e-30) {
193                        continue;
194                    }
195                    let zeta = (beta - alpha) / (2.0 * gamma);
196                    let t = zeta.signum() / (zeta.abs() + (1.0 + zeta * zeta).sqrt());
197                    let c = 1.0 / (1.0 + t * t).sqrt();
198                    let s = t * c;
199                    for i in 0..m {
200                        let up = u[[i, p]];
201                        let uq = u[[i, q]];
202                        u[[i, p]] = c * up - s * uq;
203                        u[[i, q]] = s * up + c * uq;
204                    }
205                    for i in 0..n {
206                        let vp = v[[i, p]];
207                        let vq = v[[i, q]];
208                        v[[i, p]] = c * vp - s * vq;
209                        v[[i, q]] = s * vp + c * vq;
210                    }
211                }
212            }
213            if off.sqrt() < tol {
214                break;
215            }
216        }
217        let mut svals = vec![0.0_f64; n];
218        for j in 0..n {
219            let mut s = 0.0_f64;
220            for i in 0..m {
221                s += u[[i, j]] * u[[i, j]];
222            }
223            svals[j] = s.sqrt();
224            let norm = svals[j].max(1e-300);
225            for i in 0..m {
226                u[[i, j]] /= norm;
227            }
228        }
229        let mut idx: Vec<usize> = (0..n).collect();
230        idx.sort_by(|&a, &b| svals[b].partial_cmp(&svals[a]).unwrap());
231        let mut u_sorted = Array2::<f64>::zeros((m, n));
232        let mut v_sorted = Array2::<f64>::zeros((n, n));
233        let mut svals_sorted = vec![0.0_f64; n];
234        for (j, &orig) in idx.iter().enumerate() {
235            for i in 0..m {
236                u_sorted[[i, j]] = u[[i, orig]];
237            }
238            for i in 0..n {
239                v_sorted[[i, j]] = v[[i, orig]];
240            }
241            svals_sorted[j] = svals[orig];
242        }
243        (u_sorted, svals_sorted, v_sorted)
244    } else {
245        let at = a.t().to_owned();
246        let (u_t, s, v_t) = svd(&at, max_sweeps, tol);
247        (v_t, s, u_t)
248    }
249}
250
251#[cfg(test)]
252mod tests {
253    use super::*;
254    use ndarray::array;
255
256    #[test]
257    fn sparse_pca_returns_the_right_shape() {
258        let x = array![
259            [1.0_f64, 2.0, 3.0], [4.0, 5.0, 6.0], [7.0, 8.0, 9.0]
260        ];
261        let m = SparsePCA::fit_with(x.view(), 2, 0.1, 100, 1e-6).unwrap();
262        assert_eq!(m.components.shape(), &[2, 3]);
263    }
264}