用 Rust 从头实现机器学习算法·分类②:k 近邻——最懒的分类器

70 分钟阅读 rust-ml-from-scratch · 10
Rust机器学习

用 Rust 从头实现机器学习算法·分类②:k 近邻——最懒的分类器

逻辑回归用一条直线切分世界,前提是数据真的”线性可分”。本篇的数据一上来就打破这个前提——同心圆:类 0 住在半径 1.5 的圆盘里,类 1 住在 1.8~3.5 的圆环上。任何直线都无能为力,而 k 近邻(kNN)几乎不费吹灰之力。这是清单里”低维数据快速建立基线”的典型案例。

核心思想

kNN 是机器学习里最”懒”的算法——没有训练阶段。它的全部假设只有一句谚语:近朱者赤。预测新样本时,在训练集里找离它最近的 kk 个点,多数投票决定类别(回归问题则取平均)。清单的三个要素:距离度量(本文用欧氏)、kk 的选择、投票规则。

两个细节决定成败:

  • kk 是偏差-方差的旋钮:k=1k=1 时决策边界扭出无数小岛(每个训练点都是一座孤岛,过拟合);kk 增大边界趋于平滑(方差下降),但太大时连真实形状都抹掉(偏差上升)——注意这个 k 与 k 均值聚类的 k 毫无关系(清单附录 A.3 专门提醒);
  • 量纲敏感:距离由数值大的特征主导,用 kNN 前必须标准化(04 篇的纪律在此是硬要求)。

数学:距离与多数投票

欧氏距离:d(x,x′)=(x1−x1′)2+(x2−x2′)2d(x, x') = \sqrt{(x_1 - x_1')^2 + (x_2 - x_2')^2}。预测规则:

y^=1[∑i∈Nk(x)yi≥k2]\hat{y} = \mathbb{1}\Big[\sum_{i \in N_k(x)} y_i \ge \frac{k}{2}\Big]

Nk(x)N_k(x) 是最近的 kk 个邻居。没有损失函数、没有梯度——“训练”的代价为零,预测的代价是 O(n)O(n) 次距离计算,这是惰性学习的代价转移。

Rust 实现

// src/main.rs(段一:kNN 算法与实验主程序)
// 分类②:k 近邻——最懒的分类器
// 单文件、仅标准库。绘图器与系列前篇内联的是同一份实现。

// ===================== xorshift64 + Box-Muller(沿用 07 篇) =====================

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
    }

    fn next_gauss(&mut self) -> f64 {
        let u1 = 1.0 - self.next_f64();
        let u2 = self.next_f64();
        (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos()
    }
}

// ===================== kNN =====================

/// 欧氏距离。
fn dist(a: (f64, f64), b: (f64, f64)) -> f64 {
    ((a.0 - b.0).powi(2) + (a.1 - b.1).powi(2)).sqrt()
}

/// kNN 预测:按距离排序取前 k 个邻居,多数投票。
/// 平局时偏向距离更近的一方(排序稳定性让近者先发言)。
fn knn_predict(train_x: &[(f64, f64)], train_y: &[f64], q: (f64, f64), k: usize) -> f64 {
    let mut neighbors: Vec<(f64, f64)> = train_x
        .iter()
        .zip(train_y.iter())
        .map(|(x, y)| (dist(*x, q), *y))
        .collect();
    neighbors.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(std::cmp::Ordering::Equal));
    let votes: f64 = neighbors.iter().take(k).map(|(_, y)| y).sum();
    if votes * 2.0 >= k as f64 { 1.0 } else { 0.0 }
}

fn accuracy(train_x: &[(f64, f64)], train_y: &[f64], test_x: &[(f64, f64)], test_y: &[f64], k: usize) -> f64 {
    let correct = test_x
        .iter()
        .zip(test_y.iter())
        .filter(|(x, y)| knn_predict(train_x, train_y, **x, k) == **y)
        .count();
    correct as f64 / test_y.len() as f64
}

