qualia_core_db/solvers/learning/classification/
naive_bayes.rs1use 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#[derive(Debug, Clone)]
14pub struct GaussianNb {
15 pub classes: Vec<usize>,
17 log_priors: Vec<f64>,
19 means: Vec<f64>,
21 variances: Vec<f64>,
22 p: usize,
23}
24
25impl GaussianNb {
26 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 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 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 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 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 let x = [0.0, 0.0, 0.0, 0.0, 1.0, 1.0]; let y = [0, 0, 1];
119 let nb = GaussianNb::fit(&x, &y, 3, 2).unwrap();
120 assert_eq!(nb.classes.len(), 2);
121 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}