共计 1961 个字符,预计需要花费 5 分钟才能阅读完成。
决策树基础扫盲
决策树就像人类做决策的过程,通过一系列 if-else 规则对数据进行分类。举个例子:判断水果是苹果还是橘子,可能会先问『颜色是红色吗?』,再根据重量、形状等特征逐步细分。

常见的决策树算法有:
- ID3:使用信息增益选择特征,只能处理离散值,容易过拟合
- C4.5:改进版,用信息增益率选择特征,支持连续值处理
- CART(本文主角):使用基尼系数,能同时处理分类和回归任务
CART 算法核心原理
基尼系数计算
基尼系数衡量数据的不纯度,公式很简单:
Gini(D) = 1 - Σ(p_i)^2 # p_i 是第 i 类样本的比例
比如一个袋子有 3 红球 + 7 蓝球:
Gini = 1 - (0.3² + 0.7²) = 0.42
特征选择策略
CART 采用二分法:
1. 对每个特征的所有可能分割点计算基尼系数
2. 选择使基尼系数下降最大的特征作为分裂点
数学表达式:
ΔGini = Gini(D) - (|D1|/|D|)*Gini(D1) - (|D2|/|D|)*Gini(D2)
Python 手把手实现
先定义决策树节点结构:
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 gini(y):
_, counts = np.unique(y, return_counts=True)
probabilities = counts / len(y)
return 1 - np.sum(probabilities**2)
递归建树主逻辑:
def build_tree(X, y, depth=0, max_depth=5):
# 终止条件:纯度达标 / 达到最大深度 / 样本数太少
if (gini(y) < 0.01) or (depth == max_depth) or (len(y) < 5):
return Node(value=np.argmax(np.bincount(y)))
best_gini = float('inf')
best_feature, best_thresh = None, None
# 遍历所有特征和可能的分割点
for feature in range(X.shape[1]):
thresholds = np.unique(X[:, feature])
for thresh in thresholds:
left_idx = X[:, feature] <= thresh
g = (len(y[left_idx])/len(y))*gini(y[left_idx]) + \
(len(y[~left_idx])/len(y))*gini(y[~left_idx])
if g < best_gini:
best_gini = g
best_feature = feature
best_thresh = thresh
# 递归构建子树
left_idx = X[:, best_feature] <= best_thresh
left = build_tree(X[left_idx], y[left_idx], depth+1)
right = build_tree(X[~left_idx], y[~left_idx], depth+1)
return Node(feature=best_feature, threshold=best_thresh, left=left, right=right)
算法性能分析
- 时间复杂度 :O(mnlog(n)),其中 m 是特征数,n 是样本数
- 空间复杂度 :O(深度) 递归栈开销
与 sklearn 对比测试(鸢尾花数据集):
| 指标 | 自实现 CART | sklearn |
|---|---|---|
| 训练时间 (s) | 0.12 | 0.008 |
| 测试准确率 | 93.3% | 96.7% |
避坑指南
- 连续值处理 :
- 先排序,取相邻值中点作为候选分割点
-
对于大数据集可采用近似分位数
-
剪枝策略 :
- 预剪枝:限制最大深度 / 最小样本数
-
后剪枝:通过验证集评估剪枝收益
-
类别不平衡 :
- 使用加权基尼系数
- 对少数类样本过采样
思考进阶
- 当特征之间存在强相关性时,CART 会如何选择分裂特征?
- 如何修改算法使其支持回归任务(预测连续值)?
- 在百万级数据集上,有哪些优化计算效率的方法?
实现建议
建议先用小数据集(如鸢尾花)跑通整个流程,再尝试在 UCI 的成人收入数据集上实践。遇到问题时,可以:
- 可视化决策过程:用 graphviz 绘制树结构
- 打印中间变量:观察特征选择过程
- 对比 sklearn 结果:定位差异点
决策树是理解机器学习的最佳起点,希望本文能帮你打下坚实基础。在实际业务中,它常作为特征选择工具或集成学习的基模型,后续可以继续探索随机森林、GBDT 等进阶算法。
正文完
