决策树剪枝实战:西瓜书4.5节代码资源拆解与避坑指南
简介这是一份对应周志华《机器学习》西瓜书第4.5节内容的代码与数据包主要面向正在学习决策树剪枝处理、对照课本公式做实验的读者。包内共4个文件包含两个Jupyter Notebook脚本、一个Python脚本和一个csv格式的数据集其中ipynb文件便于分步查看计算过程并可视化决策树py文件适合直接运行复现结果csv数据则用于导入样本进行训练与验证整体压缩包仅18KB轻量易用。已有436人学习下载适合配合教材逐行阅读、动手调试也可作为理解剪枝前后泛化性能对比的入门示例。通过运行这些代码读者可快速完成从数据加载、树构建到剪枝评估的完整流程直观体会西瓜书4.5节中预剪枝与后剪枝的差异并能将代码迁移到自己的数据集上做进一步试验。1. 这份压缩包的定位很明确西瓜书4.5代码.zip 是为周志华《机器学习》第四章 4.5 节“剪枝处理”准备的代码资源我一开始也以为是某个大项目的一部分解压后才发现里面只有 main.py、heart.csv 和一个 main.ipynb外加一个 .ipynb_checkpoints 目录。文件不多但每个都有明确用途main.py 是可运行的决策树剪枝脚本heart.csv 是拿来演示的小规模心脏病分类数据集main.ipynb 是同一套逻辑的 Notebook 版方便你一格一格看中间过程。这份资源解决的是很多人读西瓜书第四章时最难受的地方你知道信息增益怎么算也背得出“预剪枝”“后剪枝”的定义但一落到代码里就不知道怎么组织数据、怎么对比剪枝前后的效果。它适合正在啃西瓜书、需要完成课后作业或课程设计的学生也适合想拿手写剪枝逻辑和 sklearn 默认行为对照一遍的从业者。下面我按实际拆包顺序把文件结构、运行方式、参数含义和踩坑点讲清楚。2. 剪枝逻辑与代码包结构先搞懂 4.5 节再拆 main.py 和 main.ipynb2.1 4.5 节到底在讲什么预剪枝和后剪枝决策树生成过程中最大的问题是过拟合分支越多越容易把训练集里的噪声也学进去。西瓜书 4.5 节给出的两个解法是预剪枝和后剪枝它们的共同点都是靠验证集精度来决定要不要“砍掉”某个分支。预剪枝的做法是每次准备划分一个节点前先用验证集算一次精度。如果划分后验证集精度比划分前低就放弃这次划分直接把当前节点变成叶节点。它的优点是训练时间短边建树边剪枝缺点是只看当前一步可能丢掉后面几步带来的收益最后容易欠拟合。后剪枝正好反过来先把决策树完整地生成出来不做任何限制然后自底向上逐个考察内部节点尝试把这个节点替换成叶节点再看验证集精度有没有提升。有提升就剪没有就保留。它的优点是“事后改正”所以保留的信息更多精度一般也更高缺点是先要建一棵完整树再慢慢修训练时间更长。这本书配套代码里预剪枝和后剪枝通常都会落到一个独立的判断函数上。预剪枝对应“划分前 vs 划分后”后剪枝对应“替换节点前 vs 替换节点后”。你后面读 main.py 时只要盯住这两个对比点就能看懂它到底实现了哪一种剪枝。对比项预剪枝后剪枝判断时机节点划分之前完整树生成之后依据划分前后验证集精度替换前后验证集精度计算开销小大风险欠拟合相对更小西瓜书结论速度更快但可能欠拟合效果好但时间更长有了这个底子再去看具体文件就不容易晕。2.2 main.py 与 main.ipynb 的分工把压缩包解压后会看到这几个文件我一般先按文件清单过一次防止自己漏看什么文件用途关注重点main.py命令行直接运行的完整脚本load_data 和剪枝判断函数heart.csv心脏病分类数据集target 列分布main.ipynbJupyter Notebook 交互版逐步输出和画图.ipynb_checkpoints/Jupyter 自动保存的备份目录可以忽略不参与运行main.py 是代码入口直接python main.py就能跑。main.ipynb 是同一套逻辑的 Notebook 版适合在 Jupyter 里一格一格看中间结果。很多第一次接触的同学会把两个文件当成两个版本其实内容基本一致二选一跑就行。一般 main.py 的结构会是这样这也是我拆解类似西瓜书配套代码时常看到的分层# main.py 典型骨架与西瓜书 4.5 节对应 import pandas as pd from sklearn.model_selection import train_test_split def load_data(): # heart.csv 是 UCI Heart Disease 整理后的常见版本 df pd.read_csv(heart.csv) X df.drop(columns[target]) y df[target] # 先留出 30% 做测试集再从训练集里切 20% 作为剪枝用的验证集 X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.3, random_state42) X_train, X_val, y_train, y_val train_test_split( X_train, y_train, test_size0.2, random_state42) return X_train, X_val, X_test, y_train, y_val, y_test这段代码里random_state42是为了保证可复现换台机器跑结果也一样去掉它的话每次运行训练集划分都会变剪枝对比就会失去参照。test_size0.3表示把 30% 的数据先冻结起来做最终评估后续剪枝决策不碰这一部分否则结果会虚高。再往下通常是决策树训练和不剪枝的基准结果from sklearn.tree import DecisionTreeClassifier clf DecisionTreeClassifier(criterionentropy, random_state42) clf.fit(X_train, y_train) train_acc clf.score(X_train, y_train) test_acc clf.score(X_test, y_test) print(不剪枝训练集 %.4f测试集 %.4f % (train_acc, test_acc))这里的criterionentropy对应西瓜书里的信息增益计算方式而不是默认的基尼指数。很多从 sklearn 入门的人习惯用gini但看西瓜书配套代码时要先确认用的是哪个指标否则以后改了数据同一棵树的划分点会完全对不上。main.ipynb 里通常会在差不多同样的位置把每次划分前后的精度打印出来方便你对照书里那个“划分后精度反而下降”的例子。这个对比写得好不好直接决定你最后看到的是“剪枝更差”还是“剪枝更好”。这也是下一章运行时要重点盯的地方。3. 在本地跑通这份代码从解压到换用自己的数据3.1 环境准备四个依赖包和一条命令先确认机器上有 Python 3.8 以上版本然后安装依赖pip install pandas numpy scikit-learn matplotlib jupyterpandas 管 CSV 读取numpy 管数组计算scikit-learn 提供决策树和数据集划分函数matplotlib 负责画图jupyter 用来打开 main.ipynb。如果只想跑 main.pyjupyter 可以不装但你既然下载了这个资源大概率会用到 Notebook所以建议全部装齐。解压时别直接双击拖到桌面我建议在命令行里解压路径可控cd ~/Downloads unzip 西瓜书4.5代码.zip -d ~/xigua45 cd ~/xigua45 ls -l执行ls -l后能看到 main.py、heart.csv 和 main.ipynb 就已经正常。如果解压后发现 main.py 被套了一层子目录用find . -name main.py找到实际位置再把它和 heart.csv 放到同一个目录下。很多FileNotFoundError: heart.csv的报错本质是脚本和 CSV 不在同一目录不是代码本身有问题。文件布局确认后直接跑python main.py如果正常终端会输出剪枝前后的精度有的整合版本还会生成一张 PNG 图。如果没有任何输出大概率是作者把结果全部放在了 Notebook 里此时打开 main.ipynb 逐个执行单元格即可。3.2 heart.csv 到底长什么样不了解数据就调参等于盲人摸象。先用一条命令看前五行python -c import pandas as pd; dfpd.read_csv(heart.csv); print(df.shape); print(df.head())常见 heart.csv 是 UCI Heart Disease 数据集整理成的 CSV大约 300 行、14 列最后一列叫target0 表示没有心脏病1 表示有心脏病。前面 13 列包括年龄、性别、胸痛类型、静息血压、胆固醇、空腹血糖、心电图结果、最大心率、运动诱发心绞痛、ST 段压低、ST 段斜率、血管造影数量和地中海贫血特征。不同渠道拿到的版本列名会有些差异但target这个列名出现概率最高。拿到数据后先看类别分布import pandas as pd df pd.read_csv(heart.csv) print(df[target].value_counts())如果正负样本差异很大比如 250:50决策树会偏向多数类。这时候原来的train_test_split最好加上stratifyy否则剪枝对比没有说服力X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.3, random_state42, stratifyy )stratifyy的意思是让训练集和测试集里的正负比例保持和原始数据一致。这个参数在样本量小的时候尤其重要因为它直接影响后面剪枝判断用的验证集是否可靠。3.3 把 heart.csv 换成你自己的 CSV这份代码不一定非得跑心脏数据你完全可以把数据集换成自己课程设计的数据。但有三处必须改否则运行结果不可信。第一是分隔符。heart.csv 是逗号分隔但很多中文数据集可能用分号或制表符。读取时改成df pd.read_csv(your_data.csv, sep;, encodingutf-8)第二是目标列。假设你的目标列叫label就把df.drop(columns[target])改成df.drop(columns[label])y df[target]同步改成y df[label]。这里最容易出错的是只改了 x 没改 y代码会先报 KeyError很好发现真正危险的是目标列名恰好也叫target但语义完全不是心脏病标签这种情况要格外确认。第三是缺失值和文本列。heart.csv 是整理好的自己的数据基本都有缺失。建议在训练前先跑一个检查print(df.isnull().sum()) # 看每列缺失数量 df df.dropna() # 缺失不多时直接删行如果你的数据里全是“是/否”“高/中/低”这类文本sklearn 决策树不能直接处理需要先编码from sklearn.preprocessing import LabelEncoder for col in df.select_dtypes(includeobject).columns: df[col] LabelEncoder().fit_transform(df[col])这段代码把每一列文本转成 0、1、2 这样的数字。需要注意的是LabelEncoder对取值顺序不敏感比如“高/中/低”会被编码成 0/1/2但字符排序可能变成“低0、中1、高2”也可能是别的顺序。严格来说有序类别应该用OrdinalEncoder无序类别用OneHotEncoder但对演示剪枝逻辑来说LabelEncoder已经足够先跑通再优化。4. 避坑指南代码包常见的五个翻车点与排查方法4.1 解压提示损坏或需要密码现象用 WinRAR 或 7-Zip 解压时弹出“文件损坏”或“需要密码”但项目描述里没有提到密码。原因网上下载的 zip 经常被转载站套壳或者下载不完整。还有一种情况是 zip 伪加密压缩包的文件头被标记成加密实际数据没有加密解压工具误以为需要密码。解决先看文件大小如果只有几十 KB但有三个文件多半是下载不完整重新下。如果大小正常用 Python 打开看结构import zipfile with zipfile.ZipFile(西瓜书4.5代码.zip) as z: print(z.namelist())只要namelist()能列出 main.py 和 heart.csv说明 zip 文件本身完整换个解压工具或者去掉伪加密标志位就能解压。这份资源本身没有密码如果你下的版本要密码那是转载站加的壳不是作者加的。4.2 直接跑 main.py 报 ModuleNotFoundError现象ModuleNotFoundError: No module named sklearn或 pandas 找不到。原因当前 Python 环境是干净的缺第三方库这不是代码问题。解决执行第 3.1 节的安装命令。这里有个细节如果电脑里同时装了 Python 2 和 Python 3pip 可能装到了旧环境。用python -m pip install pandas numpy scikit-learn matplotlib代替pip install确保装到当前python解释器对应的环境。更稳妥的做法是建一个虚拟环境再装避免把系统 Python 弄乱。4.3 main.ipynb 打开是空的或和 main.py 不一致现象Jupyter 里打开 main.ipynb 只看到标题没有代码或者代码和 main.py 完全不对应。原因压缩包里同时存在main.ipynb和.ipynb_checkpoints/main-checkpoint.ipynb后者是 Jupyter 自动保存的备份文件不是双份代码。部分解压软件或中转站点会覆盖原文件导致 Notebook 损坏。解决以 main.py 为准。main.ipynb 只是交互版不值得花太多时间修复。如果两个文件都打不开可以直接把 main.py 的内容复制到新的 Notebook 里自己加断点效果一样。4.4 剪枝前后测试集精度几乎不变现象输出显示不剪枝、预剪枝、后剪枝的精度都差不多比如 0.82、0.81、0.83看起来剪枝没有意义。原因heart.csv 只有 300 行左右如果验证集只有几十条样本一次划分产生的精度波动很容易被几个样本的误判掩盖。另一个更常见的原因是代码拿测试集去做剪枝决策而不是验证集这属于实现错误。解决把test_size调大一点比如测试集 0.4验证集 0.3再看趋势。更关键的是确认代码里是否真的存在独立的验证集打印X_train.shape、X_val.shape、X_test.shape三个集合都必须有数据。如果代码里根本没有验证集说明这份实现只展示了“不剪枝 vs 剪枝后的最终结果”剪枝决策本身没有参与把它当结果对比看就行不用强行理解成完整的西瓜书 4.5 实现。4.5 画图时中文乱码现象图上的中文标签变成方块或者 matplotlib 报RuntimeWarning: Glyph missing。原因matplotlib 默认字体不包含中文字符。解决在代码开头加三行import matplotlib.pyplot as plt plt.rcParams[font.sans-serif] [SimHei, Arial Unicode MS, DejaVu Sans] plt.rcParams[axes.unicode_minus] FalseSimHei是 Windows 黑体macOS 下建议改成Arial Unicode MSLinux 下装fonts-wqy-microhei。第二行的unicode_minus是为了解决坐标轴负号显示成方块的问题这个和中文乱码是两回事但经常一起出现。5. 把剪枝结果画出来用 sklearn 对照西瓜书 4.5 的结论最后分享一个我每次拿到这类决策树代码包都会做的事把树画出来再用 sklearn 的剪枝结果和西瓜书里的结论对一遍。只看精度数字很难看出剪枝到底砍掉了哪个节点画图一眼就明白了。假设 main.py 里已经训练好一棵不剪枝的决策树clf直接追加这段代码from sklearn.tree import plot_tree import matplotlib.pyplot as plt plt.figure(figsize(16, 10)) plot_tree(clf, filledTrue, feature_namesdf.columns[:-1], class_names[0, 1]) plt.savefig(tree_no_prune.png, dpi150, bbox_inchestight) plt.show()filledTrue会按类别占比给节点上色分类更纯的节点颜色更深。feature_names必须和训练时的列顺序一致顺序错了整张图的含义就反了。bbox_inchestight是防止树太大时图片被截断这个参数每次都值得写上。如果你想再看剪枝后的树我给一个快速做法用 sklearn 的成本复杂度剪枝代替手写后剪枝找验证集精度最高的那棵树。这个参数和西瓜书后剪枝的“是否替换成叶节点”不完全等价但思路相近更适合新手对照。from sklearn.tree import DecisionTreeClassifier from sklearn.model_selection import train_test_split # 如果 main.py 里没有单独的验证集可以从训练集再切一次 X_train, X_val, y_train, y_val train_test_split( X_train, y_train, test_size0.2, random_state42) path clf.cost_complexity_pruning_path(X_train, y_train) ccp_alphas path.ccp_alphas best_alpha None best_score 0 for alpha in ccp_alphas: model DecisionTreeClassifier( criterionentropy, random_state42, ccp_alphaalpha) model.fit(X_train, y_train) score model.score(X_val, y_val) if score best_score: best_score score best_alpha alpha best_model DecisionTreeClassifier( criterionentropy, random_state42, ccp_alphabest_alpha) best_model.fit(X_train, y_train) plt.figure(figsize(12, 8)) plot_tree(best_model, filledTrue, feature_namesdf.columns[:-1], class_names[0, 1]) plt.savefig(tree_pruned.png, dpi150, bbox_inchestight) plt.show()cost_complexity_pruning_path会算出一系列候选 alpha从 0 开始逐渐增大。对每个 alpha 重新建树再用验证集打分挑出最好的那个。ccp_alpha越大剪掉的节点越多树越矮。我从那以后养成一个习惯拿到这类 zip 代码包先跑一遍 main.py 确认能复现结果再打开 main.ipynb 对照变量名最后才把数据集换成自己手上的。换数据前先用df.info()和value_counts()排查一遍永远比直接改模型参数省时间。希望这份拆包记录能帮到你少踩几个我已经不会再踩的坑。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联
返回资讯列表 →