共计 5534 个字符,预计需要花费 14 分钟才能阅读完成。
决策树基础概念
决策树是一种模仿人类决策过程的机器学习方法。CART(Classification and Regression Trees)是其中最经典的算法之一,由 Breiman 等人于 1984 年提出。与 ID3 和 C4.5 算法相比,CART 有以下几个显著特点:

- 二叉树结构:每个节点只分裂为两个子节点
- 基尼系数:使用基尼不纯度 (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}")
延伸思考题
- 如何将 CART 扩展为随机森林?
- 通过 Bagging 集成多个决策树
- 在每次分裂时随机选择特征子集
-
通过投票或平均得到最终预测
-
针对高维稀疏数据应如何优化分裂效率?
- 使用稀疏矩阵存储数据结构
- 对连续特征进行离散化预处理
- 使用近似算法加速最佳分裂点搜索
- 考虑特征哈希或降维技术
总结
本文详细介绍了 CART 决策树的核心原理和 Python 实现。我们从基尼系数出发,逐步实现了数据预处理、树生长、预测和可视化等关键功能。通过 MNIST 数据集的测试,我们的实现与 sklearn 的性能相近。最后我们还讨论了生产环境中的优化技巧和扩展方向。
决策树是很多强大模型的基础,理解其原理和实现细节对深入掌握机器学习非常重要。希望本文能帮助初学者更好地理解和使用决策树算法。
