简介:这份决策树算法实现包收录了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.py:ID3 算法主脚本,入口文件,运行后输出一棵多叉树。
- c45.py:C4.5 算法主脚本,处理连续特征和缺失值。
- cart2.py:CART 算法主脚本,二叉分裂,可做分类或回归。
- 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(boxstyle="round,pad=0.8", fc="white", ec="black") ax.annotate(node_text, xy=parent_pt, xytext=center_pt, ha="center", va="center", bbox=bbox, arrowprops=dict(arrowstyle="<-"))这段代码里的boxstyle="round,pad=0.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', header=None) 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, bins=10, labels=False)。分箱数可以在 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_size=0.3, random_state=42): X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=test_size, random_state=random_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、第二步跑通原始数据、第三步才去读建树逻辑,这个顺序能省掉至少半小时的排查时间。希望帮到你。
本文还有配套的精品资源,点击获取