多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

决策树三巨头实战复现:ID3、C4.5、CART代码拆解笔记

决策树三巨头实战复现:ID3、C4.5、CART代码拆解笔记 简介这是一份面向Python初学者的决策树三种经典算法实现代码包覆盖ID3、C4.5与CART的完整实现与调用示例适合正在学习数据挖掘、机器学习分类模型或需要参照经典算法源码的开发者。压缩包共9个文件其中6个Python脚本对应三个算法的核心逻辑与绘图辅助工具2个pyc为编译缓存1个iris.csv提供实验数据集整体仅14KB轻量易读便于快速运行调试。三套算法各有侧重ID3基于信息增益选择特征C4.5引入信息增益比并支持连续属性与缺失值处理CART采用基尼不纯度并可同时处理分类与回归任务。通过对照阅读源码与运行结果读者能够直观理解不同分裂标准、剪枝策略与树结构差异掌握数据导入、模型构建、训练预测的基本流程。目前已有389人学习下载适合作为课堂实验、课程设计或自学的参考实现。1. 决策树三种经典算法实现代码包拆解与复现笔记做机器学习的人迟早会遇到决策树而决策树绕不开的坎就是 ID3、C4.5、CART 这三个经典算法。理论看过无数遍信息熵、信息增益、基尼系数背得滚瓜烂熟但真到自己动手写的时候很多人还是懵。这份「决策树三种经典算法实现.rar」压缩包正好补上这个缺口里面是纯 Python 手写实现的三种算法源码配合 iris.csv 数据集和一个树形可视化脚本从建树到画图一条龙。我把它下载下来拆了一遍代码结构清楚注释也到位既能当算法课的课后作业参考也能拿来对比理解 sklearn 里 DecisionTreeClassifier 背后的计算逻辑。这篇笔记就从文件结构、算法差异、运行方式到常见坑完整过一遍新手照着跑能出图熟手能直接改代码做实验。2. 压缩包内部结构与代码数据流先搞清楚每个文件干什么拿到压缩包别急着跑代码先把文件清单过一遍。这种手写算法实现的包文件之间的调用关系往往比算法本身更容易让人栽跟头。2.1 文件清单与角色定位解压之后你会看到下面这些文件我先按角色给你分个类文件角色定位说明id3.pyID3 算法主实现基于信息增益多叉树划分c45.pyC4.5 算法主实现基于信息增益比支持连续属性cart.pyCART 分类树实现基于基尼不纯度二叉树结构cart2.pyCART 回归树实现基于平方误差最小化做回归用的treePlotter.py树结构可视化调用 matplotlib 绘制决策树图形iris.csv标准测试数据150 条鸢尾花样本4 个特征 3 个类别__pycache__/Python 缓存目录不用管跑过代码后自动生成注意cart2.py和cart.py的区别前者解决回归问题后者解决分类问题。这个压缩包表面上是三种算法实际上是 2 个分类算法加一个扩展变体CART 分类和 CART 回归分开写了这对你理解同一算法在不同任务下的分裂标准差异非常有帮助。2.2 代码入口与多文件协作方式我读这几个文件的时候发现它们不是统一通过一个入口文件调用的而是每个算法文件都自带if __name__ __main__测试块。也就是说你要跑哪个算法就单独执行哪个文件数据加载逻辑在各自文件里重写了一遍没有做公共模块抽取。这种结构对学习是友好的每个文件都能独立运行但对工程化是反面教材你读代码时要有这个判断。treePlotter.py是公共可视化模块先看它的核心函数签名因为这个被三个算法文件共同引用# treePlotter.py 核心接口 def createPlot(inTree): 输入决策树字典结构嵌套dict 说明递归遍历字典生成带箭头和节点的matplotlib图形 fig plt.figure(1, facecolorwhite) fig.clf() axprops dict(xticks[], yticks[]) createPlot.ax1 plt.subplot(111, frameonFalse, **axprops) plotTree.totalW float(getNumLeafs(inTree)) # 叶子数决定横向宽度 plotTree.totalD float(getTreeDepth(inTree)) # 深度决定纵向高度 plotTree(inTree, (0.5, 1.0), ) # 根节点在顶部中心 plt.show()这段代码最常见的用法是配合json或直接传 Python 字典树结构用嵌套字典表达键是特征名:切分点值是一个子字典或最终的类别标签。getNumLeafs和getTreeDepth分别递归计算叶子数和树深度用来确定画布尺寸。如果你自己写新的树构建逻辑只要最终输出也是这个嵌套字典格式createPlot就能直接复用不需要改任何东西。2.3 iris.csv 与三个算法文件的数据接口约定iris.csv 是 Fisher 的经典鸢尾花数据集150 行数据每行 5 个字段花萼长度、花萼宽度、花瓣长度、花瓣宽度、类别。三个算法文件读取这份数据的逻辑基本一致差异主要在特征处理细节上。ID3 版本通常会做离散化预处理C4.5 版本会实时排序寻找最优切分点CART 两个版本都是直接在原始连续值上选阈值。调用核心逻辑分布id3.py加载数据 → 特征离散化 → 计算信息增益 → 递归建树 → 输出字典树 → 调 treePlotter 画图c45.py加载数据 → 连续属性和离散属性分别处理 → 计算信息增益比 → 递归建树 → 输出字典树 → 画图cart.py加载数据 → 遍历所有特征和所有切分点 → 计算基尼不纯度 → 递归建二叉 → 画图# 三个算法文件的独立运行方式 python id3.py python c45.py python cart.py每个文件运行后都会在终端打印树结构文本同时弹出 matplotlib 窗口显示可视化树形。终端文本和图像是两个输出通道验证结果时记得两个都要看图像可能因为节点重叠看不清深层细节文本才是准确的。3. 三种算法的分裂标准与实现细节从公式到代码逐行对照这一章是核心中的核心。三种算法最本质的区别就是它们回答同一个问题的答案不同如何从当前数据集中选出一个最优特征进行 split。ID3 用信息增益C4.5 用信息增益比CART 用基尼不纯度。代码里的差别就在这几十行。3.1 ID3 的信息熵与信息增益实现ID3 的分裂标准是选择信息增益最大的特征。先看信息熵的计算这一块几乎所有实现都长一个样# id3.py 信息熵计算 def calcShannonEnt(dataSet): numEntries len(dataSet) labelCounts {} for featVec in dataSet: currentLabel featVec[-1] # 最后一列是标签 labelCounts[currentLabel] labelCounts.get(currentLabel, 0) 1 shannonEnt 0.0 for key in labelCounts: prob float(labelCounts[key]) / numEntries shannonEnt - prob * math.log(prob, 2) # 注意log底数是2 return shannonEnt这段代码实现了Ent(D) -Σ p_k * log2(p_k)其中p_k是类别 k 的样本占比。log(prob, 2)里的底数 2 对应信息论里比特的单位定义。这段代码有个小问题当prob为 0 时math.log会报错不过实际运行时数据集里不会出现占比为 0 的类别所以作者没做保护处理。你要是把代码改到其他数据集上建议加上if prob 0的防御判断。信息增益的计算就是分组加权后的熵差# id3.py 按特征划分后计算信息增益 def chooseBestFeatureToSplit(dataSet): baseEntropy calcShannonEnt(dataSet) # 划分前的熵 bestInfoGain 0.0 bestFeature -1 numFeatures len(dataSet[0]) - 1 for i in range(numFeatures): featList [example[i] for example in dataSet] uniqueVals set(featList) newEntropy 0.0 for value in uniqueVals: subDataSet splitDataSet(dataSet, i, value) prob len(subDataSet) / float(len(dataSet)) newEntropy prob * calcShannonEnt(subDataSet) infoGain baseEntropy - newEntropy if infoGain bestInfoGain: bestInfoGain infoGain bestFeature i return bestFeature # 返回特征下标这段代码的newEntropy就是条件熵Σ |Dv|/|D| * Ent(Dv)信息增益是两者之差。ID3 的毛病也在代码里暴露出来了uniqueVals的数量越多newEntropy倾向于越小信息增益倾向于越大所以 ID3 天然偏好取值数多的特征比如编号这种每行一个值的特征会被优先选中。这就是过拟合的一个来源。3.2 C4.5 的信息增益比与连续属性切分C4.5 针对 ID3 的两个痛点做改进一是用增益比替代增益二是在连续属性上动态找切分点。先看增益比实现# c45.py 信息增益比计算 def calcGainRatio(feature, dataSet): infoGain calcInfoGain(feature, dataSet) # 普通信息增益 featValues [example[feature] for example in dataSet] uniqueVals set(featValues) splitInfo 0.0 for value in uniqueVals: prob len(splitDataSet(dataSet, feature, value)) / float(len(dataSet)) splitInfo - prob * math.log(prob, 2) if splitInfo 0.0: return 0.0 # 防止除零 return infoGain / splitInfo # 增益除以固有值增益比在信息增益基础上除以一个固有值splitInfo固有值越大说明特征取值越分散这个惩罚项抵消了 ID3 对多取值特征的偏好。代码里splitInfo 0.0的检查不是多余的当所有样本在某个特征上取值一致时会出现除零。连续属性处理是 C4.5 和 ID3 最大的分水岭。这段代码把连续特征的所有取值排序然后在相邻值的中间点尝试切分# c45.py 连续属性最优切分点搜索 def calcContinuousFeatureBestSplit(dataSet, feature): values sorted([example[feature] for example in dataSet]) bestGainRatio 0.0 bestSplitPoint None for i in range(len(values) - 1): if values[i] values[i1]: continue # 相邻值相同则跳过 splitPoint (values[i] values[i1]) / 2.0 # 取中点 # 按 splitPoint 把数据集分成左右两组 leftSet [e for e in dataSet if e[feature] splitPoint] rightSet [e for e in dataSet if e[feature] splitPoint] # 计算这个二分方案下的信息增益比 ... # 后续与离散特征类似这里的关键参数是splitPoint (values[i] values[i1]) / 2.0也就是相邻取值的中点。C4.5 不是直接拿每个原始值当切分点而是用中点这样在连续空间里切分边界永远不会落在真实样本值上对噪声更鲁棒。代码里continue跳过相同值也是必要的防御操作原始数据里同一个值出现多次时中点会重合等于重复计算。3.3 CART 的基尼系数与二叉递归划分CART 走上了一条不同的路用基尼不纯度替代熵同时强制二叉树结构。基尼的计算比熵要简单很多不需要对数运算# cart.py 基尼不纯度计算 def calcGini(dataSet): labelCounts {} for featVec in dataSet: label featVec[-1] labelCounts[label] labelCounts.get(label, 0) 1 gini 1.0 for key in labelCounts: prob float(labelCounts[key]) / len(dataSet) gini - prob * prob return gini # 1 - Σ p^2基尼系数的计算复杂度远低于熵没有对数这是 CART 在工程上更受欢迎的原因之一。cart2.py的回归树版本则用了另一个指标——平方误差# cart2.py 回归树分裂指标 def calcMSE(dataSet): labels [e[-1] for e in dataSet] meanVal sum(labels) / float(len(labels)) mse sum((v - meanVal) ** 2 for v in labels) / float(len(labels)) return mse回归树每次切分选择使左右子集加权 MSE 最小的特征和切分点叶子节点的预测值就是该区域样本标签的均值。cart2.py的意义在于让你理解CART 不是一个分类专用算法它的树结构框架完全不变只需要替换分裂指标就摇身一变成为回归树。3.4 三个算法在同一份数据上的行为差异对照用 iris.csv 跑三个算法会出现一个有代表性的差异ID3 倾向于在花瓣长度和花瓣宽度上先分裂而且因为 iris 特征是连续值ID3 的离散化方式会直接影响建树结果C4.5 会自动搜索最优切分点树的深度通常会比 ID3 浅CART 在分类任务上产生的树是二叉的比 ID3 的多叉树在可视化时更清晰。算法分裂标准树形态连续特征支持缺失值支持ID3信息增益多叉树需提前离散化不支持C4.5信息增益比多叉树动态找切分点支持CART基尼系数 / MSE二叉树动态找切分点有替代方案选型逻辑很直观小数据集上手学原理用 ID3理解熵是怎么工作的数据里有大量连续特征且不想做太多预处理优先看 C4.5工程落地对效率有要求、或者需要回归能力直接用 CART。这几个算法的代码都在压缩包里你可以把同一份 iris 数据喂给三个文件对比生成的树结构差异。4. 动手跑通代码环境准备、运行参数与可视化输出代码读完了接下来把环境搭好把三棵树跑出来。这一章的所有操作我都按自己常用的方式走一遍你照着敲就行。4.1 运行环境最低配置treePlotter.py依赖 matplotlib三个算法文件依赖 numpy 和 pandas读 csv 用。Python 版本建议用 3.8 到 3.10太新的版本有些老代码的 API 会报废弃警告但不影响运行。# 创建虚拟环境并安装依赖 python -m venv dt_env source dt_env/bin/activate # Windows 下用 dt_env\Scripts\activate pip install numpy pandas matplotlib scikit-learnscikit-learn不是这个压缩包必需的但建议装上后面验证手工计算结果时可以用来交叉比对。装依赖这一步如果网络慢可以用国内镜像源pip install -i https://pypi.tuna.tsinghua.edu.cn/simple numpy pandas matplotlib。4.2 运行三个算法并观察输出进入解压目录直接运行注意观察终端输出和弹窗的图形python id3.py如果终端打印出类似这样的结构说明代码跑通了{petal length: {0.3: Iris-setosa, ...}}有一点要提前打预防针treePlotter.py里如果用了中文节点文本在 matplotlib 默认字体下会显示成方块乱码。这个坑后面避坑章节细讲现在先知道有这个问题就行。4.3 修改 split 阈值参数观察过拟合代码里默认会一直分裂到叶子节点纯净为止这会导致训练集上 100% 准确率但新数据上表现很差。我一般会在算法文件里临时加一个最小样本数参数来观察剪枝的效果# 在递归建树函数的入口加一个参数 def createTree(dataSet, labels, minSamples2): # 当样本数少于阈值时停止分裂 if len(dataSet) minSamples: return majorityVote([e[-1] for e in dataSet]) ...minSamples设得越大树越浅训练准确率越低但泛化能力通常更好。你可以分别设 2、5、10 跑三遍对比树的可视化效果和测试集上的准确率。这个手动改参数的过程比直接调 sklearn 更能加深对剪枝的理解。4.4 可视化样式与保存图片createPlot默认是弹窗显示如果你想把图保存下来做笔记可以用plt.savefig替代或补充plt.show()# treePlotter.py 尾部追加保存逻辑 plt.savefig(tree_output.png, dpi150, bbox_inchestight)dpi150保证图片放大后不糊bbox_inchestight会自动裁掉多余的空白边。这个参数组合是我最常用的出来的图放进文档里直接能用。5. 避坑决策树复现中的五条血泪记录拆这个包的过程中我踩了不止一个坑有些是代码本身的问题有些是换成自己的数据集后才会暴露的通用问题。整理五条最有代表性的按「现象 → 原因 → 解决」的顺序写清楚。5.1 信息熵计算结果不一致跑id3.py和我自己在 sklearn 里交叉验证的结果对不上计算出来的信息熵总是略有偏差。排查下来问题是 root 节点的熵算的是全体样本的类别分布但 sklearn 在criterionentropy时用的公式在数学上是等价的只是浮点运算顺序不同导致末尾几位不同。更隐蔽的原因是math.log传入的prob如果是 0Python 会抛ValueError代码里没有防御逻辑。解决方式是对splitDataSet返回空集的场景加try-except或者if prob 0判断同时不要在不同实现之间做完全等值比较用abs(ent1 - ent2) 1e-6做容差比较。5.2 C4.5 连续属性切分点报错c45.py在 iris 数据上一切正常我把数据换成有重复值的企业估值数据后报错float division by zero。原因是calcGainRatio里splitInfo算出来是 0因为某个特征的所有取值都一样信息熵为 0增益比的分母就没了。解决方式是加一个判断if splitInfo 0.0: return 0.0。这个坑的本质是算法没做边界防御遇到特征全等时应当直接跳过该特征。看原代码时发现有这行防御但只在主分支里连续属性子分支没有覆盖到换数据后立刻暴露。5.3 树可视化中文乱码运行treePlotter.py时节点里的特征名如果是中文图形窗口里显示成一个个方框。原因是 matplotlib 默认字体是 DejaVu Sans不支持中文。解决方式是在createPlot函数开头加两行from matplotlib import rcParams rcParams[font.sans-serif] [SimHei, Arial Unicode MS] rcParams[axes.unicode_minus] FalseSimHei是 Windows 下的黑体macOS 用Arial Unicode MS。如果你 Linux 服务器上既没有 SimHei 也没有 Arial Unicode MS就用fc-list :langzh查一下系统装了哪些中文字体把名字填进去就行。这个坑和数据无关纯环境问题但几乎人人都会遇到。5.4 同一数据集上三次运行结果不一样我连续跑三遍cart.py发现每次生成的树不完全一样。原代码里train_test_split没有固定随机种子。分类树在 iris 这种干净数据集上差别不明显但换到有噪点的数据上叶子节点的标签归属就会漂移。解决方式很简单在文件头部加random.seed(42)或np.random.seed(42)保证可复现。这也是任何机器学习代码的通用习惯跑实验不固定种子结果对不上是必然的。5.5 传入新数据集时出现维度不匹配把 iris.csv 换成自己的数据两个特征以上时splitDataSet返回的子集可能丢失维度信息。原因是特征筛选时用了for i in range(numFeatures)遍历但对离散特征的取值做了硬编码假设每个特征取值数不超过某个阈值。换数据后取值数变大或变小都会触发越界。解决方式是检查splitDataSet实现里是否用了featVec[i]直接索引如果是确认数据格式是 list of list且每行长度一致。最稳妥的做法是把数据加载部分改成 pandas DataFrame 之后再做values.tolist()转换统一格式。6. 验证代码正确性用手工复算一棵树的完整过程最后一个技巧是验证这些手写算法到底算得对不对不用 sklearn 交叉比对也能做拿原始训练数据里的一小部分手工算一次分裂然后和代码输出对比。这个手工验证法是我多年养成的习惯任何自动化的代码都值得做一次白盒核验。以 ID3 为例取 iris.csv 里前 6 条数据两个类别特征选花瓣长度。原始数据如下为了演示做了简化5.1,3.5,1.4,0.2,Iris-setosa 4.9,3.0,1.4,0.2,Iris-setosa 4.7,3.2,1.3,0.2,Iris-setosa 4.6,3.1,1.5,0.2,Iris-setosa 5.0,3.6,1.4,0.2,Iris-setosa 5.4,3.9,1.7,0.4,Iris-setosa 6.4,3.2,4.5,1.5,Iris-versicolor前 6 条是 setosa第 7 条是 versicolor。按 ID3 的离散化逻辑把所有值排序后自动找切分点切分点在 1.7 和 4.5 之间取 3.1分成左右两组。手工算一下根节点的熵Ent(D) -7/8 * log2(7/8) - 1/8 * log2(1/8) ≈ 0.5436。按切分点 3.1 划分后左侧全是 setosa熵为 0右侧 1 条 versicolor熵为 0。加权条件熵 6/8 * 0 2/8 * 0 0信息增益达到最大。代码里输出的第一次分裂位置就是花瓣长度 ≤ 3.1和手算一致说明熵计算和切分逻辑都是对的。# 手工验证脚本 import math from collections import Counter labels [setosa]*6 [versicolor] ent_d 0.0 for k, v in Counter(labels).items(): p v / len(labels) ent_d - p * math.log(p, 2) print(f根节点熵: {ent_d:.4f}) # 期望约 0.5436这个方法的价值在于你不用信任任何一个第三方库也不用盲信压缩包里的代码只凭信息论的原始定义就能验证模型对不对。从那以后我每次拿到手写机器学习算法的代码包都会先取 10 条以内的样本做一轮手工复算确认分裂标准和代码实现一致再放心地用到正式数据上。这套流程帮我筛掉过不少表面能用、换个数据就出错的代码包希望也能帮你在学这三棵树的路上少走几步弯路。盯准一个细节拆透一个资源比囫囵吞枣跑完三个脚本要有用得多。本文还有配套的精品资源点击获取
返回列表