stats_claw/algorithms/classification/gaussian.rs
1//! Gaussian Naive Bayes, matching `sklearn.naive_bayes.GaussianNB`.
2
3use super::{argmax, classification_result_from, normalize_log, sorted_unique, validate_dims};
4use crate::algorithms::classification::ClassificationResult;
5use crate::algorithms::count_to_f64;
6use crate::error::{Error, Result};
7
8/// `ln(2π)`, precomputed for the per-feature Gaussian log density.
9const LN_2PI: f64 = 1.837_877_066_409_345_5;
10
11/// `scikit-learn`'s default `GaussianNB` variance-smoothing coefficient.
12const VAR_SMOOTHING: f64 = 1e-9;
13
14/// Fitted Gaussian Naive Bayes model.
15///
16/// Holds the sorted class labels, the per-class log priors, and the per-class
17/// per-feature Gaussian mean and (smoothed) variance. Construct one with
18/// [`gaussian_nb_fit`].
19#[derive(Debug, Clone)]
20pub struct GaussianNbModel {
21 /// Sorted, de-duplicated class labels, in scoring/column order.
22 classes: Vec<usize>,
23 /// `ln P(class)` per class, in `classes` order.
24 log_priors: Vec<f64>,
25 /// Per-class per-feature MLE means (`means[c][j]`).
26 means: Vec<Vec<f64>>,
27 /// Per-class per-feature smoothed variances (`variances[c][j]`).
28 variances: Vec<Vec<f64>>,
29 /// Number of features expected in every input row.
30 n_features: usize,
31}
32
33/// Returns each feature's biased (`/n`) variance over all rows of `x`.
34fn feature_variances(x: &[Vec<f64>], n_features: usize) -> Vec<f64> {
35 let n = count_to_f64(x.len());
36 let mut mean = vec![0.0_f64; n_features];
37 for row in x {
38 for (m, &v) in mean.iter_mut().zip(row.iter()) {
39 *m += v;
40 }
41 }
42 for m in &mut mean {
43 *m /= n;
44 }
45 let mut var = vec![0.0_f64; n_features];
46 for row in x {
47 for ((acc, &v), &m) in var.iter_mut().zip(row.iter()).zip(mean.iter()) {
48 let diff = v - m;
49 *acc = diff.mul_add(diff, *acc);
50 }
51 }
52 for acc in &mut var {
53 *acc /= n;
54 }
55 var
56}
57
58/// Fits a Gaussian Naive Bayes model to continuous features.
59///
60/// Reproduces `sklearn.naive_bayes.GaussianNB`: per-class per-feature means and
61/// biased (`/n`) variance MLEs, a shared variance floor
62/// `var_smoothing = 1e-9 · max feature variance` added to every variance, and
63/// class-frequency priors. Scoring (see [`GaussianNbModel::predict`]) is done
64/// entirely in log space.
65///
66/// # Arguments
67///
68/// * `x` — training design matrix, one inner `Vec` of feature values per
69/// observation; every row must share a length.
70/// * `y` — one class label per observation.
71///
72/// # Returns
73///
74/// The fitted [`GaussianNbModel`].
75///
76/// # Errors
77///
78/// * [`Error::EmptyInput`] if `x` or `y` is empty.
79/// * [`Error::InvalidInput`] on an `x`/`y` length mismatch, zero features, or
80/// ragged rows.
81/// * [`Error::InsufficientData`] if fewer than two distinct classes are present.
82///
83/// # Examples
84///
85/// ```
86/// use stats_claw::algorithms::classification::naive_bayes::gaussian_nb_fit;
87///
88/// let x = vec![vec![0.0], vec![0.5], vec![9.0], vec![9.5]];
89/// let y = vec![0, 0, 1, 1];
90/// let model = gaussian_nb_fit(&x, &y)?;
91/// assert_eq!(model.predict(&[vec![0.1], vec![9.2]])?, vec![0, 1]);
92/// # Ok::<(), stats_claw::error::Error>(())
93/// ```
94pub fn gaussian_nb_fit(x: &[Vec<f64>], y: &[usize]) -> Result<GaussianNbModel> {
95 let n_features = validate_dims(x, y)?;
96 let classes = sorted_unique(y);
97 if classes.len() < 2 {
98 return Err(Error::InsufficientData);
99 }
100 let n_total = count_to_f64(y.len());
101 let global_var = feature_variances(x, n_features);
102 let max_var = global_var.iter().copied().fold(0.0_f64, f64::max);
103 let epsilon = VAR_SMOOTHING * max_var;
104
105 let mut log_priors = Vec::with_capacity(classes.len());
106 let mut means = Vec::with_capacity(classes.len());
107 let mut variances = Vec::with_capacity(classes.len());
108 for &cls in &classes {
109 let rows: Vec<&Vec<f64>> = x
110 .iter()
111 .zip(y)
112 .filter_map(|(row, &label)| (label == cls).then_some(row))
113 .collect();
114 let n_c = count_to_f64(rows.len());
115 let mut mean = vec![0.0_f64; n_features];
116 for row in &rows {
117 for (m, &v) in mean.iter_mut().zip(row.iter()) {
118 *m += v;
119 }
120 }
121 for m in &mut mean {
122 *m /= n_c;
123 }
124 let mut var = vec![0.0_f64; n_features];
125 for row in &rows {
126 for ((acc, &v), &m) in var.iter_mut().zip(row.iter()).zip(mean.iter()) {
127 let diff = v - m;
128 *acc = diff.mul_add(diff, *acc);
129 }
130 }
131 for acc in &mut var {
132 *acc = (*acc / n_c).max(0.0) + epsilon;
133 }
134 log_priors.push((n_c / n_total).ln());
135 means.push(mean);
136 variances.push(var);
137 }
138 Ok(GaussianNbModel {
139 classes,
140 log_priors,
141 means,
142 variances,
143 n_features,
144 })
145}
146
147impl GaussianNbModel {
148 /// Returns the sorted class labels, in [`Self::predict_log_proba`] column
149 /// order.
150 #[must_use]
151 pub fn classes(&self) -> &[usize] {
152 &self.classes
153 }
154
155 /// Computes `ln P(class) + Σⱼ ln N(xⱼ; μ, σ²)` for every sample and class.
156 fn joint_log_likelihoods(&self, x: &[Vec<f64>]) -> Result<Vec<Vec<f64>>> {
157 let mut out = Vec::with_capacity(x.len());
158 for row in x {
159 if row.len() != self.n_features {
160 return Err(Error::InvalidInput(
161 "sample feature count differs from the fitted model".to_owned(),
162 ));
163 }
164 let mut scores = Vec::with_capacity(self.classes.len());
165 for ((&log_prior, mean), var) in self
166 .log_priors
167 .iter()
168 .zip(self.means.iter())
169 .zip(self.variances.iter())
170 {
171 let mut ll = log_prior;
172 for ((&value, &m), &v) in row.iter().zip(mean.iter()).zip(var.iter()) {
173 let diff = value - m;
174 let quad = diff.mul_add(diff, 0.0) / (2.0 * v);
175 ll += (-0.5_f64).mul_add(LN_2PI + v.ln(), -quad);
176 }
177 scores.push(ll);
178 }
179 out.push(scores);
180 }
181 Ok(out)
182 }
183
184 /// Predicts the class label for every sample in `x` (ties break low).
185 ///
186 /// # Arguments
187 ///
188 /// * `x` — rows to classify; each must have the fitted feature count.
189 ///
190 /// # Returns
191 ///
192 /// One predicted class label per input row, in input order.
193 ///
194 /// # Errors
195 ///
196 /// [`Error::InvalidInput`] if any row's length differs from the fitted
197 /// feature count.
198 ///
199 /// # Examples
200 ///
201 /// ```
202 /// use stats_claw::algorithms::classification::naive_bayes::gaussian_nb_fit;
203 ///
204 /// let x = vec![vec![0.0], vec![0.5], vec![9.0], vec![9.5]];
205 /// let y = vec![0, 0, 1, 1];
206 /// let model = gaussian_nb_fit(&x, &y)?;
207 /// assert_eq!(model.predict(&[vec![9.1]])?, vec![1]);
208 /// # Ok::<(), stats_claw::error::Error>(())
209 /// ```
210 pub fn predict(&self, x: &[Vec<f64>]) -> Result<Vec<usize>> {
211 let joints = self.joint_log_likelihoods(x)?;
212 Ok(joints
213 .iter()
214 .map(|row| self.classes.get(argmax(row)).copied().unwrap_or(0))
215 .collect())
216 }
217
218 /// Predicts the normalized log posteriors for every sample in `x`.
219 ///
220 /// Each returned row has one entry per class (in [`Self::classes`] order)
221 /// whose exponentials sum to 1.
222 ///
223 /// # Arguments
224 ///
225 /// * `x` — rows to score; each must have the fitted feature count.
226 ///
227 /// # Returns
228 ///
229 /// One log-posterior row per input row.
230 ///
231 /// # Errors
232 ///
233 /// [`Error::InvalidInput`] if any row's length differs from the fitted
234 /// feature count.
235 ///
236 /// # Examples
237 ///
238 /// ```
239 /// use stats_claw::algorithms::classification::naive_bayes::gaussian_nb_fit;
240 ///
241 /// let x = vec![vec![0.0], vec![0.5], vec![9.0], vec![9.5]];
242 /// let y = vec![0, 0, 1, 1];
243 /// let model = gaussian_nb_fit(&x, &y)?;
244 /// let lp = model.predict_log_proba(&[vec![0.1]])?;
245 /// let total: f64 = lp[0].iter().map(|v| v.exp()).sum();
246 /// assert!((total - 1.0).abs() < 1e-12, "posteriors summed to {total}");
247 /// # Ok::<(), stats_claw::error::Error>(())
248 /// ```
249 pub fn predict_log_proba(&self, x: &[Vec<f64>]) -> Result<Vec<Vec<f64>>> {
250 let joints = self.joint_log_likelihoods(x)?;
251 Ok(joints.iter().map(|row| normalize_log(row)).collect())
252 }
253
254 /// Builds a populated [`ClassificationResult`] from predictions on
255 /// `(x, y_true)`: accuracy plus macro-averaged precision / recall / F1.
256 ///
257 /// # Arguments
258 ///
259 /// * `x` — evaluation design matrix.
260 /// * `y_true` — the true class labels, one per row of `x`.
261 ///
262 /// # Returns
263 ///
264 /// A [`ClassificationResult`] scored on `(x, y_true)`.
265 ///
266 /// # Errors
267 ///
268 /// * [`Error::InvalidInput`] on an `x`/`y_true` length mismatch or a row
269 /// whose feature count differs from the fitted model.
270 /// * [`Error::EmptyInput`] if `x` is empty.
271 ///
272 /// # Examples
273 ///
274 /// ```
275 /// use stats_claw::algorithms::classification::naive_bayes::gaussian_nb_fit;
276 ///
277 /// let x = vec![vec![0.0], vec![0.5], vec![9.0], vec![9.5]];
278 /// let y = vec![0, 0, 1, 1];
279 /// let model = gaussian_nb_fit(&x, &y)?;
280 /// let result = model.classification_result(&x, &y)?;
281 /// assert!((result.accuracy - 1.0).abs() < 1e-12, "accuracy {}", result.accuracy);
282 /// # Ok::<(), stats_claw::error::Error>(())
283 /// ```
284 pub fn classification_result(
285 &self,
286 x: &[Vec<f64>],
287 y_true: &[usize],
288 ) -> Result<ClassificationResult> {
289 let predictions = self.predict(x)?;
290 classification_result_from(&self.classes, &predictions, y_true, "Gaussian Naive Bayes")
291 }
292}