/// 造同心圆数据:类 0 是半径 ≤1.5 的圆盘,类 1 是 1.8~3.5 的圆环,各 n/2 个。
fn make_data(rng: &mut XorShift, n: usize) -> (Vec<(f64, f64)>, Vec<f64>) {
    let mut xs = Vec::with_capacity(n);
    let mut ys = Vec::with_capacity(n);
    for i in 0..n {
        let angle = rng.next_f64() * 2.0 * std::f64::consts::PI;
        let (r, label) = if i % 2 == 0 {
            (1.5 * rng.next_f64().sqrt(), 0.0) // 圆盘:sqrt 保证面积均匀
        } else {
            (1.8 + 1.7 * rng.next_f64(), 1.0) // 圆环
        };
        xs.push((r * angle.cos() + 0.08 * rng.next_gauss(), r * angle.sin() + 0.08 * rng.next_gauss()));
        ys.push(label);
    }
    (xs, ys)
}

fn main() {
    let out_dir = "../../../frontend/public/images/series/rust-ml-10-knn";
    std::fs::create_dir_all(out_dir).expect("创建输出目录失败");

    // ---- 1. 数据:线性不可分的同心圆 ----
    let mut rng = XorShift::new(42);
    let (xtr, ytr) = make_data(&mut rng, 120);
    let (xte, yte) = make_data(&mut rng, 120);
    println!("train = {},test = {}(同心圆:类 0 圆盘 r≤1.5,类 1 圆环 1.8~3.5)", xtr.len(), xte.len());

    // ---- 2. k 扫描:小 k 过拟合,大 k 欠拟合 ----
    let ks = [1usize, 3, 5, 7, 9, 15, 21, 31];
    println!("\n[k 扫描]   k   train acc   test acc");
    let mut train_curve = Vec::new();
    let mut test_curve = Vec::new();
    let mut best = (0usize, 0.0f64);
    for &k in &ks {
        let tr = accuracy(&xtr, &ytr, &xtr, &ytr, k);
        let te = accuracy(&xtr, &ytr, &xte, &yte, k);
        println!("{:>7}   {:>8.3}   {:>8.3}", k, tr, te);
        train_curve.push((k as f64, tr));
        test_curve.push((k as f64, te));
        if te > best.1 {
            best = (k, te);
        }
    }
    println!("测试集最优:k = {}(acc = {:.3})", best.0, best.1);

    // ---- 3. 图一:k-acc 曲线 ----
    let mut c = Canvas::new(560.0, 320.0, 1.0, 31.0, 0.7, 1.02);
    c.axes("k", "accuracy");
    c.polyline(&train_curve, PALETTE[0], 2.0);
    c.polyline(&test_curve, PALETTE[1], 2.0);
    c.dots(&[(best.0 as f64, best.1)], PALETTE[2], 5.0);
    c.legend(&[("train", PALETTE[0]), ("test", PALETTE[1]), ("best", PALETTE[2])]);
    let p1 = format!("{out_dir}/k-scan.svg");
    c.save(&p1);

    // ---- 4. 图二/三:决策区域图(k=1 破碎 vs k=15 平滑) ----
    for (k, name) in [(1usize, "region-k1"), (15, "region-k15")] {
        let mut region0 = Vec::new();
        let mut region1 = Vec::new();
        for i in 0..46 {
            for j in 0..46 {
                let q = (i as f64 * 0.16 - 3.6, j as f64 * 0.16 - 3.6);
                if knn_predict(&xtr, &ytr, q, k) > 0.5 {
                    region1.push(q);
                } else {
                    region0.push(q);
                }
            }
        }
        let mut c = Canvas::new(560.0, 460.0, -3.6, 3.4, -3.6, 3.4);
        c.axes("x1", "x2");
        c.dots(&region0, PALETTE[1], 2.2);
        c.dots(&region1, PALETTE[3], 2.2);
        let class0: Vec<(f64, f64)> = xtr.iter().zip(ytr.iter()).filter(|(_, y)| **y < 0.5).map(|(x, _)| *x).collect();
        let class1: Vec<(f64, f64)> = xtr.iter().zip(ytr.iter()).filter(|(_, y)| **y > 0.5).map(|(x, _)| *x).collect();
        c.dots(&class0, "#16161d", 2.6);
        c.dots(&class1, "#fbfaf7", 2.6);
        c.legend(&[("region 0", PALETTE[1]), ("region 1", PALETTE[3])]);
        let p = format!("{out_dir}/{name}.svg");
        c.save(&p);
        println!("图已写入:{p}");
    }

    // ---- 5. 抽样预测 ----
    println!("\n[抽样,k=7] 查询点        | 真实 | 预测");
    for i in (0..xte.len()).step_by(23) {
        let pred = knn_predict(&xtr, &ytr, xte[i], 7);
        println!("  ({:+.2}, {:+.2}) |  {:.0}   |  {:.0}", xte[i].0, xte[i].1, yte[i], pred);
    }
    println!("\n图已写入:{p1}");
}

