Python实战:CART决策树算法原理与高效实现

1次阅读
没有评论

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

image.webp

决策树基础与 CART 算法特点

决策树通过一系列规则对数据进行分类或回归,而 CART(Classification and Regression Trees)是其中最经典的算法之一。与 ID3 和 C4.5 不同,CART 有两大特点:

Python 实战:CART 决策树算法原理与高效实现

  • 始终采用二叉树结构(即使特征有多个取值也通过二分法处理)
  • 分类任务使用基尼系数,回归任务使用方差作为分裂标准

核心数学原理详解

基尼系数计算

基尼系数衡量数据的不纯度,计算公式为:

def gini_index(groups, classes):
    n_instances = sum(len(group) for group in groups)
    gini = 0.0
    for group in groups:
        size = len(group)
        if size == 0:
            continue
        score = 0.0
        for class_val in classes:
            p = [row[-1] for row in group].count(class_val) / size
            score += p * p
        gini += (1.0 - score) * (size / n_instances)
    return gini

最佳分裂点选择

对于连续特征,CART 的二分过程:
1. 对特征值排序
2. 计算相邻值的中间点作为候选分割阈值
3. 选择基尼系数最小的分割方案

对比其他决策树算法

算法 分裂标准 树结构 是否支持回归 连续值处理
ID3 信息增益 多叉树 不支持
C4.5 信息增益比 多叉树 支持
CART 基尼系数 / 方差 二叉树 支持

Python 完整实现

树节点结构

from dataclasses import dataclass
from typing import Any, Union

@dataclass
class DecisionNode:
    feature_idx: int = None    # 用于分裂的特征索引
    threshold: float = None    # 分裂阈值
    left: Any = None           # 左子树
    right: Any = None          # 右子树
    value: Any = None          # 叶节点存储的预测值 

递归建树核心逻辑

def build_tree(X, y, max_depth, min_samples_split, depth=0):
    # 终止条件检查
    if len(set(y)) == 1 or len(y) <= min_samples_split or depth >= max_depth:
        return DecisionNode(value=most_common_label(y))

    # 寻找最佳分裂
    best_gini = float('inf')
    best_feature, best_thresh = None, None

    for feature_idx in range(X.shape[1]):
        thresholds = np.unique(X[:, feature_idx])
        for threshold in thresholds:
            left_idx = X[:, feature_idx] <= threshold
            groups = [y[left_idx], y[~left_idx]]
            gini = gini_index(groups, np.unique(y))

            if gini < best_gini:
                best_gini = gini
                best_feature = feature_idx
                best_thresh = threshold

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

    return DecisionNode(best_feature, best_thresh, left, right)

过拟合解决方案

预剪枝策略

  • 最大深度限制(max_depth)
  • 节点最小样本数(min_samples_split)
  • 叶节点最小样本数(min_samples_leaf)

后剪枝实现(代价复杂度剪枝)

def prune_tree(node, alpha, X_val, y_val):
    if node.value is not None:
        return node

    # 先剪枝子树
    node.left = prune_tree(node.left, alpha, X_val, y_val)
    node.right = prune_tree(node.right, alpha, X_val, y_val)

    # 如果两个子节点都是叶节点,考虑合并
    if node.left.value is not None and node.right.value is not None:
        # 计算合并前后的误差
        merged_error = compute_error(merge_nodes(node), X_val, y_val)
        original_error = compute_error(node, X_val, y_val)

        if merged_error + alpha <= original_error:
            return DecisionNode(value=node.left.value if len(node.left.value) > len(node.right.value) else node.right.value)

    return node

性能优化实践

特征离散化技巧

对于连续特征,等频分箱可提升分裂效率:

from sklearn.preprocessing import KBinsDiscretizer

discretizer = KBinsDiscretizer(n_bins=10, encode='ordinal', strategy='quantile')
X_discrete = discretizer.fit_transform(X_continuous)

并行化处理

对大型数据集,可并行计算不同特征的分裂质量:

from joblib import Parallel, delayed

results = Parallel(n_jobs=-1)(delayed(calc_best_split)(X, y, feature_idx) 
    for feature_idx in range(X.shape[1])
)

实战案例:鸢尾花分类

from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split

# 数据准备
iris = load_iris()
X_train, X_test, y_train, y_test = train_test_split(iris.data, iris.target, test_size=0.2)

# 模型训练
tree = build_tree(X_train, y_train, max_depth=3, min_samples_split=5)

# 评估
predictions = [predict_tree(tree, x) for x in X_test]
accuracy = np.mean(predictions == y_test)
print(f"Test Accuracy: {accuracy:.2f}")

生产环境注意事项

  1. 内存管理 :深度过大的树会导致栈溢出,建议:
  2. 限制 max_depth
  3. 改用迭代方式实现建树

  4. 类别特征处理

  5. 对高基数类别特征采用目标编码(Target Encoding)
  6. 或使用特殊的二分策略(如按类别出现频率排序)

  7. 模型持久化

  8. 使用 pickle 保存训练好的树结构
  9. 对于超大模型,考虑按节点分块存储

延伸思考

  1. 如何处理特征缺失值?可以尝试:
  2. 分裂时将缺失值单独作为一类
  3. 按照当前节点样本的分布比例随机分配

  4. 对于回归任务,如何修改分裂标准?

  5. 将基尼系数改为方差减少量
  6. 叶节点输出改为子节点样本的均值

  7. 如何扩展实现随机森林?

  8. 通过 bootstrap 采样创建多棵决策树
  9. 预测时采用投票或平均机制
正文完
 0
评论(没有评论)