ARTICLE DETAIL

资讯详情

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

西瓜书4.5决策树代码实战:从信息增益率到剪枝调参

西瓜书4.5决策树代码实战:从信息增益率到剪枝调参 简介这份代码包是配合《机器学习》西瓜书第4.5节内容整理的实战示例主要面向正在学习决策树算法的初学者、考研学生及需要复现书中案例的研究者。资源共4个文件包含两个Jupyter Notebook含检查点版本、一个可直接运行的Python脚本和一个CSV格式的心脏病数据集整体压缩包仅18KB目录结构简单清晰轻量且便于下载。Notebook与Python脚本互为补充Notebook适合在Jupyter环境中交互式逐行观察决策树的构建、划分与剪枝细节脚本则适合批量运行或二次修改自带的数据集让读者不必额外寻找数据即可完成一次完整的实验闭环通过运行代码还能直观体会不同参数对模型表现的影响。目前已有436人学习代码结构清晰、注释友好可直接复制到本地环境运行帮助读者将书中抽象的归纳偏好、信息增益等概念落实到具体代码中更适合边阅读边动手调试的学习方式。1. 西瓜书4.5代码解压后先看这三个文件拿到西瓜书4.5代码.zip解压出来不是一堆PPT而是三个文件main.py、heart.csv、main.ipynb。第一次用的人常把它当成“课后题答案一键生成器”其实它是一个能直接跑的决策树分类示例。代码用heart.csv这个心血管疾病数据集演示了从数据清洗、决策树训练到可视化评估的完整流程正好对应周志华《机器学习》第四章决策树的4.5节。想理解C4.5或CART到底怎么用或者想把自己手头的CSV套进决策树拿这份代码起步比从头看理论更顺手。注意zip里那个.ipynb_checkpoints是Jupyter自动生成的目录不影响运行但会干扰你对代码结构的第一眼判断。2. 为什么这份代码值得当教材从ID3到C4.5的切换逻辑2.1 决策树划分选择的三种策略周志华老师在西瓜书里把决策树划分选择分成ID3、C4.5、CART三条路线。ID3用信息增益对取值较多的特征有天然偏好C4.5用信息增益率避免了这个偏差CART用基尼指数默认生成二叉树。这份代码里heart.csv的特征既包含连续型数值年龄、胆固醇也包含离散型类别性别、胸痛类型如果直接用ID3的信息增益会把年龄这种连续值当成离散特征处理每个不同值都切开树会迅速膨胀。所以main.py里常见的做法是让sklearn的DecisionTreeClassifier内部自动处理连续特征或者自己先做分箱。更实际的原因是heart.csv只有13个特征但年龄、胆固醇这种连续特征的可能取值有几十个。ID3算信息增益时倾向于选择取值更多的特征年龄几乎能完美切分每个样本信息增益虚高。C4.5引入信息增益率后用特征本身的信息量做分母把这个虚高拉回来。这也是为什么西瓜书的课后题要求用信息熵划分选择而不是直接用默认的gini。2.2 信息增益率计算与代码验证信息增益率的公式不复杂但只看公式很难理解它对ID3的修正作用。我一般会在Jupyter里手动算一遍配合heart.csv的某一列做验证# 手动计算信息增益率复现西瓜书4.2节公式 import numpy as np import pandas as pd def entropy(y): _, counts np.unique(y, return_countsTrue) p counts / counts.sum() return -np.sum(p * np.log2(p)) def gain_rate(feature, y): data pd.DataFrame({f: feature, y: y}) groups data.groupby(f)[y] cond_ent sum((len(g) / len(data)) * entropy(g) for _, g in groups) info_gain entropy(data[y]) - cond_ent iv entropy(data[f]) return info_gain / iv这段代码先通过entropy计算经验熵再按特征取值分组算条件熵最后用iv做分母。iv是特征自身的信息量特征可取的值越多iv越大信息增益率被压得越低。你可以把heart.csv里的age和cp分别喂进这个函数跑一遍会发现cp的信息增益率反而比age高这与医学直觉一致胸痛类型对心脏病判断更重要年龄只是辅助信号。提示如果只是想跑通流程不需要自己实现C4.5sklearn已经封装好。但如果你想通过西瓜书4.5的课后要求建议先按上面的函数手算两步再看main.py会瞬间理解代码里每个参数在干什么。2.3 heart.csv 字段说明与预处理要点回到数据集本身。heart.csv是UCI Heart Disease数据集的常见导出通常包含13个特征加1个目标列。用这份代码做实验前先建立一张字段表避免后面画决策树时看不懂节点名。字段名类型说明决策树中的处理age连续年龄直接喂给树模型sklearn自动找切分点sex离散性别0/1数值无需额外处理cp离散胸痛类型取值1~4类别值chol连续胆固醇数值型注意离群点thalach连续最大心率与年龄相关性高target离散是否患病分类标签0/1这里要特别注意main.py读CSV时如果直接用pd.read_csv(heart.csv)默认第一行是列名。换成你自己的数据时先检查有没有表头、有没有空值用data.isnull().sum()。对决策树来说空值不是致命问题sklearn的树模型支持NaN但不支持字符串类别所以cp这类离散列如果变成文本需要先用sklearn.preprocessing.LabelEncoder转成数值。把这一步提前写进main.py后面换数据就不会在训练时报错。3. main.py 拆解训练、剪枝、可视化一条龙3.1 数据加载与训练集划分main.py 的第一步通常是读取heart.csv然后用train_test_split按7:3分成训练集和测试集。这里有个容易踩的坑random_state不固定每次跑出来的准确率都不一样换了个随机种子模型结构也完全变样。所以项目里固定随机种子是必须的。# main.py 数据准备段 import pandas as pd from sklearn.model_selection import train_test_split df pd.read_csv(heart.csv) X df.drop(target, axis1) # heart.csv 的类别列多数叫 target y df[target] X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.3, random_state42, stratifyy )drop(target, axis1)把标签列剔除剩下全部作特征stratifyy保证训练集和测试集里患病/不患病的比例与原数据一致。对心脏病这种本身分布不极端的数据分层抽样影响不大但换成二分类正样本只有10%的业务数据时不加stratify很容易让测试集全是负样本。test_size0.3表示30%做测试7:3是入门常用比例样本量超过一万时可以调成0.2。3.2 决策树训练参数max_depth 和 min_samples_leaf进入核心训练段。sklearn里的DecisionTreeClassifier虽然默认走CART但参数对C4.5思路同样适用。main.py里最常见的关键参数是max_depth和min_samples_leaf。一个是限制树的深度一个是限制叶子节点的最小样本数这俩就是西瓜书4.3剪枝处理的程序化表达。# main.py 模型训练段常见写法 from sklearn.tree import DecisionTreeClassifier from sklearn.metrics import accuracy_score clf DecisionTreeClassifier( criterionentropy, # entropy 近似 C4.5 的信息增益 max_depth4, # 预剪枝限制最大深度 min_samples_leaf5, # 叶子节点至少 5 个样本 random_state0 ) clf.fit(X_train, y_train) y_pred clf.predict(X_test) print(Accuracy:, accuracy_score(y_test, y_pred))criterionentropy让树按信息熵选切分点比默认的gini更贴近西瓜书的理论推导max_depth4不是拍脑袋拍的heart.csv只有13个特征深度4到6之间模型最稳。太深会把每个病人的个体差异都记下来也就是过拟合。min_samples_leaf5强制每个叶子上至少5个样本避免某个叶子只包含一个极端样本。运行后准确率一般在0.75到0.85之间波动不要因为某些博客写0.95就怀疑自己heart数据集本身有噪声准确率不是唯一指标。3.3 剪枝策略对比预剪枝和后剪枝剪枝是西瓜书4.3的重头戏。main.py里的参数属于预剪枝在构建过程中就判断是否分裂。后剪枝则是先长出完整树再自底向上合并那些对验证集没有提升的节点。sklearn没有直接暴露后剪枝API所以代码包里通常用预剪枝参数代替。这三种策略在实践里分别对应不同的代码位置策略实现方式优点缺点预剪枝max_depth/min_samples_leaf训练快避免无用分支可能欠拟合只看局部后剪枝sklearn 最低成本复杂度路径保留更多结构精度更稳训练慢代码多不剪枝默认参数训练集拟合好测试集容易崩如果你在main.ipynb里看到ccp_alpha这个参数那就是sklearn提供的最小成本复杂度剪枝属于后剪枝。要上手的话先调出clf.cost_complexity_pruning_path(X_train, y_train)拿到不同alpha下的树信息再选验证集精度最高的alpha重新训练。这部分代码在main.py中不一定是默认开启的但能看懂就说明你已经超过单纯调包的水平。3.4 可视化画出决策树才能向别人解释训练完模型main.py一般会调用plot_tree或export_graphviz画图。画图是判断树是否合理的第一步。重点看根节点的分裂特征和分裂阈值如果第一个分裂特征是thal而不是cp说明数据分布的偶然性引导了模型这时候要怀疑是不是特征太多、样本太少。# main.py 可视化段 import matplotlib.pyplot as plt from sklearn.tree import plot_tree plt.figure(figsize(16, 8)) plot_tree( clf, feature_namesX_train.columns, class_names[No Disease, Disease], filledTrue, roundedTrue, impurityTrue ) plt.savefig(tree.png, dpi150)feature_names控制节点里显示的特征名不设的话节点里都是x[3]这种索引谁也看不懂class_names把类别从0/1转成语义化标签filledTrue会用颜色深浅表示多数类别一眼看出每个叶子偏向哪一侧。保存为PNG后放到项目说明里比贴一堆指标更直观。4. 把 heart.csv 换成自己的 CSV三个坑要绕开4.1 数据格式要求与调整主程序拿这份代码去跑自己的数据最省事的做法是只把heart.csv替换成你的文件然后改列名。但main.py里写死了target这一列。假设你的分类列叫label同时前两列是id和date不能直接当特征训练。修改方式很简单# main.py 自定义数据加载 df pd.read_csv(your_data.csv) drop_cols [id, date, label] # 不参与训练的列 X df.drop(columnsdrop_cols) y df[label]drop_cols这个列表很关键。id和时间列对分类任务没有泛化能力决策树还会试图拿它们做切分点产生一条“日期2023-03-15则患病”这种无意义路径。去掉之后再检查X.dtypes确认全是数值类型。如果X里混入object列sklearn会直接抛ValueError此时要么pd.get_dummies做one-hot要么用LabelEncoder。一般情况下离散特征取值少于5个时直接保留数值编码就行不需要one-hot因为决策树能处理无序类别。4.2 连续特征截断点分箱与直接切分的边界heart.csv里的年龄、胆固醇、最大心率都是连续值。sklearn的决策树会自动搜索最优切分点不需要你提前分箱。但西瓜书4.4讲连续属性处理时反复强调一个细节切分点不是特征取值本身而是相邻取值的中点。比如年龄19和20之间切分阈值是19.5。这份代码里如果出现age 52.5这种分裂条件就是这么算出来的。自己处理连续特征时可选的方案有两个方案做法适用场景让树模型自动切分不预处理直接训练特征分布正常样本量足手动分箱pd.cut 分成3~5个区间特征有长尾/离群点常见做法是先把age用pd.cut(age, bins5)离散化然后丢进决策树。这会让树失去寻找更细切分点的能力但能显著减少过拟合。heart.csv里胆固醇的分布偏右直接用原始值树很可能只在大于或小于一个极高阈值处分裂手动分箱后会更稳定。如果你试了两种方式对比测试集准确率分箱不会比自动切分差太多但可解释性更强。4.3 缺失值处理fillna 放哪里要小心这份代码里heart.csv是比较干净的数据但你自己数据集大概率有缺失值。决策树本身能处理NaNsklearn的实现里DecisionTreeClassifier的splitter会自动把NaN分到较优的一侧。因此你可以直接把带NaN的DataFrame喂进去不需要先fillna。这个行为很多教程没提到导致很多人习惯性先df.fillna(0)结果把缺失值变成了一个数值特征。但有一个例外如果你在4.2中手动分箱pd.cut遇到NaN会直接报错此时必须先做按类别分组填充df[age] df.groupby(target)[age].transform( lambda s: s.fillna(s.median()) )用每个类别自己的中位数填充而不是全量填一个值。按target分组很重要否则会破坏两类样本在年龄上的分布差异。替换文件后用df.isnull().sum().sum()检查剩余缺失量如果仍不为0再决定是删除还是填充特定列。4.4 类别不平衡准确率骗人的时候看什么heart.csv里患病样本大概占45%到55%比较均衡。如果换成真实业务数据正样本可能只有5%此时accuracy_score会失真。比如全部预测为负类准确率也能到95%。main.py里只打印了Accuracy一个指标换数据后建议扩成多指标# 多指标评估替代单一的 accuracy from sklearn.metrics import precision_score, recall_score, f1_score print(Precision:, precision_score(y_test, y_pred)) print(Recall:, recall_score(y_test, y_pred)) print(F1:, f1_score(y_test, y_pred))recall对医疗场景最重要它衡量真正生病的病人被找出多少。f1是precision和recall的调和平均类别不平衡下比accuracy靠谱。如果你发现F1只有0.3而accuracy是0.9说明模型把所有样本都推给了多数类这时候要调整class_weightbalanced参数让少数类在分裂时获得更高权重。在DecisionTreeClassifier里加上这个参数再训练F1通常会往上涨但整体accuracy会下降这是合理的trade-off。5. 用 main.ipynb 做增量调试在分裂点看数据长什么样5.1 在Jupyter里单步看每层分裂main.ipynb 和 main.py 是同一套代码但notebook的优势是能把中间过程留下来。调试时可以在clf.fit(X_train, y_train)之前插入一个%debug或者在plot_tree之前写clf.tree_.feature直接输出每个节点的分裂特征索引。更实用的技巧是先把max_depth改成2重新训练然后用clf.apply(X_train)拿到每个样本落在哪个叶子节点再按节点分组打印样本均值。你会发现同一个叶子里的样本在原始特征上确实有相似性比如都满足“年龄大于60且最大心率小于120”这就验证了树学习到了有意义的模式。5.2 判断这颗树学废了的快速方法有一个很土但有效的方法分别看训练集和测试集的准确率差值。在main.py末尾加上train_acc accuracy_score(y_train, clf.predict(X_train)) test_acc accuracy_score(y_test, y_pred) print(fTrain: {train_acc:.3f}, Test: {test_acc:.3f})当两者差值超过0.15说明过拟合树把训练集的特征细节当成普适规律当两者都低于0.6说明特征本身和标签关系太弱或者预处理出了问题。记住不是所有数据都适合决策树heart.csv能跑出0.8你自己的表可能永远只有0.55这不一定是代码的错可能是特征集本身信息量不够。5.3 把树结构和规则导出来做解释最后分享一个导出规则的技巧用export_text把决策树转成纯文本规则不需要装graphviz也能看到完整分裂逻辑。在main.py里加from sklearn.tree import export_text print(export_text(clf, feature_nameslist(X_train.columns)))输出里每一行就是一条if-else路径。把这些规则按照根到叶子整理出来就是一份不需要写代码的解释文档。如果你要给非技术同事讲模型把这几十行规则贴到文档里远比一张散点图来得清楚。注意max_depth4时规则约16条深度增到10后规则会指数增长所以导出前先控制树的深度。配合pd.Series(clf.feature_importances_, indexX_train.columns).sort_values(ascendingFalse)看特征重要性排名如果某个无关特征的位置太靠前就得回头检查预处理。所以下次遇到新数据先跑一下export_text看看根节点分裂得合不合理再决定要不要继续调参或换模型。本文还有配套的精品资源点击获取
返回列表