Skip to main content

stats_claw/algorithms/decomposition/
factor_analysis.rs

1//! Factor analysis via the SVD-based EM of `sklearn.decomposition.FactorAnalysis`.
2//!
3//! Factor analysis models each observation as `x = W·z + μ + ε` with `z` a
4//! low-dimensional latent factor and `ε` diagonal Gaussian noise of variance `Ψ`.
5//! The loadings `W` and noise `Ψ` are fit by the deterministic EM that scikit-learn
6//! uses: each iteration scales the centered data by `1/√Ψ`, takes the top-`k`
7//! singular directions, forms `W = √max(s²−1, 0)·Vᵀ·√Ψ`, and updates `Ψ` to the
8//! residual per-feature variance. The right singular vectors and squared singular
9//! values come from the eigen-decomposition of `MᵀM`, reusing the family's Jacobi
10//! solver.
11//!
12//! The factor loadings are defined only up to sign and rotation, so the equivalence
13//! suite compares the rotation-invariant reconstruction error rather than `W`
14//! itself.
15
16use crate::algorithms::decomposition::{
17    at, count_to_f64, jacobi_eigen, mean_center, reconstruction_error, symmetric_inverse,
18};
19
20/// Lower bound applied to noise variances and `√Ψ` for numerical stability,
21/// matching scikit-learn's `SMALL` constant.
22const SMALL: f64 = 1e-12;
23
24/// Outcome of a factor-analysis fit.
25#[derive(Debug, Clone)]
26pub struct FactorResult {
27    /// The `k × dim` factor loading matrix `W`, one row per latent factor.
28    pub loadings: Vec<Vec<f64>>,
29    /// Per-feature diagonal noise variance `Ψ`, length `dim`.
30    pub noise_variance: Vec<f64>,
31    /// Mean squared error of reconstructing the input from its latent factors.
32    pub reconstruction_error: f64,
33}
34
35/// Fits a `k`-factor model to `data` by SVD-based EM.
36///
37/// # Arguments
38///
39/// * `data` — observations; each inner slice is one row of equal dimension. Empty
40///   input yields an empty result.
41/// * `n_components` — number of latent factors `k`, clamped to the feature dimension.
42/// * `max_iter` — maximum EM iterations.
43/// * `tol` — convergence tolerance on the log-likelihood increase between iterations.
44///
45/// # Returns
46///
47/// A [`FactorResult`] with the loadings, noise variance, and reconstruction error.
48///
49/// # Examples
50///
51/// ```
52/// use stats_claw::algorithms::decomposition::factor_analysis;
53///
54/// let data = vec![
55///     vec![1.0, 1.1],
56///     vec![2.0, 2.2],
57///     vec![3.0, 2.9],
58///     vec![4.0, 4.1],
59/// ];
60/// let r = factor_analysis(&data, 1, 1000, 1e-2);
61/// // A single factor captures the shared trend, so reconstruction is tight.
62/// assert!(r.reconstruction_error < 0.1, "error was {}", r.reconstruction_error);
63/// ```
64#[must_use]
65pub fn factor_analysis(
66    data: &[Vec<f64>],
67    n_components: usize,
68    max_iter: usize,
69    tol: f64,
70) -> FactorResult {
71    let dim = data.first().map_or(0, Vec::len);
72    let k = n_components.min(dim);
73    let n = data.len();
74    if dim == 0 || k == 0 || n == 0 {
75        return FactorResult {
76            loadings: Vec::new(),
77            noise_variance: Vec::new(),
78            reconstruction_error: 0.0,
79        };
80    }
81    let (centered, means) = mean_center(data, dim);
82    let var = column_variance(&centered, dim);
83    let nsqrt = count_to_f64(n).sqrt();
84    let llconst = count_to_f64(dim).mul_add((2.0 * std::f64::consts::PI).ln(), count_to_f64(k));
85
86    let mut psi = vec![1.0_f64; dim];
87    let mut loadings = vec![vec![0.0_f64; dim]; k];
88    let mut old_ll = f64::NEG_INFINITY;
89
90    for _ in 0..max_iter {
91        let sqrt_psi: Vec<f64> = psi.iter().map(|&p| p.sqrt() + SMALL).collect();
92        let scaled = scale_columns(&centered, &sqrt_psi, nsqrt);
93        let (sq_singular, right_vectors, unexplained) = top_k_svd(&scaled, dim, k);
94
95        loadings = build_loadings(&sq_singular, &right_vectors, &sqrt_psi, dim, k);
96
97        let ll = log_likelihood(llconst, &sq_singular, &psi, unexplained, n);
98        if (ll - old_ll) < tol {
99            break;
100        }
101        old_ll = ll;
102        update_noise(&mut psi, &var, &loadings, dim);
103    }
104
105    let reconstructed = reconstruct(&centered, &loadings, &psi, &means, dim, k);
106    let error = reconstruction_error(data, &reconstructed);
107    FactorResult {
108        loadings,
109        noise_variance: psi,
110        reconstruction_error: error,
111    }
112}
113
114/// Computes the per-feature population variance of mean-centered rows (the `n`
115/// denominator, matching `numpy.var`).
116fn column_variance(centered: &[Vec<f64>], dim: usize) -> Vec<f64> {
117    let mut var = vec![0.0_f64; dim];
118    for row in centered {
119        for (v, &x) in var.iter_mut().zip(row) {
120            *v = x.mul_add(x, *v);
121        }
122    }
123    let n = count_to_f64(centered.len());
124    if n > 0.0 {
125        for v in &mut var {
126            *v /= n;
127        }
128    }
129    var
130}
131
132/// Returns `centered[i][j] / (sqrt_psi[j] * nsqrt)` — the EM-step scaling.
133fn scale_columns(centered: &[Vec<f64>], sqrt_psi: &[f64], nsqrt: f64) -> Vec<Vec<f64>> {
134    centered
135        .iter()
136        .map(|row| {
137            row.iter()
138                .zip(sqrt_psi)
139                .map(|(&x, &sp)| x / (sp * nsqrt))
140                .collect()
141        })
142        .collect()
143}
144
145/// Computes the top-`k` squared singular values, right singular vectors, and the
146/// unexplained variance of the `n × dim` matrix `scaled`, via the eigen-
147/// decomposition of `scaledᵀ·scaled`.
148///
149/// The eigenvalues of the Gram matrix are the squared singular values; the
150/// eigenvectors are the right singular vectors. The unexplained variance is the sum
151/// of the squared singular values beyond the top `k` (scikit-learn's `unexp_var`),
152/// which the log-likelihood requires so the stopping criterion matches exactly.
153///
154/// Returns `(squared_singular_values, right_vectors, unexplained_variance)` ordered
155/// by descending value.
156fn top_k_svd(scaled: &[Vec<f64>], dim: usize, k: usize) -> (Vec<f64>, Vec<Vec<f64>>, f64) {
157    let mut gram = vec![0.0_f64; dim * dim];
158    for row in scaled {
159        for i in 0..dim {
160            let ri = row.get(i).copied().unwrap_or(0.0);
161            for j in 0..dim {
162                let rj = row.get(j).copied().unwrap_or(0.0);
163                if let Some(slot) = gram.get_mut(i * dim + j) {
164                    *slot = ri.mul_add(rj, *slot);
165                }
166            }
167        }
168    }
169    let (values, vectors) = jacobi_eigen(&gram, dim);
170    let mut order: Vec<usize> = (0..dim).collect();
171    order.sort_by(|&a, &b| {
172        let va = values.get(a).copied().unwrap_or(f64::NEG_INFINITY);
173        let vb = values.get(b).copied().unwrap_or(f64::NEG_INFINITY);
174        vb.partial_cmp(&va).unwrap_or(std::cmp::Ordering::Equal)
175    });
176    let mut sq_singular = Vec::with_capacity(k);
177    let mut right_vectors = Vec::with_capacity(k);
178    for &col in order.iter().take(k) {
179        sq_singular.push(values.get(col).copied().unwrap_or(0.0).max(0.0));
180        right_vectors.push((0..dim).map(|row| at(&vectors, dim, row, col)).collect());
181    }
182    let unexplained: f64 = order
183        .iter()
184        .skip(k)
185        .map(|&col| values.get(col).copied().unwrap_or(0.0).max(0.0))
186        .sum();
187    (sq_singular, right_vectors, unexplained)
188}
189
190/// Builds `W = √max(s² − 1, 0) · Vᵀ · √Ψ`, one loading row per component.
191fn build_loadings(
192    sq_singular: &[f64],
193    right_vectors: &[Vec<f64>],
194    sqrt_psi: &[f64],
195    dim: usize,
196    k: usize,
197) -> Vec<Vec<f64>> {
198    (0..k)
199        .map(|c| {
200            let scale = (sq_singular.get(c).copied().unwrap_or(0.0) - 1.0)
201                .max(0.0)
202                .sqrt();
203            let vector = right_vectors.get(c);
204            (0..dim)
205                .map(|j| {
206                    let v = vector.and_then(|row| row.get(j)).copied().unwrap_or(0.0);
207                    scale * v * sqrt_psi.get(j).copied().unwrap_or(0.0)
208                })
209                .collect()
210        })
211        .collect()
212}
213
214/// Evaluates scikit-learn's factor-analysis log-likelihood
215/// `−n/2 · (llconst + Σ log s² + unexplained + Σ log Ψ)`.
216///
217/// The exact value (including `unexplained`) is needed so the `(ll − old_ll) < tol`
218/// stopping iteration matches scikit-learn's, which determines the final loadings.
219fn log_likelihood(
220    llconst: f64,
221    sq_singular: &[f64],
222    psi: &[f64],
223    unexplained: f64,
224    n: usize,
225) -> f64 {
226    let log_s: f64 = sq_singular.iter().map(|&s| s.max(SMALL).ln()).sum();
227    let log_psi: f64 = psi.iter().map(|&p| p.ln()).sum();
228    let ll = llconst + log_s + unexplained + log_psi;
229    ll * (-count_to_f64(n) / 2.0)
230}
231
232/// Updates the diagonal noise to the residual per-feature variance
233/// `Ψ_j = max(var_j − Σ_c W_{c,j}², SMALL)`.
234fn update_noise(psi: &mut [f64], var: &[f64], loadings: &[Vec<f64>], dim: usize) {
235    for j in 0..dim {
236        let explained: f64 = loadings
237            .iter()
238            .map(|row| {
239                let w = row.get(j).copied().unwrap_or(0.0);
240                w * w
241            })
242            .sum();
243        if let Some(slot) = psi.get_mut(j) {
244            *slot = (var.get(j).copied().unwrap_or(0.0) - explained).max(SMALL);
245        }
246    }
247}
248
249/// Reconstructs the data: latent means `z = X_c·(W/Ψ)ᵀ·(I + (W/Ψ)·Wᵀ)⁻¹`, then
250/// `x̂ = z·W + μ`, mirroring scikit-learn's `transform` followed by the linear map.
251fn reconstruct(
252    centered: &[Vec<f64>],
253    loadings: &[Vec<f64>],
254    psi: &[f64],
255    means: &[f64],
256    dim: usize,
257    k: usize,
258) -> Vec<Vec<f64>> {
259    // Wpsi[c][j] = W[c][j] / psi[j]
260    let wpsi: Vec<Vec<f64>> = loadings
261        .iter()
262        .map(|row| {
263            row.iter()
264                .zip(psi)
265                .map(|(&w, &p)| w / p)
266                .collect::<Vec<f64>>()
267        })
268        .collect();
269    // gram[a][b] = (Wpsi · Wᵀ)[a][b], plus identity → I + Wpsi·Wᵀ.
270    let mut posterior = vec![0.0_f64; k * k];
271    for a in 0..k {
272        for b in 0..k {
273            let dot: f64 = wpsi.get(a).zip(loadings.get(b)).map_or(0.0, |(wa, lb)| {
274                wa.iter().zip(lb).map(|(&x, &y)| x * y).sum()
275            });
276            let eye = if a == b { 1.0 } else { 0.0 };
277            if let Some(slot) = posterior.get_mut(a * k + b) {
278                *slot = dot + eye;
279            }
280        }
281    }
282    let cov_z = symmetric_inverse(&posterior, k);
283
284    centered
285        .iter()
286        .map(|row| {
287            // tmp[c] = row · Wpsi[c]ᵀ
288            let tmp: Vec<f64> = (0..k)
289                .map(|c| {
290                    wpsi.get(c)
291                        .map_or(0.0, |w| row.iter().zip(w).map(|(&x, &y)| x * y).sum())
292                })
293                .collect();
294            // z[c] = Σ_a tmp[a] · cov_z[a][c]
295            let z: Vec<f64> = (0..k)
296                .map(|c| {
297                    (0..k)
298                        .map(|a| tmp.get(a).copied().unwrap_or(0.0) * at(&cov_z, k, a, c))
299                        .sum()
300                })
301                .collect();
302            // x̂[j] = Σ_c z[c] · W[c][j] + mean[j]
303            (0..dim)
304                .map(|j| {
305                    let acc: f64 = (0..k)
306                        .map(|c| {
307                            z.get(c).copied().unwrap_or(0.0)
308                                * loadings
309                                    .get(c)
310                                    .and_then(|w| w.get(j))
311                                    .copied()
312                                    .unwrap_or(0.0)
313                        })
314                        .sum();
315                    acc + means.get(j).copied().unwrap_or(0.0)
316                })
317                .collect()
318        })
319        .collect()
320}
321
322#[cfg(test)]
323mod tests {
324    use super::*;
325
326    #[test]
327    fn single_factor_captures_shared_trend() {
328        let data = vec![
329            vec![1.0, 1.1],
330            vec![2.0, 2.2],
331            vec![3.0, 2.9],
332            vec![4.0, 4.1],
333        ];
334        let r = factor_analysis(&data, 1, 1000, 1e-2);
335        assert!(
336            r.reconstruction_error < 0.1,
337            "error was {}",
338            r.reconstruction_error
339        );
340    }
341
342    #[test]
343    fn empty_input_is_empty() {
344        let r = factor_analysis(&[], 2, 100, 1e-2);
345        assert!(r.loadings.is_empty(), "loadings not empty");
346    }
347
348    #[test]
349    fn k_is_clamped_to_dimension() {
350        let data = vec![vec![1.0, 2.0], vec![3.0, 4.0], vec![5.0, 6.0]];
351        let r = factor_analysis(&data, 9, 100, 1e-2);
352        assert!(r.loadings.len() <= 2, "factor count exceeded dim");
353    }
354}