【机器学习入门】7.3 剪枝:一文读懂决策树剪枝的核心逻辑与实践
对于刚入门机器学习的同学来说,“如何让模型在新数据上表现更好” 是绕不开的核心问题。而决策树作为经典且易理解的模型,其 “剪枝” 操作正是解决模型泛化能力不足的关键手段。今天我们就从基础概念出发,一步步拆解剪枝的原理、策略及对比,帮你彻底搞懂这一重要技术。
一、先搞懂:剪枝的 “前置知识”
在学习剪枝前,我们需要先明确几个核心概念 —— 它们是判断模型好坏、理解剪枝必要性的基础。
1. 错误率与精度:模型 “准不准” 的直观指标
- 错误率:当我们用模型对
m个样本进行分类时,若有α个样本分类错误,错误率E = α/m。比如用模型判断 10 个 “好瓜”,有 3 个判错,错误率就是 30%。 - 精度:精度是错误率的互补指标,计算公式为
精度 = 1 - 错误率。上面的例子中,精度就是 70%,代表模型分类正确的样本占比。
2. 误差:模型 “差在哪” 的关键衡量
- 误差定义:学习器的实际预测输出与样本真实输出的差异,就是误差。比如样本实际是 “好瓜”,模型却判为 “坏瓜”,这就产生了误差。
- 训练误差(经验误差):模型在训练集上产生的误差。比如用 100 个训练样本训练后,模型在这 100 个样本上的误差就是训练误差。
- 泛化误差:模型在新样本(没见过的样本)上产生的误差。机器学习的终极目标,是让泛化误差尽可能小 —— 毕竟我们希望模型能处理实际场景中的新数据,而不是只 “死记硬背” 训练数据。
二、为什么需要剪枝?先解决 “过拟合” 问题
剪枝的核心目的,是解决决策树的 “过拟合” 问题。要理解剪枝,就得先搞懂过拟合和欠拟合的区别。
1. 过拟合:模型 “学偏了”
过拟合是指模型把训练样本自身的 “特殊特点”,当成了所有潜在样本的 “普遍规律”,最终导致泛化性能下降。比如用 “树叶” 样本训练模型时:
- 训练样本里的树叶都有锯齿,模型就误以为 “有锯齿的才是树叶”,遇到没锯齿的树叶(新样本)就会判错;
- 训练样本里的树叶都是绿色,模型就误以为 “绿色的都是树叶”,遇到绿色的非树叶(如绿色纸片)也会判错。这种 “学太细、记太死” 的情况,就是过拟合。
2. 欠拟合:模型 “没学会”
与过拟合相反,欠拟合是指模型对训练样本的 “普遍规律” 都没学好。比如用树叶样本训练时,模型只学会了 “绿色的可能是树叶”,但连 “有叶脉、有叶柄” 这些关键特征都没掌握,遇到绿色的非树叶(如绿色塑料片)就会误判为树叶。
3. 剪枝的作用:给过拟合 “降温”
决策树在生成过程中,若不加以控制,会不断细分特征、生成复杂的分支,很容易陷入过拟合(比如把训练样本里某个特殊树叶的 “斑点” 当成判断树叶的关键特征)。而剪枝就是通过 “剪掉” 不必要的分支,让模型回归到更通用的规律上,从而降低泛化误差。
三、剪枝的两种核心策略:预剪枝与后剪枝
剪枝的核心逻辑是 “判断分支是否提升泛化性能”,但根据判断时机的不同,分为预剪枝和后剪枝两种策略。我们结合 “好瓜分类” 的案例(特征包括色泽、根蒂、敲声、纹理、脐部、触感),具体拆解两种策略的操作流程。
1. 预剪枝:“先判断,再生长”
预剪枝的思路是 “防患于未然”—— 在决策树生成过程中,每个结点在划分特征前,先通过验证集判断:如果划分后模型的泛化性能(用验证集精度衡量)没有提升,就停止划分,把当前结点标记为叶结点。
预剪枝的实际操作案例
以 “好瓜分类” 的决策树生成为例:
- 案例 1:某结点划分前,验证集精度为 42.9%;划分后(比如按 “脐部 = 凹陷 / 平坦 / 微凹” 划分),验证集精度提升到 71.4%。此时判断 “划分有价值”,继续生成分支。
- 案例 2:另一结点划分前,验证集精度已达 71.4%;划分后(比如按 “根蒂 = 蜷缩 / 稍蜷 / 硬挺” 划分),验证集精度降至 57.1%。此时判断 “划分无价值”,停止划分,将该结点标记为叶结点。
预剪枝的核心特点
- 优点:提前停止分支生长,避免生成过于复杂的决策树,计算成本低、效率高。
- 缺点:可能存在 “欠剪枝” 风险 —— 有些分支当前划分精度没提升,但后续划分可能带来性能改善,预剪枝会直接切断这种可能,导致模型欠拟合。
2. 后剪枝:“先长全,再修剪”
后剪枝的思路是 “先把树长完整,再回头优化”—— 先基于训练集生成一棵完整的决策树(不管是否过拟合),然后自底向上遍历每个非叶结点,判断:如果把该结点对应的子树替换成叶结点,验证集精度能提升,就执行剪枝;否则保留子树。
后剪枝的实际操作案例
同样以 “好瓜分类” 的完整决策树为例:
- 案例 1:某非叶结点(如 “纹理 = 清晰 / 稍糊 / 模糊” 对应的子树)剪枝前,验证集精度为 42.9%;将子树替换为叶结点后,精度提升到 57.1%。此时判断 “剪枝有价值”,执行剪枝。
- 案例 2:另一非叶结点(如 “色泽 = 青绿 / 乌黑 / 浅白” 对应的子树)剪枝前,验证集精度为 71.4%;剪枝后精度仍为 71.4%(无提升)。此时判断 “剪枝无价值”,保留子树。
- 案例 3:某顶层结点(如 “脐部 = 凹陷” 对应的子树)剪枝前,验证集精度为 57.1%;剪枝后精度提升到 71.4%,执行剪枝,最终得到更简洁的决策树。
后剪枝的核心特点
- 优点:基于完整决策树优化,能更精准地剪掉过拟合的分支,泛化性能通常比预剪枝更好,很少出现欠拟合。
- 缺点:需要先生成完整决策树,再遍历修剪,计算成本比预剪枝高,效率较低。
四、预剪枝与后剪枝的全面对比
为了帮大家更清晰地选择剪枝策略,我们从 4 个关键维度对两种策略进行对比:
| 对比维度 | 预剪枝(Pre-pruning) | 后剪枝(Post-pruning) |
|---|---|---|
| 决策时机 | 决策树生成过程中(结点划分前) | 完整决策树生成后(自底向上遍历) |
| 核心判断依据 | 划分后验证集精度是否提升 | 替换子树为叶结点后验证集精度是否提升 |
| 对模型的影响 | 易生成简单树,可能欠拟合 | 生成较优树,泛化性能更稳定 |
| 计算成本与效率 | 成本低、效率高 | 成本高、效率低 |
五、入门总结与实践建议
- 核心逻辑:剪枝的本质是 “平衡模型复杂度与泛化性能”,通过去掉过拟合的分支,让模型更关注数据的普遍规律。
- 策略选择:
- 若数据量小、追求效率,优先尝试预剪枝;
- 若数据量足够、追求更高泛化性能,后剪枝是更优选择(尽管计算成本高)。
- 关键提醒:无论是预剪枝还是后剪枝,都必须基于验证集判断性能 —— 不能用训练集判断,否则会陷入 “用训练误差衡量泛化误差” 的误区,导致剪枝无效。
对于刚入门的同学,建议先通过简单数据集(如本文的 “好瓜分类”)手动模拟两种剪枝过程,再用 Python 的 sklearn 库(如DecisionTreeClassifier的max_depth(预剪枝)、ccp_alpha(后剪枝)参数)进行实践,这样能更直观地感受剪枝对模型性能的影响。
后续我们还会讲解剪枝的具体评估指标(如损失函数)和进阶技巧,关注我,带你一步步吃透机器学习的核心知识点!
决策树剪枝 Python 实践代码(含数据集与可视化)
以下代码基于 sklearn 实现决策树的预剪枝(max_depth 控制)与后剪枝(ccp_alpha 控制),配套 “好瓜分类” 简化数据集,包含完整的数据预处理、模型训练、性能评估与可视化流程,可直接复制到 CSDN 推文的实践部分使用。
一、环境依赖
首先确保安装以下 Python 库(若未安装,执行 pip install 库名 即可):
pip install numpy pandas scikit-learn matplotlib
二、完整代码实现
1. 导入所需库
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, classification_report
2. 构建 “好瓜分类” 数据集(贴合文档案例特征)
参考文档中 “好瓜” 的核心特征(色泽、根蒂、敲声、纹理、脐部、触感),整理结构化数据集,标签 “好瓜” 为 1,“坏瓜” 为 0:
# 构建好瓜分类数据集(基于文档中训练集/测试集特征整理)
data = {
"色泽": ["青绿", "乌黑", "乌黑", "青绿", "乌黑", "青绿", "浅白", "乌黑", "浅白", "青绿",
"青绿", "浅白", "乌黑", "乌黑", "浅白", "浅白", "青绿"],
"根蒂": ["蜷缩", "蜷缩", "蜷缩", "稍蜷", "稍蜷", "硬挺", "稍蜷", "稍蜷", "蜷缩", "蜷缩",
"蜷缩", "蜷缩", "稍蜷", "稍蜷", "硬挺", "蜷缩", "稍蜷"],
"敲声": ["浊响", "沉闷", "浊响", "浊响", "浊响", "清脆", "沉闷", "浊响", "浊响", "沉闷",
"沉闷", "浊响", "浊响", "沉闷", "清脆", "浊响", "浊响"],
"纹理": ["清晰", "清晰", "清晰", "清晰", "稍糊", "清晰", "稍糊", "清晰", "模糊", "稍糊",
"清晰", "清晰", "清晰", "稍糊", "模糊", "模糊", "稍糊"],
"脐部": ["凹陷", "凹陷", "凹陷", "稍凹", "稍凹", "平坦", "凹陷", "稍凹", "平坦", "稍凹",
"凹陷", "凹陷", "稍凹", "稍凹", "平坦", "平坦", "凹陷"],
"触感": ["硬滑", "硬滑", "硬滑", "软粘", "软粘", "软粘", "硬滑", "软粘", "硬滑", "硬滑",
"硬滑", "硬滑", "硬滑", "硬滑", "硬滑", "软粘", "硬滑"],
"好瓜": [1, 1, 1, 1, 0, 0, 0, 1, 0, 0, 1, 1, 1, 0, 0, 0, 0] # 1=好瓜,0=坏瓜
}
# 转换为 DataFrame
df = pd.DataFrame(data)
# 查看数据集基本信息
print("数据集形状(样本数, 特征数):", df.shape)
print("\n数据集前 5 行:")
print(df.head())
3. 数据预处理(特征编码:文本转数值)
决策树无法直接处理文本特征,需将 “色泽”“根蒂” 等分类特征转为数值(用 LabelEncoder 编码):
from sklearn.preprocessing import LabelEncoder
# 分离特征(X)和标签(y)
X = df.drop("好瓜", axis=1) # 所有特征列
y = df["好瓜"] # 标签列
# 对每个文本特征列进行 LabelEncoder 编码
label_encoders = {} # 存储每个特征的编码器,方便后续解释
for col in X.columns:
le = LabelEncoder()
X[col] = le.fit_transform(X[col])
label_encoders[col] = le # 保存编码器(如:色泽-青绿=0,乌黑=1,浅白=2)
# 查看编码后的特征
print("\n编码后的特征前 5 行:")
print(X.head())
# 划分训练集(70%)和验证集(30%)(验证集用于评估剪枝效果)
X_train, X_val, y_train, y_val = train_test_split(
X, y, test_size=0.3, random_state=42, stratify=y # stratify 保证标签分布一致
)
print(f"\n训练集样本数:{len(X_train)},验证集样本数:{len(X_val)}")
4. 预剪枝实现(用 max_depth 控制树深度)
预剪枝通过限制决策树的最大深度(max_depth),避免树过度生长导致过拟合。我们对比不同 max_depth 下的模型性能,选择最优深度:
# 1. 训练无预剪枝的基准决策树(max_depth=None,树会完全生长)
dt_base = DecisionTreeClassifier(random_state=42)
dt_base.fit(X_train, y_train)
# 2. 训练不同 max_depth 的预剪枝决策树(对比深度 1~4)
max_depths = [1, 2, 3, 4, None] # None 表示无预剪枝
pre_prune_acc = [] # 存储每个深度的验证集精度
for depth in max_depths:
dt_pre = DecisionTreeClassifier(
max_depth=depth, # 预剪枝核心参数:最大深度
random_state=42
)
dt_pre.fit(X_train, y_train)
# 预测并计算验证集精度
y_val_pred = dt_pre.predict(X_val)
acc = accuracy_score(y_val, y_val_pred)
pre_prune_acc.append(acc)
print(f"\n预剪枝 - max_depth={depth}:")
print(f" 训练集精度:{accuracy_score(y_train, dt_pre.predict(X_train)):.2f}")
print(f" 验证集精度:{acc:.2f}")
# 3. 选择最优 max_depth(验证集精度最高的深度)
best_depth = max_depths[pre_prune_acc.index(max(pre_prune_acc))]
print(f"\n预剪枝最优 max_depth:{best_depth},对应验证集精度:{max(pre_prune_acc):.2f}")
# 4. 训练最优预剪枝模型(用于后续可视化)
dt_pre_best = DecisionTreeClassifier(max_depth=best_depth, random_state=42)
dt_pre_best.fit(X_train, y_train)
5. 后剪枝实现(用 ccp_alpha 控制剪枝强度)
后剪枝通过 ccp_alpha(代价复杂度剪枝参数)控制剪枝强度:ccp_alpha 越大,剪枝越彻底;ccp_alpha=0 表示无后剪枝。我们先计算最优 ccp_alpha,再训练后剪枝模型:
# 1. 计算决策树的 ccp_alpha 候选值(从树中自动生成)
dt_for_ccp = DecisionTreeClassifier(random_state=42)
path = dt_for_ccp.cost_complexity_pruning_path(X_train, y_train)
ccp_alphas = path.ccp_alphas[:-1] # 去除最后一个 alpha(会导致树只剩根节点)
# 2. 训练不同 ccp_alpha 的后剪枝决策树
post_prune_acc = [] # 存储每个 alpha 的验证集精度
dt_post_list = [] # 存储每个 alpha 对应的模型
for alpha in ccp_alphas:
dt_post = DecisionTreeClassifier(
ccp_alpha=alpha, # 后剪枝核心参数:剪枝强度
random_state=42
)
dt_post.fit(X_train, y_train)
dt_post_list.append(dt_post)
# 预测并计算验证集精度
y_val_pred = dt_post.predict(X_val)
acc = accuracy_score(y_val, y_val_pred)
post_prune_acc.append(acc)
print(f"\n后剪枝 - ccp_alpha={alpha:.4f}:")
print(f" 训练集精度:{accuracy_score(y_train, dt_post.predict(X_train)):.2f}")
print(f" 验证集精度:{acc:.2f}")
# 3. 选择最优 ccp_alpha(验证集精度最高的 alpha)
if ccp_alphas.size > 0:
best_alpha = ccp_alphas[post_prune_acc.index(max(post_prune_acc))]
best_post_idx = post_prune_acc.index(max(post_prune_acc))
dt_post_best = dt_post_list[best_post_idx] # 最优后剪枝模型
print(f"\n后剪枝最优 ccp_alpha:{best_alpha:.4f},对应验证集精度:{max(post_prune_acc):.2f}")
else:
# 若没有候选 alpha,使用无后剪枝模型
dt_post_best = DecisionTreeClassifier(random_state=42)
dt_post_best.fit(X_train, y_train)
best_alpha = 0.0
print("\n后剪枝无有效候选 alpha,使用无后剪枝模型(ccp_alpha=0)")
6. 模型性能对比(预剪枝 vs 后剪枝 vs 无剪枝)
# 计算三种模型的最终验证集精度
base_acc = accuracy_score(y_val, dt_base.predict(X_val)) # 无剪枝
pre_best_acc = accuracy_score(y_val, dt_pre_best.predict(X_val)) # 最优预剪枝
post_best_acc = accuracy_score(y_val, dt_post_best.predict(X_val)) # 最优后剪枝
# 打印对比结果
print("\n" + "="*60)
print("模型性能最终对比(验证集精度)")
print("="*60)
print(f"无剪枝模型:{base_acc:.2f}")
print(f"最优预剪枝模型(max_depth={best_depth}):{pre_best_acc:.2f}")
print(f"最优后剪枝模型(ccp_alpha={best_alpha:.4f}):{post_best_acc:.2f}")
print("="*60)
# 输出最优模型的分类报告(更详细的性能指标)
best_model = dt_post_best if post_best_acc >= pre_best_acc else dt_pre_best
print(f"\n最优模型({'后剪枝' if post_best_acc >= pre_best_acc else '预剪枝'})分类报告:")
print(classification_report(
y_val, best_model.predict(X_val),
target_names=["坏瓜", "好瓜"] # 对应标签 0 和 1
))
7. 结果可视化(决策树结构 + 精度对比图)
(1)决策树结构可视化(最优预剪枝 vs 最优后剪枝)
# 设置中文字体(避免 matplotlib 中文乱码)
plt.rcParams['font.sans-serif'] = ['SimHei', 'DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False
# 创建画布(2个子图:预剪枝树 + 后剪枝树)
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(20, 8))
# 1. 绘制最优预剪枝决策树
plot_tree(
dt_pre_best,
ax=ax1,
feature_names=X.columns, # 特征名称(如:色泽、根蒂)
class_names=["坏瓜", "好瓜"], # 类别名称
filled=True, # 填充颜色(不同类别颜色不同)
rounded=True, # 圆角矩形
fontsize=10
)
ax1.set_title(f"预剪枝决策树(max_depth={best_depth})", fontsize=14, pad=20)
# 2. 绘制最优后剪枝决策树
plot_tree(
dt_post_best,
ax=ax2,
feature_names=X.columns,
class_names=["坏瓜", "好瓜"],
filled=True,
rounded=True,
fontsize=10
)
ax2.set_title(f"后剪枝决策树(ccp_alpha={best_alpha:.4f})", fontsize=14, pad=20)
# 保存图片(可直接插入 CSDN 推文)
plt.tight_layout()
plt.savefig("decision_tree_pruning_structure.png", dpi=300, bbox_inches='tight')
plt.show()
(2)精度对比图(不同剪枝参数的性能趋势)
# 创建画布(2个子图:预剪枝深度-精度 + 后剪枝 alpha-精度)
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))
# 1. 预剪枝:max_depth 与验证集精度关系
ax1.plot(
[str(d) for d in max_depths], # x轴:深度(转为字符串,避免 None 报错)
pre_prune_acc,
marker='o', linewidth=2, markersize=8, color='#1f77b4'
)
ax1.set_xlabel("预剪枝 max_depth", fontsize=12)
ax1.set_ylabel("验证集精度", fontsize=12)
ax1.set_title("预剪枝:不同 max_depth 的验证集精度", fontsize=14)
ax1.grid(True, alpha=0.3)
# 标注最优精度点
best_pre_idx = pre_prune_acc.index(max(pre_prune_acc))
ax1.annotate(
f"最优:{max(pre_prune_acc):.2f}",
xy=(best_pre_idx, max(pre_prune_acc)),
xytext=(best_pre_idx, max(pre_prune_acc)+0.05),
ha='center', fontsize=10,
arrowprops=dict(arrowstyle='->', color='red')
)
# 2. 后剪枝:ccp_alpha 与验证集精度关系
if ccp_alphas.size > 0:
ax2.plot(
ccp_alphas,
post_prune_acc,
marker='s', linewidth=2, markersize=8, color='#ff7f0e'
)
ax2.set_xlabel("后剪枝 ccp_alpha", fontsize=12)
ax2.set_ylabel("验证集精度", fontsize=12)
ax2.set_title("后剪枝:不同 ccp_alpha 的验证集精度", fontsize=14)
ax2.grid(True, alpha=0.3)
# 标注最优精度点
best_post_idx = post_prune_acc.index(max(post_prune_acc))
ax2.annotate(
f"最优:{max(post_prune_acc):.2f}",
xy=(ccp_alphas[best_post_idx], max(post_prune_acc)),
xytext=(ccp_alphas[best_post_idx]+0.01, max(post_prune_acc)+0.05),
ha='center', fontsize=10,
arrowprops=dict(arrowstyle='->', color='red')
)
else:
ax2.text(0.5, 0.5, "无有效 ccp_alpha 候选值", ha='center', va='center', fontsize=12)
ax2.set_xlabel("后剪枝 ccp_alpha", fontsize=12)
ax2.set_ylabel("验证集精度", fontsize=12)
ax2.set_title("后剪枝:不同 ccp_alpha 的验证集精度", fontsize=14)
# 保存图片
plt.tight_layout()
plt.savefig("decision_tree_pruning_accuracy.png", dpi=300, bbox_inches='tight')
plt.show()
三、代码说明与使用建议
- 数据集适配:代码中 “好瓜分类” 数据集基于文档中的特征整理,若需使用自己的数据集,只需替换
data字典中的特征与标签即可(注意文本特征需保留LabelEncoder编码步骤)。 - 参数调整:
- 预剪枝可调整
max_depths列表(如增加5或6),探索更优深度; - 后剪枝的
ccp_alphas由模型自动生成,无需手动设置,若需更精细控制,可手动添加alpha值(如[0.001, 0.005, 0.01])。
- 预剪枝可调整
- 结果解读:
- 决策树结构可视化图中,每个节点会显示 “特征名称 + 阈值”“样本数量”“类别分布”,可直观理解模型的决策逻辑;
- 精度对比图可帮助判断 “是否过拟合”:若无剪枝模型训练集精度很高但验证集精度低,说明存在过拟合,剪枝后验证集精度提升则证明剪枝有效。
更多推荐


所有评论(0)