用 Rust 从头实现机器学习算法·分类③:朴素贝叶斯——条件独立假设下的极速分类器
用 Rust 从头实现机器学习算法·分类③:朴素贝叶斯——条件独立假设下的极速分类器
kNN 靠距离投票,朴素贝叶斯走另一条路:先给每个类别画像,再看新样本像谁。它是清单里”训练预测都极快、天然输出概率、可增量学习”的代言人,垃圾邮件过滤器就是它的经典战场。我们用一个人造的体育/科技新闻分类任务,从零实现完整的多项式 NB。
核心思想
贝叶斯公式翻过来用:想知道 ,但 和 更好算——
“朴素”就朴素在一个大胆的假设上:特征之间条件独立,。真实文本里”team”和”win”显然相关,假设几乎必错——但分类只需比较大小,各维度的偏差往往相互抵消(清单原话),实践中出奇地好用。文本的”特征”是词,文档是词袋(只管词频,不管语序)——顺序信息的牺牲换来了计数的简洁。
数学:计数、平滑与对数空间
训练即计数:每类维护每个词的累计出现次数 与总词数 ,似然估计
是拉普拉斯平滑:未见过的词()不能给出零概率——一篇含生词的体育稿不至于被判成科技稿。 是词表大小。
预测时直接乘 会下溢(几百个小数相乘),所以搬到对数空间:——乘法变加法,连平滑都自动包含。
Rust 实现
// src/main.rs(段一:多项式朴素贝叶斯与实验主程序)
// 分类③:朴素贝叶斯——条件独立假设下的极速分类器
// 单文件、仅标准库。绘图器与系列前篇内联的是同一份实现。
// ===================== xorshift64(沿用 03 篇) =====================
struct XorShift(u64);
impl XorShift {
fn new(seed: u64) -> Self {
assert!(seed != 0, "种子不能为 0");
XorShift(seed)
}
fn next_u64(&mut self) -> u64 {
let mut x = self.0;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
self.0 = x;
x
}
fn next_f64(&mut self) -> f64 {
(self.next_u64() >> 11) as f64 / (1u64 << 53) as f64
}
}
// ===================== 朴素贝叶斯(多项式模型) =====================
/// 词表:6 个体育词、6 个科技词、2 个共有词。文档是词袋(词的索引序列)。
const VOCAB: [&str; 14] = [
"team", "goal", "coach", "match", "win", "player", // 体育类词
"chip", "code", "server", "data", "ai", "software", // 科技类词
"the", "new", // 共有词(不提供区分度)
];
const SPORTS: usize = 0;
const TECH: usize = 1;
/// 两类各自的真实词分布(生成模拟数据用)。
fn true_profile(class: usize) -> [f64; 14] {
if class == SPORTS {
[0.20, 0.18, 0.12, 0.15, 0.15, 0.12, 0.01, 0.01, 0.02, 0.02, 0.005, 0.01, 0.04, 0.03]
} else {
[0.01, 0.01, 0.02, 0.02, 0.005, 0.01, 0.20, 0.18, 0.12, 0.15, 0.15, 0.12, 0.04, 0.03]
}
}
/// 按类别分布采样一篇文档(8~12 个词)。
fn gen_doc(rng: &mut XorShift, class: usize) -> Vec<usize> {
let p = true_profile(class);
let mut cum = [0.0f64; 14];
let mut acc = 0.0;
for i in 0..14 {
acc += p[i];
cum[i] = acc;
}
let len = 8 + (rng.next_f64() * 5.0) as usize;
(0..len)
.map(|_| {
let u = rng.next_f64();
cum.iter().position(|&c| u <= c).unwrap_or(13)
})
.collect()
}
/// 多项式朴素贝叶斯:每类维护词频计数与文档数。
struct NaiveBayes {
word_counts: [[usize; 14]; 2], // 每类每个词的累计出现次数
token_totals: [usize; 2], // 每类总词数
doc_counts: [usize; 2], // 每类文档数
}
impl NaiveBayes {
fn new() -> Self {
NaiveBayes {
word_counts: [[0; 14]; 2],
token_totals: [0; 2],
doc_counts: [0; 2],
}
}
/// 增量学习:来一篇文档,更新一次计数——没有"重新训练"的概念。
fn update(&mut self, doc: &[usize], class: usize) {
for &w in doc {
self.word_counts[class][w] += 1;
self.token_totals[class] += 1;
}
self.doc_counts[class] += 1;
}
/// 对数后验:log P(class) + Σ count_w · log P(w | class),拉普拉斯平滑 α=1。
fn log_posterior(&self, doc: &[usize], class: usize) -> f64 {
let n_docs = (self.doc_counts[0] + self.doc_counts[1]) as f64;
let log_prior = (self.doc_counts[class] as f64 + 1.0).ln() - (n_docs + 2.0).ln();
let v = VOCAB.len() as f64;
let log_norm = (self.token_totals[class] as f64 + v).ln();
// 文档转词频(fold 的典型用例)
let freq: Vec<usize> = doc.iter().fold(vec![0usize; 14], |mut f, &w| {
f[w] += 1;
f
});
let log_lik: f64 = freq
.iter()
.enumerate()
.map(|(w, &c)| {
let p = (self.word_counts[class][w] as f64 + 1.0) / log_norm.exp();
c as f64 * p.ln()
})
.sum();
log_prior + log_lik
}
fn predict(&self, doc: &[usize]) -> usize {
let a = self.log_posterior(doc, SPORTS);
let b = self.log_posterior(doc, TECH);
if a >= b { SPORTS } else { TECH }
}
}
fn main() {
let out_dir = "../../../frontend/public/images/series/rust-ml-11-naive-bayes";
std::fs::create_dir_all(out_dir).expect("创建输出目录失败");
// ---- 1. 数据:每类 150 篇文档,前 100 篇训练,后 50 篇测试 ----
let mut rng = XorShift::new(42);
let mut train: Vec<(Vec<usize>, usize)> = Vec::new();
let mut test: Vec<(Vec<usize>, usize)> = Vec::new();
for class in [SPORTS, TECH] {
for i in 0..150 {
let doc = gen_doc(&mut rng, class);
if i < 100 {
train.push((doc, class));
} else {
test.push((doc, class));
}
}
}
println!("train = {} 篇(每类 100),test = {} 篇(每类 50)", train.len(), test.len());
// ---- 2. 逐篇增量训练,顺带记录学习曲线 ----
let mut model = NaiveBayes::new();
let mut acc_curve = vec![(0.0, 0.5)];
for (i, (doc, class)) in train.iter().enumerate() {
model.update(doc, *class);
if (i + 1) % 20 == 0 {
let acc = test_acc(&model, &test);
acc_curve.push(((i + 1) as f64, acc));
}
}
// ---- 3. 评估 ----
let acc = test_acc(&model, &test);
let (mut tp, mut tn, mut fp, mut fn_) = (0, 0, 0, 0);
for (doc, class) in &test {
let pred = model.predict(doc);
let class = *class;
match (pred, class) {
(SPORTS, SPORTS) => tp += 1,
(TECH, TECH) => tn += 1,
(SPORTS, TECH) => fp += 1,
(TECH, SPORTS) => fn_ += 1,
_ => unreachable!(),
}
}
println!("test acc = {:.3}", acc);
println!("混淆矩阵: 预测体育 预测科技");
println!(" 真实体育 ({} 篇) {:>6} {:>6}", 50, tp, fn_);
println!(" 真实科技 ({} 篇) {:>6} {:>6}", 50, fp, tn);
// ---- 4. 泄密词排行榜:每词的对数似然比 ----
println!("\n[词的对数似然比 log P(w|体育) − log P(w|科技)](绝对值大 = 泄密多)");
let v = VOCAB.len() as f64;
let mut ratios: Vec<(f64, usize)> = (0..14)
.map(|w| {
let ls = ((model.word_counts[SPORTS][w] as f64 + 1.0) / (model.token_totals[SPORTS] as f64 + v)).ln();
let lt = ((model.word_counts[TECH][w] as f64 + 1.0) / (model.token_totals[TECH] as f64 + v)).ln();
(ls - lt, w)
})
.collect();
ratios.sort_by(|a, b| b.0.abs().partial_cmp(&a.0.abs()).unwrap());
for (r, w) in ratios.iter().take(8) {
println!(" {:<9} {:+.3} {}", VOCAB[*w], r, if *r > 0.0 { "→ 体育" } else { "→ 科技" });
}
let ratios_full: Vec<(f64, f64)> = (0..14).map(|w| (w as f64, ratios.iter().find(|(_, x)| *x == w).unwrap().0)).collect();
let mut c = Canvas::new(560.0, 320.0, -0.5, 13.5, -3.0, 3.0);
c.axes("词表索引", "log 似然比");
let zero: Vec<(f64, f64)> = vec![(-0.5, 0.0), (13.5, 0.0)];
c.polyline(&zero, "#16161d", 1.0);
let pos: Vec<(f64, f64)> = ratios_full.iter().filter(|(_, r)| *r >= 0.0).map(|(w, r)| (*w, *r)).collect();
let neg: Vec<(f64, f64)> = ratios_full.iter().filter(|(_, r)| *r < 0.0).map(|(w, r)| (*w, *r)).collect();
c.dots(&pos, PALETTE[0], 5.0);
c.dots(&neg, PALETTE[1], 5.0);
c.legend(&[("偏向体育", PALETTE[0]), ("偏向科技", PALETTE[1])]);
let p1 = format!("{out_dir}/word-logratio.svg");
c.save(&p1);
// ---- 5. 增量学习曲线 ----
let mut c = Canvas::new(560.0, 320.0, 0.0, 200.0, 0.5, 1.02);
c.axes("已见训练文档数", "test acc");
c.polyline(&acc_curve, PALETTE[0], 2.0);
let p2 = format!("{out_dir}/learning-curve.svg");
c.save(&p2);
// ---- 6. 手写两篇新文档做演示 ----
println!("\n[演示] 新文档预测:");
let demo_sports: Vec<usize> = ["team", "coach", "win", "match", "player", "goal"].iter().map(|w| VOCAB.iter().position(|x| x == w).unwrap()).collect();
let demo_tech: Vec<usize> = ["chip", "server", "data", "ai", "software", "code"].iter().map(|w| VOCAB.iter().position(|x| x == w).unwrap()).collect();
for (name, doc) in [("team coach win match player goal", &demo_sports), ("chip server data ai software code", &demo_tech)] {
let c_pred = model.predict(doc);
println!(" \"{}\" → {}", name, if c_pred == SPORTS { "体育" } else { "科技" });
}
println!("图已写入:{p1}、{p2}");
}
fn test_acc(model: &NaiveBayes, test: &[(Vec<usize>, usize)]) -> f64 {
let correct = test.iter().filter(|(doc, class)| model.predict(doc) == *class).count();
correct as f64 / test.len() as f64
}
// ===================== 迷你 SVG 绘图器(与系列前篇一致) =====================
(与系列前篇一致) =====================
/// 系列默认配色:朱橙 / 青绿 / 蓝 / 琥珀(取自站点设计令牌)。
const PALETTE: [&str; 4] = ["#d6491f", "#0f766e", "#2563eb", "#d97706"];
struct Canvas {
w: f64,
h: f64,
xmin: f64,
xmax: f64,
ymin: f64,
ymax: f64,
pad_l: f64,
pad_r: f64,
pad_t: f64,
pad_b: f64,
body: String,
}
impl Canvas {
fn new(w: f64, h: f64, xmin: f64, xmax: f64, ymin: f64, ymax: f64) -> Canvas {
let (xmin, xmax) = spread(xmin, xmax);
let (ymin, ymax) = spread(ymin, ymax);
Canvas {
w,
h,
xmin,
xmax,
ymin,
ymax,
pad_l: 52.0,
pad_r: 16.0,
pad_t: 16.0,
pad_b: 40.0,
body: String::new(),
}
}
fn px(&self, x: f64) -> f64 {
self.pad_l + (x - self.xmin) / (self.xmax - self.xmin) * (self.w - self.pad_l - self.pad_r)
}
fn py(&self, y: f64) -> f64 {
self.h - self.pad_b - (y - self.ymin) / (self.ymax - self.ymin) * (self.h - self.pad_t - self.pad_b)
}
fn polyline(&mut self, pts: &[(f64, f64)], color: &str, width: f64) {
let mut d = String::new();
for (i, (x, y)) in pts.iter().enumerate() {
let _ = write!(d, "{}{:.1},{:.1}", if i == 0 { "M" } else { "L" }, self.px(*x), self.py(*y));
}
let _ = write!(
self.body,
r##"<path d="{d}" fill="none" stroke="{color}" stroke-width="{width}" stroke-linejoin="round"/>"##
);
}
fn dots(&mut self, pts: &[(f64, f64)], color: &str, r: f64) {
for (x, y) in pts {
let _ = write!(
self.body,
r##"<circle cx="{:.1}" cy="{:.1}" r="{r}" fill="{color}" fill-opacity="0.65"/>"##,
self.px(*x),
self.py(*y)
);
}
}
/// 置信带:upper/lower 两条曲线围成的区域,半透明填充。
#[allow(dead_code)] // 本篇未用,保持与系列其他篇的绘图器逐字一致
fn band(&mut self, upper: &[(f64, f64)], lower: &[(f64, f64)], color: &str) {
let mut d = String::new();
for (i, (x, y)) in upper.iter().enumerate() {
let _ = write!(d, "{}{:.1},{:.1}", if i == 0 { "M" } else { "L" }, self.px(*x), self.py(*y));
}
for (x, y) in lower.iter().rev() {
let _ = write!(d, "L{:.1},{:.1}", self.px(*x), self.py(*y));
}
let _ = write!(
self.body,
r##"<path d="{d}Z" fill="{color}" fill-opacity="0.15" stroke="none"/>"##
);
}
fn axes(&mut self, xlabel: &str, ylabel: &str) {
let x0 = self.px(self.xmin);
let x1 = self.px(self.xmax);
let y0 = self.py(self.ymin);
let y1 = self.py(self.ymax);
for t in nice_ticks(self.xmin, self.xmax, 6) {
let x = self.px(t);
let _ = write!(
self.body,
r##"<line x1="{x:.1}" y1="{y0:.1}" x2="{x:.1}" y2="{y1:.1}" stroke="#16161d" stroke-opacity="0.08"/>"##
);
let _ = write!(
self.body,
r##"<text x="{x:.1}" y="{:.1}" text-anchor="middle" fill="#16161d" fill-opacity="0.45">{t:.6}</text>"##,
y0 + 16.0,
t = trim(t)
);
}
for t in nice_ticks(self.ymin, self.ymax, 6) {
let y = self.py(t);
let _ = write!(
self.body,
r##"<line x1="{x0:.1}" y1="{y:.1}" x2="{x1:.1}" y2="{y:.1}" stroke="#16161d" stroke-opacity="0.08"/>"##
);
let _ = write!(
self.body,
r##"<text x="{:.1}" y="{:.1}" text-anchor="end" fill="#16161d" fill-opacity="0.45">{t:.6}</text>"##,
x0 - 6.0,
y + 4.0,
t = trim(t)
);
}
let _ = write!(
self.body,
r##"<path d="M{x0:.1},{y0:.1}H{x1:.1}M{x0:.1},{y0:.1}V{y1:.1}" fill="none" stroke="#16161d" stroke-opacity="0.6"/>"##
);
let _ = write!(
self.body,
r##"<text x="{:.1}" y="{:.1}" text-anchor="middle" fill="#16161d" fill-opacity="0.7">{xlabel}</text>"##,
(x0 + x1) / 2.0,
self.h - 8.0
);
let _ = write!(
self.body,
r##"<text x="14" y="{:.1}" text-anchor="middle" transform="rotate(-90 14 {:.1})" fill="#16161d" fill-opacity="0.7">{ylabel}</text>"##,
(y0 + y1) / 2.0,
(y0 + y1) / 2.0
);
}
/// 图例:右上角依次画「色线 + 文字」。
fn legend(&mut self, entries: &[(&str, &str)]) {
let sample_w = 22.0;
let line_h = 18.0;
let x0 = self.w - self.pad_r - 78.0;
let x_text = x0 + sample_w + 6.0;
for (i, (label, color)) in entries.iter().enumerate() {
let y = self.pad_t + 14.0 + i as f64 * line_h;
let _ = write!(
self.body,
r##"<line x1="{x0:.1}" y1="{y:.1}" x2="{:.1}" y2="{y:.1}" stroke="{color}" stroke-width="2.5"/>"##,
x0 + sample_w
);
let _ = write!(
self.body,
r##"<text x="{x_text:.1}" y="{:.1}" fill="#16161d" fill-opacity="0.7">{label}</text>"##,
y + 4.0
);
}
}
fn save(&self, path: &str) {
let svg = format!(
r##"<svg xmlns="http://www.w3.org/2000/svg" viewBox="0 0 {w} {h}" font-family="'JetBrains Mono',monospace" font-size="11"><rect width="{w}" height="{h}" fill="#fbfaf7"/>{body}</svg>"##,
w = self.w,
h = self.h,
body = self.body
);
std::fs::write(path, svg).expect("写 SVG 失败");
}
}
fn trim(v: f64) -> String {
let s = format!("{v:.6}");
s.trim_end_matches('0').trim_end_matches('.').to_string()
}
fn spread(min: f64, max: f64) -> (f64, f64) {
if min == max {
(min - 1.0, max + 1.0)
} else {
let pad = (max - min) * 0.05;
(min - pad, max + pad)
}
}
fn nice_ticks(min: f64, max: f64, n: usize) -> Vec<f64> {
let raw = (max - min) / n.max(1) as f64;
if raw <= 0.0 || !raw.is_finite() {
return vec![min];
}
let exp = raw.abs().log10().floor() as i32;
let base = 10f64.powi(exp);
let frac = raw / base;
let step = if frac <= 1.0 {
base
} else if frac <= 2.0 {
2.0 * base
} else if frac <= 2.5 {
2.5 * base
} else if frac <= 5.0 {
5.0 * base
} else {
10.0 * base
};
let mut ticks = Vec::new();
let mut t = (min / step).ceil() * step;
while t <= max + 1e-9 {
ticks.push(if t.abs() < step * 1e-9 { 0.0 } else { t });
t += step;
}
ticks
}
实现要点:
NaiveBayes的全部状态是三个计数数组——没有权重、没有梯度,“模型”就是一张频次表。这也是它可增量学习的原因:update一篇文档只需累加。log_posterior里p.ln()的p由平滑保证 ,对数空间安全。- 文档转词频用了
fold:从空计数器出发,一个词一个词地”折叠”出直方图——计数问题的标准 idiomatic 写法。 - 生成数据的
r = 1.5·sqrt(U)和 kNN 篇同理:面积均匀采样,避免点挤在圆心(此处是词概率的累积分布采样,position做逆变换)。
// src/main.rs 段二:迷你 SVG 绘图器(与系列前篇同一份实现)
use std::fmt::Write as _;
// src/main.rs 末尾:单元测试(cargo test 运行)
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gen_doc_respects_profile() {
let mut rng = XorShift::new(7);
let doc = gen_doc(&mut rng, SPORTS);
assert!(!doc.is_empty() && doc.len() <= 12);
assert!(doc.iter().all(|&w| w < 14));
}
#[test]
fn unseen_word_never_crashes() {
// 平滑保证:训练集里没出现过的词(计数 0)也给非零概率
let mut model = NaiveBayes::new();
model.update(&[0, 1, 2], SPORTS);
model.update(&[6, 7, 8], TECH);
let doc = vec![13, 13, 13]; // "new" 未在训练中出现
let p = model.predict(&doc); // 不应 panic,也不应 NaN
assert!(p == SPORTS || p == TECH);
}
#[test]
fn predicts_class_by_content() {
let mut rng = XorShift::new(11);
let mut model = NaiveBayes::new();
for _ in 0..40 {
model.update(&gen_doc(&mut rng, SPORTS), SPORTS);
model.update(&gen_doc(&mut rng, TECH), TECH);
}
let sports_doc = vec![0, 1, 2, 3, 4, 5]; // 全是体育词
let tech_doc = vec![6, 7, 8, 9, 10, 11]; // 全是科技词
assert_eq!(model.predict(&sports_doc), SPORTS);
assert_eq!(model.predict(&tech_doc), TECH);
}
#[test]
fn fold_counts_word_frequency() {
let doc = vec![1, 1, 2, 3, 3, 3];
let freq: Vec<usize> = doc.iter().fold(vec![0usize; 14], |mut f, &w| {
f[w] += 1;
f
});
assert_eq!(freq[1], 2);
assert_eq!(freq[3], 3);
assert_eq!(freq[0], 0);
}
}
Rust 语法角:fold——把循环折叠成一个表达式
doc.iter().fold(vec![0usize; 14], |mut f, &w| { f[w] += 1; f }) 三要素:初始累加器(空直方图)、折叠闭包(来一个词加一个计数)、最终返回值。它把 Python 里”先建 dict、再 for 循环累加”的样板压缩成一行,且累加器的类型变化(比如想把词频折叠成总对数似然)同样适用。map/filter/sum 其实都是 fold 的特化。见《Rust 程序设计语言》ch13-02(迭代器)。
运行结果
cargo test(4 个用例:文档生成合法性、生词不 panic 也不 NaN、内容定类、fold 词频正确性)全部通过后,cargo run:
train = 200 篇(每类 100),test = 100 篇(每类 50)
test acc = 1.000
混淆矩阵: 预测体育 预测科技
真实体育 (50 篇) 50 0
真实科技 (50 篇) 0 50
[词的对数似然比 log P(w|体育) − log P(w|科技)](绝对值大 = 泄密多)
team +3.487 → 体育
goal +3.436 → 体育
chip -3.335 → 科技
code -3.253 → 科技
ai -2.860 → 科技
software -2.671 → 科技
win +2.553 → 体育
match +2.196 → 体育
[演示] 新文档预测:
"team coach win match player goal" → 体育
"chip server data ai software code" → 科技
图已写入:../../../frontend/public/images/series/rust-ml-11-naive-bayes/word-logratio.svg、../../../frontend/public/images/series/rust-ml-11-naive-bayes/learning-curve.svg
两张图:
怎么读这些数字和图
- 混淆矩阵满分是合成数据的礼物:两类词表几乎不相交,任务对 NB 来说太轻松。真实垃圾邮件的混淆矩阵会有非零的 off-diagonal——那正是平滑和先验发挥作用的地方。
- 泄密词榜是可解释性的范本:
team的对数似然比 +3.49 意味着”看到 team,体育的可信度凭空加 3.49 个对数单位”。词表混进的两共有词(the、new)比值趋近 0,安静地躺在零线附近——不提供区分度的特征自动失效,这是 NB 的隐藏美德。 - 学习曲线 20 篇就逼近满分:计数型模型收敛极快,小数据可用;且每 20 篇评估一次的代价只是增量更新——“在线学习”不需要任何框架支持。
- 手写演示两篇纯词文档全对:注意这是确定性推断(没有随机),训练完的模型就是一个可复现的查表计算器。
优缺点与适用场景
(抄清单 2.2 原文)
- 优点:训练预测都极快、小样本也可用、天然输出概率、可增量学习。
- 缺点:条件独立假设过强、概率估计粗糙(但分类排序常没问题)。
- 适用场景:文本分类、垃圾邮件过滤、高维稀疏数据、需要毫秒级响应的系统。
小结
NB 与 kNN 构成有趣的镜像:一个把知识全塞进训练后的参数(计数表),一个把知识留在数据里随查随算;一个需要线性特征假设,一个需要距离假设;一个生成式(对每类建模),一个判别式邻居投票。下一个分类器继续判别式路线,但放弃”线性/近邻”,改问”怎样切一刀信息量最大”——分类④:决策树,用熵和信息增益递归切分特征空间。