共计 2406 个字符,预计需要花费 7 分钟才能阅读完成。
背景介绍
决策树是机器学习中最直观且广泛使用的算法之一,而 CART(Classification and Regression Trees)算法因其简单高效的特性,成为实际项目中的常见选择。与 ID3 和 C4.5 不同,CART 既能处理分类问题也能解决回归任务,并且总是生成二叉树结构,这使得它在实现和解释上都更加简洁。

核心算法解析
基尼系数计算
基尼系数是 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. 选择基尼系数减少最多的分裂方案
递归构建过程
构建过程遵循典型的分治策略:
1. 从根节点开始,寻找最佳分裂
2. 创建左右子节点
3. 对子节点递归执行分裂过程
4. 直到满足停止条件(如达到最大深度)
Python 实现
以下是 CART 决策树的核心类框架:
class DecisionTree:
def __init__(self, max_depth=5, min_samples_split=2):
self.max_depth = max_depth
self.min_samples_split = min_samples_split
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 fit(self, X, y):
self.n_classes = len(set(y))
self.n_features = X.shape[1]
self.tree = self._grow_tree(np.column_stack((X, y)))
def _gini(self, y):
m = y.size
return 1.0 - sum((np.sum(y == c) / m) ** 2 for c in range(self.n_classes))
def _best_split(self, X, y):
# 具体实现省略...
return best_feature, best_thresh
def _grow_tree(self, data, depth=0):
# 递归构建树的具体实现...
return node
简单案例演示
以鸢尾花数据集的前两个特征为例:
from sklearn.datasets import load_iris
iris = load_iris()
X = iris.data[:, :2] # 只用前两个特征
tree = DecisionTree(max_depth=3)
tree.fit(X, iris.target)
# 可视化决策边界
x_min, x_max = X[:, 0].min() - 1, X[:, 0].max() + 1
y_min, y_max = X[:, 1].min() - 1, X[:, 1].max() + 1
xx, yy = np.meshgrid(np.arange(x_min, x_max, 0.02),
np.arange(y_min, y_max, 0.02))
Z = tree.predict(np.c_[xx.ravel(), yy.ravel()])
Z = Z.reshape(xx.shape)
plt.contourf(xx, yy, Z, alpha=0.4)
plt.scatter(X[:, 0], X[:, 1], c=iris.target, s=20, edgecolor='k')
plt.xlabel('Sepal length')
plt.ylabel('Sepal width')
生产环境注意事项
过拟合问题
- 预剪枝策略:
- 限制最大深度(max_depth)
- 设置节点最小样本数(min_samples_split)
- 定义信息增益阈值
- 后剪枝策略:
- 通过验证集评估剪枝收益
- 代价复杂度剪枝(CCP)
处理连续值
CART 天然支持连续特征:
1. 对特征值排序
2. 取相邻值的中间点作为候选分裂点
3. 选择基尼系数最小的分裂点
缺失值处理
- 替代分裂:当特征缺失时使用替代特征
- 概率分配:将样本分配到左右子节点
性能优化建议
并行化处理
- 特征选择阶段可以并行计算
- 使用 joblib 并行化特征评估:
from joblib import Parallel, delayed def evaluate_features(args): # 特征评估逻辑 return gini results = Parallel(n_jobs=-1)(delayed(evaluate_features)(feature) for feature in features)
内存优化
- 使用稀疏矩阵存储数据
- 对类别特征进行数值编码
- 限制树的深度减少内存占用
总结与扩展
CART 决策树作为基础算法,其价值不仅在于单独使用,更是随机森林和 GBDT 等集成方法的基石。在实际项目中,决策树特别适合以下场景:
– 需要模型可解释性的业务场景
– 特征包含混合类型(连续 + 离散)
– 数据存在缺失值的情况
未来改进方向可以考虑:
– 实现增量学习支持
– 加入更灵活的分裂规则
– 与线性模型结合提高表现
正文完
