Skip to main content

qualia_core_db/solvers/learning/classification/
naive_bayes.rs

1//! Gaussian Naive Bayes (ISL ch 4.4.4) — a generative classifier assuming the
2//! features are conditionally independent Gaussians given the class. Fit per-class
3//! priors and per-feature mean/variance (reusing `statistics::descriptive`);
4//! classify by the argmax log-posterior. Kernel-class `Reduction`.
5
6use crate::solvers::learning::LearningError;
7use crate::solvers::statistics::descriptive::{mean, variance};
8
9const VAR_FLOOR: f64 = 1e-9;
10const LN_2PI: f64 = 1.837_877_066_409_345_6;
11
12/// A fitted Gaussian naive-Bayes classifier.
13#[derive(Debug, Clone)]
14pub struct GaussianNb {
15    /// Distinct class labels, in ascending order.
16    pub classes: Vec<usize>,
17    /// Log class priors, aligned with `classes`.
18    log_priors: Vec<f64>,
19    /// Per-class per-feature means / variances, `n_classes × p` row-major.
20    means: Vec<f64>,
21    variances: Vec<f64>,
22    p: usize,
23}
24
25impl GaussianNb {
26    /// Fit from a row-major `n × p` matrix and integer labels.
27    pub fn fit(x: &[f64], y: &[usize], n: usize, p: usize) -> Result<Self, LearningError> {
28        if n == 0 || p == 0 || x.len() != n * p || y.len() != n {
29            return Err(LearningError::InvalidDimension);
30        }
31        let mut classes: Vec<usize> = y.to_vec();
32        classes.sort_unstable();
33        classes.dedup();
34        let c = classes.len();
35
36        let mut log_priors = vec![0.0; c];
37        let mut means = vec![0.0; c * p];
38        let mut variances = vec![0.0; c * p];
39
40        for (ci, &cls) in classes.iter().enumerate() {
41            let rows: Vec<usize> = (0..n).filter(|&i| y[i] == cls).collect();
42            if rows.is_empty() {
43                return Err(LearningError::InsufficientData);
44            }
45            log_priors[ci] = (rows.len() as f64 / n as f64).ln();
46            let mut colbuf = vec![0.0; rows.len()];
47            for j in 0..p {
48                for (t, &i) in rows.iter().enumerate() {
49                    colbuf[t] = x[i * p + j];
50                }
51                means[ci * p + j] = mean(&colbuf).ok_or(LearningError::InsufficientData)?;
52                // Population variance per class (NaN-safe via floor for n_c==1).
53                let v = variance(&colbuf, false).unwrap_or(0.0);
54                variances[ci * p + j] = v.max(VAR_FLOOR);
55            }
56        }
57
58        Ok(Self {
59            classes,
60            log_priors,
61            means,
62            variances,
63            p,
64        })
65    }
66
67    /// Log-posterior (up to the shared evidence constant) of `q` for class index `ci`.
68    fn log_score(&self, q: &[f64], ci: usize) -> f64 {
69        let mut s = self.log_priors[ci];
70        for j in 0..self.p {
71            let v = self.variances[ci * self.p + j];
72            let d = q[j] - self.means[ci * self.p + j];
73            s += -0.5 * (LN_2PI + v.ln() + d * d / v);
74        }
75        s
76    }
77
78    /// Predict the most probable class label for one query row.
79    pub fn predict_row(&self, q: &[f64]) -> usize {
80        let mut best = 0;
81        let mut best_s = f64::NEG_INFINITY;
82        for ci in 0..self.classes.len() {
83            let s = self.log_score(q, ci);
84            if s > best_s {
85                best_s = s;
86                best = ci;
87            }
88        }
89        self.classes[best]
90    }
91
92    pub fn predict(&self, x: &[f64], m: usize) -> Vec<usize> {
93        (0..m)
94            .map(|i| self.predict_row(&x[i * self.p..(i + 1) * self.p]))
95            .collect()
96    }
97}
98
99#[cfg(test)]
100mod tests {
101    use super::*;
102
103    #[test]
104    fn separates_two_gaussian_classes() {
105        // Class 0 around (0,0), class 1 around (5,5).
106        let x = [0.0, 0.1, 0.2, -0.1, -0.1, 0.0, 5.0, 5.1, 4.9, 5.0, 5.2, 4.8];
107        let y = [0, 0, 0, 1, 1, 1];
108        let nb = GaussianNb::fit(&x, &y, 6, 2).unwrap();
109        assert_eq!(nb.classes, vec![0, 1]);
110        assert_eq!(nb.predict_row(&[0.05, 0.05]), 0);
111        assert_eq!(nb.predict_row(&[5.05, 4.95]), 1);
112    }
113
114    #[test]
115    fn respects_priors() {
116        // Heavily imbalanced classes; a borderline point leans to the majority.
117        let x = [0.0, 0.0, 0.0, 0.0, 1.0, 1.0]; // 3 of class 0 (origin) ... reuse coords
118        let y = [0, 0, 1];
119        let nb = GaussianNb::fit(&x, &y, 3, 2).unwrap();
120        assert_eq!(nb.classes.len(), 2);
121        // Prior for class 0 (2/3) > class 1 (1/3).
122        assert!(nb.log_priors[0] > nb.log_priors[1]);
123    }
124
125    #[test]
126    fn guards() {
127        assert_eq!(
128            GaussianNb::fit(&[1.0, 2.0, 3.0], &[0, 1], 2, 2).unwrap_err(),
129            LearningError::InvalidDimension
130        );
131    }
132}