CART算法实战:如何用决策树解决分类问题(附完整例题解析)

1次阅读
没有评论

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

image.webp

为什么需要 CART 算法?

在机器学习项目中,我们经常遇到需要根据特征对数据进行分类的场景。比如银行要判断客户是否会违约,电商要预测用户是否会购买商品。决策树作为一种直观易懂的算法,非常适合这类问题。而 CART(Classification and Regression Trees) 算法是其中应用最广泛的一种实现。

CART 算法实战:如何用决策树解决分类问题(附完整例题解析)

不过在实际使用中,很多工程师会困惑:

  • 面对大量特征时,如何选择最佳划分点?
  • 连续型特征应该如何处理?
  • 如何避免决策树过深导致的过拟合问题?

这正是本文要解决的核心问题。

CART vs ID3/C4.5:为什么选择 Gini 系数?

决策树算法有很多变种,最著名的包括 ID3、C4.5 和 CART。它们的主要区别在于划分标准:

  • ID3 使用信息增益
  • C4.5 使用信息增益比
  • CART 使用 Gini 系数

Gini 系数的优势在于:

  1. 计算更简单,不需要计算对数
  2. 对类别分布更敏感,能更好处理类别不平衡
  3. 更倾向于产生平衡的划分

举个例子,对于一个二分类问题,Gini 系数的计算公式是:

Gini = 1 - (p₁² + p₂²)

其中 p₁和 p₂分别是两个类别的比例。当数据完全纯净时(只有一个类别),Gini 系数为 0;当两个类别各占 50% 时,Gini 系数达到最大值 0.5。

实战:完整例题解析

让我们通过一个真实案例来理解 CART 算法的工作流程。假设我们有如下数据集,用于预测客户是否会购买产品:

年龄 收入 学生 信用评级 购买
青年 一般
青年
中年 一般
老年 一般
老年 一般
老年

第一步:计算根节点的 Gini 系数

整个数据集中,购买的比例是 3 /6=0.5,不购买也是 0.5。因此根节点的 Gini 系数为:

Gini = 1 - (0.5² + 0.5²) = 0.5

第二步:评估各个特征的划分效果

我们需要计算每个特征划分后的加权 Gini 系数,选择使 Gini 系数降低最多的划分。

以 ” 年龄 ” 特征为例:

  1. 青年组:2 个不购买
    Gini = 1 – (0² + 1²) = 0
  2. 中年组:1 个购买
    Gini = 1 – (1² + 0²) = 0
  3. 老年组:2 个购买,1 个不购买
    Gini = 1 – ((2/3)² + (1/3)²) ≈ 0.444

加权 Gini = (2/6)0 + (1/6)0 + (3/6)*0.444 ≈ 0.222

类似的,我们可以计算其他特征的 Gini 系数降低情况。通过比较发现,” 年龄 ” 特征的 Gini 系数降低最多(从 0.5 降到 0.222),因此选择它作为第一个划分特征。

Python 实现代码

from sklearn.tree import DecisionTreeClassifier, plot_tree
import matplotlib.pyplot as plt
from sklearn.preprocessing import LabelEncoder

# 准备数据
data = {'age': ['青年','青年','中年','老年','老年','老年'],
    'income': ['高','高','高','中','低','低'],
    'student': ['否','否','否','否','是','是'],
    'credit': ['一般','好','一般','一般','一般','好'],
    'buy': ['否','否','是','是','是','否']
}

# 将类别特征编码
le = LabelEncoder()
X = pd.DataFrame({col: le.fit_transform(data[col]) for col in ['age','income','student','credit']})
y = le.fit_transform(data['buy'])

# 构建决策树模型
model = DecisionTreeClassifier(criterion='gini', max_depth=3)
model.fit(X, y)

# 可视化决策树
plt.figure(figsize=(12,8))
plot_tree(model, feature_names=['age','income','student','credit'], 
          class_names=['不买','买'], filled=True)
plt.show()

# 特征重要性
plt.barh(['age','income','student','credit'], model.feature_importances_)
plt.title('特征重要性')
plt.show()

优化与调参建议

剪枝策略

决策树容易过拟合,剪枝是关键。有两种主要方法:

  1. 预剪枝:在构建树时提前停止
  2. 设置最大深度 (max_depth)
  3. 设置叶节点最小样本数 (min_samples_leaf)
  4. 设置划分最小信息增益 (min_impurity_decrease)

  5. 后剪枝:先构建完整树再剪枝

  6. 代价复杂度剪枝 (ccp_alpha)

参数调优示例

下表展示了不同参数组合在验证集上的表现:

max_depth min_samples_leaf 准确率
3 1 0.82
5 2 0.85
None 5 0.83
4 3 0.86

常见问题与解决方案

  1. 样本不均衡
  2. 使用 class_weight 参数调整类别权重
  3. 对少数类样本进行过采样

  4. 高基数类别特征

  5. 考虑目标编码 (target encoding)
  6. 使用均值编码 (mean encoding)
  7. 分组稀有类别

  8. 计算资源有限

  9. 限制树的最大深度
  10. 使用随机采样减少训练数据量
  11. 考虑特征选择减少输入维度

值得思考的问题

  1. 当特征同时包含连续型和分类型时,如何设计最佳的预处理流程?
  2. 在在线学习场景中,如何增量更新决策树模型?
  3. 如何解释决策树的预测结果,使其对业务人员更友好?

通过这个完整的例子,相信你已经掌握了 CART 算法的核心要点。决策树虽然简单,但在实际业务中非常实用,特别是在需要模型可解释性的场景。下一步,你可以尝试在自己的数据集上应用这些技巧,观察不同参数对模型性能的影响。

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