决策树入门实战:从原理到代码实现CART算法

1次阅读
没有评论

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

image.webp

决策树基础扫盲

决策树就像人类做决策的过程,通过一系列 if-else 规则对数据进行分类。举个例子:判断水果是苹果还是橘子,可能会先问『颜色是红色吗?』,再根据重量、形状等特征逐步细分。

决策树入门实战:从原理到代码实现 CART 算法

常见的决策树算法有:

  • ID3:使用信息增益选择特征,只能处理离散值,容易过拟合
  • C4.5:改进版,用信息增益率选择特征,支持连续值处理
  • CART(本文主角):使用基尼系数,能同时处理分类和回归任务

CART 算法核心原理

基尼系数计算

基尼系数衡量数据的不纯度,公式很简单:

Gini(D) = 1 - Σ(p_i)^2  # p_i 是第 i 类样本的比例 

比如一个袋子有 3 红球 + 7 蓝球:

Gini = 1 - (0.3² + 0.7²) = 0.42

特征选择策略

CART 采用二分法:
1. 对每个特征的所有可能分割点计算基尼系数
2. 选择使基尼系数下降最大的特征作为分裂点

数学表达式:

ΔGini = Gini(D) - (|D1|/|D|)*Gini(D1) - (|D2|/|D|)*Gini(D2)

Python 手把手实现

先定义决策树节点结构:

class Node:
    def __init__(self, feature=None, threshold=None, left=None, right=None, value=None):
        self.feature = feature   # 分裂特征
        self.threshold = threshold  # 分裂阈值
        self.left = left         # 左子树
        self.right = right       # 右子树
        self.value = value       # 叶节点预测值 

关键函数——计算基尼系数:

def gini(y):
    _, counts = np.unique(y, return_counts=True)
    probabilities = counts / len(y)
    return 1 - np.sum(probabilities**2)

递归建树主逻辑:

def build_tree(X, y, depth=0, max_depth=5):
    # 终止条件:纯度达标 / 达到最大深度 / 样本数太少
    if (gini(y) < 0.01) or (depth == max_depth) or (len(y) < 5):
        return Node(value=np.argmax(np.bincount(y)))

    best_gini = float('inf')
    best_feature, best_thresh = None, None

    # 遍历所有特征和可能的分割点
    for feature in range(X.shape[1]):
        thresholds = np.unique(X[:, feature])
        for thresh in thresholds:
            left_idx = X[:, feature] <= thresh
            g = (len(y[left_idx])/len(y))*gini(y[left_idx]) + \
                (len(y[~left_idx])/len(y))*gini(y[~left_idx])

            if g < best_gini:
                best_gini = g
                best_feature = feature
                best_thresh = thresh

    # 递归构建子树            
    left_idx = X[:, best_feature] <= best_thresh
    left = build_tree(X[left_idx], y[left_idx], depth+1)
    right = build_tree(X[~left_idx], y[~left_idx], depth+1)

    return Node(feature=best_feature, threshold=best_thresh, left=left, right=right)

算法性能分析

  • 时间复杂度 :O(mnlog(n)),其中 m 是特征数,n 是样本数
  • 空间复杂度 :O(深度) 递归栈开销

与 sklearn 对比测试(鸢尾花数据集):

指标 自实现 CART sklearn
训练时间 (s) 0.12 0.008
测试准确率 93.3% 96.7%

避坑指南

  1. 连续值处理
  2. 先排序,取相邻值中点作为候选分割点
  3. 对于大数据集可采用近似分位数

  4. 剪枝策略

  5. 预剪枝:限制最大深度 / 最小样本数
  6. 后剪枝:通过验证集评估剪枝收益

  7. 类别不平衡

  8. 使用加权基尼系数
  9. 对少数类样本过采样

思考进阶

  1. 当特征之间存在强相关性时,CART 会如何选择分裂特征?
  2. 如何修改算法使其支持回归任务(预测连续值)?
  3. 在百万级数据集上,有哪些优化计算效率的方法?

实现建议

建议先用小数据集(如鸢尾花)跑通整个流程,再尝试在 UCI 的成人收入数据集上实践。遇到问题时,可以:

  • 可视化决策过程:用 graphviz 绘制树结构
  • 打印中间变量:观察特征选择过程
  • 对比 sklearn 结果:定位差异点

决策树是理解机器学习的最佳起点,希望本文能帮你打下坚实基础。在实际业务中,它常作为特征选择工具或集成学习的基模型,后续可以继续探索随机森林、GBDT 等进阶算法。

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