机器学习实验四:决策树剪枝实战 —— 提升贷款审批模型泛化能力
一、实验概述
1. 实验背景
在上一篇实验中,我们基于决策树构建了贷款审批预测模型,虽然在测试集上取得了 100% 的准确率,但这可能是由于数据集规模较小导致的 “过拟合” 风险。决策树算法天生容易过拟合:未限制的决策树会不断分裂节点,直到每个叶子节点仅包含单一类别样本,导致模型在训练数据上表现极好,但面对新数据时泛化能力差。
剪枝是解决决策树过拟合的核心手段,通过移除决策树中 “冗余” 的分支,降低模型复杂度,从而提升泛化能力。本实验将详细讲解决策树的预剪枝和后剪枝原理,并基于贷款审批数据集实战演示剪枝过程与效果对比。
2. 实验目标
- 理解决策树过拟合的原因及危害;
- 掌握预剪枝(Pre-pruning)的核心思想与常用方法;
- 掌握后剪枝(Post-pruning)的核心思想与常用方法;
- 实战对比剪枝前后模型的性能差异,学会选择最优剪枝策略。
3. 实验环境
- 编程语言:Python 3.8+
- 核心库:pandas、scikit-learn、matplotlib、seaborn、graphviz
二、剪枝核心概念
1. 过拟合与剪枝的关系
- 过拟合表现:训练集准确率高,测试集准确率低,模型 “死记硬背” 训练数据的噪声而非核心规律;
- 剪枝本质:通过 “去掉不必要的分支” 降低模型复杂度,让决策树更关注数据的整体规律而非局部噪声;
- 剪枝分类:
- 预剪枝:训练过程中提前停止树的生长(“防患于未然”);
- 后剪枝:先构建完整的决策树,再通过一定规则移除冗余分支(“亡羊补牢”)。
2. 预剪枝(Pre-pruning)
核心思想
在决策树生长过程中,设置终止条件,当满足条件时停止节点分裂,避免树长得过深。
常用预剪枝策略
| 策略参数 | 作用说明 |
|---|---|
max_depth | 限制决策树的最大深度(最常用),树的深度越大,复杂度越高 |
min_samples_split | 节点分裂的最小样本数,小于该值则不分裂(避免样本过少的节点继续分裂) |
min_samples_leaf | 叶子节点的最小样本数,分裂后叶子节点样本数小于该值则不分裂 |
min_impurity_decrease | 分裂后不纯度(基尼系数 / 熵)的减少量阈值,小于该值则不分裂 |
max_leaf_nodes | 限制决策树的最大叶子节点数 |
3. 后剪枝(Post-pruning)
核心思想
先构建一棵完整的、未剪枝的决策树(通常深度较大、叶子节点较多),再从树的底部向上回溯,判断每个分支是否 “有用”,若移除该分支后模型性能未下降(或下降在可接受范围内),则移除该分支。
常用后剪枝策略
- 成本复杂度剪枝(Cost-Complexity Pruning, CCP):scikit-learn 默认支持的后剪枝方法,核心是引入 “复杂度系数 α”(ccp_alpha):
- 目标函数:
Cost(T) = 训练误差 + α * 树的复杂度(树的复杂度通常用叶子节点数衡量); - α=0:不剪枝(完整树);
- α 增大:剪枝力度增大,树的复杂度降低;
- 最优 α:通过交叉验证选择使验证集性能最优的 α 值。
- 目标函数:
三、实验步骤
1. 环境准备与数据加载
1.1 导入库
python
运行
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.metrics import accuracy_score, classification_report, confusion_matrix
from sklearn.model_selection import cross_val_score
import warnings
warnings.filterwarnings('ignore')
# 定义列名(与上一篇实验一致)
columns = ['年龄段', '有工作', '有自己的房子', '信贷情况', '是否给贷款']
1.2 加载并预处理数据
python
运行
# 加载训练集和测试集
train_data = pd.read_csv('dataset.txt', header=None, names=columns)
test_data = pd.read_csv('testset.txt', header=None, names=columns)
# 分离特征(X)和目标变量(y)
X_train = train_data.drop('是否给贷款', axis=1)
y_train = train_data['是否给贷款']
X_test = test_data.drop('是否给贷款', axis=1)
y_test = test_data['是否给贷款']
# 查看数据基本信息
print("训练集形状:", X_train.shape)
print("测试集形状:", X_test.shape)
print("\n训练集目标变量分布:")
print(y_train.value_counts())
2. 基准模型:未剪枝的决策树
首先构建一棵完全未剪枝的决策树,作为后续剪枝对比的基准:
python
运行
# 构建未剪枝的决策树(不设置任何限制参数)
dt_unpruned = DecisionTreeClassifier(random_state=42)
dt_unpruned.fit(X_train, y_train)
# 评估基准模型性能
y_pred_unpruned = dt_unpruned.predict(X_test)
train_acc_unpruned = accuracy_score(y_train, dt_unpruned.predict(X_train))
test_acc_unpruned = accuracy_score(y_test, y_pred_unpruned)
print("=== 未剪枝决策树性能 ===")
print(f"训练集准确率:{train_acc_unpruned:.2f}")
print(f"测试集准确率:{test_acc_unpruned:.2f}")
print("\n测试集分类报告:")
print(classification_report(y_test, y_pred_unpruned, target_names=['否', '是']))
# 可视化未剪枝的决策树
plt.figure(figsize=(20, 12))
plot_tree(dt_unpruned, feature_names=X_train.columns, class_names=['否', '是'],
filled=True, rounded=True, fontsize=8)
plt.title('未剪枝的贷款审批决策树(深度:{},叶子节点数:{})'.format(
dt_unpruned.get_depth(), dt_unpruned.get_n_leaves()
))
plt.show()
基准模型分析
- 未剪枝决策树的训练集准确率通常为 100%(完全拟合训练数据);
- 树的深度较大(本数据集下深度约为 4-5),叶子节点数较多(约 6-8 个);
- 虽然本实验测试集准确率仍为 100%(因测试集过小),但在大规模数据中,未剪枝树的泛化能力会显著下降。
3. 预剪枝实战
预剪枝的核心是通过参数限制树的生长,我们将重点演示常用的max_depth、min_samples_split、min_samples_leaf三种策略,并通过网格搜索选择最优参数。
3.1 单一预剪枝策略演示
(1)限制树的最大深度(max_depth)
python
运行
# 测试不同max_depth值的效果
max_depths = [1, 2, 3, 4, 5, 6]
train_accs = []
test_accs = []
leaf_counts = []
for depth in max_depths:
dt = DecisionTreeClassifier(max_depth=depth, 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)))
leaf_counts.append(dt.get_n_leaves())
# 可视化不同max_depth的性能
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.plot(max_depths, train_accs, 'o-', label='训练集准确率')
plt.plot(max_depths, test_accs, 's-', label='测试集准确率')
plt.xlabel('max_depth(树的最大深度)')
plt.ylabel('准确率')
plt.title('max_depth对模型性能的影响')
plt.legend()
plt.grid(True, alpha=0.3)
plt.subplot(1, 2, 2)
plt.plot(max_depths, leaf_counts, 'o-', color='orange')
plt.xlabel('max_depth')
plt.ylabel('叶子节点数')
plt.title('max_depth对树复杂度的影响')
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
(2)限制节点分裂的最小样本数(min_samples_split)
python
运行
# 测试不同min_samples_split值的效果
min_samples_splits = [2, 3, 4, 5, 6, 7]
train_accs_split = []
test_accs_split = []
leaf_counts_split = []
for split in min_samples_splits:
dt = DecisionTreeClassifier(min_samples_split=split, random_state=42)
dt.fit(X_train, y_train)
train_accs_split.append(accuracy_score(y_train, dt.predict(X_train)))
test_accs_split.append(accuracy_score(y_test, dt.predict(X_test)))
leaf_counts_split.append(dt.get_n_leaves())
# 可视化结果
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.plot(min_samples_splits, train_accs_split, 'o-', label='训练集准确率')
plt.plot(min_samples_splits, test_accs_split, 's-', label='测试集准确率')
plt.xlabel('min_samples_split(分裂最小样本数)')
plt.ylabel('准确率')
plt.title('min_samples_split对模型性能的影响')
plt.legend()
plt.grid(True, alpha=0.3)
plt.subplot(1, 2, 2)
plt.plot(min_samples_splits, leaf_counts_split, 'o-', color='orange')
plt.xlabel('min_samples_split')
plt.ylabel('叶子节点数')
plt.title('min_samples_split对树复杂度的影响')
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
3.2 最优预剪枝模型选择(网格搜索)
实际应用中,通常组合多个预剪枝参数,通过网格搜索寻找最优组合:
python
运行
from sklearn.model_selection import GridSearchCV
# 定义参数网格
param_grid = {
'max_depth': [2, 3, 4],
'min_samples_split': [2, 3, 4],
'min_samples_leaf': [1, 2, 3]
}
# 网格搜索(使用5折交叉验证)
grid_search = GridSearchCV(
estimator=DecisionTreeClassifier(random_state=42),
param_grid=param_grid,
cv=5,
scoring='accuracy',
n_jobs=-1
)
grid_search.fit(X_train, y_train)
# 输出最优参数和性能
print("=== 预剪枝最优参数 ===")
print(grid_search.best_params_)
print(f"交叉验证最优准确率:{grid_search.best_score_:.2f}")
# 构建最优预剪枝模型
dt_prepruned = grid_search.best_estimator_
y_pred_prepruned = dt_prepruned.predict(X_test)
train_acc_prepruned = accuracy_score(y_train, dt_prepruned.predict(X_train))
test_acc_prepruned = accuracy_score(y_test, y_pred_prepruned)
print("\n=== 最优预剪枝模型性能 ===")
print(f"训练集准确率:{train_acc_prepruned:.2f}")
print(f"测试集准确率:{test_acc_prepruned:.2f}")
print(f"树的深度:{dt_prepruned.get_depth()}")
print(f"叶子节点数:{dt_prepruned.get_n_leaves()}")
# 可视化最优预剪枝决策树
plt.figure(figsize=(15, 10))
plot_tree(dt_prepruned, feature_names=X_train.columns, class_names=['否', '是'],
filled=True, rounded=True, fontsize=10)
plt.title('最优预剪枝决策树(参数:{})'.format(grid_search.best_params_))
plt.show()
4. 后剪枝实战(CCP 剪枝)
scikit-learn 的DecisionTreeClassifier通过ccp_alpha参数支持成本复杂度剪枝,步骤如下:
4.1 计算最优 ccp_alpha 值
python
运行
# 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(对应根节点,无意义)
impurities = path.impurities[:-1]
# 2. 测试不同ccp_alpha的性能
dt_ccp_models = []
for alpha in ccp_alphas:
dt = DecisionTreeClassifier(ccp_alpha=alpha, random_state=42)
dt.fit(X_train, y_train)
dt_ccp_models.append(dt)
# 3. 计算各模型的交叉验证准确率(避免过拟合)
cv_scores = [cross_val_score(dt, X_train, y_train, cv=5).mean() for dt in dt_ccp_models]
# 4. 寻找最优ccp_alpha(交叉验证准确率最高)
best_idx = np.argmax(cv_scores)
best_ccp_alpha = ccp_alphas[best_idx]
print("=== 后剪枝(CCP)最优参数 ===")
print(f"最优ccp_alpha:{best_ccp_alpha:.4f}")
print(f"交叉验证最优准确率:{cv_scores[best_idx]:.2f}")
# 可视化ccp_alpha对性能的影响
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
plt.plot(ccp_alphas, cv_scores, 'o-', color='green')
plt.scatter(best_ccp_alpha, cv_scores[best_idx], color='red', s=100, label='最优alpha')
plt.xlabel('ccp_alpha')
plt.ylabel('5折交叉验证准确率')
plt.title('ccp_alpha对交叉验证性能的影响')
plt.legend()
plt.grid(True, alpha=0.3)
plt.subplot(1, 2, 2)
depths = [dt.get_depth() for dt in dt_ccp_models]
plt.plot(ccp_alphas, depths, 'o-', color='orange')
plt.scatter(best_ccp_alpha, depths[best_idx], color='red', s=100)
plt.xlabel('ccp_alpha')
plt.ylabel('树的深度')
plt.title('ccp_alpha对树复杂度的影响')
plt.grid(True, alpha=0.3)
plt.tight_layout()
plt.show()
4.2 构建最优后剪枝模型
python
运行
# 基于最优ccp_alpha构建后剪枝模型
dt_postpruned = DecisionTreeClassifier(ccp_alpha=best_ccp_alpha, random_state=42)
dt_postpruned.fit(X_train, y_train)
# 评估后剪枝模型性能
y_pred_postpruned = dt_postpruned.predict(X_test)
train_acc_postpruned = accuracy_score(y_train, dt_postpruned.predict(X_train))
test_acc_postpruned = accuracy_score(y_test, y_pred_postpruned)
print("\n=== 最优后剪枝模型性能 ===")
print(f"训练集准确率:{train_acc_postpruned:.2f}")
print(f"测试集准确率:{test_acc_postpruned:.2f}")
print(f"树的深度:{dt_postpruned.get_depth()}")
print(f"叶子节点数:{dt_postpruned.get_n_leaves()}")
# 可视化后剪枝决策树
plt.figure(figsize=(15, 10))
plot_tree(dt_postpruned, feature_names=X_train.columns, class_names=['否', '是'],
filled=True, rounded=True, fontsize=10)
plt.title('最优后剪枝决策树(ccp_alpha={:.4f})'.format(best_ccp_alpha))
plt.show()
5. 剪枝前后模型对比
将未剪枝、预剪枝、后剪枝模型的关键指标汇总对比:
python
运行
# 汇总所有模型的性能指标
models = {
'未剪枝': {
'训练准确率': train_acc_unpruned,
'测试准确率': test_acc_unpruned,
'树深度': dt_unpruned.get_depth(),
'叶子节点数': dt_unpruned.get_n_leaves()
},
'预剪枝': {
'训练准确率': train_acc_prepruned,
'测试准确率': test_acc_prepruned,
'树深度': dt_prepruned.get_depth(),
'叶子节点数': dt_prepruned.get_n_leaves()
},
'后剪枝': {
'训练准确率': train_acc_postpruned,
'测试准确率': test_acc_postpruned,
'树深度': dt_postpruned.get_depth(),
'叶子节点数': dt_postpruned.get_n_leaves()
}
}
# 转换为DataFrame便于查看
model_df = pd.DataFrame(models).T
print("=== 各模型性能对比 ===")
print(model_df.round(2))
# 可视化对比
plt.figure(figsize=(14, 6))
# 准确率对比
plt.subplot(1, 2, 1)
model_names = list(models.keys())
train_accs = [models[name]['训练准确率'] for name in model_names]
test_accs = [models[name]['测试准确率'] for name in model_names]
x = np.arange(len(model_names))
width = 0.35
plt.bar(x - width/2, train_accs, width, label='训练准确率', alpha=0.8)
plt.bar(x + width/2, test_accs, width, label='测试准确率', alpha=0.8)
plt.xlabel('模型类型')
plt.ylabel('准确率')
plt.title('各模型准确率对比')
plt.xticks(x, model_names)
plt.legend()
plt.grid(True, alpha=0.3, axis='y')
# 复杂度对比(叶子节点数)
plt.subplot(1, 2, 2)
leaf_counts = [models[name]['叶子节点数'] for name in model_names]
plt.bar(model_names, leaf_counts, color=['red', 'blue', 'green'], alpha=0.8)
plt.xlabel('模型类型')
plt.ylabel('叶子节点数(复杂度)')
plt.title('各模型复杂度对比')
plt.grid(True, alpha=0.3, axis='y')
plt.tight_layout()
plt.show()
四、实验结果与分析
1. 核心结论
(1)剪枝的核心作用
- 降低模型复杂度:剪枝后树的深度和叶子节点数显著减少(如未剪枝树深度 5→预剪枝树深度 3→后剪枝树深度 3);
- 平衡 “拟合” 与 “泛化”:未剪枝树训练准确率 100%(过拟合风险),剪枝后训练准确率可能略有下降,但泛化能力更稳定;
- 提升模型可解释性:剪枝后的决策树结构更简单,业务人员更容易理解决策逻辑。
(2)预剪枝 vs 后剪枝
| 对比维度 | 预剪枝(参数限制) | 后剪枝(CCP) |
|---|---|---|
| 实现难度 | 简单(直接设置参数) | 稍复杂(需计算最优 alpha) |
| 剪枝效果 | 依赖参数选择,可能欠剪枝 | 剪枝更精准,不易欠剪枝 / 过剪枝 |
| 计算成本 | 低(训练时直接限制生长) | 高(先建完整树再剪枝) |
| 适用场景 | 快速迭代、数据规模大 | 追求最优泛化性能、数据规模小 |
(3)本实验最优模型
基于贷款审批数据集,预剪枝和后剪枝模型的测试准确率均与未剪枝模型一致(100%),但复杂度显著降低,因此剪枝后的模型更优(泛化能力更强)。
2. 常见问题与解决方案
(1)预剪枝参数设置不当导致欠剪枝 / 过剪枝
- 欠剪枝:参数限制过松(如 max_depth=10),树仍复杂,过拟合风险高;
- 过剪枝:参数限制过严(如 max_depth=1),树过于简单,训练 / 测试准确率均低;
- 解决方案:通过网格搜索 + 交叉验证选择参数,避免手动设置的主观性。
(2)后剪枝 ccp_alpha 选择不当
- alpha 过小:剪枝力度不足,树仍复杂;
- alpha 过大:剪枝力度过大,树过于简单;
- 解决方案:通过交叉验证寻找最优 alpha,而非直接使用默认值 0。
五、实验总结
本实验通过贷款审批数据集,完整演示了决策树剪枝的核心流程:从理解过拟合与剪枝的关系,到预剪枝的参数调优,再到后剪枝的 CCP 算法实现,最后通过多模型对比验证了剪枝的有效性。
关键收获
- 剪枝是决策树避免过拟合的核心手段,实际应用中必须剪枝(未剪枝树几乎无法用于真实场景);
- 预剪枝适合快速开发和大规模数据,后剪枝适合对性能要求较高的场景;
- 模型评估不能仅看测试集准确率,还需关注复杂度(深度、叶子节点数),追求 “准确率 - 复杂度” 的平衡;
- 决策树的可解释性是其核心优势,剪枝后的树结构更简洁,更易落地到业务场景(如贷款审批规则梳理)。
后续拓展
- 尝试更多剪枝策略(如最小代价剪枝 MCCP);
- 结合集成学习(如随机森林、XGBoost),通过多个剪枝后的决策树进一步提升泛化能力;
- 针对更大规模的贷款数据集,验证剪枝对泛化能力的提升效果。
六、附:完整代码
python
运行
# 决策树剪枝实战完整代码
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.metrics import accuracy_score, classification_report, confusion_matrix
from sklearn.model_selection import cross_val_score, GridSearchCV
import warnings
warnings.filterwarnings('ignore')
# 1. 数据加载与预处理
columns = ['年龄段', '有工作', '有自己的房子', '信贷情况', '是否给贷款']
train_data = pd.read_csv('dataset.txt', header=None, names=columns)
test_data = pd.read_csv('testset.txt', header=None, names=columns)
X_train = train_data.drop('是否给贷款', axis=1)
y_train = train_data['是否给贷款']
X_test = test_data.drop('是否给贷款', axis=1)
y_test = test_data['是否给贷款']
print("训练集形状:", X_train.shape)
print("测试集形状:", X_test.shape)
# 2. 基准模型:未剪枝决策树
dt_unpruned = DecisionTreeClassifier(random_state=42)
dt_unpruned.fit(X_train, y_train)
y_pred_unpruned = dt_unpruned.predict(X_test)
train_acc_unpruned = accuracy_score(y_train, dt_unpruned.predict(X_train))
test_acc_unpruned = accuracy_score(y_test, y_pred_unpruned)
print("\n=== 未剪枝决策树性能 ===")
print(f"训练集准确率:{train_acc_unpruned:.2f}")
print(f"测试集准确率:{test_acc_unpruned:.2f}")
# 可视化未剪枝树
plt.figure(figsize=(20, 12))
plot_tree(dt_unpruned, feature_names=X_train.columns, class_names=['否', '是'],
filled=True, rounded=True, fontsize=8)
plt.title('未剪枝的贷款审批决策树(深度:{},叶子节点数:{})'.format(
dt_unpruned.get_depth(), dt_unpruned.get_n_leaves()
))
plt.show()
# 3. 预剪枝实战(网格搜索)
param_grid = {
'max_depth': [2, 3, 4],
'min_samples_split': [2, 3, 4],
'min_samples_leaf': [1, 2, 3]
}
grid_search = GridSearchCV(
estimator=DecisionTreeClassifier(random_state=42),
param_grid=param_grid,
cv=5,
scoring='accuracy',
n_jobs=-1
)
grid_search.fit(X_train, y_train)
print("\n=== 预剪枝最优参数 ===")
print(grid_search.best_params_)
print(f"交叉验证最优准确率:{grid_search.best_score_:.2f}")
dt_prepruned = grid_search.best_estimator_
y_pred_prepruned = dt_prepruned.predict(X_test)
train_acc_prepruned = accuracy_score(y_train, dt_prepruned.predict(X_train))
test_acc_prepruned = accuracy_score(y_test, y_pred_prepruned)
print("\n=== 最优预剪枝模型性能 ===")
print(f"训练集准确率:{train_acc_prepruned:.2f}")
print(f"测试集准确率:{test_acc_prepruned:.2f}")
print(f"树的深度:{dt_prepruned.get_depth()}")
print(f"叶子节点数:{dt_prepruned.get_n_leaves()}")
# 可视化预剪枝树
plt.figure(figsize=(15, 10))
plot_tree(dt_prepruned, feature_names=X_train.columns, class_names=['否', '是'],
filled=True, rounded=True, fontsize=10)
plt.title('最优预剪枝决策树(参数:{})'.format(grid_search.best_params_))
plt.show()
# 4. 后剪枝实战(CCP)
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]
cv_scores = [cross_val_score(DecisionTreeClassifier(ccp_alpha=alpha, random_state=42),
X_train, y_train, cv=5).mean() for alpha in ccp_alphas]
best_idx = np.argmax(cv_scores)
best_ccp_alpha = ccp_alphas[best_idx]
print("\n=== 后剪枝(CCP)最优参数 ===")
print(f"最优ccp_alpha:{best_ccp_alpha:.4f}")
print(f"交叉验证最优准确率:{cv_scores[best_idx]:.2f}")
# 构建后剪枝模型
dt_postpruned = DecisionTreeClassifier(ccp_alpha=best_ccp_alpha, random_state=42)
dt_postpruned.fit(X_train, y_train)
y_pred_postpruned = dt_postpruned.predict(X_test)
train_acc_postpruned = accuracy_score(y_train, dt_postpruned.predict(X_train))
test_acc_postpruned = accuracy_score(y_test, y_pred_postpruned)
print("\n=== 最优后剪枝模型性能 ===")
print(f"训练集准确率:{train_acc_postpruned:.2f}")
print(f"测试集准确率:{test_acc_postpruned:.2f}")
print(f"树的深度:{dt_postpruned.get_depth()}")
print(f"叶子节点数:{dt_postpruned.get_n_leaves()}")
# 可视化后剪枝树
plt.figure(figsize=(15, 10))
plot_tree(dt_postpruned, feature_names=X_train.columns, class_names=['否', '是'],
filled=True, rounded=True, fontsize=10)
plt.title('最优后剪枝决策树(ccp_alpha={:.4f})'.format(best_ccp_alpha))
plt.show()
# 5. 模型对比
models = {
'未剪枝': {
'训练准确率': train_acc_unpruned,
'测试准确率': test_acc_unpruned,
'树深度': dt_unpruned.get_depth(),
'叶子节点数': dt_unpruned.get_n_leaves()
},
'预剪枝': {
'训练准确率': train_acc_prepruned,
'测试准确率': test_acc_prepruned,
'树深度': dt_prepruned.get_depth(),
'叶子节点数': dt_prepruned.get_n_leaves()
},
'后剪枝': {
'训练准确率': train_acc_postpruned,
'测试准确率': test_acc_postpruned,
'树深度': dt_postpruned.get_depth(),
'叶子节点数': dt_postpruned.get_n_leaves()
}
}
model_df = pd.DataFrame(models).T
print("\n=== 各模型性能对比 ===")
print(model_df.round(2))
# 可视化对比
plt.figure(figsize=(14, 6))
model_names = list(models.keys())
train_accs = [models[name]['训练准确率'] for name in model_names]
test_accs = [models[name]['测试准确率'] for name in model_names]
leaf_counts = [models[name]['叶子节点数'] for name in model_names]
# 准确率对比
plt.subplot(1, 2, 1)
x = np.arange(len(model_names))
width = 0.35
plt.bar(x - width/2, train_accs, width, label='训练准确率', alpha=0.8)
plt.bar(x + width/2, test_accs, width, label='测试准确率', alpha=0.8)
plt.xlabel('模型类型')
plt.ylabel('准确率')
plt.title('各模型准确率对比')
plt.xticks(x, model_names)
plt.legend()
plt.grid(True, alpha=0.3, axis='y')
# 复杂度对比
plt.subplot(1, 2, 2)
plt.bar(model_names, leaf_counts, color=['red', 'blue', 'green'], alpha=0.8)
plt.xlabel('模型类型')
plt.ylabel('叶子节点数(复杂度)')
plt.title('各模型复杂度对比')
plt.grid(True, alpha=0.3, axis='y')
plt.tight_layout()
plt.show()
更多推荐

所有评论(0)