用 Rust 从头实现机器学习算法·分类③:朴素贝叶斯——条件独立假设下的极速分类器

82 分钟阅读 rust-ml-from-scratch · 11
Rust机器学习

用 Rust 从头实现机器学习算法·分类③:朴素贝叶斯——条件独立假设下的极速分类器

kNN 靠距离投票,朴素贝叶斯走另一条路:先给每个类别画像,再看新样本像谁。它是清单里”训练预测都极快、天然输出概率、可增量学习”的代言人,垃圾邮件过滤器就是它的经典战场。我们用一个人造的体育/科技新闻分类任务,从零实现完整的多项式 NB。

核心思想

贝叶斯公式翻过来用:想知道 P(y∣x)P(y \mid x),但 P(x∣y)P(x \mid y) 和 P(y)P(y) 更好算——

P(y∣x)∝P(y)⋅P(x∣y)P(y \mid x) \propto P(y) \cdot P(x \mid y)

“朴素”就朴素在一个大胆的假设上:特征之间条件独立,P(x∣y)=∏jP(xj∣y)P(x \mid y) = \prod_j P(x_j \mid y)。真实文本里”team”和”win”显然相关,假设几乎必错——但分类只需比较大小,各维度的偏差往往相互抵消(清单原话),实践中出奇地好用。文本的”特征”是词,文档是词袋(只管词频,不管语序)——顺序信息的牺牲换来了计数的简洁。

数学:计数、平滑与对数空间

训练即计数:每类维护每个词的累计出现次数 Nw,cN_{w,c} 与总词数 NcN_c,似然估计

P^(w∣c)=Nw,c+αNc+αV\hat{P}(w \mid c) = \frac{N_{w,c} + \alpha}{N_c + \alpha V}

α\alpha 是拉普拉斯平滑:未见过的词(Nw,c=0N_{w,c}=0)不能给出零概率——一篇含生词的体育稿不至于被判成科技稿。VV 是词表大小。

预测时直接乘 ∏P^(w∣c)\prod \hat{P}(w \mid c) 会下溢(几百个小数相乘),所以搬到对数空间:arg⁡max⁡c[ln⁡P(c)+∑wcountw⋅ln⁡P^(w∣c)]\arg\max_c \big[\ln P(c) + \sum_w \text{count}_w \cdot \ln \hat{P}(w \mid c)\big]——乘法变加法,连平滑都自动包含。

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 由平滑保证 >0> 0,对数空间安全。
  • 文档转词频用了 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 构成有趣的镜像:一个把知识全塞进训练后的参数(计数表),一个把知识留在数据里随查随算;一个需要线性特征假设,一个需要距离假设;一个生成式(对每类建模),一个判别式邻居投票。下一个分类器继续判别式路线,但放弃”线性/近邻”,改问”怎样切一刀信息量最大”——分类④:决策树,用熵和信息增益递归切分特征空间。