1use ndarray::{Array1, Array2, ArrayView2};
9use solow_core::{Error, Result};
10
11#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
13#[derive(Clone, Debug, PartialEq)]
14pub struct SparsePCA {
15 pub components: Array2<f64>,
17 pub mean: Array1<f64>,
19 pub n_components: usize,
21 pub alpha: f64,
23 pub n_iter: usize,
25}
26
27impl SparsePCA {
28 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 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 let (u0, s0, v0) = svd(¢red, 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 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 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 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 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 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}