共计 2324 个字符,预计需要花费 6 分钟才能阅读完成。
决策树与 CART 算法概述
决策树是机器学习中经典的分类与回归方法,通过树形结构模拟人类决策过程。相比其他算法(如 ID3 使用信息增益、C4.5 使用增益率),CART(Classification and Regression Trees)的最大特点是:

- 二叉树结构:每个节点最多分裂为两个子节点,简化决策路径
- 基尼系数(Gini Index):作为分类任务的分裂标准,计算纯度损失更高效
- 动态类型处理:同时支持分类(Gini/ 错分率)和回归(平方误差)任务
痛点分析与业务影响
1. 高维特征计算效率
当特征维度超过 1000 时,传统递归分割可能消耗数小时。曾在一个用户画像项目中,5000 维特征使训练时间达到 8 小时。
2. 过拟合现象
某金融风控案例显示,未剪枝的决策树在训练集准确率 99%,但测试集仅 72%,导致实际业务误判率飙升。
3. 类别不平衡问题
医疗诊断数据中正样本仅占 5%,模型会倾向于预测多数类,造成漏诊风险。
Python 完整实现(含注释)
# 数据预处理示例
import pandas as pd
from sklearn.model_selection import train_test_split
data = pd.read_csv('sample_data.csv')
X = data.drop('target', axis=1)
y = data['target']
# 类别标签编码
from sklearn.preprocessing import LabelEncoder
le = LabelEncoder()
y_encoded = le.fit_transform(y)
# 基尼系数计算函数
def gini_impurity(y):
_, counts = np.unique(y, return_counts=True)
prob = counts / len(y)
return 1 - np.sum(prob ** 2)
# 节点分裂逻辑(核心代码段)class Node:
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 # 叶节点预测值
可视化实现
from sklearn.tree import export_graphviz
import graphviz
dot_data = export_graphviz(
model,
out_file=None,
feature_names=X.columns,
class_names=['class0', 'class1'],
filled=True
)
graph = graphviz.Source(dot_data)
graph.render('decision_tree')
优化方案实战
剪枝对比实验
| 方法 | 测试集准确率 | 树深度 |
|---|---|---|
| 预剪枝 | 85.2% | 5 |
| 后剪枝 | 86.7% | 7 |
| 不剪枝 | 72.1% | 32 |
# 后剪枝实现示例
from sklearn.tree._tree import TREE_LEAF
def prune_index(inner_tree, index):
if inner_tree.children_left[index] == TREE_LEAF:
return
prune_index(inner_tree, inner_tree.children_left[index])
prune_index(inner_tree, inner_tree.children_right[index])
# 计算剪枝前后误差变化
before_prune = calculate_error(...)
after_prune = calculate_error(...)
if after_prune <= before_prune:
inner_tree.children_left[index] = TREE_LEAF
inner_tree.children_right[index] = TREE_LEAF
特征重要性评估
importances = model.feature_importances_
indices = np.argsort(importances)[::-1]
plt.figure()
plt.title("Feature Importances")
plt.bar(range(X.shape[1]), importances[indices])
plt.xticks(range(X.shape[1]), X.columns[indices], rotation=90)
plt.show()
生产环境实践
- 内存监控方案:
- 使用
memory_profiler包记录训练过程内存消耗 -
设置
max_depth参数控制内存峰值 -
模型持久化:
import joblib joblib.dump(model, 'cart_model.pkl') # 比 pickle 更高效 loaded_model = joblib.load('cart_model.pkl') -
增量训练:
- 通过
warm_start=True参数复用已有树结构 - 配合
partial_fit方法实现在线学习
开放思考题
- 如何将 CART 与 GBDT 等集成方法结合发挥更大价值?
- 在实时推理场景下,如何优化决策树的查询效率?
- 面对非结构化数据(如文本),CART 需要怎样的特征工程改造?
通过这个案例可以看到,CART 决策树在保持可解释性的同时,通过合理的优化手段完全可以达到生产级精度。关键是根据业务场景选择合适的剪枝策略和特征处理方法。
正文完
