从零构建CART决策树模型:原理详解与Python实战指南

1次阅读
没有评论

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

image.webp

决策树基础概念

决策树是一种模仿人类决策过程的机器学习方法。CART(Classification and Regression Trees)是其中最经典的算法之一,由 Breiman 等人于 1984 年提出。与 ID3 和 C4.5 算法相比,CART 有以下几个显著特点:

从零构建 CART 决策树模型:原理详解与 Python 实战指南

  • 二叉树结构:每个节点只分裂为两个子节点
  • 基尼系数:使用基尼不纯度 (Gini Impurity) 作为分裂标准
  • 支持回归:既可以处理分类问题也可以处理回归问题

基尼系数 vs 信息增益

基尼系数和信息增益都是衡量数据集不纯度的指标,但计算方式不同:

基尼系数公式:
$$ Gini(p) = 1 – \sum_{k=1}^{K} p_k^2 $$

信息增益公式:
$$ IG(D_p, f) = I(D_p) – \sum_{j=1}^{m} \frac{N_j}{N_p} I(D_j) $$

其中 $p_k$ 是第 k 类样本的比例,$I$ 可以是不纯度指标(如熵)。

Python 实现关键步骤

1. 数据预处理

import pandas as pd
from sklearn.preprocessing import LabelEncoder

def preprocess_data(df):
    # 处理缺失值
    for col in df.columns:
        if df[col].dtype == 'object':
            df[col].fillna(df[col].mode()[0], inplace=True)
        else:
            df[col].fillna(df[col].median(), inplace=True)

    # 类别变量编码
    categorical_cols = df.select_dtypes(include=['object']).columns
    for col in categorical_cols:
        le = LabelEncoder()
        df[col] = le.fit_transform(df[col])

    return df

2. 基尼系数计算函数

import numpy as np

def gini_impurity(y):
    """计算基尼不纯度"""
    if len(y) == 0:
        return 0

    # 计算每个类别的比例
    p = np.bincount(y) / len(y)
    return 1 - np.sum(p ** 2)

# 向量化优化版本
def gini_impurity_vectorized(y):
    """向量化计算的基尼不纯度"""
    _, counts = np.unique(y, return_counts=True)
    p = counts / len(y)
    return 1 - np.sum(p ** 2)

3. 递归分裂终止条件

决策树生长需要设置合理的停止条件,常见的有:

  • 最大深度(max_depth)
  • 最小样本数(min_samples_split)
  • 节点纯度阈值(min_impurity_decrease)
class DecisionNode:
    """决策树节点类"""
    def __init__(self, feature_idx=None, threshold=None,
                 left=None, right=None, value=None):
        self.feature_idx = feature_idx  # 分裂特征索引
        self.threshold = threshold      # 分裂阈值
        self.left = left                # 左子树
        self.right = right              # 右子树
        self.value = value              # 叶节点值

完整模型实现

class DecisionTreeClassifier:
    def __init__(self, max_depth=None, min_samples_split=2):
        self.max_depth = max_depth
        self.min_samples_split = min_samples_split
        self.tree_ = None

    def fit(self, X, y):
        self.n_classes_ = len(np.unique(y))
        self.n_features_ = X.shape[1]
        self.tree_ = self._grow_tree(X, y)

    def _grow_tree(self, X, y, depth=0):
        n_samples, n_features = X.shape
        n_classes = len(np.unique(y))

        # 停止条件
        if (self.max_depth is not None and depth >= self.max_depth) or \
           n_samples < self.min_samples_split or \
           n_classes == 1:
            leaf_value = self._most_common_label(y)
            return DecisionNode(value=leaf_value)

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

        for feature_idx in range(n_features):
            thresholds = np.unique(X[:, feature_idx])
            for threshold in thresholds:
                left_idx = X[:, feature_idx] <= threshold
                gini = self._gini_split(y, left_idx)

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

        # 递归生长子树
        left_idx = X[:, best_feature] <= best_threshold
        left = self._grow_tree(X[left_idx], y[left_idx], depth+1)
        right = self._grow_tree(X[~left_idx], y[~left_idx], depth+1)

        return DecisionNode(feature_idx=best_feature,
                           threshold=best_threshold,
                           left=left, right=right)

    def _gini_split(self, y, left_idx):
        """计算分裂后的加权基尼系数"""
        n = len(y)
        n_left, n_right = sum(left_idx), sum(~left_idx)

        if n_left == 0 or n_right == 0:
            return float('inf')

        gini_left = gini_impurity(y[left_idx])
        gini_right = gini_impurity(y[~left_idx])

        return (n_left/n)*gini_left + (n_right/n)*gini_right

    def _most_common_label(self, y):
        """返回最常见的类别"""
        return np.argmax(np.bincount(y))

    def predict(self, X):
        return np.array([self._predict_tree(x, self.tree_) for x in X])

    def _predict_tree(self, x, node):
        if node.value is not None:
            return node.value

        if x[node.feature_idx] <= node.threshold:
            return self._predict_tree(x, node.left)
        else:
            return self._predict_tree(x, node.right)

