深入解析CART决策树:从算法原理到工程实践

1次阅读
没有评论

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

image.webp

深入解析 CART 决策树:从算法原理到工程实践

背景痛点

决策树在工业界有着广泛的应用,比如风控评分卡、推荐系统、医疗诊断等。然而,传统的决策树实现常常面临以下问题:

深入解析 CART 决策树:从算法原理到工程实践

  • 过拟合:模型在训练集上表现很好,但在测试集上表现差,泛化能力不足。
  • 特征重要性漂移:随着数据分布的变化,特征的重要性可能发生漂移,影响模型的稳定性。
  • 计算复杂度高:尤其是在高维特征场景下,决策树的训练和推理可能变得非常耗时。

算法对比

决策树的核心在于如何选择最优的特征进行分裂。常见的算法有 ID3、C4.5 和 CART,它们在特征选择上有所不同:

  1. ID3(信息增益)
    $$ IG(D, A) = H(D) – H(D|A) $$
    其中,$H(D)$ 是数据集 D 的熵,$H(D|A)$ 是在特征 A 条件下的条件熵。

  2. C4.5(信息增益比)
    $$ IG_{ratio}(D, A) = \frac{IG(D, A)}{H_A(D)} $$
    其中,$H_A(D)$ 是特征 A 的熵,用于解决 ID3 对多值特征的偏好问题。

  3. CART(Gini 系数)
    $$ Gini(D) = 1 – \sum_{i=1}^{k} p_i^2 $$
    Gini 系数越小,数据集的纯度越高。CART 采用二分法,适用于连续值和离散值特征。

核心实现

特征离散化处理

在 CART 中,连续特征需要进行离散化处理,常见的方法有等频分箱和等宽分箱:

  • 等频分箱:将数据分成 n 个区间,每个区间包含相同数量的样本。
  • 等宽分箱:将数据分成 n 个区间,每个区间的宽度相同。

递归分裂终止条件

决策树的构建是一个递归过程,终止条件通常包括:

  1. 当前节点的样本数少于预设阈值。
  2. 当前节点的 Gini 系数低于预设阈值。
  3. 树的深度达到预设最大值。
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_)

生产建议

  1. 类别不平衡时的 class_weight 配置

    tree = DecisionTreeClassifier(class_weight='balanced')

  2. 超参数 max_depth 与 min_samples_leaf 的权衡

  3. max_depth控制树的深度,防止过拟合。
  4. min_samples_leaf控制叶子节点的最小样本数,增加模型的稳定性。

  5. 在线学习时的增量更新策略

  6. 使用 warm_start=True 参数,允许模型在新增数据上继续训练。

延伸思考

  1. 如何将 CART 与 GBDT 结合提升效果?
  2. GBDT(Gradient Boosting Decision Tree)通过迭代训练多个决策树,每棵树学习前序树的残差,可以显著提升模型的预测能力。

  3. 在实时推理场景下如何优化决策路径查询?

  4. 可以通过预计算决策路径,或者使用更高效的数据结构(如位图)来加速查询。

总结

CART 决策树是一种强大且灵活的机器学习算法,适用于各种分类和回归任务。通过合理设置超参数、使用剪枝技术以及结合其他算法(如 GBDT),可以进一步提升模型的性能。希望本文能帮助你在实际项目中更好地应用 CART 决策树。

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