ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

西瓜书决策树代码全解析:从ID3到剪枝与多变量实现

西瓜书决策树代码全解析:从ID3到剪枝与多变量实现 简介《机器学习》西瓜书第4章决策树4.5节的配套示例代码包面向想要结合源码理解C4.5决策树原理的初学者也适合需要快速复现算法的开发者。压缩包共4个文件其中包含2个ipynb交互式笔记本、1个Python脚本和1个csv数据集。笔记本部分按数据读取、特征离散化、建树、剪枝、评估的顺序分步演示每一步都保留中间结果便于对照教材逐行学习Python脚本封装了核心训练函数可直接复用csv数据集提供一份心脏疾病样本加载后无需额外预处理即可运行。整个包仅18KB代码与数据高度精简没有过多依赖方便在本地环境快速试验也可以作为课后作业或课程设计的参考模板。目前已有436人学习下载内容覆盖信息增益计算、递归划分、剪枝处理与准确率验证并带有中间过程输出能帮助读者厘清C4.5算法在真实数据上的处理细节是深入理解西瓜书4.5节内容不可多得的实用材料。1. 西瓜书4.5代码.zip这份压缩包装的是“最后一公里”的例子不是现成的预测器很多人把周志华《机器学习》的配套代码压缩包下载下来名字往往就是“西瓜书4.5代码.zip”这种风格。解压之后第一眼看到一堆 Python 文件和 CSV 数据集第一反应是“跑一个 main.py 看看结果”结果要么报错要么画出来的树和图例对不上。这份代码对应的其实是第 4 章决策树的 4.24.5 节实现尤其 4.5 节的多变量决策树是树模型从“能用”到“理解原理”的分水岭。它解决的核心问题是把信息增益、剪枝处理、连续属性离散化这些公式变成能跑、能改、能输出图像的工程代码。适合正在复现教材算法、做课程设计、准备面试手撕树模型的人。下文从原理讲到运行参数再给出搭树时的参数习惯和踩坑记录。2. 先把4.5节的底摸清楚多变量决策树为什么值得单独写一段代码2.1 单变量决策树的局限轴平行边界在什么场景下翻车前几节生成的树每次只拿一个属性做判断比如“纹理是否清晰”“根蒂是否蜷缩”。从几何角度看这种划分边界是轴平行的也就是跟坐标轴完全平行或垂直的直线段。对西瓜数据集这类天然离散属性多的问题这通常够用。但一旦特征分布呈明显的斜线单变量树就必须用很多层阶梯去逼近。一个典型情况是二维特征空间里真实边界是“密度大于 0.38 且含糖率大于 0.2”构成的斜线。单变量树每层只能取一个特征做阈值切分斜线只能被一步步拆成横竖接替的长条树深很快超过 10 层中间节点的样本量越来越少叶子越来越碎。这时的训练集准确率可能很高验证集效果却差得让人想退货——这就是树模型典型过拟合也是我见过不少人在 4.2 节算完增益、兴冲冲画树之后翻车的起点。2.2 多变量决策树的建树流程把“属性选择”换成“线性分类器”第 4.5 节的思路很直接既然单属性切分不够就让每个非叶节点学一个线性分类器形如w1*x1 w2*x2 ... wn*xn t样本进入某个节点后先计算这个线性组合的值再跟阈值 t 比较决定走左子树还是右子树。“选哪个特征做分裂”变成了“学一组权重和阈值”。树的结构本质上没变仍然是递归划分只是分裂的规则从离散取值比较变成了线性模型的预测值比较。这个改动带来的收益是边界不再限制在轴平行方向。树可以用更少的分裂次数把倾斜边界切出来模型深度下降叶子数量减少泛化能力往往更好。代价也直接每个节点多了一次线性拟合的开销而且特征的量纲、数值范围、异常点都会影响权重学习这比算一次信息熵要敏感得多。因此很多“西瓜书4.5代码.zip”在实现这个节点时直接调 sklearn 的 LogisticRegression 或 Perceptron外层自己写建树循环不然纯手写线性模型代码工作量和 debug 成本会翻好几倍。2.3 自己写树还是套 sklearn三种做法对比网传各种版本的西瓜书决策树代码大体上分三类。第一种是自己实现熵、信息增益、递归建树整体逻辑可控能严格复现书上表格。第二种是直接调 sklearn 的 DecisionTreeClassifier代码短但默认的 CART 实现用的是基尼指数特征还要做数值编码跟书上的 ID3、C4.5 对不上。第三种是多变量决策树常见做法是在每个节点包一个 sklearn 线性模型外层自己控制递归和终止条件。做法复现书上公式支持多变量节点写起来适合场景自写 ID3/C4.5完全复现需要自己扩展100~200 行教学、面试、改算法sklearn DecisionTreeClassifier接近但不一致不支持几十行快速出结果、工业部署自写树骨架 sklearn 线性节点部分复现支持150~250 行研究多变量决策树我一般会优先选择第三种树的骨架自己写节点的分裂器替换成可插拔对象。这样既能跑第 4.2 节的信息增益也能在第 4.5 节把分裂器换成感知机同一个 zip 包里的代码能覆盖一整章内容。2.4 示例代码讲解一个不依赖 sklearn 的最小建树循环无论 zip 包内部怎么组织核心函数最终都会收敛到这段逻辑。这是一个只支持离散属性的最小 ID3 建树代码批量包内经常以decision_tree.py的形式出现import numpy as np import pandas as pd from collections import Counter def entropy(y): 信息熵底数取 2对应书上式(4.1) counter Counter(y) total len(y) return -sum((cnt / total) * np.log2(cnt / total) for cnt in counter.values()) def info_gain(X, y, feature): 信息增益对应书上式(4.2) base entropy(y) total len(y) weighted 0.0 for value in X[feature].unique(): subset_y y[X[feature] value] weighted len(subset_y) / total * entropy(subset_y) return base - weighted def build_tree(X, y, features): 递归建树返回字典结构叶子用 leaf 标记 if len(set(y)) 1: return {leaf: y.iloc[0]} if not features: return {leaf: y.mode()[0]} best_feature max(features, keylambda f: info_gain(X, y, f)) node {feature: best_feature, branches: {}} for value in X[best_feature].unique(): sub_X X[X[best_feature] value] sub_y y[X[best_feature] value] node[branches][value] build_tree( sub_X, sub_y, [f for f in features if f ! best_feature] ) return node这段代码最容易理解entropy把书里的连加号翻译成 Counter 循环info_gain先算根节点熵再按每个特征取值把样本切开加权求和子节点熵两者相减就是增益build_tree每次挑增益最大的特征按取值递归建树。需要关注的是features列表在每层移除已用特征保证不会无限递归叶子终止条件有两个一个是样本同一类别一个是特征用完。这个版本遇到连续属性会卡死因为X[feature].unique()拿到的是一堆浮点数每一层会按浮点数精确分叉没有任何泛化能力这正是下一节连续属性二分要解决的问题。3. 把西瓜书4.5代码.zip跑起来解压、文件识别与运行参数3.1 先看 zip 包里的文件清单哪些是核心哪些是绘图辅助解压前先别急着双击。先看压缩包内部结构常见做法是解压后直接列出文件unzip 西瓜书4.5代码.zip -d watermelon_ch4 cd watermelon_ch4 find . -maxdepth 2 -type f | sort一般会看到如下几类文件。我建议只关心带tree、prune、dataset字样的绘图文件排在最后。文件作用是否核心decision_tree.pyID3/C4.5 建树和预测主逻辑核心prune.py预剪枝、后剪枝实现核心tree_plotter.pymatplotlib 画树结构辅助watermelon.csv西瓜数据集 2.016 条离散样本数据watermelon3.0.csv西瓜数据集 3.0含连续属性数据如果 find 结果里只有watermelon.csv而少了watermelon3.0.csv说明压缩包可能不完整第 4.4 节连续属性示例大概率跑不了后面第 5 章会展开说。3.2 解压与 Python 环境准备两个容易忽略的步骤解压这一步看着简单但 Windows 自带解压工具碰上带中文文件名的 zip 时偶尔会把内部路径解出乱码。我一般用命令行解压并且在虚拟环境里装依赖python -m venv .venv # Windows 用户执行 .venv\Scripts\activatemacOS/Linux 执行下面这行 source .venv/bin/activate pip install pandas numpy matplotlib依赖其实只缺这三个。里面不会用到 sklearn除非 zip 包里已经实现了多变量决策树节点。如果有requirements.txt可以直接pip install -r requirements.txt没有的话就用上面三件套。注意 Python 版本建议 3.8 以上太老的版本对 pandas 的mode()返回结果处理会有兼容问题。3.3 运行主程序命令行参数与代码内部参数怎么对上跑通全流程最稳的命令是这样python decision_tree.py --dataset watermelon.csv --criterion info_gain --pruning after --max-depth 5三个参数分别对应书上的三块内容。--criterion选info_gain就是第 4.2 节的信息增益选gain_ratio就是 C4.5 的增益率选gini就是 CART 基尼指数。--pruning控制第 4.3 节剪枝处理none不剪枝before预剪枝after后剪枝。--max-depth是树的最大深度限制过深可以有效对抗过拟合但也会让训练集准确率偏低。参数可选值对应书上内容默认建议--criterioninfo_gain/gain_ratio/gini4.2 划分选择info_gain--pruningnone/before/after4.3 剪枝处理none--max-depth正整数防止过深不传或传 5如果主程序没提供命令行解析代码里通常会有一片全局变量比如CRITERION info_gain直接打开文件改常量也行。这两种方式本质一样只是运行入口不同。3.4 打印树结构把输出和图对应起来跑完命令控制台通常会输出一棵用缩进表示的树类似这样纹理清晰 ├── 根蒂蜷缩 → 好瓜 ├── 根蒂稍蜷 │ ├── 色泽青绿 → 好瓜 │ └── 色泽乌黑 → 好瓜 └── 根蒂硬挺 → 坏瓜这是典型的递归打印结果比较好的是它能直接和书上剪枝小节的示例做人工对照。有一点必须提醒如果--pruning after输出树大概率比不剪枝时更短因为验证集里不少分支不会带来准确率提升会被替换成叶子。这属于正常现象不是代码坏了。4. 从4.5往前后各退一步预剪枝、后剪枝与连续属性的代码实现4.1 为什么划分选择之后必须接剪枝过拟合的血泪经验决策树不加限制地生长信息增益会不断偏袒“取值数目多”的属性。比如“编号”这个属性每个样本一个取值按它切分每个子节点都只有一个样本熵直接归零信息增益直接拉满。在真实任务里虽然没有“编号”但类似的高基数特征很容易钻进树里导致训练集全对、验证集全错。预剪枝的做法是在建树时对每个节点先估算“如果当前就停验证集准确率是多少”跟“继续分裂之后是多少”做比较不提升就停止。后剪枝则是树建完以后从下往上扫描把某些子树替换成叶子替换后验证集准确率不降就保留替换。我在真实项目里很少只盯准确率还要观察树的深度和叶子数。记住一个规律预剪枝快、树矮、但可能欠拟合后剪枝慢一点但保留的分支结构更自然也更贴近书上“先画整棵树再剪枝”的过程。这属于教科书里的公开结论我自己的经验是后剪枝对边界样本更友好付出的时间成本在小数据集上可忽略。4.2 在现有代码里加一个后剪枝函数把后悔药写进决策树直接在 build_tree 之后接一个递归后处理函数这是包里最常见的扩展方式def tree_accuracy(tree, X_val, y_val): 逐条样本走完整棵树统计验证集准确率 correct 0 for i in range(len(X_val)): node tree while leaf not in node: feature node[feature] value X_val.iloc[i][feature] if value not in node[branches]: break node node[branches][value] if leaf in node and node[leaf] y_val.iloc[i]: correct 1 return correct / len(X_val) def post_prune(tree, X_val, y_val): 自底向上剪枝验证集上更优就用叶子替换当前子树 if leaf in tree: return tree for value in tree[branches]: idx X_val[tree[feature]] value tree[branches][value] post_prune( tree[branches][value], X_val[idx], y_val[idx] ) leaf_label y_val.mode()[0] leaf_tree {leaf: leaf_label} if tree_accuracy(leaf_tree, X_val, y_val) tree_accuracy(tree, X_val, y_val): return leaf_tree return tree这段代码有几个细节。tree_accuracy里如果验证集出现训练时没见过的特征取值会直接 break当前样本不计入正确数这是最保守的做法。post_prune递归到每个节点后先把子节点全部处理完再比较整棵当前子树和“一个叶子”谁在验证集上更准。叶子标签不是全局多数类而是“当前节点接收到的验证样本”的多数类这样能避免剪枝后标签偏移。参数上唯一需要调的是X_val、y_val的切分比例zip 包内部如果已经切成 70/30就直接沿用它不要自己再随机分一次否则和书上的结果没法对齐。4.3 连续属性二分切分阈值选择是关键西瓜数据集 3.0 里“密度”“含糖率”是连续值建树时必须先做离散化。常见做法是排序后取相邻样本的中间值作为候选阈值然后挑信息增益最高的切分点def best_threshold(X, y, feature): 对连续属性找最佳二分阈值书中 4.4 节思路 values np.sort(X[feature].unique()) candidates [ (values[i] values[i 1]) / 2 for i in range(len(values) - 1) ] def gain_of_split(t): left y[X[feature] t] right y[X[feature] t] n len(y) if len(left) 0 or len(right) 0: return 0.0 return entropy(y) - ( len(left) / n * entropy(left) len(right) / n * entropy(right) ) return max(candidates, keygain_of_split)这里每次生成len(values) - 1个候选点复杂度等于排序加线性扫描对小样本无所谓但特征维度大、样本多了以后会很慢。更快的做法是用二分搜索或者直方图分箱来缩小候选集但 zip 包里的教学代码通常不会做这种优化。注意候选阈值选择的是相邻点均值所以阈值永远在原数据的取值区间内不会跑到边界外。这个函数替换掉 2.4 节info_gain里的离散取值循环就能让树支持连续属性。我建议把阈值也存进节点否则预测时没法知道当时切在哪个点。5. 避坑指南与常见问题排查从zip解压报错到模型过拟合5.1 解压报错“文件损坏”或提示需要密码先怀疑 zip 伪加密现象Windows 自带解压工具提示“压缩文件已损坏”或者unzip命令弹出password required但压缩包作者从没提过密码。原因这类教学压缩包常被二次打包工具处理过其中一部分在 zip 的通用位标记上置了加密标志位但数据区并没有真正加密也就是常说的 zip 伪加密。另一种情况是中文文件名编码问题Windows 用 GBK 记录文件名Linux 按 UTF-8 解析unzip 会误判文件损坏。解决别急着重新下载。先用 7-Zip 打开如果能看到文件名直接用 7-Zip 解压即可。命令行下可以这样处理7z x 西瓜书4.5代码.zip7-Zip 对伪加密兼容得更好能自动忽略无效密码位。如果解出来文件名乱码就把解压参数改成指定编码再试一次。注意如果你发现包里确实带了password.txt之类的说明优先按说明处理没有明确提示时才走这条路。5.2 画树中文乱码matplotlib 字体设置的三个位置现象树图画出来方框里的“好瓜”“坏瓜”“纹理”全是小方块或者英文正常中文消失。原因matplotlib 默认字体不包含中文字形Linux 服务端尤其常见Windows 上装了 Anaconda 也偶尔中招。解决在tree_plotter.py头部加这段import matplotlib.pyplot as plt plt.rcParams[font.sans-serif] [SimHei, WenQuanYi Micro Hei, Microsoft YaHei] plt.rcParams[axes.unicode_minus] False三个字体依次是黑体、文泉驿微米黑、微软雅黑覆盖 Windows、Linux、macOS 三个平台。如果加了还乱码检查系统里有没有中文字体Linux 可以用fc-list :langzh查没有就apt install fonts-wqy-microhei。保存图片时再加dpi150避免在论文里插图显得模糊。5.3 信息增益算出来比书上大或者全是零log 底数和编码方式现象打印每次分裂选出的信息增益值和书上的数字对不上有的偏大有的干脆全是 0。原因最常见的三个。第一代码里写的是np.log而信息熵公式要求log2自然对数和以 2 为底的数值差一个常数倍排序可能不变但绝对值全偏。第二pandas 读 CSV 时把“色泽”这类离散列自动读成了字符串如果手动转成数值比如青绿0、乌黑1、浅白2再用unique()切分会把连续型数值当成离散枚举计算出来的熵没问题但语义变了。第三数据集被 shuffle 过计算的增益自然和书上顺序结果略有出入。解决统一改用np.log2读数据时不要用LabelEncoder重新编码保持字符串类型确认代码里对特征值做的是unique()枚举切分而不是阈值切分。每次分裂把当前节点的增益值打印出来对照一遍能省很多困惑。5.4 自己的树为什么比 sklearn 差一大截数据切分方向反了现象post_prune之后模型验证集准确率反而下降甚至比对 sklearn 的默认树低了 20 个百分点。原因sklearn 的train_test_split默认做了分层抽样保证训练集和验证集里正负样本比例接近原始分布。而很多手写代码习惯直接取前 70% 做训练、后 30% 做验证。西瓜数据集里类别顺序如果刚好被排过序前 70% 全是“好瓜”验证集全是“坏瓜”后剪枝必然翻车。解决在跑剪枝前用分层抽样重新切数据from sklearn.model_selection import train_test_split X_train, X_val, y_train, y_val train_test_split( X, y, test_size0.3, random_state42, stratifyy )stratifyy是关键这一行能让训练和验证两边的类别分布一致。做完这步再跑post_prune才可能看到论文里后剪枝略优于预剪枝的效果。5.5 一运行就报 ModuleNotFoundError 或 KeyError依赖和列名编码现象import matplotlib报错或者KeyError: 好瓜。原因环境缺依赖以及 CSV 文件列名编码不对。西瓜数据集如果是从教材主页直接下载的列名一般是中文“色泽 根蒂 敲声 纹理 脐部 触感 好瓜”CSV 文件可能是 ANSI/GBK 编码Python 默认 UTF-8 读不出来。解决读文件时显式指定编码data pd.read_csv(watermelon.csv, encodinggbk)如果还报错把encodinggbk换成encodingutf-8或encodingansi三个都试一次看哪个能读通。依赖缺失就用前面 3.2 节的三件套补齐。顺手在 IDE 里开启代码诊断插件比如 PyCharm 自带的检查或 pylint能一次性扫出漏 import 和变量未定义比报错后挨个查更快。6. 进阶把西瓜书4.5代码.zip改造成自己的决策树工具箱6.1 把 ID3 的划分函数改成 C4.5增益率实现在 2.4 节info_gain基础上增加一个函数就是 C4.5 的增益率对应书上式(4.3)def gain_ratio(X, y, feature): C4.5 增益率增益除以固有值 IV base entropy(y) total len(y) weighted 0.0 split_info 0.0 for value in X[feature].unique(): subset_y y[X[feature] value] ratio len(subset_y) / total weighted ratio * entropy(subset_y) if ratio 0: split_info - ratio * np.log2(ratio) gain base - weighted if split_info 0: return 0.0 return gain / split_info注意增益率有个玄学问题某些特征的固有值很小会导致增益率虚高实际操作中常先筛出增益高于平均水平的特征再从中选增益率最大的。把info_gain换成gain_ratio就能让整棵树从 ID3 变成 C4.5。很多面试题就考这一层这段代码会告诉你为什么 sklearn 的criterionentropy依然不是严格 C4.5。6.2 和 sklearn 对照验证自己的实现我通常会把自写树和 sklearn 放在同一份验证集上做对照用分层五折交叉验证看两个指标一个是均值准确率一个是标准差from sklearn.tree import DecisionTreeClassifier from sklearn.model_selection import StratifiedKFold sk_model DecisionTreeClassifier(criterionentropy, random_state0) sk_scores cross_val_score(sk_model, X_num, y, cvStratifiedKFold(5))自写代码的 data 必须编码成数值sklearn 不支持字符串特征。如果两边分数差在 3 个百分点以内说明建树逻辑基本没有错差太多优先怀疑数据切分和特征编码而不是算法实现。这个对照步骤值得一直留着后面把 4.5 节线性节点替换进来时还能用它确认线性组合产生了实际收益。我从这套代码里养成的习惯是拿到任何压缩包先解压看文件清单再跑一次--help或直接看主程序参数最后才动数据。现在想改树模型我也会先打印出一棵未剪枝的树看它学的东西是否合理再谈调参。希望帮到你。本文还有配套的精品资源点击获取
返回列表