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

不过在实际使用中,很多工程师会困惑:
- 面对大量特征时,如何选择最佳划分点?
- 连续型特征应该如何处理?
- 如何避免决策树过深导致的过拟合问题?
这正是本文要解决的核心问题。
CART vs ID3/C4.5:为什么选择 Gini 系数?
决策树算法有很多变种,最著名的包括 ID3、C4.5 和 CART。它们的主要区别在于划分标准:
- ID3 使用信息增益
- C4.5 使用信息增益比
- CART 使用 Gini 系数
Gini 系数的优势在于:
- 计算更简单,不需要计算对数
- 对类别分布更敏感,能更好处理类别不平衡
- 更倾向于产生平衡的划分
举个例子,对于一个二分类问题,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 系数降低最多的划分。
以 ” 年龄 ” 特征为例:
- 青年组:2 个不购买
Gini = 1 – (0² + 1²) = 0 - 中年组:1 个购买
Gini = 1 – (1² + 0²) = 0 - 老年组: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()
优化与调参建议
剪枝策略
决策树容易过拟合,剪枝是关键。有两种主要方法:
- 预剪枝:在构建树时提前停止
- 设置最大深度 (max_depth)
- 设置叶节点最小样本数 (min_samples_leaf)
-
设置划分最小信息增益 (min_impurity_decrease)
-
后剪枝:先构建完整树再剪枝
- 代价复杂度剪枝 (ccp_alpha)
参数调优示例
下表展示了不同参数组合在验证集上的表现:
| max_depth | min_samples_leaf | 准确率 |
|---|---|---|
| 3 | 1 | 0.82 |
| 5 | 2 | 0.85 |
| None | 5 | 0.83 |
| 4 | 3 | 0.86 |
常见问题与解决方案
- 样本不均衡 :
- 使用 class_weight 参数调整类别权重
-
对少数类样本进行过采样
-
高基数类别特征 :
- 考虑目标编码 (target encoding)
- 使用均值编码 (mean encoding)
-
分组稀有类别
-
计算资源有限 :
- 限制树的最大深度
- 使用随机采样减少训练数据量
- 考虑特征选择减少输入维度
值得思考的问题
- 当特征同时包含连续型和分类型时,如何设计最佳的预处理流程?
- 在在线学习场景中,如何增量更新决策树模型?
- 如何解释决策树的预测结果,使其对业务人员更友好?
通过这个完整的例子,相信你已经掌握了 CART 算法的核心要点。决策树虽然简单,但在实际业务中非常实用,特别是在需要模型可解释性的场景。下一步,你可以尝试在自己的数据集上应用这些技巧,观察不同参数对模型性能的影响。