可视化决策树

可以使用 graphviz 库可视化决策树:

from graphviz import Digraph

def visualize_tree(tree, feature_names=None):
    dot = Digraph()
    _add_nodes(dot, tree, feature_names)
    return dot

def _add_nodes(dot, node, feature_names, parent=None, edge_label=None):
    if node.value is not None:
        dot.node(str(id(node)), label=f'Class {node.value}', shape='box')
    else:
        if feature_names is not None:
            feature = feature_names[node.feature_idx]
        else:
            feature = f'Feature {node.feature_idx}'

        dot.node(str(id(node)), 
                label=f'{feature} <= {node.threshold:.2f}')

    if parent is not None:
        dot.edge(str(id(parent)), str(id(node)), label=edge_label)

    if node.left is not None:
        _add_nodes(dot, node.left, feature_names, node, 'True')
    if node.right is not None:
        _add_nodes(dot, node.right, feature_names, node, 'False')

生产环境注意事项

1. 连续特征分箱

对于连续特征,直接使用所有可能值作为分割点可能效率低下。可以考虑:

  • 等宽分箱:将特征值范围均匀划分为 N 个区间
  • 等频分箱:每个区间包含相同数量的样本
  • 基于决策树的分箱

2. 类别不平衡处理

可以通过加权基尼系数来处理类别不平衡问题:

def weighted_gini(y, sample_weight=None):
    if sample_weight is None:
        sample_weight = np.ones(len(y))

    classes = np.unique(y)
    total_weight = np.sum(sample_weight)
    weighted_p = []

    for c in classes:
        class_weight = np.sum(sample_weight[y == c])
        weighted_p.append(class_weight / total_weight)

    return 1 - np.sum(np.array(weighted_p) ** 2)

3. 后剪枝

后剪枝可以防止过拟合,代价复杂度剪枝是常用方法:

def cost_complexity_pruning(tree, X_val, y_val, alpha):
    """
    代价复杂度剪枝
    alpha: 复杂度参数,控制剪枝强度
    """
    # 计算剪枝前后的代价复杂度
    # 这里简化实现,实际需要递归计算所有可能剪枝

    # 计算验证集准确率
    def accuracy(y_true, y_pred):
        return np.mean(y_true == y_pred)

    # 剪枝前准确率
    acc_before = accuracy(y_val, tree.predict(X_val))

    # 尝试剪枝每个内部节点
    # 实际实现需要考虑所有可能的剪枝组合

    return pruned_tree

MNIST 数据集测试

我们可以用 MNIST 数据集测试我们的实现:

from sklearn.datasets import load_digits
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score

# 加载数据
digits = load_digits()
X, y = digits.data, digits.target

# 划分训练测试集
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

# 训练我们的决策树
our_tree = DecisionTreeClassifier(max_depth=5)
our_tree.fit(X_train, y_train)
our_pred = our_tree.predict(X_test)
print(f"Our Tree Accuracy: {accuracy_score(y_test, our_pred):.4f}")

# 对比 sklearn
from sklearn.tree import DecisionTreeClassifier

sk_tree = DecisionTreeClassifier(max_depth=5, random_state=42)
sk_tree.fit(X_train, y_train)
sk_pred = sk_tree.predict(X_test)
print(f"Sklearn Tree Accuracy: {accuracy_score(y_test, sk_pred):.4f}")

延伸思考题

  1. 如何将 CART 扩展为随机森林?
  2. 通过 Bagging 集成多个决策树
  3. 在每次分裂时随机选择特征子集
  4. 通过投票或平均得到最终预测

  5. 针对高维稀疏数据应如何优化分裂效率?

  6. 使用稀疏矩阵存储数据结构
  7. 对连续特征进行离散化预处理
  8. 使用近似算法加速最佳分裂点搜索
  9. 考虑特征哈希或降维技术

总结

本文详细介绍了 CART 决策树的核心原理和 Python 实现。我们从基尼系数出发,逐步实现了数据预处理、树生长、预测和可视化等关键功能。通过 MNIST 数据集的测试,我们的实现与 sklearn 的性能相近。最后我们还讨论了生产环境中的优化技巧和扩展方向。

决策树是很多强大模型的基础,理解其原理和实现细节对深入掌握机器学习非常重要。希望本文能帮助初学者更好地理解和使用决策树算法。

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