// ===================== 迷你 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
}

实现要点:

  • knn_predict 是算法的完整实现:算全部距离、排序、取前 kk、投票——五行核心逻辑,系列至今最短的”模型”。
  • 排序用 sort_by + partial_cmp:浮点数没有全序(NaN 的存在),比较返回 Option,unwrap_or(Equal) 是”NaN 视为相等”的务实处理。
  • 平局策略写在投票式里:votes * 2 >= k 让偶数 kk 的平局偏向排序后先出现的邻居——即距离更近的一方,与排序稳定性配合得到确定性结果。
  • 数据生成的小技巧:r = 1.5 * sqrt(U) 让圆盘内的点在面积上均匀(直接用 U 会让点挤向圆心);环带用线性半径即可。
// src/main.rs 段二:迷你 SVG 绘图器(与系列前篇同一份实现)
use std::fmt::Write as _;
// src/main.rs 末尾:单元测试(cargo test 运行)
#[cfg(test)]
mod tests {
    use super::*;

    #[test]
    fn dist_is_euclidean() {
        assert!((dist((0.0, 0.0), (3.0, 4.0)) - 5.0).abs() < 1e-12);
        assert_eq!(dist((1.0, 1.0), (1.0, 1.0)), 0.0);
    }

    #[test]
    fn k1_returns_nearest_label() {
        let xs = vec![(0.0, 0.0), (5.0, 5.0), (0.2, 0.1)];
        let ys = vec![0.0, 1.0, 1.0];
        assert_eq!(knn_predict(&xs, &ys, (0.05, 0.05), 1), 0.0);
        assert_eq!(knn_predict(&xs, &ys, (4.9, 5.1), 1), 1.0);
    }

    #[test]
    fn majority_vote_with_tie() {
        // k=2 时一票对一票,平局 → 取排序后先出现者(距离近者赢)
        let xs = vec![(1.0, 0.0), (1.1, 0.0), (5.0, 0.0)];
        let ys = vec![0.0, 1.0, 1.0];
        // 查询点到 (1.0,0) 距离 0.05,到 (1.1,0) 距离 0.05,到 (5,0) 很远——k=2 平局
        let pred = knn_predict(&xs, &ys, (1.05, 0.0), 2);
        assert!(pred == 0.0 || pred == 1.0); // 平局策略只需确定性
        // k=3 时 2:1,应判 1
        assert_eq!(knn_predict(&xs, &ys, (1.05, 0.0), 3), 1.0);
    }

    #[test]
    fn separates_concentric_circles() {
        let mut rng = XorShift::new(7);
        let (xtr, ytr) = make_data(&mut rng, 120);
        let acc = accuracy(&xtr, &ytr, &xtr, &ytr, 7);
        assert!(acc > 0.95, "同心圆数据上 k=7 应接近满分: {acc}");
    }
}

