机器学习实验--决策树剪枝
·
一、决策树剪枝介绍
1.1 剪枝的核心意义
决策树在无约束条件下训练时,易因过度拟合训练集细节出现过拟合现象,剪枝是解决该问题的核心手段,本次实验聚焦更常用的预剪枝。
1.2 预剪枝核心原理
预剪枝通过设置决策树生长的约束参数,在树形构建过程中提前终止分支生长,避免生成过于复杂的树结构,核心参数说明:
max_depth:决策树的最大深度(限制树的纵向生长,本次设置为 3);min_samples_split:节点分裂所需的最小样本数(样本数不足则不分裂,本次设置为 2);min_samples_leaf:叶节点所需的最小样本数;- 核心逻辑:通过参数限制树的复杂度,在 “拟合能力” 与 “泛化能力” 之间达到平衡,同时保持决策树的可解释性。
二、实践代码部分
1. 数据加载与预处理
# 导入所需库
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
from sklearn.tree import DecisionTreeClassifier, plot_tree
from sklearn.metrics import accuracy_score, confusion_matrix
plt.rcParams["font.family"] = ["SimHei", "Microsoft YaHei"]
# 解决负号显示为方块的问题
plt.rcParams["axes.unicode_minus"] = False
def load_and_preprocess_data():
"""
功能:加载贷款训练集与测试集,完成数据预处理(分离特征与标签)
返回:训练集特征、测试集特征、训练集标签、测试集标签
备注:数据集与上一实验一致,无需修改
"""
# 定义中文列名(4个特征+1个标签)
columns = ["年龄段", "有工作", "有自己的房子", "信贷情况", "类别"]
# 加载本地数据集(逗号分隔,无表头)
train_data = pd.read_csv("dataset.txt", sep=",", header=None, names=columns)
test_data = pd.read_csv("testset.txt", sep=",", header=None, names=columns)
# 分离特征(X)与标签(y)
X_train = train_data.drop("类别", axis=1) # 训练集特征
y_train = train_data["类别"] # 训练集标签(0=不给贷款,1=给贷款)
X_test = test_data.drop("类别", axis=1) # 测试集特征
y_test = test_data["类别"] # 测试集标签
# 打印数据基本信息,验证数据加载正确性
print("-"*30) # 短分隔线,保证终端排版整洁
print("数据基本信息")
print("-"*30)
print(f"训练集样本数:{len(train_data)},测试集样本数:{len(test_data)}")
print("特征编码规则:")
print("- 年龄段:青年(0)/中年(1)/老年(2)")
print("- 有工作/有自己的房子:否(0)/是(1)")
print("- 信贷情况:一般(0)/好(1)/非常好(2)")
print("- 类别:不给贷款(0)/给贷款(1)")
return X_train, X_test, y_train, y_test
2. 带预剪枝的模型构建与训练
def build_pruned_tree_model(X_train, y_train):
"""
功能:初始化带预剪枝参数的ID3决策树模型,完成训练
参数:X_train-训练集特征,y_train-训练集标签
返回:训练完成的预剪枝决策树模型
核心:添加预剪枝参数,限制树的过度生长
"""
# 初始化带预剪枝的ID3决策树分类器
pruned_dt_model = DecisionTreeClassifier(
# 基础参数(与上一实验一致)
criterion="entropy", # 划分准则:信息熵
random_state=42, # 固定随机种子,保证结果可复现
# 预剪枝核心参数(关键改动)
max_depth=3, # 树的最大深度:限制纵向生长,避免过深
min_samples_split=2, # 节点分裂最小样本数:样本不足2个则不分裂
min_samples_leaf=1 # 叶节点最小样本数:保证叶节点有代表性
)
# 使用训练集拟合预剪枝模型
pruned_dt_model.fit(X_train, y_train)
# 打印训练完成信息与训练集拟合准确率
print("\n" + "-"*30)
print("预剪枝模型训练完成")
print("-"*30)
# 输出训练集准确率,对比剪枝前后差异
print(f"预剪枝模型训练集拟合准确率:{pruned_dt_model.score(X_train, y_train):.4f}")
return pruned_dt_model
3. 指标计算与模型评估
def calculate_metrics(y_true, y_pred, class_names):
"""
功能:手动计算分类指标,避免sklearn英文残留
参数:y_true-真实标签,y_pred-预测标签,class_names-中文类别名
返回:各类别指标与总体指标(与上一实验一致,复用即可)
"""
# 计算混淆矩阵(核心中间数据)
cm = confusion_matrix(y_true, y_pred)
n_classes = len(class_names)
# 初始化指标列表
precision = [] # 精确率
recall = [] # 召回率
f1_score = [] # F1分数
support = [] # 支持数
# 遍历每个类别计算指标
for i in range(n_classes):
tp = cm[i, i] # 真阳性
fp = cm[:, i].sum() - tp # 假阳性
fn = cm[i, :].sum() - tp # 假阴性
sup = cm[i, :].sum() # 真实样本数
support.append(sup)
# 避免分母为0的异常
cls_precision = tp / (tp + fp) if (tp + fp) != 0 else 0.0
cls_recall = tp / (tp + fn) if (tp + fn) != 0 else 0.0
cls_f1 = 2 * (cls_precision * cls_recall) / (cls_precision + cls_recall) if (cls_precision + cls_recall) != 0 else 0.0
precision.append(cls_precision)
recall.append(cls_recall)
f1_score.append(cls_f1)
# 计算总体指标
macro_prec = np.mean(precision)
macro_rec = np.mean(recall)
macro_f1 = np.mean(f1_score)
weighted_prec = np.average(precision, weights=support)
weighted_rec = np.average(recall, weights=support)
weighted_f1 = np.average(f1_score, weights=support)
accuracy = accuracy_score(y_true, y_pred)
return (precision, recall, f1_score, support,
macro_prec, macro_rec, macro_f1,
weighted_prec, weighted_rec, weighted_f1,
accuracy)
def evaluate_pruned_model(pruned_dt_model, X_test, y_test):
"""
功能:评估预剪枝模型,输出纯中文分类报告
参数:pruned_dt_model-预剪枝模型,X_test-测试集特征,y_test-测试集真实标签
返回:测试集准确率
"""
# 使用预剪枝模型预测测试集
y_pred = pruned_dt_model.predict(X_test)
# 中文类别名
class_names = ["不给贷款", "给贷款"]
# 获取所有评估指标
(precision, recall, f1_score, support,
macro_prec, macro_rec, macro_f1,
weighted_prec, weighted_rec, weighted_f1,
accuracy) = calculate_metrics(y_test, y_pred, class_names)
# 打印评估结果
print("\n" + "-"*30)
print("预剪枝模型评估结果")
print("-"*30)
print(f"预剪枝模型测试集准确率:{accuracy:.4f}")
# 格式化输出分类报告
print("\n分类报告:")
header = f"{'':<12} {'精确率':<8} {'召回率':<8} {'F1分数':<8} {'支持数'}"
print(header)
print("-" * 45)
for i, name in enumerate(class_names):
line = f"{name:<12} {precision[i]:<8.2f} {recall[i]:<8.2f} {f1_score[i]:<8.2f} {support[i]}"
print(line)
print()
print(f"{'准确率':<12} {'':<8} {'':<8} {'':<8} {accuracy:.2f} {sum(support)}")
print(f"{'宏观平均':<12} {macro_prec:<8.2f} {macro_rec:<8.2f} {macro_f1:<8.2f} {sum(support)}")
print(f"{'加权平均':<12} {weighted_prec:<8.2f} {weighted_rec:<8.2f} {weighted_f1:<8.2f} {sum(support)}")
return accuracy
4. 预剪枝决策树可视化
def visualize_pruned_tree(pruned_dt_model, X_train):
"""
功能:绘制预剪枝后的全中文决策树,保存可视化图片
参数:pruned_dt_model-预剪枝模型,X_train-训练集特征(获取特征名)
备注:对比剪枝前,树形更简洁
"""
# 创建画布,保证树形清晰展示
plt.figure(figsize=(15, 10))
# 绘制预剪枝决策树
plot_tree(
pruned_dt_model,
filled=True, # 节点填充颜色
rounded=True, # 圆角矩形,更美观
feature_names=X_train.columns, # 中文特征名
class_names=["不给贷款", "给贷款"], # 中文类别名
fontsize=10 # 字体大小,避免文字重叠
)
# 设置图表标题,突出预剪枝特性
plt.title("贷款审批决策树(ID3算法+预剪枝)", fontsize=15, pad=20)
# 保存高分辨率图片,自适应边界
plt.savefig("贷款审批决策树_预剪枝.png", dpi=300, bbox_inches="tight")
# 显示图片
plt.show()
# 提示图片保存完成
print("\n预剪枝决策树可视化文件已保存:贷款审批决策树_预剪枝.png")
# ========== 主函数:串联预剪枝实验全流程 ==========
if __name__ == "__main__":
# 步骤1:数据加载与预处理
X_train, X_test, y_train, y_test = load_and_preprocess_data()
# 步骤2:构建并训练预剪枝模型
pruned_dt_model = build_pruned_tree_model(X_train, y_train)
# 步骤3:评估预剪枝模型
pruned_accuracy = evaluate_pruned_model(pruned_dt_model, X_test, y_test)
# 步骤4:可视化预剪枝决策树
visualize_pruned_tree(pruned_dt_model, X_train)
三、实验结果
3.1 终端输出
------------------------------
数据基本信息
------------------------------
训练集样本数:16,测试集样本数:7
特征编码规则:
- 年龄段:青年(0)/中年(1)/老年(2)
- 有工作/有自己的房子:否(0)/是(1)
- 信贷情况:一般(0)/好(1)/非常好(2)
- 类别:不给贷款(0)/给贷款(1)
------------------------------
预剪枝模型训练完成
------------------------------
预剪枝模型训练集拟合准确率:1.0000
------------------------------
预剪枝模型评估结果
------------------------------
预剪枝模型测试集准确率:1.0000
分类报告:
精确率 召回率 F1分数 支持数
---------------------------------------------
不给贷款 1.00 1.00 1.00 4
给贷款 1.00 1.00 1.00 3
准确率 1.00 7
宏观平均 1.00 1.00 1.00 7
加权平均 1.00 1.00 1.00 7
预剪枝决策树可视化文件已保存:贷款审批决策树_预剪枝.png
3.2 可视化结果对比
| 对比项 | 无剪枝决策树 | 预剪枝决策树(max_depth=3) |
|---|---|---|
| 树形复杂度 | 分支较多,深度略深 | 分支简洁,深度不超过 3,无冗余节点 |
| 节点数量 | 较多(含细微划分节点) | 较少(仅保留核心划分节点) |
| 业务可解释性 | 良好 | 更优(核心逻辑更突出) |
| 图片文件 | 贷款审批决策树_全中文.png | 贷款审批决策树_预剪枝.png |
| 核心划分逻辑 | 一致(以 “有自己的房子” 为根节点) | 一致(保留核心业务逻辑) |
四、实验报告
4.1 实验目的
- 理解决策树过拟合的成因与预剪枝的核心原理;
- 掌握
scikit-learn中决策树预剪枝参数的设置与调优方法; - 完成带预剪枝的贷款审批分类实验,对比剪枝前后的模型差异;
4.2 实验环境
- Windows 10/11
- Visual Studio Code
- Python
4.3 实验原理
本次实验在 ID3 算法基础上,通过预剪枝参数限制树的生长:
max_depth=3:限制树的最大深度为 3,避免树纵向过度延伸min_samples_split=2:只有当节点样本数≥2 时才允许分裂,避免对少量样本进行无效划分;min_samples_leaf=1:保证叶节点至少有 1 个样本,维持分类的有效性;- 预剪枝在树形构建过程中提前终止无效分支,既保留 “有自己的房子> 有工作 > 信贷情况 > 年龄段” 的核心业务逻辑,又去除冗余划分,提升模型稳健性。
4.4 实验步骤
- 环境配置:复用原有 Python 依赖库,无需额外配置;
- 数据准备:将
dataset.txt、testset.txt放入代码同级目录; - 数据加载:通过
pandas读取数据,分离特征与标签,验证数据正确性; - 模型构建:初始化带
max_depth等预剪枝参数的 ID3 决策树; - 模型训练:用训练集拟合预剪枝模型,输出训练集准确率;
- 模型评估:手动计算指标,输出纯中文分类报告与测试集准确率;
- 可视化对比:绘制预剪枝决策树,与无剪枝模型对比树形差异;
- 结果分析:总结预剪枝对模型复杂度与可解释性的优化效果。
4.5 实验结果与分析
- 准确率分析:预剪枝模型训练集与测试集准确率均为 100%,与无剪枝模型一致,说明在小样本数据集上,预剪枝未损失拟合能力,同时优化了树形结构;
- 泛化能力分析:预剪枝通过限制树复杂度,避免了未来数据扩充后可能出现的过拟合,模型稳健性更优;
- 参数有效性分析:
max_depth=3在本次实验中为最优取值
4.6 实验结论
- 预剪枝可有效简化决策树结构,在不损失分类准确率的前提下,提升模型的可解释性与稳健性;
max_depth是预剪枝的核心参数,合理设置可平衡模型拟合能力与泛化能力;- 带预剪枝的 ID3 决策树更适配贷款审批等金融场景,既保证分类准确性,又提供清晰的决策依据;
五、实验总结
本次实验完整实现了贷款审批决策树的预剪枝流程,明确了预剪枝参数的作用与调优方法,通过对比验证了预剪枝对模型结构的优化效果。实验结果表明,预剪枝在小样本场景下可保持分类准确率,同时为大样本场景提供抗过拟合能力,是决策树工程落地的关键步骤。
更多推荐

所有评论(0)