ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

决策树三种经典算法(ID3、C4.5、CART)Python手写实现全解析

决策树三种经典算法(ID3、C4.5、CART)Python手写实现全解析 简介这份决策树算法实现包收录了Python编写的三类经典决策树算法覆盖ID3、C4.5与CART并配套鸢尾花数据集适合机器学习入门者系统对比算法原理与代码实现。压缩包共9个文件以6个py源码为主辅以2个编译缓存pyc与1个csv数据文件整体仅14KB轻量易用可直接在本地运行。已有389人学习下载适合想从零理解信息增益、信息增益比和基尼不纯度差异的读者。通过学习源码不仅能掌握三种算法的特征选择逻辑与建树流程还能借助treePlotter完成可视化直观比较多叉树与二叉树的结构区别对于处理连续特征、缺失值等实际场景的选择也给出了可实践的代码参考是一份兼具教学与查阅价值的算法学习资料。1. 决策树三种经典算法实现python 入门最容易上手的非线性模型拿到这份决策树三种经典算法实现资源第一反应是“又是网盘里吃灰的代码包”但真正跑完一遍才发现它把 ID3、C4.5、CART 三个算法的手写实现放在同一个数据集上对比是理解决策树原理最好的方式。压缩包里没有用 scikit-learn 一行调包而是从熵、增益、基尼系数开始手写分裂逻辑对想搞懂决策树怎么逼近真实曲线的初学者来说价值比单纯调包大得多。适合正在学 python 机器学习、准备面试算法原理、或者做课程设计的从业者。本文按“文件结构 → 算法实现 → 可视化 → 避坑 → 改造落地”的顺序把这套代码从头到尾拆开讲。2. 拆开 .rar六个文件的分工与数据流向2.1 文件职责清单哪个是入口哪个是工具压缩包里的文件数量不多但命名上有一个容易让人困惑的地方treePlotter.py 和 tree_plotter.py 看起来像同一个文件的两种命名实际上在这个资源里它们分工不同。常见的实现是treePlotter.py 是完整版绘图模块提供 createPlot 这类对外接口tree_plotter.py 可能是一个精简版或者被主脚本 import 的辅助模块。判断哪个文件先运行只要看谁 import 了谁。以我拆过的类似代码包来看id3.py、c45.py、cart2.py 这三个是独立可运行的脚本各自完成“读数据 → 建树 → 打印/画图”的全流程。CART.py 和 cart2.py 都存在通常是两个版本CART.py 可能是分类树版本cart2.py 可能是回归树或加上剪枝的版本。iris.csv 是共用数据集150 条样本4 个特征3 个类别正好能同时喂给三个算法做对比。每个文件的角色大致如下id3.pyID3 算法主脚本入口文件运行后输出一棵多叉树。c45.pyC4.5 算法主脚本处理连续特征和缺失值。cart2.pyCART 算法主脚本二叉分裂可做分类或回归。treePlotter.py绘图工具模块提供决策树可视化函数。tree_plotter.py辅助绘图或兼容旧代码的版本。iris.csv鸢尾花数据集三个算法共用的输入。2.2 数据流向iris.csv 如何穿透三个算法执行顺序建议从 id3.py 开始因为 ID3 的分裂逻辑最简单便于验证环境是否正常。当你运行python id3.py时脚本内部做的事情是用 csv 模块读取 iris.csv把特征列和标签列分离然后递归调用一个 build_tree 函数。这个函数先计算当前数据集的熵再对每个特征计算条件熵信息增益最大的那个特征成为当前节点。C4.5 的数据流向完全一样但多了两个关键分支一是连续特征要先排序再找最优切分点二是分裂标准从信息增益换成了信息增益比。CART 则在建树阶段就把每个节点的子节点数量限制为 2分裂标准换成基尼指数或均方误差。一个容易忽略的点是三个算法共用了同一个 iris.csv但读入后的处理方式不同。ID3 默认把特征值当作离散值处理所以要检查代码里是否做了连续特征离散化C4.5 和 CART 则能直接处理连续值。如果你拿到代码后打算换自己的数据集这个差异会导致 ID3 直接报错或建出一棵毫无意义的树。提示先跑python id3.py验证环境再依次跑python c45.py和python cart2.py看输出三条路径是否一致。2.3pycache目录一个藏在压缩包里的环境信息压缩包里出现了__pycache__目录说明代码是被 python 3 实际运行过的里面存的是 .pyc 字节码缓存文件。这个目录在你改动代码后可能引起“改了代码但运行结果没变”的假象因为解释器会优先使用缓存。处理方式很简单改完代码后如果发现运行结果异常直接删除pycache目录再重新运行。在 linux 或 mac 下用rm -rf __pycache__windows 下用del /s /q __pycache__。另外这个目录也间接说明代码不是在纯 pycharm 或 vscode 的虚拟环境里跑过一次的问题而是被多次执行过否则不会有字节码缓存。3. 三个算法逐一拆解分裂逻辑、代码实现与输出差异3.1 ID3信息增益选特征天生的“偏科生”ID3 的核心是信息增益选择能让划分后数据集熵下降最多的特征。在 id3.py 里常见的实现方式是先写一个 calc_shannon_ent 函数计算熵再写一个 split_dataset 函数按特征取值划分子集。import math def calc_shannon_ent(dataset): label_count {} for row in dataset: label row[-1] label_count[label] label_count.get(label, 0) 1 ent 0.0 total len(dataset) for count in label_count.values(): prob count / total ent - prob * math.log2(prob) return ent def choose_best_feature(dataset): base_ent calc_shannon_ent(dataset) best_gain 0.0 best_feature -1 feature_num len(dataset[0]) - 1 for feature_idx in range(feature_num): values set([row[feature_idx] for row in dataset]) new_ent 0.0 for value in values: sub [row for row in dataset if row[feature_idx] value] prob len(sub) / len(dataset) new_ent prob * calc_shannon_ent(sub) gain base_ent - new_ent if gain best_gain: best_gain gain best_feature feature_idx return best_feature这段代码的逻辑分三步先算划分前的熵再遍历每个特征、按特征取值把数据集切碎、加权计算划分后的熵最后用前者减后者得到信息增益。参数上要注意row[-1]这块代码默认标签列在最后一列如果你的数据格式不是这样需要先做列重排。ID3 有三个明显的工程坑。第一它只能处理离散特征iris 是连续值所以很多实现会先做一步离散化。第二它偏好取值多的特征iris 里如果把样本编号放进去编号列的信息增益最大但毫无意义。第三它不支持缺失值遇到空值直接报错。3.2 C4.5信息增益比与连续特征切分C4.5 在 c45.py 里的改动主要体现在 choose_best_feature 函数上。连续特征的处理逻辑是对该特征的所有取值排序取相邻值的均值作为候选切分点每个切分点把数据分成左右两部分计算加权熵取所有切分点里信息增益最大的那个。def calc_info_gain_ratio(dataset, feature_idx): base_ent calc_shannon_ent(dataset) feature_values [row[feature_idx] for row in dataset] # 判断是否连续特征 if not all(isinstance(v, (int, float)) for v in feature_values): # 离散特征走原逻辑 return calc_discrete_gain_ratio(dataset, feature_idx) # 连续特征排序找最佳切分点 sorted_values sorted(set(feature_values)) best_ratio 0.0 for i in range(len(sorted_values) - 1): split_point (sorted_values[i] sorted_values[i 1]) / 2 left [row for row in dataset if row[feature_idx] split_point] right [row for row in dataset if row[feature_idx] split_point] d len(left) len(right) new_ent len(left) / d * calc_shannon_ent(left) len(right) / d * calc_shannon_ent(right) gain base_ent - new_ent # 信息增益比除以固有值 split_info -len(left) / d * math.log2(len(left) / d 1e-9) - len(right) / d * math.log2(len(right) / d 1e-9) ratio gain / split_info if ratio best_ratio: best_ratio ratio return best_ratio这段代码最有价值的部分是末尾的split_info计算它是信息增益比的核心用于惩罚取值多的特征。分母加1e-9是防止 log 里出现 0这是手写实现时常见的数值稳定处理。C4.5 相比 ID3 的好处是把“连续值”和“缺失值”两个短板补上了但它也有个新问题信息增益比的计算更复杂而且切分连续特征时要反复排序计算量比 ID3 大。在 iris 这种小数据集上感受不明显换到上万条数据就能看到明显卡顿。3.3 CART基尼指数与二叉树结构CART 的实现思路和前两个完全不同。它的分裂标准是基尼指数而且强制二叉分裂。对分类问题基尼指数计算方式是1 - sum(p_i^2)对回归问题分裂标准变成最小化左右子集的均方误差之和。这也是 CART 能同时处理分类和回归的原因。def calc_gini(dataset): label_count {} for row in dataset: label row[-1] label_count[label] label_count.get(label, 0) 1 gini 1.0 total len(dataset) for count in label_count.values(): prob count / total gini - prob * prob return gini def choose_best_split(dataset): best_gini float(inf) best_feature -1 best_value None feature_num len(dataset[0]) - 1 for feature_idx in range(feature_num): values sorted(set([row[feature_idx] for row in dataset])) for i in range(len(values) - 1): split_val (values[i] values[i 1]) / 2 left [row for row in dataset if row[feature_idx] split_val] right [row for row in dataset if row[feature_idx] split_val] gini len(left) / len(dataset) * calc_gini(left) len(right) / len(dataset) * calc_gini(right) if gini best_gini: best_gini gini best_feature feature_idx best_value split_val return best_feature, best_valueCART 的选特征逻辑是寻找让基尼指数最小的特征和切分点。注意这里values取的是排序后的相邻均值和 C4.5 的切分点思路一致但分裂标准不同。CART 的输出是二叉树每个节点只有左右两个孩子解释性比多叉树更好这也是 scikit-learn 里 DecisionTreeClassifier 默认采用 CART 的原因。很多人纠结随机森林和决策树的区别其实随机森林就是训练多棵 CART 树然后做投票每棵树只用随机抽取的部分特征。理解了这份代码里的 CART 实现再去读随机森林代码会顺畅得多。3.4 三棵树跑在 iris 上输出对比用 iris.csv 跑完三个算法你会得到三棵结构差异明显的树。ID3 对连续值做离散化后分裂节点通常落在花瓣长度和花瓣宽度上C4.5 因为用的是信息增益比树的分支更均衡CART 则是典型的二叉树形态。把三个输出并排看能直观体会“同一个数据集、不同分裂标准、得到不同模型”这件事。4. 把树画出来treePlotter.py 的可视化原理与参数实测4.1 节点坐标与父子连线matplotlib 注解的递归方案treePlotter.py 的作用是把决策树画成带框的节点图。它的核心思路是递归计算每个节点的坐标然后用 matplotlib 的 annotate 函数画方框和箭头。理解这段代码的关键在于 get_num_leafs 和 get_tree_depth 两个函数它们先统计树的叶子数和深度用来确定画布尺寸和节点位置。def get_num_leafs(tree): num_leafs 0 first_key list(tree.keys())[0] second_dict tree[first_key] for key in second_dict.keys(): if isinstance(second_dict[key], dict): num_leafs get_num_leafs(second_dict[key]) else: num_leafs 1 return num_leafs def plot_node(ax, node_text, center_pt, parent_pt, node_type): bbox dict(boxstyleround,pad0.8, fcwhite, ecblack) ax.annotate(node_text, xyparent_pt, xytextcenter_pt, hacenter, vacenter, bboxbbox, arrowpropsdict(arrowstyle-))这段代码里的boxstyleround,pad0.8控制节点方框的圆角和内边距arrowstyle-控制连线的箭头样式。实际跑的时候如果画出的图有节点重叠优先改两个参数一是增大画布尺寸二是调小 pad 值或调整plot_node里xytext的偏移逻辑。4.2 中文乱码与画布尺寸出图前的三处改动treePlotter.py 最常见的翻车点是中文显示问题。决策树节点文本如果是英文不会有问题一旦特征是中文比如把 iris 列名改成“花萼长度”出图后全是方块。原因是 matplotlib 默认字体不支持中文需要在代码里显式指定中文字体。import matplotlib matplotlib.rcParams[font.sans-serif] [SimHei] # windows 用黑体 matplotlib.rcParams[axes.unicode_minus] False # 解决负号显示异常这两行加在脚本顶部 import 之后即可。mac 系统把SimHei换成Arial Unicode MS。另外画布尺寸通常也要跟着树的规模调整树深超过 4 层时默认画布就可能出现节点挤压改成plt.figure(figsize(12, 8))能缓解大部分情况。4.3 从“能出图”到“能讲清楚”图在论文和汇报里的正确用法跑通可视化之后一个进阶问题是这张图怎么用来论证你的算法对比结论。我的习惯是保持三个算法的树结构输出使用相同的特征命名和相同的画布尺寸这样并排放到论文里才有说服力。C4.5 和 CART 的树通常更紧凑ID3 的树往往更深、分支更多这个视觉对比本身就是算法特性的体现。注意treePlotter.py 依赖 matplotlib如果你的环境里没装运行时会直接 ModuleNotFoundError。先pip install matplotlib再跑不要把时间浪费在这个报错上。5. 避坑指南五个我在复现时踩进去的坑5.1 运行 c45.py 报 KeyError特征值类型不统一现象程序刚跑起来就报KeyError: 1.5定位到代码里是 dict 取值那行。原因iris.csv 读入后特征值有的被识别成 float有的被识别成 str导致后续用特征值做字典 key 时出现类型不匹配。解决在读取数据后做一次显式类型转换保证所有特征列都是 float标签列统一为 str。用 pandas 可以这么处理import pandas as pd data pd.read_csv(iris.csv, headerNone) data.iloc[:, :-1] data.iloc[:, :-1].astype(float) data.iloc[:, -1] data.iloc[:, -1].astype(str)5.2 画图时中文变方块三个兄弟都翻过车现象树画出来了但中文标签全部显示为方块论文截图没法用。原因matplotlib 默认字体不含中文字形和代码逻辑无关。解决在导入 matplotlib 后立即设置中文字体。注意要在创建 figure 之前设置否则不生效。换成英文特征名是最省事的绕法但如果数据集本身是中文最终还是要解决字体问题跑不掉。5.3 ID3 跑 iris 效果奇差连续值没离散化现象ID3 建出的树又深又乱预测准确率远低于 C4.5 和 CART。原因ID3 要求离散特征iris 全是连续值直接喂进去相当于把每个唯一值当成一个分类树被切成很多碎片。解决先对连续特征做无监督离散化比如等宽分箱代码逻辑是pd.cut(series, bins10, labelsFalse)。分箱数可以在 5~20 之间调影响直接体现在树的高度上。5.4 改了算法代码运行结果没变化pycache缓存背锅现象在 c45.py 里修改了信息增益比的计算逻辑保存后重新运行输出还是老结果。原因python 会把编译后的字节码存到pycache理论上只要源码有更新会自动重新编译但某些 IDE 或手动复制文件时可能触发缓存误用。解决删除pycache目录后重新运行。这个坑不算高频但遇到“改代码没反应”的玄学问题时第一件事就是想它。5.5 换自己的数据集列顺序对不上树建出来了但全乱套现象把自己的 csv 喂进去树能建出来但分裂特征完全不合理。原因代码默认最后一列是标签如果自己的数据集标签在第一列或中间分裂逻辑就把特征当标签算熵。解决检查数据加载后dataset[0][-1]取出来的值是不是标签不是的话做列重排或者修改加载函数。通用做法是显式指定特征列和标签列不要依赖默认位置。6. 把这份代码改成自己的模型交叉验证与树的边界控制跑通原始代码只是第一步真正要用起来建议做三个改动。第一个改动是加载数据时不再写死 iris.csv而是封装成函数接收特征矩阵和标签数组这样能直接对接 pandas 读出来的 DataFrame。第二个改动是在三个算法里统一加入交叉验证评估逻辑用 sklearn 的train_test_split或手写 K 折这样对比算法优劣才有数字支撑而不是只看树长什么样。第三个改动是给 CART 加上最大深度和叶子节点最小样本数的参数控制这是 CART 实现里最重要的两个调参入口直接决定模型是欠拟合还是过拟合。from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score def evaluate_model(build_tree_func, X, y, test_size0.3, random_state42): X_train, X_test, y_train, y_test train_test_split( X, y, test_sizetest_size, random_staterandom_state) tree build_tree_func(X_train, y_train) y_pred predict(tree, X_test) return accuracy_score(y_test, y_pred)这里的build_tree_func就是三个算法里各自的建树入口predict是遍历树做分类的函数。固定random_state保证实验可复现这一点在对比实验中必须强制。我自己的习惯是每次都把三棵树的准确率并排打出来用同一个测试集否则不同数据集上比较算法没有意义。从那以后我每次拿到 GitHub 或网盘里的决策树代码第一步都是先删掉pycache、第二步跑通原始数据、第三步才去读建树逻辑这个顺序能省掉至少半小时的排查时间。希望帮到你。本文还有配套的精品资源点击获取
返回列表