决策树剪枝实战:用西瓜数据集手把手教你预剪枝与后剪枝(附Python代码)

如果你刚开始接触机器学习,决策树模型可能是你最早学会的几个算法之一。它直观、易于理解,就像我们做决策时画流程图一样。但很快你就会发现,一个不加限制的决策树会“长”得过于茂盛——它对训练数据中的每一个细节都了如指掌,以至于把噪声也当成了规律。结果呢?它在训练集上表现完美,遇到新数据却一塌糊涂。这就是我们常说的过拟合

解决过拟合的关键技术之一,就是剪枝。剪枝不是简单地砍掉树枝,而是一门平衡的艺术:如何在保持模型强大学习能力的同时,防止它钻牛角尖。预剪枝像一位严格的园丁,在树苗生长时就判断哪些分支不该长;后剪枝则像一位经验丰富的园艺师,等树长成后再回头修剪,让整体形态更优美。这两种策略没有绝对的优劣,选择哪一种,往往取决于你的数据、算力以及对模型可解释性的要求。

今天,我们就抛开复杂的数学公式,直接动手。我将带你使用经典的西瓜数据集,在Jupyter Notebook里,一步步用Python代码构建决策树,并亲自实施预剪枝与后剪枝。你会看到代码如何运行,参数如何影响结果,以及最终如何得到一个既简洁又强大的模型。无论你是想巩固理论知识的初学者,还是希望优化实际项目中模型性能的开发者,这篇实战指南都将为你提供清晰的路径。

1. 环境准备与数据理解

在开始任何机器学习项目之前,搭建一个清晰、可复现的工作环境是第一步。我们选择Jupyter Notebook,因为它能完美地结合代码、运行结果和文字说明,非常适合这种循序渐进的教程。

首先,确保你安装了必要的Python库。我们将主要依赖scikit-learn,它是机器学习领域的瑞士军刀,同时也需要pandasmatplotlib进行数据处理和可视化。

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-learnDecisionTreeClassifier中,控制树生长的核心参数是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%,但测试集准确率很可能会显著提升。这就是预剪枝的功劳:它牺牲了一点对训练数据的拟合度,换来了模型在新数据上更强的泛化能力。

提示:寻找最佳的超参数组合是一个迭代过程。你可以使用GridSearchCVRandomizedSearchCV来自动搜索。例如,尝试不同的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。

这个过程可以分为三步:

  1. 生成完整的树并计算CCP路径。
  2. 针对不同的alpha值生成一系列剪枝后的子树。
  3. 用一个独立的验证集(或交叉验证)评估这些子树,选择最优的。
# 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()

这张图会非常直观。你几乎总能观察到:

  1. 完整树:训练准确率接近100%,测试准确率最低,过拟合差距最大。
  2. 预剪枝树:训练和测试准确率都处于中间水平,差距缩小。
  3. 后剪枝树:训练准确率可能略低于完整树,但测试准确率通常是三者中最高的,实现了更好的泛化。

那么,在实际项目中该如何选择?

我的经验是,可以遵循一个简单的决策流程:

  • 第一步,永远先试试后剪枝。如果计算资源和时间允许,后剪枝通常能给出泛化能力最强的模型。使用cost_complexity_pruning_path配合验证集来选择ccp_alpha是非常有效的方法。
  • 第二步,如果效率是关键。比如你的数据集非常大,或者需要快速迭代多个模型,那么预剪枝是你的好朋友。通过GridSearchCV快速搜索max_depth, min_samples_leaf等几个关键参数,也能得到一个相当不错的模型,而且训练速度飞快。
  • 第三步,考虑模型的可解释性。如果需要向业务方展示清晰的决策规则,一棵深度较浅、节点较少的预剪枝树可能比一棵经过复杂后剪枝的树更受欢迎。毕竟,一个只有3-5条规则的模型比一个有20个节点的树更容易讲清楚。

最后,别忘了结合其他技术。剪枝是解决决策树过拟合的核心手段,但不是唯一手段。在实际项目中,我通常会:

  1. 使用集成学习方法,如随机森林(Random Forest)或梯度提升树(Gradient Boosting),它们天生具有更强的抗过拟合能力。
  2. 在集成模型中,仍然会对其中的每一棵基学习器(决策树)进行剪枝或深度限制,这就是“双重保障”。
  3. 对于后剪枝,那个“最优alpha”的寻找过程,完全可以整合到交叉验证(Cross-Validation)的循环中去,确保选择的参数更加稳健。

纸上得来终觉浅,绝知此事要躬行。最好的学习方式就是动手把代码跑一遍,调整参数,观察树结构的变化,感受准确率的波动。你可以尝试更换其他数据集(如Scikit-learn自带的鸢尾花iris或乳腺癌breast_cancer数据集),重复这个流程。很快你就会发现,面对不同的数据分布,最优的剪枝策略和参数也会不同,而这种“手感”正是从初学者迈向实践者的关键一步。

Logo

这里是“一人公司”的成长家园。我们提供从产品曝光、技术变现到法律财税的全栈内容,并连接云服务、办公空间等稀缺资源,助你专注创造,无忧运营。

更多推荐