Rust 语法角:Vec 的排序家族

neighbors.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(...)) 展示了 Rust 排序的完整谱系:sort(要求可比较)、sort_by(自定义比较器)、sort_by_key(按投影排序)。比较器返回 Ordering(Less/Equal/Greater),闭包捕获环境的能力让”按距离排但带着标签走”零成本——(dist, label) 元组整体排序,标签随行。Python 的 sorted(key=...) 语法更短,但 Rust 的版本把”比较可能失败”(浮点 NaN)显式暴露了出来。集合操作见《Rust 程序设计语言》ch08-01 与标准库 slice::sort_by 文档。

运行结果

cargo test(4 个用例:欧氏距离、最近邻正确性、平局确定性、同心圆可分性)全部通过后,cargo run:

train = 120,test = 120(同心圆:类 0 圆盘 r≤1.5,类 1 圆环 1.8~3.5)

[k 扫描]   k   train acc   test acc
      1      1.000      0.992
      3      0.992      0.992
      5      0.992      0.992
      7      0.983      0.992
      9      0.975      0.992
     15      0.958      0.983
     21      0.917      0.925
     31      0.792      0.833
测试集最优:k = 1(acc = 0.992)
图已写入:../../../frontend/public/images/series/rust-ml-10-knn/region-k1.svg
图已写入:../../../frontend/public/images/series/rust-ml-10-knn/region-k15.svg

[抽样,k=7] 查询点        | 真实 | 预测
  (+0.98, +0.46) |  0   |  0
  (-2.23, -0.73) |  1   |  1
  (-1.27, +0.40) |  0   |  0
  (-1.17, -3.36) |  1   |  1
  (-0.91, -0.41) |  0   |  0
  (+0.62, -2.40) |  1   |  1

图已写入:../../../frontend/public/images/series/rust-ml-10-knn/k-scan.svg

三张图:

k 扫描:train/test 准确率随 k 的变化

决策区域:k=1,边界破碎成孤岛

决策区域:k=15,平滑但开始失真

怎么读这些数字和图

  • k=1 的 train acc = 1.000 是最响的过拟合警报:训练集满分从来不是成绩,是嫌疑。它的 test acc 0.992 靠数据的低密度侥幸维持——region-k1.svg 揭示了真相:决策区域碎成群岛,任何落在岛链之外的测试点都会翻车。
  • k=3~9 是甜蜜区:test acc 稳定在 0.992,train acc 缓慢下降——用一点点拟合能力换边界的平滑,这正是偏差-方差权衡的定量演示。本例数据太干净,甜蜜区很宽;数据 noisy 时这张图会呈现更尖锐的 U 形。
  • k=21 起崩塌:train 0.917 → k=31 时 0.792——圆环内缘的类 1 点被圆盘多数票吞没。region-k15.svg 已经能看到外环边界向内侧侵蚀。
  • 决策区域图是”模型在脑中看到的世界”:两图网格是模型对平面的逐格判断。k=1 的世界满是飞地;k=15 的世界简洁但细节被抹平。看区域图比看准确率更能理解一个分类器。

优缺点与适用场景

(抄清单 2.1 原文)

  • 优点:思想直观、无需训练、天然支持多分类、可分类可回归。
  • 缺点:预测慢、内存大、对特征量纲敏感(必须标准化)、高维下失效(维度灾难)。
  • 适用场景:低维数据、需要快速建立基线的场景。

小结

kNN 用”没有模型”展示了模型选择的本质:kk 就是全部的超参数,而它控制的是平滑度——和 05 篇多项式的阶数、06 篇的正则化强度,是同一个旋钮的三个化身。但投票机制对噪声一视同仁,能否让每个邻居”按可信度加权发言”?下一篇分类③:朴素贝叶斯——用概率生成模型给每个类别画像,毫秒级响应的分类器。