← Gym/Gaussian Naive Bayes from Scratch
00:00/ 30 min

🧭 Do not search for the first 15 minutes. When stuck: re-read the requirements → define I/O → choose the data structure → trace a small example by hand → write code.

Build a classifier out of a single (mostly false) assumption: that the features are independent. What makes this model distinctive is that training is one pass of aggregation — no iterations, no learning rate.

Input

python
X  # (n, d) real-valued features
y  # (n,)  integer labels

If X arrives one-dimensional, treat it as (n, 1).

Training

Handle classes in ascending sorted order. For class c:

  • prior P(c)=nc/nP(c) = n_c / n
  • per-feature mean μc,i\mu_{c,i}
  • per-feature population variance (ddof=0, i.e. divide by n_c) σc,i2\sigma^2_{c,i}

A zero variance would divide by zero during prediction, so add 1e-6 to every variance.

The usual choice is 1e-9, and there is a reason it is 1e-6 here. Rounding to 6 decimal places — as the return rule below requires — turns 1e-9 straight into 0.0, the smoothing disappears, and prediction blows up on log(0). The constant has to survive the rounding.

Prediction

Pick the class with the largest log posterior.

log⁡P(c∣x)∝log⁡P(c)+∑ilog⁡N(xi;μc,i,σc,i2)\log P(c \mid x) \propto \log P(c) + \sum_i \log \mathcal{N}(x_i; \mu_{c,i}, \sigma^2_{c,i})

log⁡N(x;μ,σ2)=−12log⁡(2πσ2)−(x−μ)22σ2\log \mathcal{N}(x; \mu, \sigma^2) = -\frac{1}{2}\log(2\pi\sigma^2) - \frac{(x - \mu)^2}{2\sigma^2}

On a tie, pick the smaller class value.

Return

fit_gaussian_nb(X, y) → a dictionary with these four keys. Round every float to 6 decimal places.

python
{"classes": [...], "priors": [...], "means": [[...], ...], "vars": [[...], ...]}

Level 1 · Per-class aggregation

Implement fit_gaussian_nb(X, y).

  • classes is the class list in ascending sorted order.
  • priors[j] = n_c / n
  • means[j][i] and vars[j][i] — per-class, per-feature mean and population variance (ddof=0). numpy's .var() already defaults to ddof=0, so it works as-is.
  • Add 1e-6 to every variance — with 1e-9 the smoothing vanishes under round(v, 6).
  • Round every float with round(v, 6).