决策树之剪枝介绍
1 课程概览
本课讲解决策树剪枝。剪枝是防止决策树过拟合的正则化方法,类似L1/L2正则化。分为预剪枝和后剪枝两种方法。预剪枝在生成过程中剪枝,后剪枝在生成完成后自底向上剪枝。
2 核心概念与定义
- 剪枝:防止决策树过拟合的正则化方法。
- 预剪枝:在决策树生成过程中,对每个节点在划分前先进行估计。
- 后剪枝:从训练集生成完整树后,自底向上考察非叶节点。
- 泛化性能:模型对未知数据的预测能力。
3 算法与模型详解
3.1 为什么需要剪枝
原因:
- 树过于茂密/繁琐
- 模型过于复杂
- 容易导致过拟合
目的:
- 防止过拟合
- 提高泛化能力
- 简化模型
3.2 剪枝原理
操作:把子树的节点全部删除,使用叶子节点替换
效果:
- 树结构变简单
- 模型更简单
- 类似L1正则化(L1可以让权重归零)
与正则化对比:
- 剪枝 ≈ L1正则化(删除节点)
- L2正则化:权重无限趋向于零
3.3 预剪枝
定义:在决策树生成过程中,对每个节点在划分前先进行估计
流程:
- 在划分前估计
- 若当前节点的划分不能带来决策树泛化性能提升
- 则停止划分
- 将当前节点标记为叶子节点
特点:
- 预先处理
- 凭经验做处理
- 节省资源
缺点:
- 可能过早停止
- 某一层对模型提升不大,但下一层可能有质的提升
- 不能考虑到每一层
3.4 后剪枝
定义:先从训练集生成一颗完整的树,然后自底向上考察非叶节点
流程:
- 生成完整决策树
- 自底向上(从叶子到根)
- 考察非叶节点
- 若该节点对应的子树替换为叶节点能带来决策树泛化提升
- 则用子树替换为叶节点
特点:
- 树完全长成后处理
- 自底向上裁剪
- 结果更精准
缺点:
- 更消耗资源
- 需要先生成完整树
3.5 预剪枝 vs 后剪枝
| 特性 | 预剪枝 | 后剪枝 |
|---|---|---|
| 时机 | 生成过程中 | 生成完成后 |
| 方向 | 自顶向下 | 自底向上 |
| 资源消耗 | 少 | 多 |
| 精准度 | 较低 | 较高 |
| 风险 | 过早停止 | 无 |
3.6 剪枝示例
数据集:好瓜判断
- 特征:色泽、根蒂、敲声、纹理、脐部、触感
- 标签:是否好瓜
预剪枝过程:
- 计算每个特征的信息增益(或基尼指数)
- 选择最优特征作为根节点
- 在划分前估计是否带来泛化提升
- 若不提升,停止划分,标记为叶子节点
后剪枝过程:
- 生成完整决策树
- 从叶子节点开始考察
- 若替换为叶节点能带来提升,则替换
- 自底向上直到根节点
3.7 剪枝方法选择
选择标准:
- 资源有限:预剪枝
- 追求精度:后剪枝
- 实际应用:通常结合使用
4 代码示例
import numpy as np
import matplotlib.pyplot as plt
from sklearn.tree import DecisionTreeClassifier
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score
from sklearn.datasets import make_classification
# 1. 生成数据
X, y = make_classification(
n_samples=1000,
n_features=20,
n_informative=5,
n_redundant=5,
n_classes=2,
random_state=42
)
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.3, random_state=42
)
print(f"训练集: {X_train.shape}")
print(f"测试集: {X_test.shape}")
# 2. 不剪枝的决策树(过拟合)
print("\n=== 不剪枝的决策树 ===")
dt_no_prune = DecisionTreeClassifier(random_state=42)
dt_no_prune.fit(X_train, y_train)
train_acc = accuracy_score(y_train, dt_no_prune.predict(X_train))
test_acc = accuracy_score(y_test, dt_no_prune.predict(X_test))
print(f"训练集准确率: {train_acc:.4f}")
print(f"测试集准确率: {test_acc:.4f}")
print(f"树的深度: {dt_no_prune.get_depth()}")
print(f"叶子节点数: {dt_no_prune.get_n_leaves()}")
# 3. 预剪枝(限制参数)
print("\n=== 预剪枝(限制参数)===")
dt_pre_prune = DecisionTreeClassifier(
max_depth=5, # 最大深度
min_samples_split=10, # 内部节点再划分最小样本数
min_samples_leaf=5, # 叶子节点最少样本数
random_state=42
)
dt_pre_prune.fit(X_train, y_train)
train_acc_pre = accuracy_score(y_train, dt_pre_prune.predict(X_train))
test_acc_pre = accuracy_score(y_test, dt_pre_prune.predict(X_test))
print(f"训练集准确率: {train_acc_pre:.4f}")
print(f"测试集准确率: {test_acc_pre:.4f}")
print(f"树的深度: {dt_pre_prune.get_depth()}")
print(f"叶子节点数: {dt_pre_prune.get_n_leaves()}")
# 4. 后剪枝(ccp_alpha参数)
print("\n=== 后剪枝(ccp_alpha)===")
dt_post_prune = DecisionTreeClassifier(
ccp_alpha=0.01, # 复杂度参数,越大剪枝越多
random_state=42
)
dt_post_prune.fit(X_train, y_train)
train_acc_post = accuracy_score(y_train, dt_post_prune.predict(X_train))
test_acc_post = accuracy_score(y_test, dt_post_prune.predict(X_test))
print(f"训练集准确率: {train_acc_post:.4f}")
print(f"测试集准确率: {test_acc_post:.4f}")
print(f"树的深度: {dt_post_prune.get_depth()}")
print(f"叶子节点数: {dt_post_prune.get_n_leaves()}")
# 5. 对比不同ccp_alpha
print("\n=== 不同ccp_alpha对比 ===")
ccp_alphas = [0, 0.001, 0.01, 0.05, 0.1]
train_accs = []
test_accs = []
depths = []
n_leaves = []
for alpha in ccp_alphas:
dt = DecisionTreeClassifier(ccp_alpha=alpha, random_state=42)
dt.fit(X_train, y_train)
train_accs.append(accuracy_score(y_train, dt.predict(X_train)))
test_accs.append(accuracy_score(y_test, dt.predict(X_test)))
depths.append(dt.get_depth())
n_leaves.append(dt.get_n_leaves())
print(f"alpha={alpha:.3f}: 训练={train_accs[-1]:.4f}, 测试={test_accs[-1]:.4f}, "
f"深度={depths[-1]}, 叶子={n_leaves[-1]}")
# 6. 可视化
plt.figure(figsize=(15, 5))
# 准确率对比
plt.subplot(1, 3, 1)
plt.plot(ccp_alphas, train_accs, 'bo-', label='训练集')
plt.plot(ccp_alphas, test_accs, 'ro-', label='测试集')
plt.xlabel('ccp_alpha')
plt.ylabel('准确率')
plt.title('准确率 vs ccp_alpha')
plt.legend()
plt.grid(True)
# 深度对比
plt.subplot(1, 3, 2)
plt.plot(ccp_alphas, depths, 'go-')
plt.xlabel('ccp_alpha')
plt.ylabel('树的深度')
plt.title('树的深度 vs ccp_alpha')
plt.grid(True)
# 叶子节点数对比
plt.subplot(1, 3, 3)
plt.plot(ccp_alphas, n_leaves, 'mo-')
plt.xlabel('ccp_alpha')
plt.ylabel('叶子节点数')
plt.title('叶子节点数 vs ccp_alpha')
plt.grid(True)
plt.tight_layout()
plt.show()
# 7. 三种情况对比
print("\n=== 三种情况对比 ===")
print(f"{'方法':<15} {'训练准确率':<12} {'测试准确率':<12} {'深度':<8} {'叶子数':<8}")
print("-" * 55)
print(f"{'不剪枝':<15} {train_acc:<12.4f} {test_acc:<12.4f} {dt_no_prune.get_depth():<8} {dt_no_prune.get_n_leaves():<8}")
print(f"{'预剪枝':<15} {train_acc_pre:<12.4f} {test_acc_pre:<12.4f} {dt_pre_prune.get_depth():<8} {dt_pre_prune.get_n_leaves():<8}")
print(f"{'后剪枝':<15} {train_acc_post:<12.4f} {test_acc_post:<12.4f} {dt_post_prune.get_depth():<8} {dt_post_prune.get_n_leaves():<8}")
输出示例:
训练集: (700, 20)
测试集: (300, 20)
=== 不剪枝的决策树 ===
训练集准确率: 1.0000
测试集准确率: 0.8733
树的深度: 15
叶子节点数: 45
=== 预剪枝(限制参数)===
训练集准确率: 0.9286
测试集准确率: 0.8900
树的深度: 5
叶子节点数: 23
=== 后剪枝(ccp_alpha)===
训练集准确率: 0.9429
测试集准确率: 0.8967
树的深度: 8
叶子节点数: 15
=== 不同ccp_alpha对比 ===
alpha=0.000: 训练=1.0000, 测试=0.8733, 深度=15, 叶子=45
alpha=0.001: 训练=0.9857, 测试=0.8900, 深度=12, 叶子=30
alpha=0.010: 训练=0.9429, 测试=0.8967, 深度=8, 叶子=15
alpha=0.050: 训练=0.9000, 测试=0.8833, 深度=3, 叶子=5
alpha=0.100: 训练=0.8714, 测试=0.8567, 深度=1, 叶子=2
=== 三种情况对比 ===
方法 训练准确率 测试准确率 深度 叶子数
-------------------------------------------------------
不剪枝 1.0000 0.8733 15 45
预剪枝 0.9286 0.8900 5 23
后剪枝 0.9429 0.8967 8 15
5 重难点与易错提醒
- ❗重点:剪枝是防止过拟合的正则化方法。
- ❗重点:预剪枝在生成过程中剪枝,后剪枝在生成完成后剪枝。
- ❗重点:预剪枝节省资源,后剪枝更精准。
- ❗重点:剪枝类似L1正则化(删除节点)。
- ⚠️易错:预剪枝可能过早停止。
- ⚠️易错:后剪枝资源消耗大。
- 💡深入理解:通过参数控制树的复杂度。
6 课堂问答精选
Q: 预剪枝和后剪枝有什么区别?
A:
- 预剪枝:在决策树生成过程中,对每个节点在划分前先进行估计,若不能带来泛化性能提升,则停止划分。节省资源,但可能过早停止。
- 后剪枝:先生成完整树,然后自底向上考察非叶节点,若替换为叶节点能带来提升,则替换。资源消耗大,但结果更精准。
Q: 剪枝和正则化有什么关系?
A: 剪枝是一种正则化方法,类似L1正则化。L1正则化可以让权重归零(删除特征),剪枝也是删除节点,用叶子节点替换。两者都是通过简化模型来防止过拟合。
7 本课小结
- 剪枝:防止决策树过拟合的正则化方法。
- 预剪枝:生成过程中剪枝,节省资源。
- 后剪枝:生成完成后自底向上剪枝,更精准。
- 剪枝类似L1正则化。
- API:max_depth、min_samples_split、ccp_alpha等参数。
8 延伸思考与实践
- 实践:对比不同剪枝参数的效果。
- 预习:集成学习。
- 思考:如何选择合适的剪枝参数?