Skip to main content

qualia_core_db/solvers/learning/classification/
discriminant.rs

1//! Discriminant analysis (ISL ch 4.4) — LDA (shared covariance ⇒ linear boundary)
2//! and QDA (per-class covariance ⇒ quadratic boundary). Both are Gaussian
3//! generative classifiers; the covariance inverse / log-determinant come from
4//! `linear_algebra::cholesky` (no new solver). Kernel-class `DenseLinear`.
5
6use crate::solvers::learning::LearningError;
7use crate::solvers::linear_algebra::cholesky::{
8    cholesky_determinant, cholesky_factor, cholesky_solve,
9};
10
11/// Per-class means + priors + the index map, shared setup for LDA and QDA.
12struct ClassStats {
13    classes: Vec<usize>,
14    priors: Vec<f64>,
15    means: Vec<f64>, // n_classes × p
16    counts: Vec<usize>,
17}
18
19fn class_stats(x: &[f64], y: &[usize], n: usize, p: usize) -> Result<ClassStats, LearningError> {
20    if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
21        return Err(LearningError::InvalidDimension);
22    }
23    let mut classes: Vec<usize> = y.to_vec();
24    classes.sort_unstable();
25    classes.dedup();
26    let c = classes.len();
27    let mut priors = vec![0.0; c];
28    let mut means = vec![0.0; c * p];
29    let mut counts = vec![0usize; c];
30    let cls_idx = |cls: usize| classes.iter().position(|&v| v == cls).unwrap();
31    for i in 0..n {
32        let ci = cls_idx(y[i]);
33        counts[ci] += 1;
34        for j in 0..p {
35            means[ci * p + j] += x[i * p + j];
36        }
37    }
38    for ci in 0..c {
39        if counts[ci] == 0 {
40            return Err(LearningError::InsufficientData);
41        }
42        priors[ci] = counts[ci] as f64 / n as f64;
43        for j in 0..p {
44            means[ci * p + j] /= counts[ci] as f64;
45        }
46    }
47    Ok(ClassStats {
48        classes,
49        priors,
50        means,
51        counts,
52    })
53}
54
55// ── LDA ───────────────────────────────────────────────────────────────────────
56
57/// Linear Discriminant Analysis: one shared within-class covariance. The
58/// discriminant is linear, `δ_c(x) = xᵀwc + bc`.
59#[derive(Debug, Clone)]
60pub struct LdaModel {
61    pub classes: Vec<usize>,
62    w: Vec<f64>, // n_classes × p : Σ⁻¹ μ_c
63    b: Vec<f64>, // n_classes      : −½ μcᵀ Σ⁻¹ μc + ln π_c
64    p: usize,
65}
66
67impl LdaModel {
68    /// Fit LDA. Needs `n > n_classes` for a non-degenerate pooled covariance.
69    pub fn fit(x: &[f64], y: &[usize], n: usize, p: usize) -> Result<Self, LearningError> {
70        let cs = class_stats(x, y, n, p)?;
71        let c = cs.classes.len();
72        if n <= c {
73            return Err(LearningError::InsufficientData);
74        }
75        // Pooled within-class scatter → covariance Σ = S/(n−C).
76        let mut cov = vec![0.0; p * p];
77        let cls_idx = |cls: usize| cs.classes.iter().position(|&v| v == cls).unwrap();
78        for i in 0..n {
79            let ci = cls_idx(y[i]);
80            for a in 0..p {
81                let da = x[i * p + a] - cs.means[ci * p + a];
82                for bb in 0..p {
83                    let db = x[i * p + bb] - cs.means[ci * p + bb];
84                    cov[a * p + bb] += da * db;
85                }
86            }
87        }
88        let denom = (n - c) as f64;
89        for v in cov.iter_mut() {
90            *v /= denom;
91        }
92        let l = cholesky_of(&cov, p)?;
93
94        let mut w = vec![0.0; c * p];
95        let mut b = vec![0.0; c];
96        let mut mu = vec![0.0; p];
97        let mut wc = vec![0.0; p];
98        for ci in 0..c {
99            mu.copy_from_slice(&cs.means[ci * p..(ci + 1) * p]);
100            cholesky_solve(p, &l, &mu, &mut wc)?; // wc = Σ⁻¹ μ_c
101            w[ci * p..(ci + 1) * p].copy_from_slice(&wc);
102            let quad: f64 = mu.iter().zip(wc.iter()).map(|(m, w)| m * w).sum();
103            b[ci] = -0.5 * quad + cs.priors[ci].ln();
104        }
105        Ok(Self {
106            classes: cs.classes,
107            w,
108            b,
109            p,
110        })
111    }
112
113    pub fn predict_row(&self, q: &[f64]) -> usize {
114        let mut best = 0;
115        let mut best_s = f64::NEG_INFINITY;
116        for ci in 0..self.classes.len() {
117            let s: f64 = q
118                .iter()
119                .zip(&self.w[ci * self.p..(ci + 1) * self.p])
120                .map(|(x, w)| x * w)
121                .sum::<f64>()
122                + self.b[ci];
123            if s > best_s {
124                best_s = s;
125                best = ci;
126            }
127        }
128        self.classes[best]
129    }
130
131    pub fn predict(&self, x: &[f64], m: usize) -> Vec<usize> {
132        (0..m)
133            .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
134            .collect()
135    }
136}
137
138// ── QDA ───────────────────────────────────────────────────────────────────────
139
140/// Quadratic Discriminant Analysis: a separate covariance per class. The
141/// discriminant `δ_c(x) = −½ ln|Σc| − ½ (x−μc)ᵀ Σc⁻¹ (x−μc) + ln π_c`.
142#[derive(Debug, Clone)]
143pub struct QdaModel {
144    pub classes: Vec<usize>,
145    means: Vec<f64>,   // n_classes × p
146    chol: Vec<f64>,    // n_classes × p × p : Cholesky factors of Σ_c
147    log_det: Vec<f64>, // n_classes
148    log_prior: Vec<f64>,
149    p: usize,
150}
151
152impl QdaModel {
153    /// Fit QDA. Needs each class to have `> p` samples for a non-singular covariance.
154    pub fn fit(x: &[f64], y: &[usize], n: usize, p: usize) -> Result<Self, LearningError> {
155        let cs = class_stats(x, y, n, p)?;
156        let c = cs.classes.len();
157        let cls_idx = |cls: usize| cs.classes.iter().position(|&v| v == cls).unwrap();
158        let mut chol = vec![0.0; c * p * p];
159        let mut log_det = vec![0.0; c];
160        let mut log_prior = vec![0.0; c];
161        for ci in 0..c {
162            if cs.counts[ci] <= p {
163                return Err(LearningError::InsufficientData); // covariance would be singular
164            }
165            // Per-class covariance Σ_c = S_c/(n_c−1).
166            let mut cov = vec![0.0; p * p];
167            for i in 0..n {
168                if cls_idx(y[i]) != ci {
169                    continue;
170                }
171                for a in 0..p {
172                    let da = x[i * p + a] - cs.means[ci * p + a];
173                    for bb in 0..p {
174                        let db = x[i * p + bb] - cs.means[ci * p + bb];
175                        cov[a * p + bb] += da * db;
176                    }
177                }
178            }
179            let denom = (cs.counts[ci] - 1) as f64;
180            for v in cov.iter_mut() {
181                *v /= denom;
182            }
183            let l = cholesky_of(&cov, p)?;
184            log_det[ci] = cholesky_determinant(p, &l).max(1e-300).ln();
185            chol[ci * p * p..(ci + 1) * p * p].copy_from_slice(&l);
186            log_prior[ci] = cs.priors[ci].ln();
187        }
188        Ok(Self {
189            classes: cs.classes,
190            means: cs.means,
191            chol,
192            log_det,
193            log_prior,
194            p,
195        })
196    }
197
198    pub fn predict_row(&self, q: &[f64]) -> usize {
199        let p = self.p;
200        let mut diff = vec![0.0; p];
201        let mut sol = vec![0.0; p];
202        let mut best = 0;
203        let mut best_s = f64::NEG_INFINITY;
204        for ci in 0..self.classes.len() {
205            for j in 0..p {
206                diff[j] = q[j] - self.means[ci * p + j];
207            }
208            // Mahalanobis² = diffᵀ Σ⁻¹ diff via the Cholesky factor.
209            let l = &self.chol[ci * p * p..(ci + 1) * p * p];
210            if cholesky_solve(p, l, &diff, &mut sol).is_err() {
211                continue;
212            }
213            let maha: f64 = diff.iter().zip(sol.iter()).map(|(d, s)| d * s).sum();
214            let s = -0.5 * self.log_det[ci] - 0.5 * maha + self.log_prior[ci];
215            if s > best_s {
216                best_s = s;
217                best = ci;
218            }
219        }
220        self.classes[best]
221    }
222
223    pub fn predict(&self, x: &[f64], m: usize) -> Vec<usize> {
224        (0..m)
225            .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
226            .collect()
227    }
228}
229
230/// Cholesky factor of a `p×p` covariance, mapping a solver failure to `Singular`.
231fn cholesky_of(cov: &[f64], p: usize) -> Result<Vec<f64>, LearningError> {
232    let mut l = vec![0.0; p * p];
233    cholesky_factor(p, cov, &mut l).map_err(|_| LearningError::Singular)?;
234    Ok(l)
235}
236
237#[cfg(test)]
238mod tests {
239    use super::*;
240
241    // Two 2-D classes, well separated, enough points for per-class covariance.
242    fn data() -> (Vec<f64>, Vec<usize>, usize) {
243        let mut x = Vec::new();
244        let mut y = Vec::new();
245        // class 0 around (0,0)
246        for &(a, b) in &[
247            (0.0, 0.0),
248            (0.5, -0.3),
249            (-0.4, 0.2),
250            (0.2, 0.4),
251            (-0.3, -0.2),
252        ] {
253            x.push(a);
254            x.push(b);
255            y.push(0);
256        }
257        // class 1 around (5,5)
258        for &(a, b) in &[(5.0, 5.0), (5.4, 4.7), (4.6, 5.3), (5.1, 4.9), (4.8, 5.2)] {
259            x.push(a);
260            x.push(b);
261            y.push(1);
262        }
263        (x, y, 10)
264    }
265
266    #[test]
267    fn lda_separates_classes() {
268        let (x, y, n) = data();
269        let m = LdaModel::fit(&x, &y, n, 2).unwrap();
270        assert_eq!(m.classes, vec![0, 1]);
271        assert_eq!(m.predict_row(&[0.1, 0.1]), 0);
272        assert_eq!(m.predict_row(&[5.0, 5.0]), 1);
273        // Training accuracy is perfect on this separable set.
274        let preds = m.predict(&x, n);
275        assert!(preds.iter().zip(&y).all(|(a, b)| a == b));
276    }
277
278    #[test]
279    fn qda_separates_classes() {
280        let (x, y, n) = data();
281        let m = QdaModel::fit(&x, &y, n, 2).unwrap();
282        assert_eq!(m.predict_row(&[0.0, 0.0]), 0);
283        assert_eq!(m.predict_row(&[5.2, 4.8]), 1);
284        let preds = m.predict(&x, n);
285        assert!(preds.iter().zip(&y).all(|(a, b)| a == b));
286    }
287
288    #[test]
289    fn qda_fails_closed_on_too_few_per_class() {
290        // class 0 has 3 non-collinear points (fine); class 1 has only 2, p=2 →
291        // n_c ≤ p → its covariance would be singular → fail closed.
292        let x = [0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 5.0, 5.0, 6.0, 5.0];
293        let y = [0, 0, 0, 1, 1];
294        assert_eq!(
295            QdaModel::fit(&x, &y, 5, 2).unwrap_err(),
296            LearningError::InsufficientData
297        );
298    }
299}