资讯动态

手写决策树:从信息增益到预剪枝后剪枝的Python实现

发布时间:2026/9/12 23:26:57 来源:尧图企业网站定制
简介围绕《机器学习》西瓜书第四章决策树这份资源提供基于信息熵与基尼指数的决策树算法Python实现适合正在阅读该书并希望动手实践的机器学习初学者、高校学生及研究人员。压缩包共9个文件含4个Python脚本和5个CSV数据集脚本覆盖决策树构建、CART剪枝与树的可视化绘图数据集包括西瓜数据集3.0、2.0、英文版西瓜数据集以及iris、adult-stretch等UCI数据可直接运行做对比实验。资源包仅16KB内容精炼已有10318人学习下载。通过研读代码可理清信息熵划分选择、基尼指数划分、预剪枝与后剪枝的完整实现逻辑并能针对多个数据集比较不同剪枝策略的决策树效果、进行统计显著性检验是配套书本章节实践的便捷代码参考。1. 从西瓜书第4章到可运行的Python决策树代码很多人学周志华《机器学习》第四章时卡在了明明看懂了信息增益的公式一合上书自己写代码就不知道从哪下手。更常见的做法是直接调sklearn的DecisionTreeClassifier把西瓜书当理论看看等面试被问到手撕决策树或者ID3和CART的区别就露馅。这篇博客就按西瓜书第4章的脉络用Python把决策树的划分选择、递归建树、预剪枝和后剪枝一步步实现出来并在西瓜数据集上跑通。新手能跟着代码复现熟手能重点看剪枝的实现细节和与sklearn参数的对应关系。2. 决策树划分选择信息增益、增益率与基尼指数的Python计算决策树的本质是递归地把数据集按属性划分让每个子集尽可能纯。西瓜书第4章给出了三种衡量纯度的准则信息增益ID3、增益率C4.5、基尼指数CART。很多人直接背公式但这里我更建议先把准则写成函数因为后续建树代码里所有属性选择都会调用它们。2.1 先用代码把信息熵算出来信息熵是划分选择的起点。西瓜书的定义是Ent(D) -sum(p_k * log2(p_k))p_k是第k类样本所占比例。写出来很直接import numpy as np from collections import Counter def calc_entropy(labels): 计算数据集D的信息熵labels是样本标签列表 counter Counter(labels) total len(labels) entropy 0.0 for count in counter.values(): p count / total entropy - p * np.log2(p) return entropy # 示例西瓜书表4.1前6条样本的正例/反例分布 labels [是, 是, 是, 是, 否, 否] print(calc_entropy(labels)) # 0.918这里有一个容易算错的地方log2(p)在p趋近0时结果趋近0但Python里直接算0 * np.log2(0)会得到nan所以要么先判断p为0就跳过该项要么直接用Counter只统计已出现的类别天然避开了零概率问题。上面这个写法用Counter遍历类别就是规避这个坑的常见做法。2.2 信息增益ID3怎么选属性信息增益的定义是父节点熵减去按属性划分后各子节点熵的加权和。西瓜书的公式是Gain(D, a) Ent(D) - sum(|Dv|/|D| * Ent(Dv))。实现时最需要注意的是按哪个特征取值分组def split_by_feature(X, y, feature_idx): 按第feature_idx列的特征值把数据集分组返回 {特征值: (子集X, 子集y)} groups {} for x, label in zip(X, y): value x[feature_idx] if value not in groups: groups[value] ([], []) groups[value][0].append(x) groups[value][1].append(label) return groups def calc_info_gain(X, y, feature_idx): 计算西瓜书Gain(D,a)父熵 - 加权子熵 base_entropy calc_entropy(y) groups split_by_feature(X, y, feature_idx) total len(y) weighted_entropy 0.0 for sub_X, sub_y in groups.values(): weight len(sub_y) / total weighted_entropy weight * calc_entropy(sub_y) return base_entropy - weighted_entropy参数说明feature_idx是当前候选属性的列下标split_by_feature用字典按特征值分组返回的子集在后续递归建树时直接传给下一层。ID3每轮选择信息增益最大的属性但它的缺陷是偏好取值多的属性——比如编号这个属性每个样本一个值算出来信息增益最大但完全没有泛化能力。这就是西瓜书紧接着引出增益率的原因。2.3 增益率C4.5为什么加惩罚项增益率在信息增益基础上除以固有值IV(a)IV(a) -sum(|Dv|/|D| * log2(|Dv|/|D|))属性的取值越多IV越大惩罚越重。但增益率又会反过来偏好取值少的属性所以C4.5的实际做法是先从候选属性里挑出信息增益高于平均水平的再从中选增益率最高的而不是直接选增益率最大。def calc_gain_ratio(X, y, feature_idx): 计算增益率信息增益 / 固有值IV gain calc_info_gain(X, y, feature_idx) groups split_by_feature(X, y, feature_idx) total len(y) iv 0.0 for sub_X, sub_y in groups.values(): weight len(sub_y) / total iv - weight * np.log2(weight) return gain / iv if iv 0 else 0.0注意iv为0的情况说明这个属性只有一个取值此时直接用除零保护返回0。实际写的时候建议参照C4.5的两阶段选择先算出所有候选属性的信息增益并求平均候选集只保留增益大于平均值的再在这个候选集里比较增益率。2.4 基尼指数CART的划分逻辑CART用的是基尼值Gini(D) 1 - sum(p_k^2)基尼指数则是按属性划分后各子集基尼值的加权和。CART选择使基尼指数最小的属性而且CART是二叉树所以对于多取值离散属性需要计算二分子集的组合这是它和ID3、C4.5最大的区别。def calc_gini(labels): 计算基尼值 Gini(D) counter Counter(labels) total len(labels) gini 1.0 for count in counter.values(): p count / total gini - p ** 2 return gini def calc_gini_index(X, y, feature_idx): 计算基尼指数按离散属性取值加权求和 groups split_by_feature(X, y, feature_idx) total len(y) gini_index 0.0 for sub_X, sub_y in groups.values(): weight len(sub_y) / total gini_index weight * calc_gini(sub_y) return gini_indexCART在每次划分时把数据集切成两个部分离散属性要遍历所有二分组合连续属性则用后文会讲的二分法找切分点。西瓜书第4章的表4.2对纹理等属性算过基尼指数可以拿那些数值核对上面代码的输出。下表把三种准则的优化方向和适用场景做了对比准则对应算法选择方向偏好常见适用场景信息增益ID3最大取值多的属性属性取值较均衡的小数据集增益率C4.5最大两阶段取值少的属性含大量离散属性的数据集基尼指数CART最小计算快天然二叉树sklearn默认工业界最常用3. 手写决策树递归建树、预剪枝与后剪枝的代码实现上一章的三个准则解决了选哪个属性的问题但一棵树还要解决什么时候停叶子标什么标签以及怎么防止过拟合。这才是西瓜书第四章代码实现里最容易被忽略的部分。3.1 节点结构与递归终止条件先把树的节点定义成一个简单的类或者直接用字典。我习惯用类因为后剪枝时要递归修改子树class DecisionNode: 决策树节点leafTrue表示叶子节点 def __init__(self, feature_idxNone, valueNone, childrenNone, labelNone, leafFalse): self.feature_idx feature_idx # 当前节点划分用的属性下标 self.value value # 父节点到本节点经过的取值 self.children children or {} # {特征值: 子节点} self.label label # 叶子节点的类别 self.leaf leaf # 是否为叶子递归建树的终止条件有三个对应西瓜书第4章的递归返回三种情形第一当前子集样本全部属于同一类别直接返回叶子第二候选属性集为空返回叶并标为样本数最多的类别第三当前子集在剩余属性上取值完全相同无法继续划分同样返回多数类叶子。多数类可以用Counter一行搞定def majority_label(labels): 返回样本数最多的类别用于无法继续划分时的叶子标记 return Counter(labels).most_common(1)[0][0]3.2 用信息增益选出最优划分属性建树的核心逻辑是每轮对当前子集计算所有候选属性的信息增益选出最大那个按它的取值把数据集分组然后对每个子集递归建树。代码里有个细节容易被忽略——每深入一层候选属性集合就要剔除已经用过的属性否则同一个属性会被反复用来划分。def build_tree(X, y, feature_names, remain_features): 递归建树remain_features是当前剩余可用的属性下标列表 if len(set(y)) 1: return DecisionNode(labely[0], leafTrue) if not remain_features or len(set(tuple(x[i] for i in remain_features) for x in X)) 1: return DecisionNode(labelmajority_label(y), leafTrue) best_idx None best_gain -1 for idx in remain_features: gain calc_info_gain(X, y, idx) if gain best_gain: best_gain gain best_idx idx node DecisionNode(feature_idxbest_idx) groups split_by_feature(X, y, best_idx) rest [i for i in remain_features if i ! best_idx] for value, (sub_X, sub_y) in groups.items(): node.children[value] build_tree(sub_X, sub_y, feature_names, rest) return node参数说明feature_names用于最后打印树结构时把下标转成属性名remain_features是列表但每层递归都会生成新列表rest所以不会污染上层。第二个终止条件里用set判断剩余属性上取值完全相同这是课本没细讲但实现时必须处理的边界情况不然递归会无限进行下去。3.3 预剪枝在分裂前判断泛化能力西瓜书把剪枝分成了预剪枝和后剪枝它们的区别在于评估时机。预剪枝在每次决定是否分裂时先用验证集评估如果分裂后的验证集精度比不分裂时高才允许分裂。实现时把构建完整树改成边构建边用验证集打分def predict(node, sample): 按特征取值递归查找叶子返回预测类别 if node.leaf: return node.label value sample[node.feature_idx] child node.children.get(value) if child is None: # 训练时没见过的取值直接返回当前节点多数类 return majority_label(get_all_labels(node)) return predict(child, sample) def calc_acc(tree, X_val, y_val): 在验证集上计算精度 correct sum(1 for x, y in zip(X_val, y_val) if predict(tree, x) y) return correct / len(y_val)预剪枝的评估逻辑是构造当前候选节点后先假设它是叶子用当前子集多数类标算验证集精度再假设它完整展开再算一次精度只有展开后精度更高才真正分裂。预剪枝的缺点是贪心它看不到两层之后的分裂可能带来的提升所以精度可能不如后剪枝。上面的predict里有一个容易忽略的分支验证集样本出现训练时没见过的属性取值直接用dict.get返回None此时需要一个兜底策略常见做法是回退到该节点的多数类标签。3.4 后剪枝先生成整棵树再自底向上替换后剪枝先按3.2节生成完整树然后从底向上遍历每个内部节点判断把这个子树替换成叶子用该节点子集的多数类标能不能提升验证集精度。实现用递归自底向上非常自然def post_prune(node, X_val, y_val): 自底向上后剪枝返回剪枝后的节点 if node.leaf: return node # 先递归处理子树 groups split_by_feature(X_val, y_val, node.feature_idx) for value, (sub_X, sub_y) in groups.items(): if value in node.children: node.children[value] post_prune(node.children[value], sub_X, sub_y) # 尝试把当前子树替换成叶子节点 labels [y for _, y in zip(X_val, y_val)] leaf_label majority_label(labels) acc_before calc_acc(node, X_val, y_val) acc_if_leaf sum(1 for y in y_val if y leaf_label) / len(y_val) if acc_if_leaf acc_before: return DecisionNode(labelleaf_label, leafTrue) return node这里用精度不低于原树作为剪枝条件是西瓜书后剪枝的经典做法宁可保持精度相当也换成更小的树换取泛化能力和可解释性。注意acc_before和acc_if_leaf的对比必须用同一个验证集如果把训练集混进去剪枝几乎不会发生因为整棵树在训练集上精度永远是1。两种剪枝方式在各个维度上的差异如下表剪枝方式评估时机优点缺点预剪枝分裂前训练开销小树小贪心可能欠拟合后剪枝树生成后保留更多结构精度通常更高训练开销大先建完整树4. 在西瓜数据集上跑通决策树并对比sklearn的参数设置理论代码写完了但光有函数还不够得在西瓜书表4.1那个经典数据集上完整跑一遍再用sklearn的DecisionTreeClassifier做对照这样才知道手写版本和工业实现差在哪。4.1 构造西瓜数据集并训练自己的决策树西瓜书表4.1是17条数据6个离散属性色泽、根蒂、敲声、纹理、脐部、触感加一个类别好瓜。我用字典直接硬编码前几条做个最小可运行示例# 简化版只取色泽、根蒂、敲声三个属性前10条样本 feature_names [色泽, 根蒂, 敲声] X [ [青绿, 蜷缩, 浊响], [乌黑, 蜷缩, 沉闷], [乌黑, 蜷缩, 浊响], [青绿, 硬挺, 清脆], [浅白, 蜷缩, 浊响], [青绿, 稍蜷, 浊响], [乌黑, 稍蜷, 浊响], [乌黑, 硬挺, 清脆], [浅白, 稍蜷, 沉闷], [青绿, 硬挺, 清脆], ] y [是, 是, 是, 否, 否, 是, 是, 否, 否, 否] tree build_tree(X, y, feature_names, list(range(len(feature_names)))) print(预测结果:, predict(tree, [乌黑, 蜷缩, 浊响]))跑通后可以对照西瓜书图4.4那棵ID3树的结构检查根节点选的是纹理而不是色泽或根蒂因为完整17条样本里纹理的信息增益最大。注意如果用我上面简化版数据选出的根节点可能不同这正常判断正确性的标准是信息增益计算有没有出错而不是树长得和书上一样。4.2 sklearn的DecisionTreeClassifier参数对照把同样的数据喂给sklearn只需要三行from sklearn.tree import DecisionTreeClassifier # sklearn的criterion支持gini和entropy对应CART基尼指数和信息增益二分版本 clf DecisionTreeClassifier(criterionentropy, random_state42) clf.fit(X_encoded, y)但直接fit会报错因为sklearn不接受字符串特征必须先用LabelEncoder或OrdinalEncoder编码。这里就体现出手写版本和工业实现的第一个差异sklearn的CART只支持数值输入且默认是基尼指数手写版本用字典分组天然支持离散字符串。第二个差异是sklearn不实现ID3那种多叉树——criterionentropy只是把信息增益用在二分上树的结构仍然是二叉树。提示用sklearn处理离散特征时建议用OrdinalEncoder统一编码训练集和测试集不要对训练集和测试集分别fit否则相同的字符串会编码成不同的整数树的分裂完全错乱。下表是我整理的手写版本与sklearn参数对照关系手写版本概念sklearn参数说明划分准则信息增益criterionentropysklearn的entropy计算是二分多叉树的混合划分准则基尼指数criteriongini默认值CART标准预剪枝最小样本数min_samples_leaf叶子至少含多少样本值越大树越小预剪枝最大深度max_depth限制树深度最常用的防过拟合参数预剪枝最小不纯度下降min_impurity_decrease对应精度提升才分裂的阈值版本随机性控制random_state固定后结果可复现4.3 调max_depth和min_samples_leaf避免过拟合决策树最容易过拟合完全生长的树在训练集上精度100%验证集上可能只剩70%。我的经验是先固定random_state然后网格搜索max_depth和min_samples_leaf的组合from sklearn.model_selection import GridSearchCV param_grid { max_depth: [2, 3, 4, 5, None], min_samples_leaf: [1, 2, 5, 10], } clf DecisionTreeClassifier(criteriongini, random_state42) grid GridSearchCV(clf, param_grid, cv5, scoringaccuracy) grid.fit(X_encoded, y) print(grid.best_params_) # 输出最优参数组合 print(grid.best_score_)逻辑说明max_depth2时树只有根节点加一层孩子几乎不会过拟合但可能欠拟合min_samples_leaf10强制每个叶子至少10个样本。这两者组合起来能有效抑制决策树对噪声样本的拟合。注意GridSearchCV的cv5用的是交叉验证比单次划分验证集更能反映泛化能力但西瓜书第4章讲的预剪枝用的是单独的验证集二者评估口径不一样做实验报告时要说明清楚。X_encoded必须用同一套编码器对训练和测试数据做变换否则特征取值错位树的分裂根本没意义。5. 决策树的连续值与缺失值处理两个值得单独实现的细节西瓜书第4章的4.4节讲连续值和缺失值处理很多代码实现都会跳过这两块但实际工作中离散属性占比其实不高连续值处理反而更常用。5.1 连续属性的二分法离散化西瓜书的做法是先把连续属性在样本上的取值排序然后取相邻取值的中位点作为候选划分点逐个计算信息增益选最优划分点。示例代码def calc_best_split_point(X, y, feature_idx): 连续属性二分返回(最优划分点, 最大信息增益) values sorted(set(x[feature_idx] for x in X)) best_point, best_gain None, -1 for i in range(len(values) - 1): mid (values[i] values[i 1]) / 2 # 把样本按 mid 和 mid 分成两组 left_y [y[j] for j, x in enumerate(X) if x[feature_idx] mid] right_y [y[j] for j, x in enumerate(X) if x[feature_idx] mid] gain calc_entropy(y) - (len(left_y) / len(y) * calc_entropy(left_y) len(right_y) / len(y) * calc_entropy(right_y)) if gain best_gain: best_gain, best_point gain, mid return best_point, best_gain参数说明候选划分点是每两个相邻值的均值所以17个样本最多产生16个候选点相对于穷举所有实数阈值这个复杂度完全可以接受。注意连续属性可以在同一棵树的多个分支上重复使用因为每次都是二分后续子节点里该属性仍然可以继续划分出更细的边界——这是连续属性和离散属性在递归时最大的区别。5.2 缺失值处理的两个要义缺失值处理有两个问题要解决一是划分属性时怎么用带缺失值的样本计算信息增益二是样本被划分到哪个分支。西瓜书的做法是给每个样本引入一个权重w初始为1计算信息增益时只统计属性未缺失的样本但要按缺失比例放缩权重。实现上最省事的做法是def calc_info_gain_with_missing(X, y, feature_idx, sample_weight): 带缺失值的信息增益sample_weight是每个样本的权重列表 # 找出该属性未缺失的样本下标 valid_idx [i for i, x in enumerate(X) if x[feature_idx] is not None] total_w sum(sample_weight[i] for i in valid_idx) # 权重版的熵需要自行实现p_k 该类权重和 / 总权重 base_entropy calc_weighted_entropy( [y[i] for i in valid_idx], [sample_weight[i] for i in valid_idx] ) # 信息增益按缺失样本比例放缩 return base_entropy * total_w / sum(sample_weight)这里calc_weighted_entropy需要自己实现权重版的熵核心是每个类别的p_k变成该类权重和 / 总权重。缺失样本不作为某个特定分支的成员而是以权重比例同时进入所有分支——这就是西瓜书说的权重调整。真要动手做建议先用无缺失的数据验证权重版熵和普通版结果一致再加缺失样本排查会容易很多。验证方法可以这样设计把西瓜数据集中某几个值手动置为None分别用完整版和缺失版实现跑一遍对比两者选出的根节点是否一致。跑完这组对比如果根节点选择和完整数据一致缺失值分支的权重计算基本可以放心不一致时优先检查样本权重有没有在传给子节点前被重置。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价