用 Rust 从头实现机器学习算法·分类②:k 近邻——最懒的分类器
用 Rust 从头实现机器学习算法·分类②:k 近邻——最懒的分类器
逻辑回归用一条直线切分世界,前提是数据真的”线性可分”。本篇的数据一上来就打破这个前提——同心圆:类 0 住在半径 1.5 的圆盘里,类 1 住在 1.8~3.5 的圆环上。任何直线都无能为力,而 k 近邻(kNN)几乎不费吹灰之力。这是清单里”低维数据快速建立基线”的典型案例。
核心思想
kNN 是机器学习里最”懒”的算法——没有训练阶段。它的全部假设只有一句谚语:近朱者赤。预测新样本时,在训练集里找离它最近的 个点,多数投票决定类别(回归问题则取平均)。清单的三个要素:距离度量(本文用欧氏)、 的选择、投票规则。
两个细节决定成败:
- 是偏差-方差的旋钮: 时决策边界扭出无数小岛(每个训练点都是一座孤岛,过拟合); 增大边界趋于平滑(方差下降),但太大时连真实形状都抹掉(偏差上升)——注意这个 k 与 k 均值聚类的 k 毫无关系(清单附录 A.3 专门提醒);
- 量纲敏感:距离由数值大的特征主导,用 kNN 前必须标准化(04 篇的纪律在此是硬要求)。
数学:距离与多数投票
欧氏距离:。预测规则:
是最近的 个邻居。没有损失函数、没有梯度——“训练”的代价为零,预测的代价是 次距离计算,这是惰性学习的代价转移。
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(®ion0, PALETTE[1], 2.2);
c.dots(®ion1, 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是算法的完整实现:算全部距离、排序、取前 、投票——五行核心逻辑,系列至今最短的”模型”。- 排序用
sort_by+partial_cmp:浮点数没有全序(NaN 的存在),比较返回Option,unwrap_or(Equal)是”NaN 视为相等”的务实处理。 - 平局策略写在投票式里:
votes * 2 >= k让偶数 的平局偏向排序后先出现的邻居——即距离更近的一方,与排序稳定性配合得到确定性结果。 - 数据生成的小技巧:
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=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 用”没有模型”展示了模型选择的本质: 就是全部的超参数,而它控制的是平滑度——和 05 篇多项式的阶数、06 篇的正则化强度,是同一个旋钮的三个化身。但投票机制对噪声一视同仁,能否让每个邻居”按可信度加权发言”?下一篇分类③:朴素贝叶斯——用概率生成模型给每个类别画像,毫秒级响应的分类器。