共计 2275 个字符,预计需要花费 6 分钟才能阅读完成。
深入解析 CART 决策树:从算法原理到工程实践
背景痛点
决策树在工业界有着广泛的应用,比如风控评分卡、推荐系统、医疗诊断等。然而,传统的决策树实现常常面临以下问题:

- 过拟合:模型在训练集上表现很好,但在测试集上表现差,泛化能力不足。
- 特征重要性漂移:随着数据分布的变化,特征的重要性可能发生漂移,影响模型的稳定性。
- 计算复杂度高:尤其是在高维特征场景下,决策树的训练和推理可能变得非常耗时。
算法对比
决策树的核心在于如何选择最优的特征进行分裂。常见的算法有 ID3、C4.5 和 CART,它们在特征选择上有所不同:
-
ID3(信息增益):
$$ IG(D, A) = H(D) – H(D|A) $$
其中,$H(D)$ 是数据集 D 的熵,$H(D|A)$ 是在特征 A 条件下的条件熵。 -
C4.5(信息增益比):
$$ IG_{ratio}(D, A) = \frac{IG(D, A)}{H_A(D)} $$
其中,$H_A(D)$ 是特征 A 的熵,用于解决 ID3 对多值特征的偏好问题。 -
CART(Gini 系数):
$$ Gini(D) = 1 – \sum_{i=1}^{k} p_i^2 $$
Gini 系数越小,数据集的纯度越高。CART 采用二分法,适用于连续值和离散值特征。
核心实现
特征离散化处理
在 CART 中,连续特征需要进行离散化处理,常见的方法有等频分箱和等宽分箱:
- 等频分箱:将数据分成 n 个区间,每个区间包含相同数量的样本。
- 等宽分箱:将数据分成 n 个区间,每个区间的宽度相同。
递归分裂终止条件
决策树的构建是一个递归过程,终止条件通常包括:
- 当前节点的样本数少于预设阈值。
- 当前节点的 Gini 系数低于预设阈值。
- 树的深度达到预设最大值。
from sklearn.tree import DecisionTreeClassifier
# 初始化决策树模型
tree = DecisionTreeClassifier(
criterion='gini',
max_depth=5,
min_samples_leaf=10
)
# 训练模型
tree.fit(X_train, y_train)
后剪枝的 CCP 代价复杂度实现
后剪枝(Post-Pruning)通过最小化代价复杂度来防止过拟合:
$$ R_{\alpha}(T) = R(T) + \alpha |\widetilde{T}| $$
其中,$R(T)$ 是树的预测误差,$|\widetilde{T}|$ 是树的叶子节点数,$\alpha$ 是复杂度参数。
# 使用 CCP 进行后剪枝
path = tree.cost_complexity_pruning_path(X_train, y_train)
ccp_alphas = path.ccp_alphas
# 选择最优的 alpha
optimal_alpha = ccp_alphas[-2]
pruned_tree = DecisionTreeClassifier(ccp_alpha=optimal_alpha)
pruned_tree.fit(X_train, y_train)
代码示例
特征重要性可视化
import matplotlib.pyplot as plt
import seaborn as sns
# 获取特征重要性
feature_importance = tree.feature_importances_
# 可视化
plt.figure(figsize=(10, 6))
sns.barplot(x=feature_importance, y=X_train.columns)
plt.title('Feature Importance')
plt.show()
GridSearchCV 调参流程
from sklearn.model_selection import GridSearchCV
# 定义参数网格
param_grid = {'max_depth': [3, 5, 7],
'min_samples_leaf': [5, 10, 15],
'criterion': ['gini', 'entropy']
}
# 初始化 GridSearchCV
grid_search = GridSearchCV(estimator=DecisionTreeClassifier(),
param_grid=param_grid,
cv=5,
scoring='accuracy'
)
# 执行网格搜索
grid_search.fit(X_train, y_train)
# 输出最优参数
print(grid_search.best_params_)
生产建议
-
类别不平衡时的 class_weight 配置:
tree = DecisionTreeClassifier(class_weight='balanced') -
超参数 max_depth 与 min_samples_leaf 的权衡:
max_depth控制树的深度,防止过拟合。-
min_samples_leaf控制叶子节点的最小样本数,增加模型的稳定性。 -
在线学习时的增量更新策略:
- 使用
warm_start=True参数,允许模型在新增数据上继续训练。
延伸思考
- 如何将 CART 与 GBDT 结合提升效果?
-
GBDT(Gradient Boosting Decision Tree)通过迭代训练多个决策树,每棵树学习前序树的残差,可以显著提升模型的预测能力。
-
在实时推理场景下如何优化决策路径查询?
- 可以通过预计算决策路径,或者使用更高效的数据结构(如位图)来加速查询。
总结
CART 决策树是一种强大且灵活的机器学习算法,适用于各种分类和回归任务。通过合理设置超参数、使用剪枝技术以及结合其他算法(如 GBDT),可以进一步提升模型的性能。希望本文能帮助你在实际项目中更好地应用 CART 决策树。
