ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

决策树手写实现:信息熵与基尼指数及剪枝策略详解

2026/9/10 2:36:20 拓冰建站 浏览量
决策树手写实现:信息熵与基尼指数及剪枝策略详解 简介西瓜书《机器学习》第四章决策树Python实现资源包面向机器学习初学者、复试备考者以及需要动手复现教材算法的读者。资源围绕信息熵与基尼指数两种划分准则完整实现决策树生成、预剪枝、后剪枝与可视化绘图并针对西瓜数据集2.0、3.0以及多个UCI数据集给出实验比较和统计显著性检验相关代码便于读者直观理解不同划分选择与剪枝策略对模型泛化能力的影响。包内共9个文件其中5个csv文件提供实验所需数据集4个py文件分别对应树生成、剪枝对比、树形绘制等核心功能压缩包仅16KB轻巧便携适合逐行研读、二次修改或课程设计复用。目前已有10318人学习下载内容紧贴周志华《机器学习》第四章课后编程题运行后即可生成决策树并完成未剪枝、预剪枝、后剪枝结果对比是入门决策树实战的高性价比素材。1. 决策树手写实现为什么信息熵和基尼指数不能只靠调库很多读者学西瓜书第四章时直接调sklearn的DecisionTreeClassifier就完事了但到面试或做课程设计时被问“预剪枝和后剪枝在代码里到底怎么实现”就卡壳。这份资源正是把第四章的算法逐一“拆开”给你看基于信息熵的ID3、基于基尼指数的CART以及手写的预剪枝/后剪枝逻辑而不是让你调用一个黑盒。项目附带西瓜数据集2.0、3.0和UCI的adult-stretch、iris等数据适合正在啃西瓜书的学生、需要复现课程实验的开发者以及想从底层理解决策树原理的机器学习初学者。我建议你先跑一遍纯Python实现的Decision_tree.py再去碰任何机器学习库。你会发现手动实现树的生长、剪枝与可视化比调库更能建立对剪枝策略的直觉。这份代码没有复杂的工程依赖核心是递归、字典和Matplotlib绘图你完全可以一行一行读通。2. 信息熵划分与ID3用Decision_tree.py生成西瓜3.0决策树西瓜书4.3要求“基于信息熵进行划分选择”并给西瓜数据集3.0生成一棵树。这里提到的信息熵就是香农熵公式为 Ent(D) -Σ p_k * log2(p_k)。ID3算法在每个节点选择信息增益最大的特征作为划分属性信息增益就是父节点熵减去子节点熵加权和。Decision_tree.py正是按这个逻辑实现的。2.1 从文本到信息熵calcShannonEnt的实现细节数据集以CSV格式存储例如西瓜数据集3.0包含色泽、根蒂、敲声等特征最后一列是类别好瓜/坏瓜。计算信息熵前需要把每列特征值和标签读成list。下面是项目中的核心熵函数我简化了文件读写部分import math from collections import Counter def calcShannonEnt(data_set): label_counts Counter(data[-1] for data in data_set) total len(data_set) entropy 0.0 for count in label_counts.values(): p count / total entropy - p * math.log2(p) return entropydata_set是一个二维列表每一行是[feature1, feature2, ..., featureN, label]最后一位是标签。Counter先统计每个类别的样本数再按香农熵公式累加。若某个类别概率为00 * log2(0)在Python里会报错但实际分类数据不会出现若要严谨可加if count 0: continue。这个函数只依赖标签列不关心特征维度所以可以被ID3和CART共用。2.2 划分选择chooseBestFeature如何计算信息增益有了熵下一步按特征划分数据。项目中的splitDataSet负责取出某特征等于某值的子集chooseBestFeature遍历每个特征计算划分前后熵的差值也就是信息增益。代码骨架如下def choose_best_feature(data_set): num_features len(data_set[0]) - 1 base_entropy calcShannonEnt(data_set) best_gain 0.0 best_feature -1 for i in range(num_features): unique_vals set(row[i] for row in data_set) new_entropy 0.0 for value in unique_vals: sub_set [row for row in data_set if row[i] value] prob len(sub_set) / len(data_set) new_entropy prob * calcShannonEnt(sub_set) info_gain base_entropy - new_entropy if info_gain best_gain: best_gain info_gain best_feature i return best_feature需要注意num_features排除了最后一列标签。unique_vals用set去重天然支持字符串类型的特征值。这里没有处理连续值因为西瓜3.0的特征全是离散的。如果遇到连续值需要先做二分排序但那是4.4的扩展内容Decision_tree.py不做。2.3 在西瓜3.0上生成决策树的完整流程createTree按“样本全同→返回类别特征用尽→返回多数类否则选最佳特征递归”三步走。你可以在main函数里这样调用import csv def load_data(file_path): with open(file_path, encodingutf-8-sig) as f: reader csv.reader(f) header next(reader) # 跳过表头 return [row for row in reader if row] data_list load_data(西瓜数据集3.0.csv) feat_labels [色泽, 根蒂, 敲声, 纹理, 脐部, 触感] # 按CSV列顺序 tree create_tree(data_list, feat_labels) print(tree)输出是一个嵌套字典形如{纹理: {清晰: {根蒂: ...}, 稍糊: 坏瓜, ...}}。字典的键是特征名值有两种若不是叶子则是一个新字典若是叶子则直接是“好瓜/坏瓜”字符串。函数输入输出关键点calcShannonEntdata_set浮点数熵值只统计最后一列choose_best_featuredata_set最优特征索引比较各特征信息增益splitDataSetdata_set, i, value子集列表用于递归建树create_treedata_set, feat_labels嵌套字典树递归终止条件如果我第一次跑会先打印根节点选的是“纹理”这正是西瓜书上的经典结果。你可以在自己的机器上验证训练集和测试集都是同一个数据集时ID3生成的树在训练集上必然100%准确但泛化能力要依赖后面剪枝。提示CSV文件若带BOMencoding要用utf-8-sig否则第一列特征名会带有\ufeff前缀。3. 基尼指数与CART剪枝西瓜2.0上的预剪枝/后剪枝对比西瓜书4.4要求改用基尼指数。CART决策树使用基尼指数选择划分属性Gini(D) 1 - Σ p_k^2基尼指数越小纯度越高。和ID3不同CART是一棵二叉树对离散特征不再按“每个取值一分枝”而是把特征取值划分成两个子集。项目中的CART.py实现后者CART_剪枝.py则加入了预剪枝和后剪枝逻辑。3.1 基尼指数计算CART.py中的纯度度量CART里的核心函数和ID3的熵函数几乎对称只是把log部分换成平方和。下面这段代码来自CART.py的核心部分def calc_gini(data_set): label_counts Counter(data[-1] for data in data_set) total len(data_set) gini 1.0 for count in label_counts.values(): p count / total gini - p * p return gini基尼指数比熵计算更省运算因为它不涉及log。当样本类别完全相同时基尼指数为0类别越混杂值越大。这也是CART偏好二元切分的原因对于取值很多的离散特征基尼指数比信息增益更不容易偏向取值多的特征。3.2 预剪枝在建树前和建树中刹车预剪枝最常见的是限制最大深度max_depth、最小样本数min_samples_leaf或者提前判断划分后验证集精度是否提升。CART_剪枝.py里通常把验证集精度作为早停依据我在项目中看到类似这样的代码def prune_pre(data_set, valid_set, max_depth, depth0): labels [row[-1] for row in data_set] if depth max_depth or len(set(labels)) 1: return majority(labels) # 尝试当前节点的最优划分 feature, value choose_best_split(data_set) if feature is None: return majority(labels) # 构造验证集预测如果划分后精度没有提升则直接返回叶子 if valid_acc_after_split(valid_set, feature, value) valid_acc_before(valid_set): return majority(labels) left_data, right_data split_binary(data_set, feature, value) left_valid, right_valid split_binary(valid_set, feature, value) tree {feature: { value: prune_pre(left_data, left_valid, max_depth, depth1), ! value: prune_pre(right_data, right_valid, max_depth, depth1)}} return tree这段代码不是完整版但表达了预剪枝的三道关卡深度限制、类别纯度、验证集精度。choose_best_split返回的feature, value是基尼指数最小的二元划分valid_acc_before是根节点直接投多数类的验证集准确率valid_acc_after_split是按划分后的两个子树预测的准确率。只有后者严格大于前者才继续生长否则就地剪成叶子。这样能显著减小树体积但有可能欠拟合因为某些没有立刻提升的划分在深层可能会带来收益。3.3 后剪枝从完整树自底向上减叶子与预剪枝不同后剪枝先生成完整的未剪枝树再用验证集自底向上判断。核心思想是如果某个子树的所有叶子替换成该子树样本的多数类验证集精度不下降那就替换。实现时常用递归函数prune_post(tree, valid_set)需要先递归到叶子再返回比较def prune_post(tree, valid_set): if not isinstance(tree, dict): return tree # 假设树结构为 {feature: {value: subtree, ...}} for feature, branch in tree.items(): for value, subtree in branch.items(): sub_valid [row for row in valid_set if row[feature] value] if isinstance(subtree, dict): tree[feature][value] prune_post(subtree, sub_valid) # 计算剪枝前后的验证集精度 if valid_acc(tree, valid_set) valid_acc(majority_label(valid_set), valid_set): return majority_label(valid_set) return tree这里的valid_acc要传入一个能预测的结构叶子节点直接返回多数类。注意递归函数必须先把子树剪完再比较当前节点因为自底向上意味着每个内部节点只有在所有子树都处理完后才考虑是否合并成叶子。后剪枝通常比预剪枝更能保持精度但训练时间更长因为要先完整生长。3.4 西瓜2.0上的三棵决策树对比在西瓜数据集2.0上可以分别运行CART.py得到未剪枝树设置不同max_depth得到预剪枝树再对未剪枝树跑prune_post得到后剪枝树。三者的可用对比维度如下对比项未剪枝预剪枝后剪枝树深度最大较小中等叶子数最多最少较少训练集精度100%较低接近100%验证集精度可能过拟合稳定但欠拟合通常最高时间开销最短中最多实际项目里CART_剪枝.py会把这三个模型都保存下来并用同一个测试集计算准确率。如果你的验证集划分规模很小预剪枝可能直接在根节点就停止这是正常现象建议用分层抽样保证训练/验证集类别比例一致否则剪枝判断会失真。4. UCI数据集横向实验双算法×剪枝策略的显著性检验4.6要求选4个UCI数据集对ID3和CART的未剪枝、预剪枝、后剪枝共6种配置做比较并做统计显著性检验。这一步是为了回答一个务实问题预剪枝和后剪枝的精度差异在统计上是否可信还是只是随机波动。4.1 实验准备数据集与模型封装项目中给出的adult-stretch.csv、iris.csv等就是UCI标准数据集。Iris是鸢尾花分类类别3个特征都是连续值而Decision_tree.py只接受离散特征所以需要先对连续特征做离散化。常见做法是按照等宽或等频分成35个区间代码类似import pandas as pd def discretize(df, column, bins4): df[column] pd.cut(df[column], bins, labelsFalse) return df对于CART.py本身支持连续值的二分划分可以不必离散化。但为了公平比较建议两种算法都使用统一预处理后的数据否则无法判断差异来自算法还是数据处理。此外每个数据集要划分成训练集、验证集、测试集三份比例可为6:2:2。验证集用于剪枝测试集用于最终评估。4.2 六种配置的评估脚本写一个循环对每个数据集跑6个模型记录测试集准确率。下面是一个脚本骨架使用自写模型和scikit-learn的准确率计算from sklearn.metrics import accuracy_score from sklearn.model_selection import train_test_split data pd.read_csv(adult-stretch.csv) X data.drop(columns[label]) y data[label] X_train, X_temp, y_train, y_temp train_test_split(X, y, test_size0.4, stratifyy) X_valid, X_test, y_valid, y_test train_test_split(X_temp, y_temp, test_size0.5, stratifyy_temp) # 转成列表形式给自写模型 train_data [list(x) [label] for x, label in zip(X_train.values, y_train)] valid_data [list(x) [label] for x, label in zip(X_valid.values, y_valid)] # ID3未剪枝 tree_id3 create_tree(train_data, list(X.columns)) pred predict_tree(tree_id3, X_test.values) print(ID3 no prune:, accuracy_score(y_test, pred)) # CART预剪枝 max_depth3 tree_cart_pre build_cart_pruned(train_data, valid_data, max_depth3)这里的predict_tree是遍历树结构的预测函数项目中没有明确提供但实现不复杂根节点匹配特征值走到叶子就返回类别。整体思路是把每行测试样本沿着树往下走。注意自写算法的输入输出格式要和sklearn准确率函数兼容即传list或numpy array。为了方便实验记录6种配置可以按下面的方式编号配置编号算法剪枝策略使用的函数1ID3未剪枝create_tree2ID3预剪枝create_tree max_depth3ID3后剪枝create_tree prune_post4CART未剪枝build_tree5CART预剪枝prune_pre6CART后剪枝prune_post4.3 统计显著性检验配对t检验与Friedman检验6组模型的测试集准确率只有一份不足以做统计检验。常见做法是在同一个数据集上按不同随机种子做10次或30次重复实验得到多组成对样本再进行配对t检验或Wilcoxon符号秩检验。下面是配对t检验示例from scipy import stats acc_cart_post [0.92, 0.90, 0.91, 0.93, 0.89] # 5次实验CART后剪枝 acc_id3_none [0.88, 0.85, 0.89, 0.86, 0.90] t_stat, p_value stats.ttest_rel(acc_cart_post, acc_id3_none) print(p , p_value)当p0.05时可以认为两种配置在测试集上的平均精度差异显著。对于多个数据集、多对算法更规范的做法是Friedman检验再配合Nemenyi后续检验画临界差图。不过手写Friedman检验并不复杂公式是在每个数据集上对算法排名再计算卡方统计量。你可以用scipy的friedmanchisquare直接得到结果。4.4 实验结果的解读顺序报告里不要只放准确率表格。建议先展示每个数据集上6种配置的平均准确率和标准差再用热力图展示排名。你会发现在数据集规模小时预剪枝往往优于后剪枝因为验证集本身噪声大数据集大时后剪枝更稳定。统计检验的意义就在于如果p值大于0.05说明两者差异很可能来自随机性不要盲目下结论。另外不同数据类型也会影响结论——Iris这种连续特征多的数据CART自然占优因为ID3离散化会丢失信息而adult-stretch这类属性离散的数据两者差距会缩小。5. plotTree.py可视化与决策树调参陷阱最后聊聊plotTree.py。很多人手写决策树能在文本里打印嵌套字典但课程报告需要图形。plotTree.py用Matplotlib完成这件事它的核心不是绘图API而是如何给每个节点分配坐标否则任意深度的树会重叠。可视化的常用技巧是先递归计算树的深度和叶子数再按叶子数平铺横坐标深度决定纵坐标。下面这段是典型的递归布局代码def plot_node(ax, tree, x, y, parent_x, parent_y, features, plot_dict): if not isinstance(tree, dict): ax.text(x, y, tree, hacenter, vacenter, fontpropertiesSimHei) ax.annotate(, xy(x, y), xytext(parent_x, parent_y), arrowpropsdict(arrowstyle-)) return feature next(iter(tree.keys())) ax.text(x, y, feature, hacenter, vacenter, fontpropertiesSimHei) total_width plot_dict[total_width] depth plot_dict[depth] # 平分横坐标给各分支 for i, (value, subtree) in enumerate(tree[feature].items()): child_x x (i - (len(tree[feature]) - 1) / 2) * total_width / (2 ** depth) child_y y - 1 ax.plot([x, child_x], [y, child_y], k-) plot_node(ax, subtree, child_x, child_y, x, y, features, plot_dict)这里的plot_dict需要预先保存整棵树的叶子数来算total_width否则子节点会越画越挤。图中文字是中文Linux/macOS上经常出现方块我通常在脚本开头加import matplotlib as mpl mpl.rcParams[font.sans-serif] [SimHei, Noto Sans CJK SC, WenQuanYi Zen Hei] mpl.rcParams[axes.unicode_minus] False调参陷阱也不可忽视。第一个陷阱是最大深度很多人为了追求测试集精度把max_depth设得很大导致训练集过拟合。我在用CART_剪枝.py时会画一条深度从1到10的验证集精度曲线选择拐点。另一个陷阱是特征值里有缺失值本项目的数据集都没有缺失但一旦换成真实数据splitDataSet会把缺失值样本直接丢弃造成偏差。常见做法是给缺失特征单独放一个分支或者在计算熵时按权重分配但Decision_tree.py没有实现这个需要你自己加。最后一个实用技巧把剪枝前后的树可视化到同一张画布上用ax1和ax2并排显示你能非常直观地看到预剪枝把哪些分支砍掉了。这种对比图写进课程报告比一张精度表格有说服力得多。本文还有配套的精品资源点击获取