从零开始理解CART决策树:手把手解析经典例题与实现细节

1次阅读
没有评论

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

image.webp

CART 决策树是机器学习中最直观的算法之一,特别适合处理分类和回归任务。比如在银行判断客户信用风险、电商预测用户购买行为时,它都能生成可解释的树形规则。本文将以鸢尾花数据集为例,带你完整走一遍特征选择、节点分裂和模型评估的全流程。

从零开始理解 CART 决策树:手把手解析经典例题与实现细节

一、从 Gini 系数理解节点分裂

假设我们有一个简化版的鸢尾花数据集,只考虑花瓣长度和宽度两个特征,目标是将花分为 Setosa 和 Versicolor 两类。决策树的核心问题是如何选择最佳分裂点:

  1. 计算父节点 Gini 系数
    原始节点有 50 个 Setosa 和 50 个 Versicolor 样本

    Gini = 1 - (\frac{50}{100})^2 - (\frac{50}{100})^2 = 0.5

  2. 评估所有可能的分裂点
    比如尝试按花瓣长度≤2.45cm 分裂:

  3. 左节点:48 个 Setosa 和 2 个 Versicolor
    Gini_{left} = 1 - (\frac{48}{50})^2 - (\frac{2}{50})^2 ≈ 0.076
  4. 右节点:2 个 Setosa 和 48 个 Versicolor
    Gini_{right} ≈ 0.076
  5. 加权平均 Gini:
    Gini_{split} = \frac{50}{100}×0.076 + \frac{50}{100}×0.076 = 0.076
  6. 信息增益:0.5 – 0.076 = 0.424

  7. 递归执行上述过程 直到满足停止条件(如 max_depth=3)

二、Python 实战:从数据到可视化

from sklearn.tree import DecisionTreeClassifier, export_graphviz
import graphviz

# 加载鸢尾花数据集(简化版)X = [[5.1, 3.5], [4.9, 3.0], ..., [6.7, 3.0]]  # 花瓣长 / 宽
Y = [0, 0, ..., 1]  # 0=Setosa, 1=Versicolor

# 关键参数说明:# max_depth: 控制树深防止过拟合
# min_samples_split: 节点最小样本数
clf = DecisionTreeClassifier(max_depth=3, criterion='gini')
clf.fit(X, Y)

# 可视化决策树
dot_data = export_graphviz(clf, feature_names=['petal_length', 'petal_width'], 
                          class_names=['Setosa', 'Versicolor'], filled=True)
graphviz.Source(dot_data)

三、避坑指南与实践技巧

  1. 连续特征处理
  2. 决策树天然支持连续值,但要注意:
  3. 排序后取相邻值的中位数作为候选分裂点
  4. 对大数据集可采用分位数减少计算量

  5. 缺失值应对策略

  6. 训练时:将缺失值单独分为一个分支
  7. 预测时:走出现概率更高的分支

  8. 剪枝参数调优

  9. 预剪枝(更常用):
    DecisionTreeClassifier(
        max_depth=5,
        min_samples_leaf=10,  # 叶节点最小样本数
        ccp_alpha=0.01  # 代价复杂度剪枝系数
    )
  10. 后剪枝:用 cost_complexity_pruning_path 找最佳 alpha

四、进阶思考与应用边界

最后留两个开放性问题:
1. 当特征大部分是类别型变量时,为什么 C4.5 可能比 CART 更合适?
2. 在金融风控场景中,决策树的哪些特性(如对单调性约束的支持不足)可能导致模型失效?

通过这个完整的案例,你应该已经掌握了 CART 的核心原理和实现方法。建议尝试用其他数据集(如泰坦尼克号生存预测)练习,观察不同参数对模型的影响。

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