Skip to main content

qualia_core_db/solvers/learning/dimensionality/
pca.rs

1//! Principal Component Analysis (ISL ch 12, PRML ch 12) — the eigendecomposition
2//! of the feature covariance, reusing `linear_algebra::{gemm, eigen}` (no new
3//! solver). Mission note: PCA is the principled way to choose the engine's 10D→5D
4//! NQuin relevance projection.
5//!
6//! Centre the data, form the `p×p` covariance `C = Xcᵀ Xc /(n−1)` with `gemm`,
7//! symmetric-eigendecompose it with `symmetric_eigen`, and sort the eigenpairs
8//! descending. Eigenvalue `k` is the variance along principal component `k`.
9//! Kernel-class `DenseLinear` (covariance GEMM), dispatch-ready.
10
11use crate::solvers::learning::LearningError;
12use crate::solvers::linear_algebra::eigen::symmetric_eigen;
13use crate::solvers::linear_algebra::gemm::{gemm, Transpose};
14use crate::solvers::statistics::descriptive::mean;
15
16/// A fitted PCA. `components` holds the principal axes as **rows** (`n_components ×
17/// p`), ordered by descending explained variance; `explained_variance[k]` is the
18/// variance (eigenvalue) along component `k`.
19#[derive(Debug, Clone)]
20pub struct Pca {
21    pub mean: Vec<f64>,
22    pub components: Vec<f64>, // n_components × p, row-major
23    pub explained_variance: Vec<f64>,
24    pub explained_variance_ratio: Vec<f64>,
25    pub n_components: usize,
26    pub p: usize,
27}
28
29impl Pca {
30    /// Project a row-major `n × p` matrix onto the first `k` components, returning
31    /// the `n × k` scores. `k` is clamped to `n_components`.
32    pub fn transform(&self, x: &[f64], n: usize, k: usize) -> Option<Vec<f64>> {
33        let p = self.p;
34        if x.len() != n * p {
35            return None;
36        }
37        let k = k.min(self.n_components);
38        let mut out = vec![0.0; n * k];
39        for i in 0..n {
40            for c in 0..k {
41                let comp = &self.components[c * p..(c + 1) * p];
42                let mut s = 0.0;
43                for j in 0..p {
44                    s += (x[i * p + j] - self.mean[j]) * comp[j];
45                }
46                out[i * k + c] = s;
47            }
48        }
49        Some(out)
50    }
51}
52
53/// Fit PCA to a row-major `n × p` matrix. `None`-equivalent failures are returned
54/// as `LearningError` (fail closed): `InvalidDimension`, `InsufficientData`
55/// (`n < 2`), or `Singular` if the eigensolver cannot decompose the covariance.
56pub fn fit(x: &[f64], n: usize, p: usize) -> Result<Pca, LearningError> {
57    if n == 0 || p == 0 || x.len() != n * p {
58        return Err(LearningError::InvalidDimension);
59    }
60    if n < 2 {
61        return Err(LearningError::InsufficientData);
62    }
63
64    // Column means and centred data.
65    let mut col = vec![0.0; n];
66    let mut means = vec![0.0; p];
67    for j in 0..p {
68        for i in 0..n {
69            col[i] = x[i * p + j];
70        }
71        means[j] = mean(&col).ok_or(LearningError::InsufficientData)?;
72    }
73    let mut xc = vec![0.0; n * p];
74    for i in 0..n {
75        for j in 0..p {
76            xc[i * p + j] = x[i * p + j] - means[j];
77        }
78    }
79
80    // Covariance C = Xcᵀ Xc / (n-1)  (p×p, symmetric).
81    let mut cov = vec![0.0; p * p];
82    gemm(
83        Transpose::Yes,
84        Transpose::No,
85        p,
86        p,
87        n,
88        1.0 / (n as f64 - 1.0),
89        &xc,
90        &xc,
91        0.0,
92        &mut cov,
93    )?;
94    // Symmetrize against round-off so the eigensolver's symmetry check passes.
95    for i in 0..p {
96        for j in (i + 1)..p {
97            let avg = 0.5 * (cov[i * p + j] + cov[j * p + i]);
98            cov[i * p + j] = avg;
99            cov[j * p + i] = avg;
100        }
101    }
102
103    // Symmetric eigendecomposition: eigenvalues on the diagonal, eigenvectors as
104    // columns of `vecs`.
105    let mut vecs = vec![0.0; p * p];
106    symmetric_eigen(p, &mut cov, &mut vecs).map_err(|_| LearningError::Singular)?;
107    let eigvals: Vec<f64> = (0..p).map(|i| cov[i * p + i].max(0.0)).collect();
108
109    // Order components by descending eigenvalue.
110    let mut order: Vec<usize> = (0..p).collect();
111    order.sort_by(|&a, &b| {
112        eigvals[b]
113            .partial_cmp(&eigvals[a])
114            .unwrap_or(core::cmp::Ordering::Equal)
115    });
116
117    let mut components = vec![0.0; p * p]; // p components (rows) × p
118    let mut explained_variance = vec![0.0; p];
119    for (rank, &e) in order.iter().enumerate() {
120        explained_variance[rank] = eigvals[e];
121        for j in 0..p {
122            // eigenvector e is column `e` of `vecs`: vecs[j*p + e].
123            components[rank * p + j] = vecs[j * p + e];
124        }
125    }
126    let total: f64 = explained_variance.iter().sum();
127    let explained_variance_ratio: Vec<f64> = explained_variance
128        .iter()
129        .map(|&v| if total > 0.0 { v / total } else { 0.0 })
130        .collect();
131
132    Ok(Pca {
133        mean: means,
134        components,
135        explained_variance,
136        explained_variance_ratio,
137        n_components: p,
138        p,
139    })
140}
141
142#[cfg(test)]
143mod tests {
144    use super::*;
145
146    #[test]
147    fn one_dominant_direction() {
148        // Variance almost entirely along x; tiny along y → PC1 ~ x-axis, ratio ~1.
149        let x = [-2.0, 0.01, -1.0, -0.01, 0.0, 0.0, 1.0, 0.01, 2.0, -0.01];
150        let m = fit(&x, 5, 2).unwrap();
151        assert!(
152            m.explained_variance_ratio[0] > 0.99,
153            "ratio {}",
154            m.explained_variance_ratio[0]
155        );
156        // PC1 aligns with the x-axis (|component_x| ~ 1, |component_y| ~ 0).
157        assert!(m.components[0].abs() > 0.99 && m.components[1].abs() < 0.05);
158        // Ratios sum to 1.
159        let s: f64 = m.explained_variance_ratio.iter().sum();
160        assert!((s - 1.0).abs() < 1e-9);
161    }
162
163    #[test]
164    fn diagonal_correlation_axis() {
165        // Perfectly correlated x=y → PC1 along the (1,1)/√2 diagonal, PC2 ~ 0 var.
166        let x = [1.0, 1.0, 2.0, 2.0, 3.0, 3.0, 4.0, 4.0, 5.0, 5.0];
167        let m = fit(&x, 5, 2).unwrap();
168        assert!(m.explained_variance_ratio[0] > 0.999);
169        let inv_sqrt2 = 1.0 / 2.0_f64.sqrt();
170        // PC1 components both ≈ ±1/√2.
171        assert!((m.components[0].abs() - inv_sqrt2).abs() < 1e-6);
172        assert!((m.components[1].abs() - inv_sqrt2).abs() < 1e-6);
173        // Second component carries ~no variance.
174        assert!(m.explained_variance[1] < 1e-9);
175    }
176
177    #[test]
178    fn transform_projects_and_decorrelates() {
179        let x = [1.0, 1.0, 2.0, 2.0, 3.0, 3.1, 4.0, 3.9, 5.0, 5.0];
180        let m = fit(&x, 5, 2).unwrap();
181        let scores = m.transform(&x, 5, 1).unwrap();
182        assert_eq!(scores.len(), 5);
183        // Scores are centered (mean ~0) because the data was centered.
184        let mean_score: f64 = scores.iter().sum::<f64>() / 5.0;
185        assert!(mean_score.abs() < 1e-9);
186    }
187
188    #[test]
189    fn total_explained_variance_equals_total_variance() {
190        let x = [2.0, 1.0, 4.0, 3.0, 6.0, 2.0, 8.0, 5.0, 10.0, 4.0];
191        let m = fit(&x, 5, 2).unwrap();
192        // Sum of eigenvalues == trace of covariance == sum of per-feature variances.
193        use crate::solvers::statistics::descriptive::variance;
194        let mut total_var = 0.0;
195        for j in 0..2 {
196            let col: Vec<f64> = (0..5).map(|i| x[i * 2 + j]).collect();
197            total_var += variance(&col, true).unwrap();
198        }
199        let sum_eig: f64 = m.explained_variance.iter().sum();
200        assert!(
201            (sum_eig - total_var).abs() < 1e-9,
202            "{sum_eig} vs {total_var}"
203        );
204    }
205
206    #[test]
207    fn guards() {
208        assert_eq!(
209            fit(&[1.0, 2.0], 1, 2).unwrap_err(),
210            LearningError::InsufficientData
211        );
212        assert_eq!(
213            fit(&[1.0, 2.0, 3.0], 2, 2).unwrap_err(),
214            LearningError::InvalidDimension
215        );
216    }
217}