简介:这是一份对应周志华《机器学习》(西瓜书)第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.ipynb
2.1 4.5 节到底在讲什么:预剪枝和后剪枝
决策树生成过程中最大的问题是过拟合:分支越多,越容易把训练集里的噪声也学进去。西瓜书 4.5 节给出的两个解法是预剪枝和后剪枝,它们的共同点都是靠验证集精度来决定要不要“砍掉”某个分支。
预剪枝的做法是:每次准备划分一个节点前,先用验证集算一次精度。如果划分后验证集精度比划分前低,就放弃这次划分,直接把当前节点变成叶节点。它的优点是训练时间短,边建树边剪枝;缺点是只看当前一步,可能丢掉后面几步带来的收益,最后容易欠拟合。
后剪枝正好反过来:先把决策树完整地生成出来,不做任何限制,然后自底向上逐个考察内部节点,尝试把这个节点替换成叶节点,再看验证集精度有没有提升。有提升就剪,没有就保留。它的优点是“事后改正”,所以保留的信息更多,精度一般也更高;缺点是先要建一棵完整树再慢慢修,训练时间更长。
这本书配套代码里,预剪枝和后剪枝通常都会落到一个独立的判断函数上。预剪枝对应“划分前 vs 划分后”,后剪枝对应“替换节点前 vs 替换节点后”。你后面读 main.py 时,只要盯住这两个对比点,就能看懂它到底实现了哪一种剪枝。
| 对比项 | 预剪枝 | 后剪枝 |
|---|---|---|
| 判断时机 | 节点划分之前 | 完整树生成之后 |
| 依据 | 划分前后验证集精度 | 替换前后验证集精度 |
| 计算开销 | 小 | 大 |
| 风险 | 欠拟合 | 相对更小 |
| 西瓜书结论 | 速度更快但可能欠拟合 | 效果好但时间更长 |
有了这个底子,再去看具体文件就不容易晕。
2.2 main.py 与 main.ipynb 的分工
把压缩包解压后会看到这几个文件,我一般先按文件清单过一次,防止自己漏看什么:
| 文件 | 用途 | 关注重点 |
|---|---|---|
| main.py | 命令行直接运行的完整脚本 | load_data 和剪枝判断函数 |
| heart.csv | 心脏病分类数据集 | target 列分布 |
| main.ipynb | Jupyter 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_size=0.3, random_state=42) X_train, X_val, y_train, y_val = train_test_split( X_train, y_train, test_size=0.2, random_state=42) return X_train, X_val, X_test, y_train, y_val, y_test这段代码里random_state=42是为了保证可复现,换台机器跑结果也一样;去掉它的话,每次运行训练集划分都会变,剪枝对比就会失去参照。test_size=0.3表示把 30% 的数据先冻结起来做最终评估,后续剪枝决策不碰这一部分,否则结果会虚高。
再往下通常是决策树训练和不剪枝的基准结果:
from sklearn.tree import DecisionTreeClassifier clf = DecisionTreeClassifier(criterion='entropy', random_state=42) 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))这里的criterion='entropy'对应西瓜书里的信息增益计算方式,而不是默认的基尼指数。很多从 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.py,jupyter 可以不装,但你既然下载了这个资源,大概率会用到 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; df=pd.read_csv('heart.csv'); print(df.shape); print(df.head())"常见 heart.csv 是 UCI Heart Disease 数据集整理成的 CSV,大约 300 行、14 列,最后一列叫target,0 表示没有心脏病,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最好加上stratify=y,否则剪枝对比没有说服力:
X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.3, random_state=42, stratify=y )stratify=y的意思是让训练集和测试集里的正负比例保持和原始数据一致。这个参数在样本量小的时候尤其重要,因为它直接影响后面剪枝判断用的验证集是否可靠。
3.3 把 heart.csv 换成你自己的 CSV
这份代码不一定非得跑心脏数据,你完全可以把数据集换成自己课程设计的数据。但有三处必须改,否则运行结果不可信。
第一是分隔符。heart.csv 是逗号分隔,但很多中文数据集可能用分号或制表符。读取时改成:
df = pd.read_csv('your_data.csv', sep=';', encoding='utf-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(include='object').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 3,pip 可能装到了旧环境。用
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 MS,Linux 下装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, filled=True, feature_names=df.columns[:-1], class_names=['0', '1']) plt.savefig('tree_no_prune.png', dpi=150, bbox_inches='tight') plt.show()filled=True会按类别占比给节点上色,分类更纯的节点颜色更深。feature_names必须和训练时的列顺序一致,顺序错了整张图的含义就反了。bbox_inches='tight'是防止树太大时图片被截断,这个参数每次都值得写上。
如果你想再看剪枝后的树,我给一个快速做法:用 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_size=0.2, random_state=42) 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( criterion='entropy', random_state=42, ccp_alpha=alpha) 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( criterion='entropy', random_state=42, ccp_alpha=best_alpha) best_model.fit(X_train, y_train) plt.figure(figsize=(12, 8)) plot_tree(best_model, filled=True, feature_names=df.columns[:-1], class_names=['0', '1']) plt.savefig('tree_pruned.png', dpi=150, bbox_inches='tight') plt.show()cost_complexity_pruning_path会算出一系列候选 alpha,从 0 开始逐渐增大。对每个 alpha 重新建树,再用验证集打分,挑出最好的那个。ccp_alpha越大,剪掉的节点越多,树越矮。
我从那以后养成一个习惯:拿到这类 zip 代码包,先跑一遍 main.py 确认能复现结果,再打开 main.ipynb 对照变量名,最后才把数据集换成自己手上的。换数据前先用df.info()和value_counts()排查一遍,永远比直接改模型参数省时间。希望这份拆包记录能帮到你,少踩几个我已经不会再踩的坑。
本文还有配套的精品资源,点击获取