简介:与西瓜书《机器学习》第四章配套的决策树Python代码实现包,适合正在学习决策树原理、需要动手复现经典算法的读者,可作为课程作业或自学实验的参考。压缩包共9个文件,包含5个csv数据文件与4个py脚本:数据文件覆盖西瓜数据集2.0、3.0及UCI数据集,脚本实现基于信息熵的划分选择、基于基尼指数的CART算法、预剪枝与后剪枝对比,以及决策树可视化绘图,包体仅16KB,轻量易用。目前已有10322人学习下载。资源对应书中4.3、4.4、4.6节编程任务,既有从零构建决策树的完整代码,也有剪枝前后效果比较与多数据集实验分析思路,能帮助读者直观理解不同划分标准和剪枝策略对模型泛化能力的影响,快速完成从算法原理到代码实现的落地。
1. 决策树代码实现的底层选择:从西瓜书第四章到可运行代码
拿到周志华《机器学习》第四章,你会发现所有公式推导都指向同一个结论:决策树的核心就是“如何选特征、何时停”。但只看公式,你依然写不出能跑通的代码。很多人第一反应是调 sklearn 的 DecisionTreeClassifier,可一旦进入调参,面对 max_depth、min_samples_split 这些参数时完全不知道内部发生了什么。与其黑匣子式调库,不如把 ID3、C4.5、CART 用 Python 自己实现一遍,再和 sklearn 的结果对照。这篇笔记就是从西瓜书第四章出发,带你从零写出一棵能画出来、能剪枝、能处理连续值和缺失值的决策树,适合正在啃西瓜书、准备机器学习期末复现算法、或者在真实数据上想搞清楚分裂逻辑的人。
2. 核心算法实现:从信息增益到递归建树,把公式变成可运行代码
2.1 信息熵与信息增益:先写对两个基础函数
决策树的第一个坑往往是信息熵没写对。信息熵公式是 -Σ p·log2(p),p 是每个类别占比,代码实现时要注意两点:log 的底数不敏感,因为同一特征下所有条件熵共享同一底数,比较顺序不受影响;但标签计数为 0 时 math.log(0) 会直接报错,必须跳过。
import math from collections import Counter def entropy(labels): """计算信息熵,labels 是样本标签列表""" counter = Counter(labels) total = len(labels) ent = 0.0 for count in counter.values(): p = count / total ent -= p * math.log2(p) return ent这个函数是所有后续计算的基石。逻辑很简单:统计每个类别的数量,算占比,累加负 p·log2(p)。注意 total 不能在循环里重复取,否则 p 算错定位时非常隐蔽。
条件熵和信息增益建立在特征划分之上。假设样本是 list of dict,每个 dict 是一条记录,key 是特征名,value 是特征值,另外单独传入一个标签列表。
def split_by_feature(data, labels, feature): """按特征的特征值划分数据,返回 {特征值: (子数据集, 子标签)}""" groups = {} for sample, label in zip(data, labels): value = sample[feature] if value not in groups: groups[value] = ([], []) groups[value][0].append(sample) groups[value][1].append(label) return groups def info_gain(data, labels, feature): """计算某特征的信息增益""" base_ent = entropy(labels) groups = split_by_feature(data, labels, feature) cond_ent = 0.0 total = len(data) for sub_data, sub_labels in groups.values(): weight = len(sub_data) / total cond_ent += weight * entropy(sub_labels) return base_ent - cond_ent这里有个容易被忽略的细节:split_by_feature 里 groups 的值是「一对列表」,一个存子数据集,一个存子标签。如果不打包而只存子数据集,后续递归时还要重新对齐标签,索引错位是血泪级事故。
2.2 增益率与基尼指数:C4.5 和 CART 的差异在哪
ID3 用信息增益选特征,缺点是偏向取值多的特征——比如把“编号”作为特征,每个样本一个值,条件熵直接为 0,信息增益最大,但泛化能力为零。C4.5 用增益率修正,CART 用基尼指数,三者选特征的逻辑代码结构一致,只是评价函数不同。
def gain_ratio(data, labels, feature): """C4.5 的增益率:信息增益 / 特征自身的固有值""" gain = info_gain(data, labels, feature) groups = split_by_feature(data, labels, feature) total = len(data) intrinsic = 0.0 for sub_data, _ in groups.values(): weight = len(sub_data) / total intrinsic -= weight * math.log2(weight) return gain / intrinsic if intrinsic > 0 else 0.0 def gini_index(data, labels, feature): """CART 的基尼指数:按特征值划分后的加权基尼""" total = len(data) groups = split_by_feature(data, labels, feature) gini = 0.0 for sub_data, sub_labels in groups.values(): weight = len(sub_data) / total counter = Counter(sub_labels) sub_gini = 1.0 for count in counter.values(): p = count / len(sub_labels) sub_gini -= p * p gini += weight * sub_gini return gini注意增益率公式里,分母是特征取值分布的熵,叫固有值。如果某个特征每个样本取值都不相同,固有值很大,增益率被压低,这就平衡了 ID3 的偏好。基尼指数越小越好,和信息增益相反,选特征时逻辑要反转。
2.3 递归建树:终止条件、叶子节点与树的存储结构
树的结构用嵌套 dict 表达:内部节点存「划分特征」,子节点按特征值映射到子树;叶子节点直接存类别。关键在递归出口,西瓜书第四章给了三个:当前节点样本都属于同一类别、当前可用特征集为空、当前样本在可用特征上取值相同。
def build_tree(data, labels, features, criterion='id3'): """递归建树,criterion 可选 id3 / c4.5 / cart""" if len(set(labels)) == 1: return labels[0] if not features: return Counter(labels).most_common(1)[0][0] # 检查当前样本在剩余特征上是否取值完全一致 first = data[0] identical = True for feat in features: if any(sample[feat] != first[feat] for sample in data): identical = False break if identical: return Counter(labels).most_common(1)[0][0] # 选特征 best_feat = None best_score = -math.inf if criterion != 'cart' else math.inf for feat in features: if criterion == 'id3': score = info_gain(data, labels, feat) if score > best_score: best_score = score best_feat = feat elif criterion == 'c4.5': score = gain_ratio(data, labels, feat) if score > best_score: best_score = score best_feat = feat else: # cart 选基尼指数最小 score = gini_index(data, labels, feat) if score < best_score: best_score = score best_feat = feat # 递归构建子节点 tree = {best_feat: {}} remaining_features = [f for f in features if f != best_feat] groups = split_by_feature(data, labels, best_feat) for value, (sub_data, sub_labels) in groups.items(): tree[best_feat][value] = build_tree( sub_data, sub_labels, remaining_features, criterion ) return tree这段代码兼顾三个算法的统一框架,只换评价函数就能在 ID3、C4.5、CART 之间切换。选特征时注意 cart 的 score 初始值是正无穷,比较符号是小于,其他两个是大于。common 的口袋里还有个常见误用:递归时 features 列表直接用原列表再删元素,会导致兄弟节点共享同一份特征列表,递归回去后特征被污染。正确做法是用列表推导式生成新列表,不修改原 features。
3. 处理西瓜数据集:离散特征、连续特征与缺失值的完整方案
3.1 数据的 Python 结构设计:list of dict 与标签对齐
西瓜书第四章用西瓜数据集训练决策树,最经典的版本是 2.0 和 3.0。3.0 版含 17 条样本、8 个特征,其中“密度”和“含糖率”是连续特征。这里直接以 3.0 数据为底,设计成 Python 内置结构,保证代码跑起来不依赖外部文件。
# 特征名和样本数据,保持和西瓜书一致 features = ['色泽', '根蒂', '敲声', '纹理', '脐部', '触感', '密度', '含糖率'] # 每条样本一个 dict,对应一个标签 raw_data = [ ({'色泽': '青绿', '根蒂': '蜷缩', '敲声': '浊响', '纹理': '清晰', '脐部': '凹陷', '触感': '硬滑', '密度': 0.697, '含糖率': 0.460}, '好瓜'), ({'色泽': '乌黑', '根蒂': '蜷缩', '敲声': '沉闷', '纹理': '清晰', '脐部': '凹陷', '触感': '硬滑', '密度': 0.774, '含糖率': 0.376}, '好瓜'), # 其余 15 条按西瓜书 3.0 表 4.1 补齐 ]设计理由:list of dict 的好处是特征名直接当 key,划分时不需要记住特征索引,后续可视化要打印节点分裂信息也很方便。标签单独一个列表,和 data 列表按下标对齐,永远不要做成一个 dict 里塞 label,否则分组时要额外剥离。
加载函数的另一重价值是处理“连续特征和离散特征同时存在”的场景。离散特征直接按键值分组,连续特征需要先离散化,这放在 3.2 节。实际工程里数据来源可能是 CSV,但为了把注意力放在决策树本身,建议先用这种内嵌数据打通流程,再换成 pandas 读取也不迟。
3.2 连续特征二分法:候选切分点的计算与递归限制
连续特征的划分和离散特征完全不同。取值 0.697、0.774 这样的密度,不能在递归时按值分组,那样一个值一个组,退化回“编号”特征。西瓜书的做法:对特征值排序,取相邻两个值的中点作为候选切分点,用信息增益或基尼指数选出最优切分点。
def best_continuous_split(data, labels, feature): """对连续特征找最优二分点,返回 (切分阈值, 信息增益或基尼)""" sorted_samples = sorted(zip(data, labels), key=lambda x: x[0][feature]) values = [sample[feature] for sample, _ in sorted_samples] total = len(data) base_ent = entropy(labels) best_threshold = None best_gain = -math.inf for i in range(total - 1): if values[i] == values[i + 1]: continue threshold = (values[i] + values[i + 1]) / 2 # 按阈值二分 left_labels = [label for sample, label in sorted_samples if sample[feature] <= threshold] right_labels = [label for sample, label in sorted_samples if sample[feature] > threshold] cond_ent = (len(left_labels) / total) * entropy(left_labels) + \ (len(right_labels) / total) * entropy(right_labels) gain = base_ent - cond_ent if gain > best_gain: best_gain = gain best_threshold = threshold return best_threshold, best_gain这段代码把排序和阈值搜索放在一起,实现里最容易翻车的是 sorted_samples 按特征值升序排列后,left 对应 <= 阈值。如果只按值排序而没把数据和标签一起移动,划分后标签和样本错位,信息增益会计算出离谱的结果。
选连续特征时要注意递归顺序:离散特征用一次就从可用特征列表里删除,但连续特征不能删。西瓜书明确说,连续特征在当前节点的最优切分点已被使用,但后续子节点还可以用同一特征、选更细的切分点。所以代码要把连续特征单独处理,递归时始终保留在候选特征里。
3.3 缺失值处理:样本权重归一化的实现细节
数据缺失在西瓜书第四章有专门讨论,核心是“样本权重”。无缺失值样本按比例加权参与特征选择,有缺失值样本按不同取值落入子节点的概率分配权重。简化版本适用于数据结构:对每个特征判断哪些样本有缺失,用有值的子集计算增益,再按有值样本占比打折。
def info_gain_with_missing(data, labels, feature): """带缺失值的信息增益:只统计有该特征值的样本""" valid_idx = [i for i, sample in enumerate(data) if sample.get(feature) is not None] if not valid_idx: return 0.0 valid_data = [data[i] for i in valid_idx] valid_labels = [labels[i] for i in valid_idx] rho = len(valid_idx) / len(data) # 用有效样本算信息增益,再乘 rho return rho * info_gain(valid_data, valid_labels, feature)这里的关键是 rho 乘在外面。如果直接拿有效样本算增益不乘比例,等于认为缺失值不存在,会高估该特征的区分能力。乘 rho 后,缺失比例高的特征在竞争中天然处于劣势,这是西瓜书公式的直观含义。
缺失值在最终预测阶段还要处理:给 dict 树增加特殊分支,当待预测样本某个特征缺失时,按该节点训练时样本的多数类别投票。真实项目里更稳妥的做法是直接填充训练集众数,但要注意这种做法会把“缺失”本身的信息抹掉,两种方案各有利弊。
4. 决策树可视化:用 matplotlib 画树与 sklearn 对照验证
4.1 matplotlib 递归画树:布局计算与节点样式
代码能跑只是第一步,能看见树长什么样才能判断分裂逻辑对不对。用 matplotlib 手工画树的关键是算出每个节点的坐标,常规做法是递归时给每个叶节点分配一个 x 坐标,内部节点的 x 坐标取两个孩子中点,y 坐标按深度递减。
import matplotlib.pyplot as plt def plot_tree(tree, x, y, parent_x, parent_y, annotate): """递归画树,x, y 为当前节点坐标,parent_* 为父节点坐标""" if not isinstance(tree, dict): # 叶子节点 plt.text(x, y, tree, bbox=dict(boxstyle='round,pad=0.3', facecolor='lightgreen', edgecolor='black')) return feature = list(tree.keys())[0] children = tree[feature] child_x = x - 2 ** (y - 2) # 用深度控制横向偏移 for value, subtree in children.items(): plot_x = child_x # 递归画子树 plt.plot([x, plot_x], [y, y - 1], color='gray', linewidth=0.8) plt.text((x + plot_x) / 2, (y + plot_x) / 2, value, fontsize=9, color='blue') plot_tree(subtree, plot_x, y - 1, x, y, value) plt.text(x, y, feature, bbox=dict(fill=False, edgecolor='black'), ha='center', fontsize=10)坐标计算的原则:y 减 1 表示深度加一层,x 的间距按 2 的负深度幂缩放,叶子节点分布在整个画布上。这里最容易踩坑的是子节点坐标计算变量名冲突,建议把当前节点坐标和父节点坐标分开变量维护。
画完树最好加一行 plt.show(),把训练时选的特征和阈值打印出来核对。第一次跑出树的时候,用 sklearn 的 DecisionTreeClassifier 在相同数据上训练,把树的结构打印出来对比,就能验证手写版和工业版差异在哪。
4.2 与 sklearn 对照:手写版和开源库的结果一致性验证
验证代码正确性最直接的办法是跑 sklearn 的决策树,传入相同离散化后的数据,对比树结构的节点分裂特征。由于 sklearn 自带连续特征最优切分点搜索和 CART 算法,和手写版直接对比时要用完全相同的数据预处理。
from sklearn.tree import DecisionTreeClassifier from sklearn.preprocessing import OrdinalEncoder # 离散特征编码 encoder = OrdinalEncoder(categories='auto') X_encoded = encoder.fit_transform( [[sample[feat] for feat in features[:6]] for sample, _ in raw_data] ) # 连续特征单独按阈值二值化 threshold, _ = best_continuous_split( [s for s, _ in raw_data], [l for _, l in raw_data], '密度' ) for i, (sample, label) in enumerate(raw_data): X_encoded[i][6] = 1 if sample['密度'] <= threshold else 0 X_encoded[i][7] = 1 if sample['含糖率'] <= /* 同样的阈值计算 */ else 0 clf = DecisionTreeClassifier(criterion='entropy', random_state=0, max_depth=5) clf.fit(X_encoded, [l for _, l in raw_data])对照时注意一个客观差异:sklearn 默认对离散特征也用 CART 的二分方式处理,而西瓜书手写版是按多分支方式划分离散特征,树形会明显不同,不表示哪方代码错了。判别的关键在于手写版在离散特征上多分支、连续特征二分,sklearn 全部分支都是二分。想得到一致结果,对手写版加限制,强制离散特征也二分化处理,或者只比较同一子树下的分裂顺序。
验证的实用建议:打印手写版和 sklearn 每层选择的分裂特征,前三层一致就可以认为实现正确。第四层开始会因为细节偏差分叉,不影响算法实现正确性。
5. 决策树实现避坑:五个让新手翻车的细节与排查方法
5.1 症状:训练完成后,某些分支从未被选中过
原因:连续特征没有正确递归,导致分裂点只选了全局最优,后续子节点无法继续用同一特征找细分点。用密度去区分“好瓜”,根节点选了 0.381 作为阈值,子节点里密度均匀分布就没法再分裂。
解决:给 build_tree 增加连续特征逻辑,递归时把连续特长保留在候选里,离散特征移除。具体做法在 3.2 节已经实现,这里提醒一点:连续特征保留后会让特征集合永远非空,必须靠剪枝或深度限制来防止树无限生长。
5.2 症状:训练集准确率 100%,测试集表现极差
现象是经典过拟合,决策树每层都把样本分尽,到后面每个叶子只剩一个样本。
原因:没有预设停止条件,或者只用“样本同类别”作为唯一出口,树会生长到完全拟合。西瓜书第四章说,决策树算法不进行剪枝就会过拟合,书里对剪枝的讨论也集中在这里。
解决:在 build_tree 里加 max_depth 参数,递归时深度超过限制就返回多数类别。另一个常见做法是 min_samples_split,少于 N 个样本就不分裂。用验证集做后剪枝的效果更好,见第 6 章。
5.3 症状:递归调用报错 RecursionError: maximum recursion depth exceeded
原因:某个特征在划分时没能减少样本数,比如离散特征某个取值只有一个样本,但该样本标签不纯或特征相同,递归卡死在原地。
解决:检查 split_by_feature 每次划分后,子数据长度是否严格小于父数据集。若某个分支的子数据长度等于父数据集,说明特征取值只有一种,此时的递归永远不会收敛。在 build_tree 的递归入口加一个保护判断:len(unique_values) < 2 时直接返回多数类别。
5.4 症状:信息增益算出来是负数或 nan
原因:math.log2 的参数必须是正数,而某个子集标签为空。常见来源是特征取值非常多,某个取值只对应零个样本,或者连续特征切分后一侧样本数为 0。
解决:在 entropy 函数里加保护,counter 里已无缺失;filter 掉空子集;连续特征二分时如果左右任意一侧为空,跳过该切分点。更隐蔽的情况是数据里混入 None 值,entropy 会按一个独立类别算,模型完全错乱,同样需要提前清洗。
5.5 症状:手写版和 sklearn 的输出完全不同
原因:离散特征处理方式不同,sklearn 的 CART 全是二分,手写版默认多分支。另一个常见原因是连续特征的切分点搜索范围不同,手写版按西瓜书取相邻均值,sklearn 用优化后的近似枚举。
解决:不要追求完全一致的树形,只对比根节点和前两层的分裂特征与阈值。若前两层一致,说明核心计算逻辑正确。若根节点都不同,优先排查信息熵或基尼指数是否算对,用简单数据集做单元测试。
6. 剪枝的代码实现与调参:预剪枝、后剪枝及其在真实数据上的表现
6.1 预剪枝实现:max_depth 与 min_samples_split 参数化
预剪枝在生成树的过程中提前终止,代码就是在 build_tree 递归入口加两个参数判断深度和样本数。
def build_tree_pruned(data, labels, features, criterion='id3', max_depth=5, min_samples_split=2, depth=0): if len(set(labels)) == 1 or depth >= max_depth or len(data) < min_samples_split: return Counter(labels).most_common(1)[0][0] # 其余逻辑和之前 build_tree 一样,递归时 depth+1 # 选特征、划分、递归预剪枝容易欠拟合,因为提前终止可能错过后续有区分度的分裂。真实项目里 max_depth 取 3 到 5 之间,min_samples_split 取 2 到 10,具体值要对着验证集调。
6.2 后剪枝实现:用测试集验证子树替换的增量效果
后剪枝在树完全长好后从底向上回缩,判断标准是:把子树替换成多数类别叶子,如果测试集精度不下降,就替换。代码比预剪枝多一步,需要遍历整棵树找出所有决策节点。
def post_prune(tree, val_data, val_labels, features): """后剪枝:返回剪枝后的树""" if not isinstance(tree, dict): return tree feature = list(tree.keys())[0] children = tree[feature] new_children = {} for value, subtree in children.items(): if isinstance(subtree, dict): new_children[value] = post_prune(subtree, val_data, val_labels, features) else: new_children[value] = subtree tree[feature] = new_children # 尝试剪掉当前节点 majority = Counter(val_labels).most_common(1)[0][0] if evaluate_tree(tree, val_data, val_labels) <= evaluate_tree(majority, val_data, val_labels): return majority return tree后剪枝效果通常比预剪枝好,但计算量更大,因为每个替换都要在验证集上评估一次精度。小数据上用全量验证没问题,大数据量常见做法是抽一个子集做后剪枝评估。
我的习惯是先跑一个 max_depth=10 的完整树,再看训练精度和验证精度的差距,差距大就先上后剪枝,不够再配合预剪枝参数。剪枝这件事没有最优解,只能靠验证集反馈反复调整,但手写实现这个过程,你对 sklearn 的 ccp_alpha、min_samples_leaf 背后的动机就全清楚了。这套代码把西瓜书第四章的公式过了一遍,希望能帮到正在读决策树、准备期末复现算法的你,动手把每个函数跑通一次,比背十遍公式有用。
本文还有配套的精品资源,点击获取