决策树剪枝实战:用西瓜数据集手把手教你预剪枝与后剪枝(附Python代码)
决策树剪枝实战:用西瓜数据集手把手教你预剪枝与后剪枝(附Python代码)
如果你刚开始接触机器学习,决策树模型可能是你最早学会的几个算法之一。它直观、易于理解,就像我们做决策时画流程图一样。但很快你就会发现,一个不加限制的决策树会“长”得过于茂盛——它对训练数据中的每一个细节都了如指掌,以至于把噪声也当成了规律。结果呢?它在训练集上表现完美,遇到新数据却一塌糊涂。这就是我们常说的过拟合。
解决过拟合的关键技术之一,就是剪枝。剪枝不是简单地砍掉树枝,而是一门平衡的艺术:如何在保持模型强大学习能力的同时,防止它钻牛角尖。预剪枝像一位严格的园丁,在树苗生长时就判断哪些分支不该长;后剪枝则像一位经验丰富的园艺师,等树长成后再回头修剪,让整体形态更优美。这两种策略没有绝对的优劣,选择哪一种,往往取决于你的数据、算力以及对模型可解释性的要求。
今天,我们就抛开复杂的数学公式,直接动手。我将带你使用经典的西瓜数据集,在Jupyter Notebook里,一步步用Python代码构建决策树,并亲自实施预剪枝与后剪枝。你会看到代码如何运行,参数如何影响结果,以及最终如何得到一个既简洁又强大的模型。无论你是想巩固理论知识的初学者,还是希望优化实际项目中模型性能的开发者,这篇实战指南都将为你提供清晰的路径。
1. 环境准备与数据理解
在开始任何机器学习项目之前,搭建一个清晰、可复现的工作环境是第一步。我们选择Jupyter Notebook,因为它能完美地结合代码、运行结果和文字说明,非常适合这种循序渐进的教程。
首先,确保你安装了必要的Python库。我们将主要依赖scikit-learn,它是机器学习领域的瑞士军刀,同时也需要pandas和matplotlib进行数据处理和可视化。
pip install scikit-learn pandas matplotlib numpy
接下来,我们导入所有需要的模块。
import pandas as pd
import numpy as np
from sklearn import tree
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score
import matplotlib.pyplot as plt
现在,让我们来认识一下今天的主角——西瓜数据集。这个数据集在机器学习教材中非常经典,它用一系列特征来描述西瓜的好坏,例如色泽、根蒂、敲声等。为了专注于剪枝技术的核心,我们这里使用一个简化版的数值化数据集。它包含了西瓜的多个特征和最终的类别标签(好瓜=1,坏瓜=0)。
# 构建一个简化的西瓜数据集(数值化版本)
# 特征:色泽(0:青绿,1:乌黑,2:浅白),根蒂(0:蜷缩,1:稍蜷,2:硬挺),敲声(0:浊响,1:沉闷,2:清脆)
# 标签:好瓜=1,坏瓜=0
data = {
'色泽': [0, 0, 0, 0, 1, 1, 1, 1, 2, 2, 2, 2, 0, 1, 2, 0, 1],
'根蒂': [0, 0, 1, 1, 0, 0, 1, 2, 0, 1, 1, 2, 1, 2, 0, 2, 1],
'敲声': [0, 1, 0, 1, 0, 2, 1, 2, 0, 1, 2, 0, 2, 1, 1, 0, 2],
'好瓜': [1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 1, 0, 1, 0, 1]
}
df = pd.DataFrame(data)
X = df[['色泽', '根蒂', '敲声']]
y = df['好瓜']
注意:在实际项目中,你的数据很可能来自CSV文件或数据库。使用
pd.read_csv()加载数据后,通常还需要进行缺失值处理、特征编码(将文字转为数字)等步骤。这里我们跳过了这些,直接使用准备好的数值数据。
任何监督学习任务都离不开将数据划分为训练集和测试集。训练集用于“教导”模型,测试集则用于客观评估模型面对从未见过数据时的表现,这是我们判断模型是否过拟合的关键。
# 划分训练集和测试集,测试集占比30%,设置随机种子确保结果可复现
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)
print(f"训练集样本数: {len(X_train)}")
print(f"测试集样本数: {len(X_test)}")
运行上述代码,你会看到类似“训练集样本数: 11, 测试集样本数: 6”的输出。我们的模型将从这11个训练样本中学习规律。
2. 构建一棵完整的决策树:过拟合的典型
在讨论剪枝之前,我们必须先看看“问题”本身是什么样子。让我们用scikit-learn训练一棵不施加任何限制的决策树,也就是让它自由生长,直到所有叶子节点都“纯净”(即只包含同一类样本)或无法继续分裂为止。
在scikit-learn的DecisionTreeClassifier中,控制树生长的核心参数是max_depth(最大深度)和min_samples_split(节点分裂所需最小样本数)。为了得到一棵完整的树,我们把这些参数设置为None,并让min_samples_leaf(叶节点最小样本数)为1。
# 创建并训练一棵未剪枝的完整决策树
full_tree = tree.DecisionTreeClassifier(criterion='entropy', # 使用信息增益作为分裂标准
splitter='best',
max_depth=None,
min_samples_split=2,
min_samples_leaf=1,
random_state=42)
full_tree.fit(X_train, y_train)
# 评估模型
y_train_pred = full_tree.predict(X_train)
y_test_pred = full_tree.predict(X_test)
train_acc = accuracy_score(y_train, y_train_pred)
test_acc = accuracy_score(y_test, y_test_pred)
print(f"完整决策树 - 训练集准确率: {train_acc:.2%}")
print(f"完整决策树 - 测试集准确率: {test_acc:.2%}")
执行这段代码,你很可能会看到一个典型的结果:训练集准确率是100%,而测试集准确率可能只有60%多甚至更低。这巨大的差距就是过拟合的鲜明标志。模型完美地记住了训练数据中的所有细节(包括噪声),但其学到的规则过于复杂和特殊,无法很好地推广到新数据。
我们可以把决策树的结构可视化出来,看看它到底有多“深”。
# 可视化决策树
plt.figure(figsize=(12, 8))
tree.plot_tree(full_tree,
feature_names=['色泽', '根蒂', '敲声'],
class_names=['坏瓜', '好瓜'],
filled=True, # 填充颜色
rounded=True)
plt.title("完整的决策树(可能过拟合)")
plt.show()
生成的树图会非常庞大,节点众多。你可以清晰地看到它为了区分每一个训练样本,生成了许多只包含单个样本的叶子节点。这样的模型虽然对训练数据了如指掌,但毫无实用价值。接下来,我们就用两种策略来修剪这棵“疯长”的树。
3. 预剪枝实战:在生长时踩刹车
预剪枝的核心思想是防患于未然。在决策树构建的每个节点,我们都会提前用验证集评估一下:如果在这里停止分裂(直接变成叶子节点),和继续分裂下去相比,哪个对模型泛化能力更有益?如果分裂不能带来提升,就立即停止。
在scikit-learn中,虽然没有一个叫“预剪枝”的独立函数,但我们可以通过设置一系列超参数来达到预剪枝的效果。这些参数就像给树的生长设置了各种“交通规则”。
| 超参数 | 作用 | 相当于预剪枝的哪种判断? |
|---|---|---|
max_depth |
树的最大深度 | 限制树的生长层数,防止过于复杂。 |
min_samples_split |
节点分裂所需的最小样本数 | 如果节点样本太少,认为其规律不可靠,停止分裂。 |
min_samples_leaf |
叶节点所需的最小样本数 | 确保每个叶子节点有足够样本支撑,避免过于特殊的规则。 |
min_impurity_decrease |
分裂所需的最小不纯度减少量 | 如果分裂带来的信息增益(纯度提升)太小,则认为不值得分裂。 |
让我们尝试组合这些参数,训练一棵经过预剪枝的树。
# 创建并训练一棵经过预剪枝的决策树
pre_pruned_tree = tree.DecisionTreeClassifier(criterion='entropy',
max_depth=3, # 限制深度
min_samples_split=4, # 节点至少4个样本才考虑分裂
min_samples_leaf=2, # 叶节点至少2个样本
min_impurity_decrease=0.01, # 纯度提升需大于0.01
random_state=42)
pre_pruned_tree.fit(X_train, y_train)
# 评估预剪枝模型
y_train_pred_pre = pre_pruned_tree.predict(X_train)
y_test_pred_pre = pre_pruned_tree.predict(X_test)
train_acc_pre = accuracy_score(y_train, y_train_pred_pre)
test_acc_pre = accuracy_score(y_test, y_test_pred_pre)
print(f"预剪枝决策树 - 训练集准确率: {train_acc_pre:.2%}")
print(f"预剪枝决策树 - 测试集准确率: {test_acc_pre:.2%}")
这次的结果会大不相同。训练集准确率可能从100%下降到80%-90%,但测试集准确率很可能会显著提升。这就是预剪枝的功劳:它牺牲了一点对训练数据的拟合度,换来了模型在新数据上更强的泛化能力。
提示:寻找最佳的超参数组合是一个迭代过程。你可以使用
GridSearchCV或RandomizedSearchCV来自动搜索。例如,尝试不同的max_depth值(2, 3, 4, 5),观察训练和测试准确率的变化曲线,找到那个测试准确率最高且模型不过于复杂的“甜蜜点”。
让我们再看看这棵被修剪过的树长什么样。
plt.figure(figsize=(10, 6))
tree.plot_tree(pre_pruned_tree,
feature_names=['色泽', '根蒂', '敲声'],
class_names=['坏瓜', '好瓜'],
filled=True,
rounded=True)
plt.title("经过预剪枝的决策树")
plt.show()
你会发现这棵树明显小了很多,结构也清晰了。它可能只用了“色泽”和“根蒂”两个特征,深度也只有2到3层。这样的模型不仅预测性能更稳定,而且更容易理解和解释——你可以直接向别人描述:“我们的模型主要看西瓜的色泽和根蒂来判断好坏”。
预剪枝的优缺点非常明显:
- 优点:训练速度快,因为树不会长得太深;能有效防止过拟合。
- 缺点:有欠拟合的风险。就像园丁过早掐掉了嫩芽,可能有些分支如果允许它再长一点,后续会分出非常有价值的枝条(即更有判别力的规则)。预剪枝是一种“贪心”算法,只看眼前一步的收益,缺乏全局视野。
4. 后剪枝实战:长成后再精修
与预剪枝的“提前干预”不同,后剪枝采取的是“先发展,后治理”的策略。它允许决策树首先不受限制地生长到完全拟合训练数据(就像我们第一节做的那样),生成一棵完整的、可能过拟合的树。然后,它从树的底部(叶子节点)开始向上回溯,逐一考察每个非叶子节点:如果把以这个节点为根的整个子树替换成一个叶子节点(类别由该节点下训练样本的多数类决定),模型的泛化性能是否会提升?如果会,就果断剪掉这整个分支。
scikit-learn提供了cost_complexity_pruning_path方法来实现后剪枝的核心思想——代价复杂度剪枝(CCP)。它会计算一系列不同的剪枝强度(alpha值)对应的子树,并给出这些子树在训练集上的性能。我们需要从中选择一个能使模型在验证集上表现最好的alpha。
这个过程可以分为三步:
- 生成完整的树并计算CCP路径。
- 针对不同的alpha值生成一系列剪枝后的子树。
- 用一个独立的验证集(或交叉验证)评估这些子树,选择最优的。
# 1. 首先,训练一棵完整的树作为起点
full_tree_for_post = tree.DecisionTreeClassifier(criterion='entropy', random_state=42)
full_tree_for_post.fit(X_train, y_train)
# 2. 获取代价复杂度剪枝路径
path = full_tree_for_post.cost_complexity_pruning_path(X_train, y_train)
ccp_alphas, impurities = path.ccp_alphas, path.impurities
# ccp_alphas是递增的剪枝强度值,alpha越大,剪枝越狠
print(f"生成的alpha值数量: {len(ccp_alphas)}")
print(f"Alpha值示例: {ccp_alphas[:5]}") # 查看前几个
接下来,我们为每一个alpha值训练一棵剪枝后的树,并评估其在测试集上的表现。
# 为每个alpha训练一棵树,并记录准确率
trees = []
train_accs = []
test_accs = []
for ccp_alpha in ccp_alphas:
pruned_tree = tree.DecisionTreeClassifier(criterion='entropy',
random_state=42,
ccp_alpha=ccp_alpha)
pruned_tree.fit(X_train, y_train)
trees.append(pruned_tree)
# 计算准确率
train_accs.append(accuracy_score(y_train, pruned_tree.predict(X_train)))
test_accs.append(accuracy_score(y_test, pruned_tree.predict(X_test)))
# 将结果转为DataFrame方便查看
results_df = pd.DataFrame({
'alpha': ccp_alphas,
'训练准确率': train_accs,
'测试准确率': test_accs,
'树节点数': [pruned_tree.tree_.node_count for pruned_tree in trees]
})
print(results_df.head(10)) # 查看前10行
通过这个表格,你可以清晰地看到,随着alpha值增大(剪枝变强):
- 树的节点数逐渐减少(模型变简单)。
- 训练准确率单调下降(对训练数据的拟合度降低)。
- 测试准确率会先上升后下降,形成一个峰值。这个峰值对应的alpha,就是我们要找的最优剪枝强度。
让我们用图表来更直观地展示这个过程。
# 可视化准确率随alpha变化的情况
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(14, 5))
# 左图:准确率 vs Alpha
ax1.plot(ccp_alphas, train_accs, marker='o', label='训练准确率', drawstyle="steps-post")
ax1.plot(ccp_alphas, test_accs, marker='s', label='测试准确率', drawstyle="steps-post")
ax1.set_xlabel('剪枝强度 (alpha)')
ax1.set_ylabel('准确率')
ax1.set_title('剪枝强度对准确率的影响')
ax1.legend()
ax1.grid(True)
# 右图:树节点数 vs Alpha
ax2.plot(ccp_alphas, results_df['树节点数'], marker='^', drawstyle="steps-post")
ax2.set_xlabel('剪枝强度 (alpha)')
ax2.set_ylabel('树节点数')
ax2.set_title('剪枝强度对模型复杂度的影响')
ax2.grid(True)
plt.tight_layout()
plt.show()
从图中,你可以轻松找到使测试准确率最高的那个alpha值。假设我们通过观察发现alpha=0.02时测试准确率最高,那么就可以用这个参数来训练最终的模型。
# 选择最优alpha(这里假设我们通过上图判断0.02是最佳值)
optimal_alpha = 0.02
final_pruned_tree = tree.DecisionTreeClassifier(criterion='entropy',
random_state=42,
ccp_alpha=optimal_alpha)
final_pruned_tree.fit(X_train, y_train)
# 评估最终的后剪枝模型
final_test_acc = accuracy_score(y_test, final_pruned_tree.predict(X_test))
print(f"最优后剪枝树 (alpha={optimal_alpha}) - 测试集准确率: {final_test_acc:.2%}")
# 可视化最终树
plt.figure(figsize=(10, 6))
tree.plot_tree(final_pruned_tree,
feature_names=['色泽', '根蒂', '敲声'],
class_names=['坏瓜', '好瓜'],
filled=True,
rounded=True)
plt.title(f"经过后剪枝的决策树 (alpha={optimal_alpha})")
plt.show()
后剪枝的优缺点同样值得权衡:
- 优点:由于是基于完整的树进行修剪,它保留了更多生成完整树时获得的信息,欠拟合风险更低,通常得到的模型泛化性能比预剪枝更好。
- 缺点:训练时间开销大。你需要先训练一棵完整的、可能非常庞大的树,然后再进行剪枝评估,计算成本更高。
5. 策略对比与实战选择指南
到现在为止,我们已经亲手实现了两种剪枝策略。是时候把它们放在一起,从多个维度进行对比,以便你在实际项目中做出明智的选择。
为了更系统地进行比较,我们可以创建一个对比表格:
| 对比维度 | 预剪枝 | 后剪枝 |
|---|---|---|
| 剪枝时机 | 在树生成过程中 | 在树生成完成后 |
| 基本策略 | 提前停止分裂 | 生成完整树后,再剪掉子树 |
| 计算效率 | 高,避免生成复杂分支 | 低,需生成完整树再回溯 |
| 过拟合控制 | 好,但可能过于严格 | 非常好,通常能找到更优平衡点 |
| 欠拟合风险 | 较高,可能过早停止 | 较低,基于全局信息修剪 |
| 模型泛化性能 | 通常不错,但可能非最优 | 通常更优 |
| 实现复杂度 | 简单(调参) | 相对复杂(需计算CCP路径) |
| 适用场景 | 数据集大、训练时间敏感、对可解释性要求高 | 追求最高模型性能、有充足计算资源 |
光看理论对比还不够,让我们用代码在同一张图上展示我们刚才训练的三种树(完整树、预剪枝树、后剪枝树)在训练集和测试集上的表现。
# 汇总三种模型的性能
models = ['完整决策树', '预剪枝决策树', '后剪枝决策树']
train_scores = [train_acc, train_acc_pre, accuracy_score(y_train, final_pruned_tree.predict(X_train))]
test_scores = [test_acc, test_acc_pre, final_test_acc]
x = np.arange(len(models))
width = 0.35
fig, ax = plt.subplots(figsize=(9, 6))
rects1 = ax.bar(x - width/2, train_scores, width, label='训练准确率', color='skyblue')
rects2 = ax.bar(x + width/2, test_scores, width, label='测试准确率', color='lightcoral')
ax.set_ylabel('准确率')
ax.set_title('不同决策树策略性能对比')
ax.set_xticks(x)
ax.set_xticklabels(models)
ax.legend()
ax.set_ylim(0, 1.1)
# 在柱子上方标注数值
def autolabel(rects):
for rect in rects:
height = rect.get_height()
ax.annotate(f'{height:.1%}',
xy=(rect.get_x() + rect.get_width() / 2, height),
xytext=(0, 3), # 3 points vertical offset
textcoords="offset points",
ha='center', va='bottom')
autolabel(rects1)
autolabel(rects2)
plt.tight_layout()
plt.show()
这张图会非常直观。你几乎总能观察到:
- 完整树:训练准确率接近100%,测试准确率最低,过拟合差距最大。
- 预剪枝树:训练和测试准确率都处于中间水平,差距缩小。
- 后剪枝树:训练准确率可能略低于完整树,但测试准确率通常是三者中最高的,实现了更好的泛化。
那么,在实际项目中该如何选择?
我的经验是,可以遵循一个简单的决策流程:
- 第一步,永远先试试后剪枝。如果计算资源和时间允许,后剪枝通常能给出泛化能力最强的模型。使用
cost_complexity_pruning_path配合验证集来选择ccp_alpha是非常有效的方法。 - 第二步,如果效率是关键。比如你的数据集非常大,或者需要快速迭代多个模型,那么预剪枝是你的好朋友。通过
GridSearchCV快速搜索max_depth,min_samples_leaf等几个关键参数,也能得到一个相当不错的模型,而且训练速度飞快。 - 第三步,考虑模型的可解释性。如果需要向业务方展示清晰的决策规则,一棵深度较浅、节点较少的预剪枝树可能比一棵经过复杂后剪枝的树更受欢迎。毕竟,一个只有3-5条规则的模型比一个有20个节点的树更容易讲清楚。
最后,别忘了结合其他技术。剪枝是解决决策树过拟合的核心手段,但不是唯一手段。在实际项目中,我通常会:
- 使用集成学习方法,如随机森林(Random Forest)或梯度提升树(Gradient Boosting),它们天生具有更强的抗过拟合能力。
- 在集成模型中,仍然会对其中的每一棵基学习器(决策树)进行剪枝或深度限制,这就是“双重保障”。
- 对于后剪枝,那个“最优alpha”的寻找过程,完全可以整合到交叉验证(Cross-Validation)的循环中去,确保选择的参数更加稳健。
纸上得来终觉浅,绝知此事要躬行。最好的学习方式就是动手把代码跑一遍,调整参数,观察树结构的变化,感受准确率的波动。你可以尝试更换其他数据集(如Scikit-learn自带的鸢尾花iris或乳腺癌breast_cancer数据集),重复这个流程。很快你就会发现,面对不同的数据分布,最优的剪枝策略和参数也会不同,而这种“手感”正是从初学者迈向实践者的关键一步。
更多推荐



所有评论(0)