Python实战:如何用CART决策树解决分类问题中的过拟合难题

1次阅读
没有评论

共计 2162 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

技术背景

CART(Classification and Regression Trees)决策树是一种基于树结构的监督学习算法,通过递归二分数据实现分类或回归。其核心优势在于:

Python 实战:如何用 CART 决策树解决分类问题中的过拟合难题

  • 可解释性强:决策路径可视化为 if-then 规则
  • 非参数特性:无需假设数据分布
  • 多类型数据处理:同时支持数值型和类别型特征

数学上,分类问题采用基尼指数(Gini Index)作为分裂标准:

$$
Gini = 1 – \sum_{i=1}^k p_i^2
$$

其中 $p_i$ 是节点中第 $i$ 类样本的比例。每次分裂选择使基尼指数下降最大的特征。

痛点分析

CART 决策树的过拟合主要表现为:

  1. 模型过度复杂:树深度过大导致创建只适用于训练数据的特殊规则
  2. 对噪声敏感:在叶节点中捕获训练数据中的随机波动
  3. 泛化能力差:测试集准确率显著低于训练集

主要原因包括:

  • 无限制生长直到所有叶节点纯净
  • 缺乏对模型复杂度的惩罚机制

解决方案

预剪枝实现

通过限制树的生长提前防止过拟合,常用 sklearn 的 DecisionTreeClassifier 参数:

from sklearn.tree import DecisionTreeClassifier

# 预剪枝参数设置示例
clf = DecisionTreeClassifier(
    max_depth=3,               # 最大树深度
    min_samples_split=10,      # 节点最小分裂样本数
    min_samples_leaf=5,        # 叶节点最小样本数
    max_features='sqrt',       # 每次分裂考虑的特征数
    random_state=42
)
clf.fit(X_train, y_train)

参数说明:
max_depth:控制树的最大层级
min_samples_split:节点可分裂的最小样本量
min_samples_leaf:确保叶节点有足够支持

后剪枝实现

代价复杂度剪枝(CCP)通过交叉验证选择最优子树:

import matplotlib.pyplot as plt
from sklearn.model_selection import cross_val_score

# 计算 CCP 路径
path = clf.cost_complexity_pruning_path(X_train, y_train)
ccp_alphas = path.ccp_alphas

# 遍历不同 alpha 值训练模型
clfs = []
for ccp_alpha in ccp_alphas:
    clf = DecisionTreeClassifier(random_state=42, ccp_alpha=ccp_alpha)
    clf.fit(X_train, y_train)
    clfs.append(clf)

# 绘制准确率随 alpha 变化曲线
train_scores = [clf.score(X_train, y_train) for clf in clfs]
test_scores = [clf.score(X_test, y_test) for clf in clfs]

plt.plot(ccp_alphas, train_scores, label="train")
plt.plot(ccp_alphas, test_scores, label="test")
plt.legend()
plt.show()

选择使测试集准确率最高的 ccp_alpha 作为最终参数。

性能优化

特征重要性评估

import pandas as pd

# 获取特征重要性
importance = pd.DataFrame({
    'feature': X_train.columns,
    'importance': clf.feature_importances_
}).sort_values('importance', ascending=False)

# 可视化
importance.plot.bar(x='feature', y='importance')

与随机森林对比

指标 单棵 CART 随机森林
抗过拟合能力
特征重要性 不稳定 更可靠
训练速度 较慢

避坑指南

类别不平衡处理

  1. 调整类别权重:

    class_weight = {0:1, 1:5}  # 少数类权重提高
    clf = DecisionTreeClassifier(class_weight=class_weight)

  2. 使用 SMOTE 过采样少数类

超参数调优建议

  1. 优先调节 max_depthmin_samples_leaf
  2. 使用网格搜索:
    from sklearn.model_selection import GridSearchCV
    
    params = {'max_depth': [3,5,7],
        'min_samples_leaf': [1,3,5]
    }
    gs = GridSearchCV(clf, params, cv=5)
    gs.fit(X_train, y_train)

总结与延伸

学习路径建议

  1. 进阶阅读《Elements of Statistical Learning》第 9 章
  2. 实践 CART 在回归问题中的应用
  3. 研究 GBDT 等基于决策树的集成方法

思考题

尝试在 UCI 的 Bank Marketing 数据集(https://archive.ics.uci.edu/ml/datasets/Bank+Marketing)上实现以下任务:
1. 比较预剪枝与后剪枝的效果差异
2. 当特征中存在大量类别型变量时,如何处理能提升模型性能?
3. 设计实验验证不同剪枝方法对模型推理速度的影响

正文完
 0
评论(没有评论)