共计 2543 个字符,预计需要花费 7 分钟才能阅读完成。
背景痛点:为什么需要后剪枝
决策树是一种直观且易于解释的机器学习模型,但在实际应用中很容易出现过拟合问题。过拟合的表现主要包括:

- 训练集准确率极高(接近 100%),但测试集准确率骤降
- 决策树深度过大,生成大量只覆盖个别样本的叶子节点
- 模型对训练数据中的噪声过度敏感
预剪枝(如限制树深度、设置叶子节点最小样本数)虽然简单,但存在明显局限性:
- 提前停止生长可能欠拟合
- 难以找到全局最优的停止条件
- 对数据分布变化敏感
后剪枝方法对比
常见的后剪枝方法主要有三种:
- REP(Reduced Error Pruning):使用验证集评估剪枝效果
- PEP(Pessimistic Error Pruning):基于统计修正的误差估计
- CCP(Cost-Complexity Pruning):基于代价复杂度平衡
其中 CCP 是 scikit-learn 实现的方法,其核心公式为:
$$R_\alpha(T) = R(T) + \alpha|T|$$
- $R(T)$:子树 T 的误差率
- $|T|$:子树 T 的叶子节点数
- $\alpha$:调节系数,控制复杂度惩罚力度
当 $\alpha=0$ 时不剪枝,随着 $\alpha$ 增大,算法会逐步剪掉对整体误差影响最小的子树。
CCP 剪枝 Python 实现
1. 训练初始决策树
from sklearn.tree import DecisionTreeClassifier
from sklearn.datasets import load_breast_cancer
from sklearn.model_selection import train_test_split
# 加载乳腺癌数据集
X, y = load_breast_cancer(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)
# 训练最大深度决策树(先不过剪枝)clf = DecisionTreeClassifier(random_state=42)
clf.fit(X_train, y_train)
print(f"初始树深度:{clf.get_depth()}, 叶子节点数:{clf.get_n_leaves()}")
2. 获取剪枝路径
# 获取 CCP 路径
path = clf.cost_complexity_pruning_path(X_train, y_train)
ccp_alphas, impurities = path.ccp_alphas, path.impurities
print(f"Alpha 取值范围:{ccp_alphas.min():.4f} 到 {ccp_alphas.max():.4f}")
3. 交叉验证选择最优 alpha
import matplotlib.pyplot as plt
from sklearn.model_selection import cross_val_score
# 测试不同 alpha 下的准确率
alpha_scores = []
for alpha in ccp_alphas:
tree = DecisionTreeClassifier(random_state=42, ccp_alpha=alpha)
scores = cross_val_score(tree, X_train, y_train, cv=5)
alpha_scores.append(scores.mean())
# 绘制准确率曲线
plt.plot(ccp_alphas, alpha_scores, marker='o')
plt.xlabel('alpha')
plt.ylabel('CV Accuracy')
plt.title('Accuracy vs alpha')
plt.show()
# 选择最佳 alpha
optimal_alpha = ccp_alphas[alpha_scores.index(max(alpha_scores))]
print(f"最优 alpha 值:{optimal_alpha:.4f}")
4. 使用最优 alpha 训练最终模型
# 训练剪枝后的决策树
pruned_tree = DecisionTreeClassifier(random_state=42, ccp_alpha=optimal_alpha)
pruned_tree.fit(X_train, y_train)
print(f"剪枝后深度:{pruned_tree.get_depth()}, 叶子节点:{pruned_tree.get_n_leaves()}")
避坑指南
类别不平衡处理
当数据类别不平衡时,建议:
- 使用 class_weight 参数平衡类别权重
- 改用 F1-score 作为剪枝评估指标
- 对少数类样本进行过采样
可解释性维护
剪枝后可能影响模型可解释性,可通过:
- 导出决策树图形可视化
- 记录重要特征的变化
- 限制最大剪枝程度
与集成学习的兼容性
RandomForest 等集成方法本身具有抗过拟合特性:
- 单个决策树不需要深度剪枝
- 可适当减小 max_depth
- 优先调整 n_estimators 参数
性能验证
在乳腺癌数据集上对比结果:
| 指标 | 原始树 | 剪枝树 |
|---|---|---|
| 测试集准确率 | 0.912 | 0.924 |
| F1-score | 0.931 | 0.941 |
| ROC-AUC | 0.963 | 0.972 |
| 树深度 | 7 | 4 |
可以看到剪枝后在保持准确率的同时显著简化了模型结构。
特征重要性可视化
import seaborn as sns
# 获取特征重要性
importances = pruned_tree.feature_importances_
feat_names = load_breast_cancer().feature_names
# 绘制条形图
plt.figure(figsize=(10,6))
sns.barplot(x=importances, y=feat_names)
plt.title('Feature Importance After Pruning')
plt.show()
总结
CART 决策树的后剪枝是解决过拟合的有效方法:
- CCP 剪枝通过平衡误差和复杂度找到最优模型
- 需要交叉验证选择最佳 alpha 参数
- 实际应用中要考虑业务场景的特殊需求
- 剪枝后的模型通常具有更好的泛化能力
建议在树深度超过 5 层或测试集表现明显差于训练集时考虑使用后剪枝技术。
正文完
