一、决策树剪枝介绍

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 实验目的

  1. 理解决策树过拟合的成因与预剪枝的核心原理;
  2. 掌握scikit-learn中决策树预剪枝参数的设置与调优方法;
  3. 完成带预剪枝的贷款审批分类实验,对比剪枝前后的模型差异;

4.2 实验环境

  • Windows 10/11
  • Visual Studio Code
  • Python 

4.3 实验原理

本次实验在 ID3 算法基础上,通过预剪枝参数限制树的生长:

  1. max_depth=3:限制树的最大深度为 3,避免树纵向过度延伸
  2. min_samples_split=2:只有当节点样本数≥2 时才允许分裂,避免对少量样本进行无效划分;
  3. min_samples_leaf=1:保证叶节点至少有 1 个样本,维持分类的有效性;
  4. 预剪枝在树形构建过程中提前终止无效分支,既保留 “有自己的房子> 有工作 > 信贷情况 > 年龄段” 的核心业务逻辑,又去除冗余划分,提升模型稳健性。

4.4 实验步骤

  1. 环境配置:复用原有 Python 依赖库,无需额外配置;
  2. 数据准备:将dataset.txt、testset.txt放入代码同级目录;
  3. 数据加载:通过pandas读取数据,分离特征与标签,验证数据正确性;
  4. 模型构建:初始化带max_depth等预剪枝参数的 ID3 决策树;
  5. 模型训练:用训练集拟合预剪枝模型,输出训练集准确率;
  6. 模型评估:手动计算指标,输出纯中文分类报告与测试集准确率;
  7. 可视化对比:绘制预剪枝决策树,与无剪枝模型对比树形差异;
  8. 结果分析:总结预剪枝对模型复杂度与可解释性的优化效果。

4.5 实验结果与分析

  1. 准确率分析:预剪枝模型训练集与测试集准确率均为 100%,与无剪枝模型一致,说明在小样本数据集上,预剪枝未损失拟合能力,同时优化了树形结构;
  2. 泛化能力分析:预剪枝通过限制树复杂度,避免了未来数据扩充后可能出现的过拟合,模型稳健性更优;
  3. 参数有效性分析:max_depth=3在本次实验中为最优取值

4.6 实验结论

  1. 预剪枝可有效简化决策树结构,在不损失分类准确率的前提下,提升模型的可解释性与稳健性;
  2. max_depth是预剪枝的核心参数,合理设置可平衡模型拟合能力与泛化能力;
  3. 带预剪枝的 ID3 决策树更适配贷款审批等金融场景,既保证分类准确性,又提供清晰的决策依据;

五、实验总结

本次实验完整实现了贷款审批决策树的预剪枝流程,明确了预剪枝参数的作用与调优方法,通过对比验证了预剪枝对模型结构的优化效果。实验结果表明,预剪枝在小样本场景下可保持分类准确率,同时为大样本场景提供抗过拟合能力,是决策树工程落地的关键步骤。

更多推荐