# 2.4. 决策树
决策树用一串 if/else 判断,把特征空间切成一个个轴对齐的方块。每个方块给出一个常数预测:分类时是类别标签(或类别分布),回归时是目标的均值。RustyML 把这一切装进一个 `DecisionTree` 类型。你用 `Algorithm` 参数选择算法:`ID3`、`C45` 或 `CART`。
在 scikit-learn 里,你选择 `DecisionTreeClassifier` 或 `DecisionTreeRegressor`,再用 `criterion=` 字符串挑选准则。RustyML 把算法作为最顶层的选择。每个 `Algorithm` 变体把不纯度度量、划分选取规则和类别列的处理策略打包在了一起。
数值特征总是按二元阈值划分。`feature <= t` 的样本走左子树。算法只改变阈值如何打分,以及类别列是否展开成多路分支。
`DecisionTree` 位于 `rustyml::machine_learning` 之下,遵循[经典机器学习](./2.0._经典机器学习.md)里介绍的 `new` -> `fit` -> `predict` 约定。
## 2.4.1. 选择算法:ID3、C4.5、CART
`Algorithm` 枚举恰好有 3 个变体。这些区别不是表面文章。每个变体都会改变这棵树能做什么,以及如何给划分打分。
| `ID3` | 仅分类 | 香农熵 | 信息增益(原始不纯度下降量) | 多路(每个取值一个分支) |
| `C45` | 仅分类 | 香农熵 | 增益**率**(增益 / 划分信息) | 多路 |
| `CART` | 分类**与**回归 | 分类用 Gini,回归用 MSE | 原始不纯度下降量 | 仅二元 |
这张表给出 2 个结论。第一,只有 `CART` 支持回归。`DecisionTree::new(Algorithm::ID3, false)` 和 `DecisionTree::new(Algorithm::C45, false)` 会立刻以 `Error::InvalidInput` 失败,而不是等到 fit 时才报错。构造函数就是这项检查的快速失败关卡。
第二,C4.5 的增益率修正了 ID3 偏向高基数特征的毛病。信息增益偏爱取值众多的特征。极端情况下,一列每行都唯一的 ID 列能拿到满分增益,却完全无法泛化。C4.5 把增益除以划分的*固有信息*,用来惩罚那种散成许多细小分支的划分。当固有信息降到零(退化成单分支的划分)时,C4.5 会拒绝这次划分。当数据集里混着基数相差很大的类别列时,增益率是更稳妥的默认选择。
对于没有类别列的纯数值数据,这 3 种算法都会退化成同一套贪心的二元阈值搜索。它们的结果通常一致。除非你需要基于熵的打分或多路类别分支,否则直接用 `CART`。
## 2.4.2. 分类:构造、训练、预测
分类走的是 `new(algorithm, true)`。标签必须是编码为 `f64` 的非负整数。标签必须从 `0` 开始连续。RustyML 按 `max(label) + 1` 推断类别数。像 `{0, 5}` 这样的标签集会分配出 6 个类别,其中 4 个是空的、用不上。请在训练前把标签编码到 `0..k-1`(见[标签编码](../Chapter-04/4.3._标签编码.md))。
```rust
use rustyml::machine_learning::{Algorithm, DecisionTree};
use ndarray::array;
fn main() {
// 2 个特征,3 个类别,沿 feature 0 干净地分开
let x = array![
[0.0, 0.0],
[0.1, 0.0],
[0.2, 0.1],
[10.0, 1.0],
[10.1, 1.0],
[20.0, 2.0],
[20.1, 2.0],
];
let y = array![0.0, 0.0, 0.0, 1.0, 1.0, 2.0, 2.0];
let mut tree = DecisionTree::new(Algorithm::CART, true)
.unwrap()
.with_max_depth(5)
.with_random_state(42);
tree.fit(&x, &y).unwrap();
let x_test = array![[0.05, 0.0], [10.2, 1.0], [20.2, 2.0]];
let labels = tree.predict(&x_test).unwrap(); // Array1<f64>,每行 1 个标签
let proba = tree.predict_proba(&x_test).unwrap(); // Array2<f64>
println!("labels: {:?}", labels);
println!("proba shape: {:?}", proba.shape()); // [3, 3] = (n_samples, n_classes)
println!("n_classes: {:?}", tree.get_n_classes()); // Some(3)
}
```
`predict` 返回一个 `Array1<f64>` 类型的标签数组。`predict_proba` 返回一个 `Array2<f64>`。每一行是叶子的经验类别分布,因此每行加和为 1.0。每行的 argmax 与对应的 `predict` 标签一致。
纯叶子上,分布是 one-hot 的。不纯叶子上,它保存的是落到该叶子的训练样本的精确类别频率。只有提前停止生长才会出现不纯叶子(见下一节)。
`DecisionTree` 还有 2 个单样本方法:`predict_one(&[f64]) -> f64` 和 `predict_proba_one(&[f64]) -> Vec<f64>`。`fit_predict` 训练树,然后用 1 次调用对同一个训练矩阵做预测。
对回归树调用 `predict_proba` 是运行时错误,会返回 `Error::Tree(TreeError::NotClassificationTree)`。
## 2.4.3. 用 CART 做回归
把 `is_classifier = false`,树就切换到 MSE 不纯度。MSE 不纯度是节点目标值的总体方差。每个叶子预测它那批训练目标的均值。
这条规则解释了回归树标志性的阶梯状输出。预测是分段常数,每个叶子对应 1 段平台。回归树永远不会外推到训练范围之外。
```rust
use rustyml::machine_learning::{Algorithm, DecisionTree};
use ndarray::array;
fn main() {
// 阶跃函数:目标值在 x = 2 到 x = 10 之间某处从 1 跳到 10
let x = array![[0.0], [1.0], [2.0], [10.0], [11.0], [12.0]];
let y = array![1.0, 1.0, 1.0, 10.0, 10.0, 10.0];
// 只有 CART 支持回归
let mut tree = DecisionTree::new(Algorithm::CART, false).unwrap();
tree.fit(&x, &y).unwrap();
// 每次查询落到某个叶子,拿到该叶子的均值目标
let preds = tree.predict(&array![[1.5], [11.5]]).unwrap();
println!("{:?}", preds); // [1.0, 10.0],2 个叶子的均值
}
```
请求非 CART 的回归器会在构造时就失败。这是 RustyML 唯一一处急切校验算法与任务组合的地方:
```rust,ignore
// ID3 和 C4.5 只能分类,所以这里会快速失败:
let err = DecisionTree::new(Algorithm::ID3, false).unwrap_err();
// -> Error::InvalidInput("Only CART algorithm is supported for regression tasks")
```
## 2.4.4. 超参数与过拟合控制
不加任何限制地生长,树会一直划分,直到每个叶子都纯(分类)或只剩 1 个样本(回归)。这样的树会把训练集连噪声一起背下来。这是决策树的经典翻车方式。
下面这 5 个生长参数,是你在训练拟合与泛化之间权衡的全部工具。其中 4 个是生长过程中生效的预剪枝停止规则。RustyML 不做后剪枝,没有 `ccp_alpha` 那类代价复杂度剪枝,所以所有的过拟合控制都发生在训练开始之前。
| `with_max_depth` | `max_depth: Option<usize>` | `None`(无限制) | 不会失败,返回 `Self` |
| `with_min_samples_split` | `min_samples_split: usize` | `2` | `>= 2`,否则 `InvalidParameter` |
| `with_min_samples_leaf` | `min_samples_leaf: usize` | `1` | `>= 1`,否则 `InvalidParameter`(还必须 `<= min_samples_split`) |
| `with_min_impurity_decrease` | `min_impurity_decrease: f64` | `0.0` | 非负且有限,否则 `InvalidParameter` |
| `with_random_state` | `random_state: Option<u64>` | `None` | 不会失败,返回 `Self` |
带取值范围的 setter 返回 `Result`。链式调用时,在它们后面加上 `.unwrap()` 或 `?`。`with_max_depth` 和 `with_random_state` 直接返回 `Self`,不涉及 `Result`。
有一条规则同时涉及 `min_samples_leaf` 和 `min_samples_split`。因为这两个值是独立设置的,单个 builder 没法单独强制这条规则。`min_samples_leaf` 不能超过 `min_samples_split`。RustyML 在 `fit` 时检查这条约束,遇到不合理的搭配就返回 `Error::InvalidParameter`。不合理的搭配会在训练时失败,而不是构造时。
每个参数否决候选划分的理由各不相同:
- **`max_depth`** 限制任意一条从根到叶路径上的边数。`Some(0)` 强制根节点本身成为叶子。这个叶子预测全局多数类(或全局均值),是一个方便的健全性基线。
- **`min_samples_split`** 当一个节点持有的样本数少于这个值时,禁止它继续划分。它会提前结束递归,得到一棵更浅的树。
- **`min_samples_leaf`** 约束的是划分**搜索**本身,而不只是最终选取。RustyML 根本不会考虑那些会让子节点低于这个下限的阈值。树不会因此塌掉这个节点,而是退而求其次,选一个两个子节点都满足下限的最佳划分。这与 scikit-learn 的语义一致。它是大家最容易误读的参数:一个罕见类别或一个孤立离群点不会因此毁掉一个本来不错的划分。
- **`min_impurity_decrease`** 否决那些不纯度下降量(按 `N_t / N_total` 缩放)低于阈值的划分。`N_t / N_total` 是到达该节点的样本占全部训练样本的比例。这种按节点权重缩放的做法沿用了 scikit-learn 的约定。树深处剩下的样本已经不多,哪怕有很大的不纯度下降,也会比根部附近同样幅度的下降打更多折扣。
下面这个例子,在 2 个故意打错标签的样本上,对比过拟合的树与受约束的树。无约束的树会多长几层来隔离噪声,把训练准确率拉到 100%。受约束的树保持浅层、让每个叶子都住着足够样本,所以它不会去拟合噪声:
```rust
use rustyml::machine_learning::{Algorithm, DecisionTree, Node, NodeType};
use ndarray::{array, Array1, Array2};
// 从根到叶最长路径的边数(光秃秃的叶子深度为 0)
fn depth(node: &Node) -> usize {
match &node.node_type {
NodeType::Leaf { .. } => 0,
NodeType::Internal { .. } => {
let mut d = 0;
if let Some(l) = &node.left {
d = d.max(depth(l));
}
if let Some(r) = &node.right {
d = d.max(depth(r));
}
if let Some(children) = &node.children {
for c in children.values() {
d = d.max(depth(c));
}
}
1 + d
}
}
}
fn train_accuracy(tree: &DecisionTree, x: &Array2<f64>, y: &Array1<f64>) -> f64 {
let preds = tree.predict(x).unwrap();
let correct = preds
.iter()
.zip(y.iter())
.filter(|(p, t)| (*p - *t).abs() < 0.5)
.count();
correct as f64 / y.len() as f64
}
fn main() {
// 底层规则:x <= 4 -> 类别 0,x >= 5 -> 类别 1。
// 2 个标签违反了它:x = 1 和 x = 8 是翻转后的噪声。
let x = array![[0.0], [1.0], [2.0], [3.0], [4.0], [5.0], [6.0], [7.0], [8.0], [9.0]];
let y = array![0.0, 1.0, 0.0, 0.0, 0.0, 1.0, 1.0, 1.0, 0.0, 1.0];
// 无约束:一直长到纯,把噪声也背下来
let mut overfit = DecisionTree::new(Algorithm::CART, true).unwrap();
overfit.fit(&x, &y).unwrap();
// 受约束:只划分 1 层,每个叶子必须保留 >= 2 个样本
let mut constrained = DecisionTree::new(Algorithm::CART, true)
.unwrap()
.with_max_depth(1)
.with_min_samples_leaf(2)
.unwrap();
constrained.fit(&x, &y).unwrap();
println!(
"unconstrained: depth {}, train acc {:.3}",
depth(overfit.get_root().unwrap()),
train_accuracy(&overfit, &x, &y)
);
println!(
"constrained: depth {}, train acc {:.3}",
depth(constrained.get_root().unwrap()),
train_accuracy(&constrained, &x, &y)
);
}
```
无约束的树给出更深的结构和完美的训练准确率。受约束的树只有 1 次划分,还刻意保留了 2 个错误。面对没见过的数据,你要的是那棵浅树。
真正的调参会拿这些参数在验证集划分上做扫描(见[训练集与测试集划分](../Chapter-04/4.1._训练集与测试集划分.md)),再用[分类指标](../Chapter-05/5.2._分类指标.md)打分。
## 2.4.5. 类别特征、缺失值与常量列
RustyML 没有单独的类别数据类型,每个特征都是一列 `f64`。你用 `set_categorical_features(vec![...])` 来声明哪些列存的是离散的类别编码。这是一个 `&mut self` 的 setter。在 `new` 和 `fit` 之间调用它。
在 `ID3` 或 `C45` 下,声明过的列按多路划分,每个不同取值对应 1 个子分支。多路划分能解出单个数值阈值搞不定的模式,比如"编码等于 1 时才是类别 1"。`CART` 天生就是二元的,完全忽略这个声明。在 CART 树上,标记某列为类别列不会有任何效果。RustyML 不会给出任何提示。
```rust
use rustyml::machine_learning::{Algorithm, DecisionTree};
use ndarray::array;
fn main() {
// 1 个类别特征(编码 0/1/2)。类别取决于编码,而非任何
// 单个数值切分:0 -> 类别 0,1 -> 类别 1,2 -> 类别 0。
let x = array![[0.0], [0.0], [1.0], [1.0], [2.0], [2.0]];
let y = array![0.0, 0.0, 1.0, 1.0, 0.0, 0.0];
let mut tree = DecisionTree::new(Algorithm::C45, true).unwrap();
tree.set_categorical_features(vec![0]); // 把第 0 列当作类别列
tree.fit(&x, &y).unwrap();
// 每个不同编码都成为自己的分支,因此模式被精确学到
println!("train preds: {:?}", tree.predict(&x).unwrap());
// 训练中从未见过的类别会走到节点的兜底叶子:不报错
println!("unseen code 99 -> {:?}", tree.predict(&array![[99.0]]).unwrap());
}
```
RustyML 通过四舍五入到 6 位小数来归一化类别取值。`1.0000001` 和 `1.0000002` 会并入同一个分支,而 `1.0` 和 `2.0` 仍然区分开。把编码写成干净的、存在 `f64` 里的整数,可以避免意外。
预测时,没见过的类别匹配不上任何分支。它会落到一个预先存好的兜底叶子,也就是父节点训练样本的多数预测。你总能得到一个有效预测,绝不会报错。
即便在 `ID3` 或 `C45` 下,`min_samples_leaf` 仍然守着类别划分。RustyML 只要至少有 2 个分支各自满足叶子下限,就会保留这次划分。某个样本太少的罕见类别,不会把整个多路划分作废。
**缺失值。** RustyML 没有针对 NaN 的路由。`fit` 和 `predict` 都会检查每一项输入是否有限。特征矩阵里任何 `NaN` 或无穷都会返回 `Error::NonFinite`。请在训练前填补或删除缺失项(见[数据预处理](../Chapter-04/4.0._数据预处理.md))。
**常量列。** 常量列是无害的。数值划分只在 2 个不同的特征值之间才成立。一列从不变化就产生不出候选阈值,于是 RustyML 会直接跳过它。一个声明为类别列、却不足 2 个不同取值的列,同样产生不出划分。常量特征会多花一点搜索时间,但绝不会破坏这棵树。
## 2.4.6. 检视训练好的树
训练好的树很容易检视。`generate_tree_structure()` 返回一份训练好的树的 ASCII 渲染,拿来就能打印:划分、阈值、叶子类别和概率向量一应俱全。训练前调用它会返回 `Error::NotFitted`。
```rust
use rustyml::machine_learning::{Algorithm, DecisionTree};
use ndarray::array;
fn main() {
let x = array![[0.0], [1.0], [2.0], [3.0]];
let y = array![0.0, 0.0, 1.0, 1.0];
let mut tree = DecisionTree::new(Algorithm::CART, true).unwrap();
tree.fit(&x, &y).unwrap();
print!("{}", tree.generate_tree_structure().unwrap());
println!("features seen: {}", tree.get_n_features());
}
```
要在程序里检视,`get_root() -> Option<&Node>` 会返回原始的树。`Node` 是公开的,带有 `node_type`、`left`、`right` 和 `children`(多路类别节点用的 `AHashMap<String, Box<Node>>`)几个字段。`NodeType` 要么是 `Internal { feature_index, threshold, categories }`,要么是 `Leaf { value, class, probabilities }`。
[2.4.4](#244-超参数与过拟合控制) 里的深度辅助函数,就是靠遍历这个结构来工作的。你也可以遍历它,来提取类似特征重要性的统计量,或把树导出成别的格式。
其余的 getter 都是朴素的访问器。它们是 `get_algorithm()`、`get_is_classifier()`、`get_n_features()`、`get_n_classes()`(回归时为 `None`)、`get_parameters()`(一个 `Copy` 的 `DecisionTreeParams`)和 `get_categorical_features()`。
## 2.4.7. 确定性与随机种子
生长是贪心且确定的,只有 1 个例外。当 2 个或更多候选划分的选取分数分毫不差地打平时,树必须打破这个平局。
在 `random_state = None` 且没有设置 crate 级种子时,打破平局是确定的。最后一个打平的候选胜出。在同一份数据上反复训练,会得到逐位一致的树。即便数据里有平局,你也不需要种子来保证可复现。
设置 `with_random_state(seed)`,则会改为从一个带种子的随机流里均匀随机地挑一个打平的候选。同一个种子会复现同一棵树。不同种子可能挑到不同的、得分相同的特征。
```rust
use rustyml::machine_learning::{Algorithm, DecisionTree};
use ndarray::array;
fn main() {
// feature 0 和 feature 1 是完全相同的列,所以它们的最佳划分正好打平
let x = array![[0.0, 0.0], [0.0, 0.0], [1.0, 1.0], [1.0, 1.0]];
let y = array![0.0, 0.0, 1.0, 1.0];
let fit_seeded = |seed: u64| {
let mut t = DecisionTree::new(Algorithm::CART, true)
.unwrap()
.with_random_state(seed);
t.fit(&x, &y).unwrap();
t.generate_tree_structure().unwrap()
};
// 相同种子 -> 相同的树,即便平局是随机打破的
assert_eq!(fit_seeded(7), fit_seeded(7));
println!("seed 7 is reproducible");
}
```
`random_state = Some(seed)` 会独立使用那个种子,忽略任何全局种子。而停在 `random_state = None` 的树,在有活跃的线程本地全局流时,会从中获取打破平局所需的随机性。调用 `set_global_seed(s)`(配合 `clear_global_seed()`),可以让一整条由无种子模型组成的流水线一起变得可复现。
[可复现性与随机种子](../Chapter-07/7.1._可复现性与随机种子.md)讲述了这套机制,以及把所有随机性都汇入同一个种子的理由。
## 2.4.8. 错误
树专属的失败都在 `TreeError` 枚举里,重新导出为 `rustyml::machine_learning::TreeError`。你通过 crate 级的 `Error::Tree` 变体触及它。它是 `#[non_exhaustive]` 的,有 2 个成员。
`NotClassificationTree` 在你对回归树调用 `predict_proba` 或 `predict_proba_one` 时返回。`CorruptStructure(&'static str)` 守护的是一次不变量违背。正常训练并使用的模型永远不会触发它。它保护的是树的遍历过程,防的是手工搭建或以其他方式损坏的节点图。
其余的错误都来自[错误处理](../Chapter-01/1.6._错误处理.md)描述的共享错误面:
- `InvalidInput`:标签有误、样本太少,或零特征。
- `InvalidParameter`:builder 取值越界,或 `min_samples_leaf > min_samples_split` 的交叉校验。
- `NonFinite`:特征里有 NaN 或无穷。
- `NotFitted`:训练前调用会返回它。
- `EmptyInput`:训练或预测数据为空。
- `DimensionMismatch { expected, found }`:预测矩阵的特征数不对。
```rust
use rustyml::machine_learning::{Algorithm, DecisionTree, TreeError};
use rustyml::error::Error;
use ndarray::array;
fn main() {
let x = array![[0.0], [1.0], [2.0], [10.0], [11.0], [12.0]];
let y = array![1.0, 1.0, 1.0, 10.0, 10.0, 10.0];
let mut reg = DecisionTree::new(Algorithm::CART, false).unwrap();
reg.fit(&x, &y).unwrap();
// 回归树没有类别概率可报告
match reg.predict_proba(&x) {
Err(Error::Tree(TreeError::NotClassificationTree)) => {
println!("predict_proba is classification-only");
}
other => panic!("unexpected: {:?}", other),
}
}
```
## 2.4.9. 持久化
训练好的 `DecisionTree` 用 `save_to_path` 和 `load_from_path` 序列化。这两个方法通过 [serde](https://serde.rs) 读写紧凑的 [postcard](https://docs.rs/postcard) 二进制格式。整个结构都能完整往返,包括多路类别节点的 `AHashMap` 子节点。加载回来的模型能精确复现预测。
文件名随你取。不管扩展名是什么,内容都是二进制。
```rust
use rustyml::machine_learning::{Algorithm, DecisionTree};
use ndarray::array;
fn main() {
let x = array![[0.0], [1.0], [2.0], [10.0], [11.0], [12.0]];
let y = array![0.0, 0.0, 0.0, 1.0, 1.0, 1.0];
let mut tree = DecisionTree::new(Algorithm::CART, true).unwrap();
tree.fit(&x, &y).unwrap();
tree.save_to_path("dt_model.bin").unwrap();
let loaded = DecisionTree::load_from_path("dt_model.bin").unwrap();
// 预测结果逐位挺过这趟往返
assert_eq!(tree.predict(&x).unwrap(), loaded.predict(&x).unwrap());
println!("round trip OK");
std::fs::remove_file("dt_model.bin").unwrap();
}
```
读写失败或数据损坏会返回 `Error::Io`。格式的保证与版本注意事项见[深入模型持久化](../Chapter-07/7.2._深入模型持久化.md)。
## 2.4.10. 复杂度与并行
生长一个节点时,会把每个特征的取值排 1 次序,再带着滚动的不纯度统计量扫一遍。一个覆盖 `n` 个样本的节点,花费 `O(n_features * n log n)`。对一棵还算平衡的树,整体累加下来大约是 `O(n_features * n log^2 n)`。这是 CART 类学习器的标准开销。这也是宽(特征众多)数据集会主导训练预算的原因。
预测是一趟从根到叶的行走,每个样本 `O(depth)`。训练好的树并不存自己的深度。并行的门槛改用大约 16 个节点的行走作为替代估计值。
RustyML 会运行 2 处独立的 rayon 并行,一旦工作量越过校准好的门槛就会生效。`fit` 期间,当 `n_samples * n_features` 越过排序-扫描门槛时,逐特征的划分搜索就会并行执行。`predict` 或 `predict_proba` 期间,当样本数越过树遍历门槛时,逐样本的遍历就会并行执行。
规模小的问题保持单线程,以免白搭并行开销。你无需做任何事来启用它。这些门槛只挑选策略。它们绝不会改变结果。[性能调优与并行](../Chapter-07/7.3._性能调优与并行.md)讲述了如何调这些阈值。