用 Rust 从头实现机器学习算法·终篇:集成学习——随机森林把决策树的方差熨平
用 Rust 从头实现机器学习算法·终篇:集成学习——随机森林把决策树的方差熨平
系列最后一篇。12 篇我们给决策树下了诊断:极易过拟合、对扰动敏感——但清单留了药方:它是集成学习的最佳基学习器。本篇把几百棵”病树”装进投票箱,看”三个臭皮匠”如何超过诸葛亮。这是从教材算法走向工业实践的桥梁(清单 5.1),也是全系列的收官之作。
核心思想
集成学习的假设朴素得可疑:多个弱学习器组合成强学习器。两条路线:
- Bagging(自助聚合):对数据有放回采样 个副本,各训一棵树,多数投票。每棵树看到的是略有不同的数据集,它们的错误互不相关;投票把随机噪声平均掉——降低方差。方差公式说得很清楚: 个相关性为 、方差 的学习器取平均,方差变为 —— 增大第二项消失,但第一项被 钉死。所以:
- 随机特征子空间:每个节点分裂时只随机看 个特征——让树与树更”不一样”(进一步压低 ),这就是随机森林之于普通 Bagging 的关键一跃。
清单里还有第三条路线 Boosting(AdaBoost 加权错分样本、GBDT 拟合负梯度、XGBoost/LightGBM 的工程巅峰)——思想是串行纠错降偏差,与 Bagging 的并行投票降方差形成对位;Stacking 则是元模型再学习一层。篇幅所限,本文把 Bagging 路线走到底。
数学:投票与置换重要性
回归森林取平均、分类森林多数投票。 棵树对 的预测 ——系列代码里最短的”模型融合”。
置换重要性给森林加上可解释性:要度量特征 的价值,把测试集第 列随机打乱(切断该特征与标签的关系),准确率掉了多少就是它的重要性。比训练期的 Gini 重要性更抗偏——不需要重新训练,任何模型都能用。
Rust 实现
数据沿用 12 篇的矩形规则,但加了一个纯噪声特征 x3——置换重要性的试金石。
// src/main.rs(段一:决策树基学习器、随机森林、置换重要性与主程序)
// 终篇:集成学习——随机森林把决策树的方差熨平
// 单文件、仅标准库。基学习器是 12 篇的决策树;绘图器与系列前篇一致。
// ===================== 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
}
}
// ===================== CART 决策树(沿用 12 篇,加随机特征子空间) =====================
enum TreeNode {
Leaf(f64),
Branch {
feature: usize,
threshold: f64,
left: Box<TreeNode>,
right: Box<TreeNode>,
},
}
fn gini(ys: &[f64]) -> f64 {
if ys.is_empty() {
return 0.0;
}
let p1 = ys.iter().sum::<f64>() / ys.len() as f64;
1.0 - p1.powi(2) - (1.0 - p1).powi(2)
}
fn majority(ys: &[f64]) -> f64 {
let ones = ys.iter().filter(|&&y| y > 0.5).count();
if ones * 2 >= ys.len() { 1.0 } else { 0.0 }
}
/// 在候选特征子集上找最优切分(随机森林的"随机特征子空间")。
fn best_split_subset(
xs: &[Vec<f64>],
ys: &[f64],
candidates: &[usize],
) -> Option<(usize, f64)> {
let parent = gini(ys);
let mut best: Option<(usize, f64, f64)> = None; // (feature, threshold, drop)
for &feature in candidates {
let mut order: Vec<usize> = (0..xs.len()).collect();
order.sort_by(|&a, &b| xs[a][feature].partial_cmp(&xs[b][feature]).unwrap());
for w in 0..order.len() - 1 {
let (a, b) = (order[w], order[w + 1]);
if xs[a][feature] == xs[b][feature] {
continue;
}
let t = (xs[a][feature] + xs[b][feature]) / 2.0;
let mut gl = Vec::new();
let mut gr = Vec::new();
for (x, y) in xs.iter().zip(ys.iter()) {
if x[feature] <= t {
gl.push(*y);
} else {
gr.push(*y);
}
}
let n = ys.len() as f64;
let score = gl.len() as f64 / n * gini(&gl) + gr.len() as f64 / n * gini(&gr);
let drop = parent - score;
if drop > 1e-12 && best.map_or(true, |(_, _, d)| drop > d) {
best = Some((feature, t, drop));
}
}
}
best.map(|(f, t, _)| (f, t))
}
/// 递归建树。m_features > 0 时每个节点随机选 m 个候选特征(随机森林模式)。
fn build_tree(
xs: &[Vec<f64>],
ys: &[f64],
depth: usize,
max_depth: usize,
min_samples: usize,
m_features: usize,
rng: &mut XorShift,
) -> TreeNode {
if ys.iter().all(|&y| (y - ys[0]).abs() < 1e-12) {
return TreeNode::Leaf(ys[0]);
}
if depth >= max_depth || ys.len() < min_samples {
return TreeNode::Leaf(majority(ys));
}
// 候选特征:全部(单树)或随机子集(森林)
let candidates: Vec<usize> = if m_features > 0 && m_features < xs[0].len() {
let mut idx: Vec<usize> = (0..xs[0].len()).collect();
// 部分 Fisher-Yates 打乱后取前 m 个
for i in 0..m_features {
let j = i + ((rng.next_u64() as usize) % (idx.len() - i));
idx.swap(i, j);
}
idx[..m_features].to_vec()
} else {
(0..xs[0].len()).collect()
};
match best_split_subset(xs, ys, &candidates) {
Some((feature, threshold)) => {
let mut lx = Vec::new();
let mut ly = Vec::new();
let mut rx = Vec::new();
let mut ry = Vec::new();
for (x, y) in xs.iter().zip(ys.iter()) {
if x[feature] <= threshold {
lx.push(x.clone());
ly.push(*y);
} else {
rx.push(x.clone());
ry.push(*y);
}
}
TreeNode::Branch {
feature,
threshold,
left: Box::new(build_tree(&lx, &ly, depth + 1, max_depth, min_samples, m_features, rng)),
right: Box::new(build_tree(&rx, &ry, depth + 1, max_depth, min_samples, m_features, rng)),
}
}
None => TreeNode::Leaf(majority(ys)),
}
}
fn predict_tree(node: &TreeNode, x: &[f64]) -> f64 {
match node {
TreeNode::Leaf(label) => *label,
TreeNode::Branch { feature, threshold, left, right } => {
if x[*feature] <= *threshold {
predict_tree(left, x)
} else {
predict_tree(right, x)
}
}
}
}
// ===================== 随机森林:Bagging + 随机子空间 + 投票 =====================
/// Bagging:有放回自助采样,每棵树看到自己的副本。
fn bootstrap(xs: &[Vec<f64>], ys: &[f64], rng: &mut XorShift) -> (Vec<Vec<f64>>, Vec<f64>) {
let n = xs.len();
let mut sx = Vec::with_capacity(n);
let mut sy = Vec::with_capacity(n);
for _ in 0..n {
let i = (rng.next_u64() as usize) % n;
sx.push(xs[i].clone());
sy.push(ys[i]);
}
(sx, sy)
}
/// 训练 B 棵树的森林。m_features = ⌊log₂d⌋ + 1 是常见默认值。
fn train_forest(xs: &[Vec<f64>], ys: &[f64], b: usize, max_depth: usize, rng: &mut XorShift) -> Vec<TreeNode> {
let d = xs[0].len();
let m = ((d as f64).log2() as usize) + 1;
(0..b)
.map(|_| {
let (sx, sy) = bootstrap(xs, ys, rng);
build_tree(&sx, &sy, 0, max_depth, 2, m, rng)
})
.collect()
}
/// 多数投票:迭代器把各树预测求和过阈值。
fn predict_forest(forest: &[TreeNode], x: &[f64]) -> f64 {
let votes: f64 = forest.iter().map(|t| predict_tree(t, x)).sum();
if votes * 2.0 >= forest.len() as f64 { 1.0 } else { 0.0 }
}
fn accuracy_forest(forest: &[TreeNode], xs: &[Vec<f64>], ys: &[f64]) -> f64 {
let correct = xs.iter().zip(ys.iter()).filter(|(x, y)| predict_forest(forest, x) == **y).count();
correct as f64 / ys.len() as f64
}
fn accuracy_tree(tree: &TreeNode, xs: &[Vec<f64>], ys: &[f64]) -> f64 {
let correct = xs.iter().zip(ys.iter()).filter(|(x, y)| predict_tree(tree, x) == **y).count();
correct as f64 / ys.len() as f64
}
/// 置换重要性:打乱第 j 列后准确率掉了多少。
fn permutation_importance(forest: &[TreeNode], xs: &[Vec<f64>], ys: &[f64], j: usize, rng: &mut XorShift) -> f64 {
let base = accuracy_forest(forest, xs, ys);
let mut shuffled: Vec<Vec<f64>> = xs.to_vec();
for i in (1..shuffled.len()).rev() {
let k = (rng.next_u64() as usize) % (i + 1);
let tmp = shuffled[i][j];
shuffled[i][j] = shuffled[k][j];
shuffled[k][j] = tmp;
}
base - accuracy_forest(forest, &shuffled, ys)
}
// ===================== 数据:矩形规则 + 一个纯噪声特征 x3 =====================
fn make_data(rng: &mut XorShift, n: usize) -> (Vec<Vec<f64>>, Vec<f64>) {
let mut xs = Vec::with_capacity(n);
let mut ys = Vec::with_capacity(n);
for _ in 0..n {
let x1 = (rng.next_f64() - 0.25) * 6.0;
let x2 = (rng.next_f64() - 0.25) * 6.0;
let x3 = rng.next_f64() * 6.0 - 1.5; // 与标签无关的噪声特征
let mut label = if (1.0..=4.0).contains(&x1) && (1.0..=4.0).contains(&x2) { 1.0 } else { 0.0 };
if rng.next_f64() < 0.1 {
label = 1.0 - label;
}
xs.push(vec![x1, x2, x3]);
ys.push(label);
}
(xs, ys)
}
fn main() {
let out_dir = "../../../frontend/public/images/series/rust-ml-17-ensemble";
std::fs::create_dir_all(out_dir).expect("创建输出目录失败");
// ---- 1. 数据(12 篇的矩形规则 + 噪声特征 x3) ----
let mut rng = XorShift::new(42);
let (xtr, ytr) = make_data(&mut rng, 300);
let (xte, yte) = make_data(&mut rng, 200);
println!("train = {},test = {}(矩形规则,10% 标签噪声,x3 为纯噪声特征)", xtr.len(), xte.len());
// ---- 2. 单树对照:深树过拟合,浅树欠拟合 ----
println!("\n[单树对照] max_depth train acc test acc");
for d in [4usize, 8, 20] {
let tree = build_tree(&xtr, &ytr, 0, d, 2, 0, &mut XorShift::new(7));
println!("{:>11} {:>8.3} {:>8.3}", d, accuracy_tree(&tree, &xtr, &ytr), accuracy_tree(&tree, &xte, &yte));
}
// ---- 3. 森林规模扫描:方差随 B 收缩 ----
println!("\n[随机森林] B train acc test acc");
let mut single_test = 0.0f64;
let mut curve = Vec::new();
let mut best_b = (0usize, 0.0f64);
for b in [1usize, 5, 11, 25, 51, 101] {
let forest = train_forest(&xtr, &ytr, b, 8, &mut XorShift::new(11));
let tr = accuracy_forest(&forest, &xtr, &ytr);
let te = accuracy_forest(&forest, &xte, &yte);
println!("{:>4} {:>8.3} {:>8.3}", b, tr, te);
curve.push((b as f64, te));
if b == 1 {
single_test = te;
}
if te > best_b.1 {
best_b = (b, te);
}
}
// ---- 4. 深度鲁棒性:森林几乎不用调深度 ----
println!("\n[森林的深度鲁棒性,B=51] max_depth test acc");
for d in [4usize, 8, 20] {
let forest = train_forest(&xtr, &ytr, 51, d, &mut XorShift::new(13));
println!("{:>18} {:>8.3}", d, accuracy_forest(&forest, &xte, &yte));
}
// ---- 5. 特征重要性:x3 应≈0 ----
let forest = train_forest(&xtr, &ytr, 51, 8, &mut XorShift::new(17));
println!("\n[置换重要性](打乱后 acc 下降越多越重要)");
for j in 0..3 {
let imp = permutation_importance(&forest, &xte, &yte, j, &mut XorShift::new(19));
println!(" x{}:{:+.4}", j + 1, imp);
}
// ---- 6. 图一:B-准确率曲线(对照单树水平线) ----
let mut c = Canvas::new(560.0, 320.0, 1.0, 101.0, 0.7, 1.02);
c.axes("树的数量 B", "test acc");
c.polyline(&curve, PALETTE[0], 2.2);
let base: Vec<(f64, f64)> = vec![(1.0, single_test), (101.0, single_test)];
c.polyline(&base, PALETTE[1], 1.4);
c.dots(&[((best_b.0) as f64, best_b.1)], PALETTE[2], 5.0);
c.legend(&[("随机森林", PALETTE[0]), ("单棵自助树基线", PALETTE[1]), ("最优", PALETTE[2])]);
let p1 = format!("{out_dir}/forest-size.svg");
c.save(&p1);
// ---- 7. 图二:特征重要性 ----
let imps: Vec<(f64, f64)> = (0..3)
.map(|j| (j as f64, permutation_importance(&forest, &xte, &yte, j, &mut XorShift::new(23))))
.collect();
let mut c = Canvas::new(560.0, 320.0, -0.5, 2.5, 0.0, 0.0);
c.axes("特征(x1 / x2 / x3)", "置换重要性(acc 下降)");
for (x, v) in &imps {
c.polyline(&[(*x, 0.0), (*x, *v)], PALETTE[0], 2.4);
}
c.dots(&imps, PALETTE[2], 5.0);
let p2 = format!("{out_dir}/importance.svg");
c.save(&p2);
println!("图已写入:{p1}、{p2}");
}
// ===================== 迷你 SVG 绘图器(与系列前篇一致) =====================
实现要点:
build_tree加一个m_features参数:m > 0时每节点先做部分 Fisher-Yates 打乱选 个候选特征再找最优切分——同一棵树,两种性格(m=0是普通树,m>0是森林模式)。bootstrap是有放回采样:每棵树约只见到 63.2% 的唯一样本(其余重复),这是 Bagging 方差缩减的来源。predict_forest的投票一行搞定:forest.iter().map(predict).sum()(语法角细说迭代器)。- 默认 ——本文 时 ,每次分裂都有一个特征被随机藏起来。
// 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 bootstrap_preserves_size_with_replacement() {
let xs: Vec<Vec<f64>> = (0..50).map(|i| vec![i as f64]).collect();
let ys: Vec<f64> = (0..50).map(|i| i as f64 % 2.0).collect();
let (sx, sy) = bootstrap(&xs, &ys, &mut XorShift::new(3));
assert_eq!(sx.len(), 50);
assert_eq!(sy.len(), 50);
}
#[test]
fn forest_beats_single_deep_tree() {
let mut rng = XorShift::new(42);
let (xtr, ytr) = make_data(&mut rng, 200);
let (xte, yte) = make_data(&mut rng, 150);
let tree = build_tree(&xtr, &ytr, 0, 20, 2, 0, &mut XorShift::new(7));
let single = accuracy_tree(&tree, &xte, &yte);
let forest = train_forest(&xtr, &ytr, 51, 8, &mut XorShift::new(11));
let ens = accuracy_forest(&forest, &xte, &yte);
assert!(ens >= single - 0.03, "森林应不弱于单树: {ens} vs {single}");
}
#[test]
fn noise_feature_has_near_zero_importance() {
let mut rng = XorShift::new(5);
let (xtr, ytr) = make_data(&mut rng, 200);
let (xte, yte) = make_data(&mut rng, 150);
let forest = train_forest(&xtr, &ytr, 31, 8, &mut XorShift::new(11));
let imp_noise = permutation_importance(&forest, &xte, &yte, 2, &mut XorShift::new(19));
assert!(imp_noise.abs() < 0.05, "噪声特征重要性应近 0: {imp_noise}");
}
#[test]
fn voting_is_deterministic() {
let mut rng = XorShift::new(9);
let (xtr, ytr) = make_data(&mut rng, 80);
let forest = train_forest(&xtr, &ytr, 11, 6, &mut XorShift::new(11));
let a = predict_forest(&forest, &xtr[0]);
let b = predict_forest(&forest, &xtr[0]);
assert_eq!(a, b);
}
}
Rust 语法角:iter / iter_mut / into_iter——三种迭代视角
predict_forest 里 forest.iter() 产生只读借用的迭代器:遍历时森林仍归你所有,预测完照常可用。三兄弟的分工:iter() 迭代 &T(只看不拿),iter_mut() 迭代 &mut T(可修改),into_iter() 消耗集合本身产出 T(拿完即弃)——投票用 iter(),移动所有权用 into_iter(),混淆三者是新手期编译器报错的重灾区。见《Rust 程序设计语言》ch13-02(迭代器)。
运行结果
cargo test(4 个用例:自助采样规模、森林不弱于单树、噪声特征重要性近零、投票确定性)全部通过后,cargo run:
train = 300,test = 200(矩形规则,10% 标签噪声,x3 为纯噪声特征)
[单树对照] max_depth train acc test acc
4 0.937 0.855
8 0.990 0.810
20 1.000 0.795
[随机森林] B train acc test acc
1 0.903 0.790
5 0.960 0.860
11 0.980 0.865
25 0.990 0.880
51 0.993 0.890
101 0.990 0.875
[森林的深度鲁棒性,B=51] max_depth test acc
4 0.870
8 0.885
20 0.875
[置换重要性](打乱后 acc 下降越多越重要)
x1:+0.2050
x2:+0.2100
x3:+0.0150
图已写入:../../../frontend/public/images/series/rust-ml-17-ensemble/forest-size.svg、../../../frontend/public/images/series/rust-ml-17-ensemble/importance.svg
两张图:
怎么读这些数字和图
- 单树对照组的过拟合昭然若揭:depth 从 4 加到 20,train acc 0.937→1.000,test acc 却从 0.855 滑到 0.795——调深度是在过拟合与欠拟合的窄缝里走钢丝。
- 森林用 B 换稳定:B=1 时就是一棵自助采样的树(0.790,比单树还差——采样本身丢信息),B 增大到 51 爬到 0.890,超过单树的最好成绩 6 个百分点。曲线在 B≈25 后趋平——方差公式里 项已被熨平,剩下的 是相关性下限。
- 深度鲁棒性是森林最省心的性质:depth 4/8/20 的 test acc 是 0.870/0.885/0.875——几乎调无可调(清单:“几乎不用调参”)。单树那条陡峭的敏感曲线,被投票抹成了直线。
- 置换重要性让噪声现形:x1、x2 各贡献约 0.21 的准确率,x3 只有 0.015——一个不携带任何标签信息的特征,重要性自动归零。工程上做特征筛选,先跑一遍置换重要性。
优缺点与适用场景
(抄清单 5.1 原文)
- Bagging / 随机森林:抗过拟合、自带特征重要性、几乎不用调参;单模型可并行训练。代表:随机森林(Bagging + 随机选特征)。
- Boosting:串行训练、每轮纠正前一轮的错误,降低偏差。AdaBoost 加大错分样本权重;GBDT 每棵新树拟合负梯度、损失可定制;XGBoost / LightGBM 是高效工程实现(正则化、列采样、并行化、直方图加速)——表格数据竞赛与工业界的默认王者。
- Stacking:不同模型的输出作为元模型的输入。
小结——系列收官
十七篇到此收官。回看整条路:工具(Matrix、绘图器)→ 回归(OLS 到高斯过程)→ 分类(逻辑回归到 SVM)→ 降维(PCA)→ 聚类(k 均值到谱聚类)→ 集成(随机森林)。所有算法共享同一副骨架:一个目标函数、一种优化、一组假设——换来换去的只是这三样的配方。清单附录的那句话可以当作全系列的注脚:线性代数是底座,概率是语言,迭代是引擎。
剩下的路在清单的模块五里延伸:Boosting 的实战巅峰、流形学习的可视化魔法、HMM 与概率图模型——当某天它们也变成 cargo new 后的一个下午,这个系列就完成了它的使命。
谢谢你读到这里。代码都在仓库的 series/rust-ml/ 下,每篇一个工程,cargo run 即可复现文中每一张图。