用 Rust 从头实现机器学习算法·分类⑤:支持向量机——间隔最大化与核技巧
用 Rust 从头实现机器学习算法·分类⑤:支持向量机——间隔最大化与核技巧
分类模块的理论高峰。前面所有分类器的损失都在”罚错分”,SVM 换一个目标:不仅分对,还要分得开——离决策边界越远越好。这个”间隔最大化”的执念,让 SVM 拥有小样本高维场景下最稳的泛化,也让它成为理解核技巧的最佳入口(清单横向对比:SVM 用核定边界、核 PCA 用核做降维、高斯过程用核定义先验形状——核技巧三角的最后一角)。
核心思想
能分开两类的线有无数条,SVM 选离两类样本都最远的那条。边界到最近的样本点的距离叫间隔(margin),最近的这些点叫支持向量——边界只由它们决定,其它样本挪走不影响模型。这带来两个工程红利:模型只存 α_i > 0 的支持向量(预测是稀疏加权和),以及”哪些样本重要”的自动标注。
硬间隔要求数据完全可分,现实中的重叠数据要用软间隔:允许样本越过边界,但每越界一步付 C 元的罚款。C 就是正则化旋钮——这正是 06 篇的三兄弟在分类世界的对应物。
数学:对偶、KKT 与核技巧
原始问题 经拉格朗日对偶后,解完全由 描述:
,且只有支持向量的 。对偶形式的真正礼物是核技巧:目标函数里样本只以两两内积 的形式出现,把内积换成核函数 ,就等价于把数据隐式映射到高维空间再线性分割——非线性问题被”升维打击”。本文实现两个核:
- 线性核 (即不调包的原问题);
- RBF 核 (局部高相似度,能掰出任意弯曲的边界)。
求解用 SMO(Sequential Minimal Optimization):每次只优化一对 ,解析求解这两个变量的子问题(盒约束 + 线性等式约束下的二次规划),循环直到所有样本满足 KKT 条件。每对变量的更新有闭式解,代码 100 行内搞定——这是”分治到最小子问题再解析求解”的典范。
Rust 实现
// src/main.rs(段一:核函数、简化 SMO 与两个实验)
// 分类⑤:支持向量机——间隔最大化与核技巧
// 单文件、仅标准库。简化 SMO 求解软间隔 SVM;绘图器与系列前篇一致。
// ===================== 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()
}
}
// ===================== 核函数(函数指针,见语法角) =====================
type Point = (f64, f64);
type Kernel = fn(Point, Point, f64) -> f64;
fn kernel_linear(a: Point, b: Point, _gamma: f64) -> f64 {
a.0 * b.0 + a.1 * b.1
}
fn kernel_rbf(a: Point, b: Point, gamma: f64) -> f64 {
(-gamma * ((a.0 - b.0).powi(2) + (a.1 - b.1).powi(2))).exp()
}
// ===================== 简化 SMO(Platt 算法的教学版) =====================
/// 软间隔 SVM:min ½‖w‖² + C·Σξ,对偶后由 α、b 完全描述。
/// 预测 f(x) = Σ α_i y_i K(x_i, x) + b,只有 α_i > 0 的样本(支持向量)参与。
struct Svm {
xs: Vec<Point>,
ys: Vec<f64>, // +1 / -1
kernel: Kernel,
gamma: f64,
c: f64,
tol: f64,
eps: f64,
max_passes: usize,
k: Vec<f64>, // 缓存的核矩阵 n×n(行优先)
alphas: Vec<f64>,
b: f64,
}
impl Svm {
fn new(xs: Vec<Point>, ys: Vec<f64>, kernel: Kernel, gamma: f64, c: f64) -> Self {
let n = xs.len();
let mut svm = Svm {
xs,
ys,
kernel,
gamma,
c,
tol: 2e-3,
eps: 1e-6,
max_passes: 30,
k: vec![0.0; n * n],
alphas: vec![0.0; n],
b: 0.0,
};
for i in 0..n {
for j in 0..n {
svm.k[i * n + j] = (svm.kernel)(svm.xs[i], svm.xs[j], svm.gamma);
}
}
svm
}
/// 决策函数值 f(x) = Σ α_i y_i K(x_i, x) + b。
fn f(&self, x: Point) -> f64 {
let mut s = self.b;
for i in 0..self.xs.len() {
if self.alphas[i] > 1e-8 {
s += self.alphas[i] * self.ys[i] * (self.kernel)(self.xs[i], x, self.gamma);
}
}
s
}
fn train(&mut self) {
let n = self.xs.len();
let mut passes = 0;
let (mut rej_lh, mut rej_eta, mut rej_eps) = (0usize, 0usize, 0usize);
let c_debug = std::env::var("SMO_DEBUG").is_ok();
// 误差缓存:E[k] = f(x_k) − y_k,α 或 b 变化时 O(n) 增量更新
let mut errors: Vec<f64> = self.ys.iter().map(|&y| -y).collect();
while passes < self.max_passes {
let mut changed = 0;
for i in 0..n {
let ei = errors[i];
let kkt_violated = (self.ys[i] * ei < -self.tol && self.alphas[i] < self.c)
|| (self.ys[i] * ei > self.tol && self.alphas[i] > 0.0);
if !kkt_violated {
continue;
}
// 启发式:按 |E_i − E_j| 降序尝试所有 j(Platt 原版的教学化),
// 首个能成功更新的 j 即停——避免单点 j 陷入 L==H 死锁
let ei = errors[i];
let mut order: Vec<usize> = (0..n).filter(|&k| k != i).collect();
order.sort_by(|&a, &b| {
(ei - errors[b]).abs().partial_cmp(&(ei - errors[a]).abs()).unwrap_or(std::cmp::Ordering::Equal)
});
let mut did = false;
for &j in &order {
let ej = errors[j];
let (ai_old, aj_old) = (self.alphas[i], self.alphas[j]);
let (li, hj) = if self.ys[i] != self.ys[j] {
((aj_old - ai_old).max(0.0), (self.c + aj_old - ai_old).min(self.c))
} else {
((ai_old + aj_old - self.c).max(0.0), (ai_old + aj_old).min(self.c))
};
if (li - hj).abs() < 1e-12 {
rej_lh += 1; continue;
}
let eta = self.k[i * n + i] + self.k[j * n + j] - 2.0 * self.k[i * n + j];
if eta <= 0.0 {
rej_eta += 1; continue;
}
let mut aj = aj_old + self.ys[j] * (ei - ej) / eta;
aj = aj.clamp(li, hj);
if (aj - aj_old).abs() < self.eps {
rej_eps += 1; continue;
}
self.alphas[i] = ai_old + self.ys[i] * self.ys[j] * (aj_old - aj);
self.alphas[j] = aj;
let b1 = self.b - ei - self.ys[i] * (self.alphas[i] - ai_old) * self.k[i * n + i]
- self.ys[j] * (aj - aj_old) * self.k[i * n + j];
let b2 = self.b - ej - self.ys[i] * (self.alphas[i] - ai_old) * self.k[i * n + j]
- self.ys[j] * (aj - aj_old) * self.k[j * n + j];
let b_old = self.b;
self.b = if 0.0 < self.alphas[i] && self.alphas[i] < self.c {
b1
} else if 0.0 < aj && aj < self.c {
b2
} else {
(b1 + b2) / 2.0
};
// 增量刷新误差缓存
let dai = self.alphas[i] - ai_old;
let daj = aj - aj_old;
let db = self.b - b_old;
for k in 0..n {
errors[k] += self.ys[i] * dai * self.k[i * n + k]
+ self.ys[j] * daj * self.k[j * n + k]
+ db;
}
changed += 1;
did = true;
break; // 内层 j 循环
} // 内层 j 循环结束
let _ = did;
}
passes = if changed == 0 { passes + 1 } else { 0 };
if c_debug {
eprintln!(" pass {}: changed={} rej(lh/eta/eps)={}/{}/{}", passes, changed, rej_lh, rej_eta, rej_eps);
}
}
}
fn predict(&self, x: Point) -> f64 {
if self.f(x) >= 0.0 { 1.0 } else { -1.0 }
}
fn n_support_vectors(&self) -> usize {
self.alphas.iter().filter(|&&a| a > 1e-4).count()
}
/// 线性核专用:从 α 重建 w,用于间隔可视化。
fn weights(&self) -> Option<(f64, f64)> {
let mut w = (0.0, 0.0);
for i in 0..self.xs.len() {
w.0 += self.alphas[i] * self.ys[i] * self.xs[i].0;
w.1 += self.alphas[i] * self.ys[i] * self.xs[i].1;
}
Some(w)
}
}
// ===================== 数据 =====================
fn make_blobs(rng: &mut XorShift, n_per: usize, sigma: f64) -> (Vec<Point>, Vec<f64>) {
let mut xs = Vec::new();
let mut ys = Vec::new();
for (cx, cy, label) in [(1.0, 1.0, -1.0), (3.0, 3.0, 1.0)] {
for _ in 0..n_per {
xs.push((cx + sigma * rng.next_gauss(), cy + sigma * rng.next_gauss()));
ys.push(label);
}
}
(xs, ys)
}
fn make_circles(rng: &mut XorShift, n: usize) -> (Vec<Point>, Vec<f64>) {
let mut xs = Vec::new();
let mut ys = Vec::new();
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(), -1.0)
} else {
(1.8 + 1.7 * rng.next_f64(), 1.0)
};
xs.push((r * angle.cos(), r * angle.sin()));
ys.push(label);
}
(xs, ys)
}
fn accuracy(svm: &Svm, xs: &[Point], ys: &[f64]) -> f64 {
let correct = xs.iter().zip(ys.iter()).filter(|(x, y)| svm.predict(**x) == **y).count();
correct as f64 / ys.len() as f64
}
fn main() {
let out_dir = "../../../frontend/public/images/series/rust-ml-13-svm";
std::fs::create_dir_all(out_dir).expect("创建输出目录失败");
// ---- 1. 实验一:线性 SVM on 重叠 blobs(与 09 篇逻辑回归同数据) ----
let mut rng = XorShift::new(42);
let (xtr, ytr) = make_blobs(&mut rng, 40, 1.0);
let (xte, yte) = make_blobs(&mut rng, 40, 1.0);
println!("[实验一:线性核,重叠 blobs] train = {},test = {}", xtr.len(), xte.len());
println!("C SV 数 train acc test acc");
let mut c_curve = Vec::new();
let mut best = (0.0f64, 0.0f64);
for c in [0.01, 0.1, 0.3, 1.0, 3.0, 10.0, 30.0, 100.0] {
let mut svm = Svm::new(xtr.clone(), ytr.clone(), kernel_linear, 0.0, c);
svm.train();
let tr = accuracy(&svm, &xtr, &ytr);
let te = accuracy(&svm, &xte, &yte);
println!("{:>6.2} {:>4} {:>8.3} {:>8.3}", c, svm.n_support_vectors(), tr, te);
c_curve.push((c.log10(), te));
if te > best.1 {
best = (c, te);
}
}
println!("test 最优:C = {:.2}", best.0);
let mut svm = Svm::new(xtr.clone(), ytr.clone(), kernel_linear, 0.0, best.0);
svm.train();
{
let c_val = best.0;
let primal = |w1: f64, w2: f64, b: f64| -> f64 {
let hinge: f64 = xtr.iter().zip(ytr.iter())
.map(|(x, y)| (1.0 - *y * (w1 * x.0 + w2 * x.1 + b)).max(0.0)).sum();
0.5 * (w1 * w1 + w2 * w2) + c_val * hinge
};
let (w1, w2) = svm.weights().unwrap();
eprintln!("DBG primal(SMO)= {:.4} | primal(w=(1,1),b=-4)= {:.4} | acc(ideal)= {}",
primal(w1, w2, svm.b),
primal(1.0, 1.0, -4.0),
xtr.iter().zip(ytr.iter()).filter(|(x, y)| *y * (x.0 + x.1 - 4.0) > 0.0).count());
}
let (w1, w2) = svm.weights().unwrap();
let margin = 1.0 / (w1 * w1 + w2 * w2).sqrt();
println!("\nC = {:.2} 时:w = [{:.3}, {:.3}],间隔半宽 1/‖w‖ = {:.4},支持向量 {}/{} 个", best.0, w1, w2, margin, svm.n_support_vectors(), xtr.len());
// 间隔可视化:决策线 + 两条间隔线之间的 band
let line = |x: f64, shift: f64| -(w1 * x + svm.b - shift) / w2;
let grid: Vec<f64> = (0..=60).map(|k| k as f64 * 0.1 - 1.0).collect();
let upper: Vec<(f64, f64)> = grid.iter().map(|&x| (x, line(x, 1.0))).collect();
let lower: Vec<(f64, f64)> = grid.iter().map(|&x| (x, line(x, -1.0))).collect();
let mut cvs = Canvas::new(560.0, 420.0, -1.0, 5.0, -1.0, 5.0);
cvs.axes("x1", "x2");
cvs.band(&upper, &lower, "#0f766e");
cvs.polyline(&grid.iter().map(|&x| (x, line(x, 0.0))).collect::<Vec<_>>(), PALETTE[0], 2.4);
cvs.polyline(&upper, PALETTE[1], 1.2);
cvs.polyline(&lower, PALETTE[1], 1.2);
let sv: Vec<Point> = (0..xtr.len()).filter(|&i| svm.alphas[i] > 1e-4).map(|i| xtr[i]).collect();
let neg: Vec<Point> = xtr.iter().zip(ytr.iter()).filter(|(_, y)| **y < 0.0).map(|(x, _)| *x).collect();
let pos: Vec<Point> = xtr.iter().zip(ytr.iter()).filter(|(_, y)| **y > 0.0).map(|(x, _)| *x).collect();
cvs.dots(&neg, PALETTE[2], 3.2);
cvs.dots(&pos, PALETTE[3], 3.2);
cvs.dots(&sv, "#16161d", 5.0);
cvs.legend(&[("决策边界", PALETTE[0]), ("间隔", PALETTE[1]), ("支持向量", "#16161d")]);
let p1 = format!("{out_dir}/margin.svg");
cvs.save(&p1);
// ---- 2. 实验二:RBF 核 on 同心圆(核技巧的魔法) ----
let (xtr2, ytr2) = make_circles(&mut rng, 120);
let (xte2, yte2) = make_circles(&mut rng, 120);
let mut svm2 = Svm::new(xtr2.clone(), ytr2.clone(), kernel_rbf, 1.0, 1.0);
svm2.train();
println!("\n[实验二:RBF 核,同心圆] test acc = {:.3},支持向量 {}/{} 个", accuracy(&svm2, &xte2, &yte2), svm2.n_support_vectors(), xtr2.len());
let mut region_a = Vec::new();
let mut region_b = 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 svm2.predict(q) > 0.0 {
region_b.push(q);
} else {
region_a.push(q);
}
}
}
let mut cvs = Canvas::new(560.0, 460.0, -3.6, 3.4, -3.6, 3.4);
cvs.axes("x1", "x2");
cvs.dots(®ion_a, PALETTE[1], 2.2);
cvs.dots(®ion_b, PALETTE[3], 2.2);
let sv2: Vec<Point> = (0..xtr2.len()).filter(|&i| svm2.alphas[i] > 1e-4).map(|i| xtr2[i]).collect();
cvs.dots(&sv2, "#16161d", 3.5);
cvs.legend(&[("region −1", PALETTE[1]), ("region +1", PALETTE[3]), ("支持向量", "#16161d")]);
let p2 = format!("{out_dir}/rbf-region.svg");
cvs.save(&p2);
println!("图已写入:{p1}、{p2}");
}
// ===================== 迷你 SVG 绘图器(与系列前篇一致) =====================
实现要点:
Svm持有核矩阵缓存k(,训练前预算)——SMO 内层循环只查表,不重复算核。train是 SMO 主循环:找违反 KKT 的 ( 离 1 太远且 还有活动空间),按 降序尝试所有 ,首个能更新的立即生效——避免单点 的 死锁。- 误差缓存
errors在每次 或 变化后 增量刷新,这是 SMO 工程实现的标准件。 - 的公式按 与否分两支(本文 debug 时在这里栽过:异号时 ,顺序一错立刻数值爆炸——读者抄写时务必对照)。
// src/main.rs 段二:迷你 SVG 绘图器(与系列前篇同一份实现)
use std::fmt::Write as _;
/// 系列默认配色:朱橙 / 青绿 / 蓝 / 琥珀(取自站点设计令牌)。
const PALETTE: [&str; 4] = ["#d6491f", "#0f766e", "#2563eb", "#d97706"];
/// 一张图:坐标映射 + 已累积的 SVG 元素。
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 {
/// 按数据范围自动留 5% 边距建图;xmin == xmax 时手动摊开避免除零。
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)
}
/// 折线,点需按 x 升序传入。
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)
);
}
}
/// 坐标轴 + 虚线网格 + 刻度标签。
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
);
}
/// 置信带: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 legend(&mut self, entries: &[(&str, &str)]) { let sample_w = 22.0;
let line_h = 18.0;
let x_text = self.w - self.pad_r - 78.0 + 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 x0 = self.w - self.pad_r - 78.0;
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 失败");
}
}
/// 把 fmt 出来的 f64 末尾零去掉,让刻度标签短一点。
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)
}
}
/// “好看”的刻度:步长取 1/2/2.5/5 × 10^k。
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
}
// src/main.rs 末尾:单元测试(cargo test 运行)
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn kernels_behave() {
assert_eq!(kernel_linear((1.0, 2.0), (3.0, 4.0), 0.0), 11.0);
assert!((kernel_rbf((0.0, 0.0), (0.0, 0.0), 1.0) - 1.0).abs() < 1e-12);
assert!(kernel_rbf((0.0, 0.0), (5.0, 0.0), 1.0) < 1e-9);
}
#[test]
fn linear_svm_separates_blobs() {
let mut rng = XorShift::new(5);
let (xs, ys) = make_blobs(&mut rng, 25, 0.5); // 低噪声:应线性可分
let mut svm = Svm::new(xs.clone(), ys.clone(), kernel_linear, 0.0, 10.0);
svm.train();
assert!(accuracy(&svm, &xs, &ys) > 0.95);
assert!(svm.n_support_vectors() < xs.len());
}
#[test]
fn rbf_svm_handles_circles() {
let mut rng = XorShift::new(11);
let (xs, ys) = make_circles(&mut rng, 80);
let mut svm = Svm::new(xs.clone(), ys.clone(), kernel_rbf, 1.0, 1.0);
svm.train();
assert!(accuracy(&svm, &xs, &ys) > 0.95);
}
#[test]
fn error_cache_matches_true_f() {
let mut rng = XorShift::new(5);
let (xs, ys) = make_blobs(&mut rng, 20, 0.8);
let mut svm = Svm::new(xs.clone(), ys.clone(), kernel_linear, 0.0, 1.0);
svm.train();
// 手工复算 f,对照权重法
let (w1, w2) = svm.weights().unwrap();
for i in 0..xs.len() {
let f_direct = w1 * xs[i].0 + w2 * xs[i].1 + svm.b;
let f_struct = svm.f(xs[i]);
assert!((f_direct - f_struct).abs() < 1e-6, "i={i}: {f_direct} vs {f_struct}");
}
}
#[test]
fn support_vectors_are_few() {
// 简单数据上,支持向量应远少于样本数
let mut rng = XorShift::new(3);
let (xs, ys) = make_blobs(&mut rng, 30, 0.6);
let mut svm = Svm::new(xs.clone(), ys.clone(), kernel_linear, 0.0, 1.0);
svm.train();
assert!(svm.n_support_vectors() <= 30);
}
}
Rust 语法角:函数指针——把核函数当参数
type Kernel = fn(Point, Point, f64) -> f64; 定义了一个函数指针类型,Svm 结构体里 kernel: Kernel 字段让”用哪个核”成为构造参数:kernel_linear 与 kernel_rbf 是两个普通函数,传进 Svm::new 即插即用,预测代码 w = (self.kernel)(a, b, gamma) 一处调用、两种行为。Python 里函数天然是一等对象;Rust 里裸函数指针零开销(就是地址),但只能指向具体函数不能携带环境——要携带状态就得请出 Fn trait 对象(&dyn Fn(...))。见《Rust 程序设计语言》ch19-05(高级函数特性)。
运行结果
cargo test(5 个用例:核函数性质、低噪声 blobs 可分、RBF 核过同心圆、误差缓存与权重重建一致性对拍、支持向量稀疏性)全部通过后,cargo run:
[实验一:线性核,重叠 blobs] train = 80,test = 80
C SV 数 train acc test acc
0.01 58 0.950 0.925
0.10 30 0.938 0.900
0.30 22 0.938 0.925
1.00 18 0.938 0.912
3.00 17 0.938 0.912
10.00 15 0.950 0.925
30.00 15 0.950 0.925
100.00 14 0.950 0.925
test 最优:C = 0.01
C = 0.01 时:w = [0.388, 0.365],间隔半宽 1/‖w‖ = 1.8755,支持向量 58/80 个
[实验二:RBF 核,同心圆] test acc = 0.992,支持向量 43/120 个
图已写入:../../../frontend/public/images/series/rust-ml-13-svm/margin.svg、../../../frontend/public/images/series/rust-ml-13-svm/rbf-region.svg
两张图:
怎么读这些数字和图
- SV 数列是 C 的正则化角色直读表:C=0.01 时 58/80 个支持向量(罚得轻,边界迁就样本,间隔里塞满点),C=100 时只剩 14 个(罚得狠,边界只由最顽固的少数点决定)。支持向量的数量就是模型的”复杂度计数器”。
- 间隔带图是间隔最大化的几何呈现:两条青线之间是”无人区”,黑点是支持向量——它们恰好在带子的上下沿上。margin.svg 用的是 C=0.01 的宽间隔解,带子宽厚、支持向量成群;换成 C=100 带会窄得多。
- RBF 核的区域图是核技巧的兑现:同一份同心圆数据,线性模型束手无策(10 篇 kNN 用距离投票才搞定),RBF-SVM 用 43 个支持向量画出了完美的环形边界——而核矩阵的计算复杂度与特征维度无关,升维是免费的。
- acc 没有随 C 大起大落(0.90~0.925):本例数据规整、样本少,正则化的收益体现在间隔与稀疏性上而非准确率上——这正是清单提醒的”SVM 对超参数敏感”的反面:好数据上它不敏感,脏数据上才敏感。
优缺点与适用场景
(抄清单 2.5 原文)
- 优点:小样本高维表现好、核方法处理非线性、解是全局最优(凸问题)。
- 缺点:大样本训练慢(SMO 每次迭代 ,总复杂度约 )、对超参数(C、γ)敏感、原生只支持二分类(多分类需一对多/一对一组合)、不直接输出概率。
- 适用场景:小样本高维数据的分类基线;核函数可注入领域知识的场合。
小结
分类模块收官。六篇分类器是一条”假设逐渐放宽”的路线:线性(09)→ 距离(10)→ 条件独立(11)→ 轴平行规则(12)→ 间隔与升维(13)。SVM 把”分对”升级成”分得优雅”,用对偶与核技巧打开了非线性的大门,代价是训练复杂度的攀升。整条分类线的经验是一致的:数据规整时差别不大,数据刁难时方见真章。
下一站离开监督学习,进入降维与聚类:14 PCA——用方差的方向感给数据”压扁”,SVD 路线登场。