Skip to main content

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}