ML ML Learning Lab
10 / 14
第 10 步 · 用"问问题"做分类

决策树Decision Tree

告别梯度下降 — 用"问问题"的方式做分类,分而治之

10 分钟
阅读 + 实操
1 个
交互演示
中级
难度

决策树 · 交互演示

节点总数0
叶子节点0
训练准确率--
最大深度 4
最小样本数 5

为什么学这步?

之前的分类器都是"算分数 + 梯度下降",但人类做决策往往是"问问题"——"年龄大于 40 吗?""收入高于 5 万吗?"决策树把这个直觉变成了算法。它可解释性强、能处理数值和类别特征,是随机森林、GBDT 等集成方法的基础。

📌 发生了什么

  • 从根节点开始,计算每个特征每个阈值的信息增益。
  • 选增益最大的特征做分裂,把数据一分为二。
  • 递归地在子节点上重复,直到深度上限或样本不足。

⚠️ 常见陷阱

  • 不限制深度就完美?错,单棵树几乎一定过拟合。
  • 决策边界是斜线?不是,只能画平行于坐标轴的阶梯。
  • 信息增益绝对公平?不,它偏向取值多的特征。

本章小结

  • 分而治之:每次选信息增益最大的特征 + 阈值分裂。
  • 不纯度衡量:Gini 或 Entropy。
  • 叶节点预测:该区域样本的多数类别(多数投票)。

📐 信息增益与基尼系数

决策树通过最大化信息增益或最小化基尼不纯度来选择分裂特征。

信息增益 = 父节点熵 - 子节点熵的加权平均:

基尼不纯度衡量从数据集中随机抽取两个样本类别不同的概率:

CART 算法用基尼系数(计算无需 log,更快),ID3/C4.5 用信息增益(或增益率):

🎛 最大深度对比

最大深度限制树的生长,是控制过拟合的关键超参数。

最大深度 表现 结果
max_depth = 1 只做一次分裂,过于粗糙 欠拟合
max_depth = 3 规则简洁可解释,泛化好 推荐
max_depth = 5 拟合更细,开始有风险 需剪枝
无限制 每个叶节点一个样本,完美记忆 严重过拟合

💡 剪枝策略:预剪枝(限制深度/叶节点最小样本数)或后剪枝(先生长再回缩)。

💻 信息增益计算与分裂

计算每个特征的信息增益并选择最优分裂(Python):

import math

# 计算熵
def entropy(labels):
    counts = {}
    for l in labels:
        counts[l] = counts.get(l, 0) + 1
    n = len(labels)
    h = 0.0
    for k in counts:
        p = counts[k] / n
        h -= p * math.log2(p)
    return h

# 计算信息增益
def info_gain(data, labels, feature, threshold):
    parent_h = entropy(labels)
    left, right = [], []
    for i in range(len(data)):
        if data[i][feature] <= threshold:
            left.append(labels[i])
        else:
            right.append(labels[i])
    child_h = (len(left) / len(labels)) * entropy(left) \
            + (len(right) / len(labels)) * entropy(right)
    return parent_h - child_h

# 选择最优分裂
def best_split(data, labels):
    best_gain, best_feat, best_thr = -1, 0, 0
    for f in range(len(data[0])):
        for thr in unique_values(data, f):
            gain = info_gain(data, labels, f, thr)
            if gain > best_gain:
                best_gain, best_feat, best_thr = gain, f, thr
    return {'feature': best_feat, 'threshold': best_thr, 'gain': best_gain}

📚 参考文献与延伸阅读

  • Breiman, L. et al. (1984). Classification and Regression Trees. — CART 算法的奠基著作
  • Quinlan, J. R. (1986). Induction of Decision Trees. Machine Learning — ID3 算法
  • Quinlan, J. R. (1993). C4.5: Programs for Machine Learning. — C4.5 算法与增益率
  • scikit-learn: Decision Trees — CART 的工业实现与可视化

📝 课后练习

检验你的理解——答对为止