共计 2744 个字符,预计需要花费 7 分钟才能阅读完成。
决策树的训练与构建
训练决策树的过程确实就是构建决策树的过程,这两个说法本质上是等价的。决策树的构建是一个递归的过程,通过不断地选择最优特征对数据集进行划分,直到满足停止条件。ID3 算法就是其中最经典的构建方法之一。

信息熵与信息增益
决策树的核心在于如何选择最优划分特征,ID3 算法使用信息增益作为特征选择标准。
-
信息熵(Entropy):表示随机变量的不确定性,定义为:
$$H(D) = -\sum_{k=1}^{K}p_k\log_2 p_k$$
其中 $p_k$ 是第 k 类样本在数据集 D 中的比例 -
信息增益(Information Gain):表示特征 A 对数据集 D 的信息增益,定义为:
$$Gain(D,A) = H(D) – \sum_{v=1}^{V}\frac{|D^v|}{|D|}H(D^v)$$
其中 V 是特征 A 的取值个数,$D^v$ 是特征 A 取值为 v 的子集
ID3 算法流程
- 从根节点开始,计算所有特征的信息增益
- 选择信息增益最大的特征作为当前节点的划分特征
- 对特征的每一个取值创建子节点
- 递归地在子节点上重复上述过程,直到:
- 所有样本属于同一类别
- 没有剩余特征可供划分
- 子节点样本数小于阈值
Python 实现示例
import math
from collections import Counter
class DecisionTreeID3:
def __init__(self, max_depth=5):
self.max_depth = max_depth
self.tree = {}
def entropy(self, labels):
"""计算信息熵"""
counts = Counter(labels)
probs = [count/len(labels) for count in counts.values()]
return -sum(p * math.log2(p) for p in probs)
def information_gain(self, data, feature, labels):
"""计算信息增益"""
total_entropy = self.entropy(labels)
# 按特征值分组
values = set([d[feature] for d in data])
subsets = [[d for d in data if d[feature] == v] for v in values]
# 计算条件熵
conditional_entropy = 0
for subset in subsets:
prob = len(subset)/len(data)
subset_labels = [labels[i] for i, d in enumerate(data) if d[feature] == subset[0][feature]]
conditional_entropy += prob * self.entropy(subset_labels)
return total_entropy - conditional_entropy
def fit(self, data, features, labels, depth=0):
"""递归构建决策树"""
# 终止条件 1:所有样本属于同一类别
if len(set(labels)) == 1:
return labels[0]
# 终止条件 2:没有特征或达到最大深度
if not features or depth >= self.max_depth:
return Counter(labels).most_common(1)[0][0]
# 选择最佳划分特征
best_feature = max(features, key=lambda f: self.information_gain(data, f, labels))
tree = {best_feature: {}}
# 递归构建子树
remaining_features = [f for f in features if f != best_feature]
for value in set([d[best_feature] for d in data]):
subset_data = [d for d in data if d[best_feature] == value]
subset_labels = [labels[i] for i, d in enumerate(data) if d[best_feature] == value]
tree[best_feature][value] = self.fit(subset_data, remaining_features, subset_labels, depth+1)
return tree
# 测试用例
if __name__ == "__main__":
data = [{'天气':'晴', '温度':'高', '湿度':'高', '风':'弱'},
{'天气':'晴', '温度':'高', '湿度':'高', '风':'强'},
{'天气':'阴', '温度':'高', '湿度':'高', '风':'弱'},
{'天气':'雨', '温度':'中', '湿度':'高', '风':'弱'},
{'天气':'雨', '温度':'低', '湿度':'正常', '风':'弱'},
{'天气':'雨', '温度':'低', '湿度':'正常', '风':'强'},
{'天气':'阴', '温度':'低', '湿度':'正常', '风':'强'},
{'天气':'晴', '温度':'中', '湿度':'高', '风':'弱'},
]
labels = ['不玩', '不玩', '玩', '玩', '玩', '不玩', '玩', '不玩']
features = ['天气', '温度', '湿度', '风']
dt = DecisionTreeID3()
tree = dt.fit(data, features, labels)
print(tree)
算法对比与演进
ID3 算法之后,决策树算法主要有两个发展方向:
-
C4.5 算法:改进 ID3 只能处理离散值的限制,引入信息增益比解决 ID3 偏向选择取值多特征的问题,支持连续值和缺失值处理。
-
CART 算法:使用基尼系数代替信息增益,可以处理回归问题,采用二叉树结构提高效率。
避坑指南
过拟合问题
- 预剪枝 :在构建过程中提前停止,如限制树的最大深度、设置最小样本数等
- 后剪枝 :先构建完整树,再自底向上剪枝,通常效果更好但计算成本更高
连续值处理
ID3 算法本身不支持连续值,实际应用中可以采用:
- 离散化处理:设定阈值将连续值分段
- 改进算法:采用 C4.5 或 CART 算法
思考题
- 当数据中存在缺失值时,该如何处理?
- 为什么 ID3 算法倾向于选择取值多的特征?有什么改进方法?
- 如何扩展决策树来处理多分类问题?
总结
决策树是机器学习中最直观的算法之一,ID3 算法作为其基础实现,虽然简单但包含了决策树的核心思想。通过本文的讲解和代码实现,相信你已经掌握了决策树的基本构建过程。在实际应用中,建议从 ID3 开始理解原理,再逐步过渡到更强大的 C4.5 和 CART 算法。
正文完
发表至: 未分类
近一天内
