rustyml 0.15.0

A high-performance machine learning & deep learning library in pure Rust, offering ML algorithms and neural network support
Documentation
# 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)讲述了如何调这些阈值。