ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

手写决策树:从信息熵到CART剪枝的Python实现与避坑指南

手写决策树:从信息熵到CART剪枝的Python实现与避坑指南 简介面向西瓜书《机器学习》第四章决策树的编程练习这套Python代码包提供了完整的实现方案适合正在学习机器学习基础、需要完成4.3至4.6课后实验的读者也便于教师或研究者快速复现决策树对比实验。代码基于信息熵与基尼指数两种划分准则分别构建决策树并实现了预剪枝、后剪枝及未剪枝三种模式的对照附带西瓜数据集2.0、3.0以及4个UCI数据集的CSV文件支持直接运行验证可观察不同剪枝策略对模型精度和过拟合倾向的影响。资源共9个文件其中5个CSV数据集负责提供实验输入4个Python脚本分别承担决策树构建、剪枝处理与树结构可视化等功能压缩包仅16KB轻量且易于部署。目前已有10322人学习下载源码结构清晰、注释到位能帮助读者逐行理解决策树递归划分、信息熵与基尼指数的计算逻辑以及预剪枝、后剪枝在不同数据集上的效果差异从而扎实掌握第四章核心算法与实验比较方法。1. 决策树 Python 实现包为什么从西瓜书第四章开始动手学《机器学习》最容易崩的地方就在第四章前面还在讲线性模型和距离度量突然冒出一堆信息熵、信息增益、基尼指数公式看懂了打开 Python 却发现连数据从哪读都不知道。这份资源解决的就是这个卡点——它是西瓜书第四章决策树的完整代码实现包含西瓜数据集2.0、西瓜数据集3.0、iris、watermelon3_0_En、adult-stretch 共 4 份数据集以及 Decision_tree.py信息熵划分版、CART.py基尼指数版、CART_剪枝.py预剪枝后剪枝和 plotTree.py 可视化脚本。适合三类人正在啃周志华《机器学习》决策树章节的学生、要交机器学习课程设计的人、以及想彻底搞懂 sklearn 里 DecisionTreeClassifier 底层到底怎么切分数据的开发者。看完这份东西你能从零手写一棵决策树并且知道模型过拟合时到底是该限制深度还是做后剪枝。2. 信息熵决策树复现西瓜书4.3为西瓜数据集3.0生成决策树2.1 信息熵与信息增益为什么选它而不是错误率决策树划分的核心问题只有一个给定一堆样本选哪个属性先分。西瓜书 4.2 节给出的答案是按信息增益来挑而信息增益的两个基础概念是信息和熵。信息量的定义是I(x) -log2 p(x)一个事件概率越小包含的信息量越大熵Ent(D) -sum(p_k * log2 p_k)则是信息量的期望也就是对数据集 D 的纯度打一个总分。注意熵和错误率的一个关键差别熵对样本分布更敏感。比如 100 个样本里好瓜和坏瓜各 50 个熵是 1.0如果分成 90 个和 10 个熵降到约 0.469。而错误率只会从 0.5 变成 0.1看不出分布越来越偏的过程。这就是为什么信息增益比错误率更适合做划分依据——它能在属性取多个值、每个取值下纯度变化很大时依然给出细粒度的区分。西瓜数据集3.0 里有 17 条样本8 个好瓜、9 个坏瓜标签分布很接近。先算根节点的熵Ent(D) -(8/17)*log2(8/17) - (9/17)*log2(9/17) ≈ 0.998。这就是 baseline。接下来每个属性算加权熵两者相减就是信息增益。哪个属性增益最大就用它做当前节点的划分属性。整个过程是每进一层换一批样本子集重算一次熵这是 ID3 算法的核心逻辑。2.2 Decision_tree.py 的核心特征选择与递归建树数据集加载直接用 pandas 读 CSV但有一个前置操作把 DataFrame 转成 Python 的 list of list最后一列是标签。这样做是为了方便后续递归时对每一行的特征列做切片操作。下面这段是原代码里最核心的熵计算函数我拆开讲import numpy as np def calc_shannon_ent(data_set): num len(data_set) if num 0: return 0.0 label_counts {} for feat_vec in data_set: current_label feat_vec[-1] label_counts[current_label] label_counts.get(current_label, 0) 1 shannon_ent 0.0 for key in label_counts: prob float(label_counts[key]) / num shannon_ent - prob * np.log2(prob) return shannon_ent参数说明data_set的每一行是一个样本行末元素是标签。代码先统计各类别数量再按熵公式累加。用np.log2而不是math.log因为后续要和 numpy 数组混用类型转换上更省事。注意开头对空数据集返回 0这个保护很重要后面递归到叶子节点的子集时经常出现为空的情况。实际跑决策树时熵计算会被调用无数次这个函数不要写得过于花哨保持简单和可读性优先。选特征时代码遍历每个特征列先把该列所有取值用set去重再对每个取值做一次split_data_set划分算子集的加权熵。这里有一个容易被忽略的点split_data_set里要保留好原始特征列的位置因为递归下一层时当前用掉的特征要从特征列表里删掉。看建树函数就明白了def create_tree(data_set, features): class_list [row[-1] for row in data_set] if class_list.count(class_list[0]) len(class_list): return class_list[0] if len(data_set[0]) 1: return majority_vote(class_list) best_feat_idx chose_best_feature_to_split(data_set) best_feat_name features[best_feat_idx] tree {best_feat_name: {}} feat_values set([row[best_feat_idx] for row in data_set]) for value in feat_values: sub_data split_data_set(data_set, best_feat_idx, value) sub_features features[:best_feat_idx] features[best_feat_idx 1:] tree[best_feat_name][value] create_tree(sub_data, sub_features) return tree递归终止有三个条件缺一不可。第一个是当前子集样本全属于同一类直接返回该类标签第二个是特征用完了但标签还不纯此时用多数投票第三个是data_set[0]长度为 1说明只剩标签列。这三个条件保护了递归不会无限循环。sub_features features[:best_feat_idx] features[best_feat_idx1:]是删掉已用特征的标准写法我第一次写时直接用features.remove(best_feat_name)结果 Python 的值传递和引用问题导致上一层列表也被改了树直接建歪。切片重组才是安全做法。2.3 连续属性怎么处理密度与含糖率的二分法西瓜数据集3.0 里除了 6 个离散属性还有密度和含糖率两个连续属性。如果直接套用上面的set([row[i] for row in data_set])每个样本的密度值都不一样划分出来的子集全是单样本点信息增益直接爆炸。常见做法是用 C4.5 的二分法离散化先把连续属性值排序相邻值取中点作为候选切分点然后逐个计算信息增益选增益最大的那个切分阈值。这段逻辑在 Decision_tree.py 里是单独拆出来的代码类似这样def split_by_threshold(data_set, feature_index, threshold): left [row for row in data_set if row[feature_index] threshold] right [row for row in data_set if row[feature_index] threshold] return left, right def choose_best_threshold(data_set, feature_index): values sorted(set(row[feature_index] for row in data_set)) best_gain 0.0 best_threshold None base_ent calc_shannon_ent(data_set) for i in range(len(values) - 1): threshold (values[i] values[i 1]) / 2.0 left, right split_by_threshold(data_set, feature_index, threshold) new_ent len(left) / len(data_set) * calc_shannon_ent(left) \ len(right) / len(data_set) * calc_shannon_ent(right) gain base_ent - new_ent if gain best_gain: best_gain gain best_threshold threshold return best_threshold, best_gain参数说明values排序后取相邻中点阈值点是唯一候选位置符合西瓜书说的把候选划分点集合中使信息增益最大的作为最优划分点。运行这段时要注意连续属性一旦被选中在当前分支的后续递归中仍然可以再次参与划分不同于离散属性用一次就删。之前有同学把连续属性也直接从特征列表里删掉结果树退化成一串按密度排行的线性序列泛化能力很差。这是 ID3 扩展到连续属性时最容易踩的坑。2.4 画图plotTree.py 把树结构变成可提交的图Decision_tree.py 生成的树本质是一个嵌套字典直接打印出来层级一多就很难看清比如{纹理: {清晰: {根蒂: {蜷缩: 好瓜, 稍蜷: {...}}}, 稍糊: 坏瓜}}。plotTree.py 做的事就是把这个嵌套字典画成树状图。核心逻辑有两个一是递归计算每个叶子节点的坐标位置二是用 matplotlib 的 annotate 画节点框和连线。由于树深度不定叶子数量不定坐标归一化是难点。plotTree.py 里用了一个全局变量来记录叶子节点数再按比例分配 x 坐标这样画出来的树不会重叠。运行前确保装了 matplotlib并且系统里有中文字体否则树节点上的色泽纹理会变成方块具体解决见第 5 章。3. CART 与基尼指数换成二叉树后代码差在哪3.1 基尼指数为什么是 CART 的选择信息增益有个隐藏的毛病它天然偏向取值多的属性。比如编号这个属性每个取值只有一条样本划分后每个子集的熵都是 0信息增益直接达到最大但结果是棵废树。CART 用基尼指数替代信息增益从两个层面规避了这个问题。第一基尼值Gini(D) 1 - sum(p_k^2)衡量的是从数据集里随机抽两个样本、类别不一致的概率计算里没有 log纯加减和平方速度比熵快很多第二CART 强制生成二叉树即使属性有多个离散取值也只会按是某个值/不是某个值来划分天然限制了一棵树的宽度。对比一组数据二分类场景下当正负样本各半时熵是 1.0基尼值是 0.5当样本全部同类时两者都是 0。基尼值的变化曲线比熵更陡意味着它更容易在纯度提升到一定程度后迅速收敛。实际使用中两者的选特征结果往往高度一致但基尼指数计算开销小、且对多取值属性的偏斜更温和这就是 sklearn 里criteriongini成为默认值的原因。如果在西瓜数据集上对比 ID3 和 CART 的树结构会发现 ID3 偏向用取值多的属性先分而 CART 会优先找区分度更强、但取值不一定多的切分点。3.2 CART 建树代码特征与阈值同时搜索CART 的分类树在建树时每次划分都要同时决定两件事选哪个特征、选什么切分点。对离散属性切分点是取值等于某类对连续属性切分点是某个阈值。CART.py 把这两件事统一成了同一个搜索逻辑核心代码如下def calc_gini(data_set): label_counts {} for row in data_set: label row[-1] label_counts[label] label_counts.get(label, 0) 1 gini 1.0 for key in label_counts: prob label_counts[key] / len(data_set) gini - prob * prob return gini def choose_best_split(data_set): num_features len(data_set[0]) - 1 best_gini float(inf) best_feature -1 best_value None for i in range(num_features): values set(row[i] for row in data_set) for value in values: left [row for row in data_set if row[i] value] right [row for row in data_set if row[i] ! value] if not left or not right: continue gini_index len(left) / len(data_set) * calc_gini(left) \ len(right) / len(data_set) * calc_gini(right) if gini_index best_gini: best_gini gini_index best_feature i best_value value return best_feature, best_value参数说明这里用left表示特征取值等于 value 的样本right表示不等于 value 的CART 永远是二路划分。best_gini初始化设为无穷大因为我们要找的是最小基尼指数。注意values set(...)与ID3 的区别——CART 不是枚举所有离散取值来生成多分叉而是把每个取值当作一次是/否判断。连续属性在 CART.py 里的处理方式和 ID3 类似但切分点是 value / value的二元划分最终选基尼指数最小的那个 value。建树函数create_tree与 ID3 版有三处不同。第一终止条件除了类别全纯、特征耗尽之外还要加一个if len(data_set) 0的防护因为某个切分可能把数据全分到一边第二返回的树结构里每个节点的分支只有两个键比如{纹理: {是: 子树, 否: 子树}}第三由于是二叉结构最大深度往往比多叉树更深同样 17 条样本的西瓜数据集ID3 可能 3 层就到底CART 可能要 5 到 6 层所以剪枝对 CART 尤其重要。这就是为什么有人在西瓜数据集上跑 CART 后发现树形和书上的 ID3 树对不上两棵树的生长逻辑本来就不同。3.3 和 sklearn 的 DecisionTreeClassifier 对照着看手写代码最大的价值不是替代 sklearn而是知道每个参数背后在做什么。这里给一张对照表方便调试时互相对照手写实现sklearn 参数对应关系calc_gini 选最小基尼指数criteriongini节点划分依据choose_best_split 遍历所有特征和取值splitterbest最优切分搜索递归直到类别纯净max_depthNone不限制时会长满多数投票返回叶子标签class_weight 默认 None类别均衡时的叶节点决策CART_剪枝.py 的后剪枝ccp_alpha代价复杂度剪枝参数调试建议把 sklearn 的 DecisionTreeClassifier 用max_depth3, random_state0训练一棵树然后把 plotTree.py 画图和 sklearn 自带决策树画图并排对比。如果手写代码的树结构在浅层和 sklearn 不一致优先检查两个点一是数据集是否做了相同的编码方式比如把青绿映射成 0、1、2二是熵或基尼计算时是否把len(data_set)写成了len(label_counts)。这两处是手写版和 sklearn 结果对不齐的常见来源。4. 预剪枝与后剪枝CART_剪枝.py 的对比实验4.1 决策树为什么必须剪枝决策树是典型的训练集越高分、测试集越容易翻车的模型。原因很简单一棵树如果一路长到叶子节点都纯为止它记住的是训练样本的个体特征而不是类别的分布规律。西瓜数据集2.0 只有 17 条样本6 个离散属性。不加任何限制地长树完全可以做到所有训练样本全部分对但泛化到新样本时表现很差。剪枝的核心思想就一句话用一部分数据当裁判树长到某个程度后问它一句这层还值不值得继续分。预剪枝和后剪枝的区别在于是边建树边问还是建完再回头改。预剪枝在每次准备划分当前节点时先用测试集算一次划分前 vs 划分后的准确率如果划分后准确率没有提升就把当前节点直接变成叶节点不再往下建。后剪枝则先让树完整长出来然后自底向上考察每个内部节点——如果把以它为根的子树整体替换成一个叶节点测试集准确率不降低就真的剪掉。预剪枝快但有欠拟合风险因为可能刚好处在一个当前层没提升、下一层会有大提升的位置被提前掐断后剪枝慢一点但通常更稳因为它基于一棵完整的树做局部修正。西瓜书 4.3 节的结论在实际跑的时候表现很稳定后剪枝保留的信息更多尤其在样本量小的时候。4.2 预剪枝实现先测试后划分CART_剪枝.py 的预剪枝逻辑不是直接改 create_tree而是在递归建树函数里加了一个判断节点。核心思路是这样的先用多数投票算法算出当前数据集的主类标签记为major_label假设当前节点不再划分用这个主类去预测测试集中的对应样本再把当前数据集按最优特征划分后用划分后的子树去预测测试集两边准确率做比较def pre_prune_tree(data_set, test_set, features, depth0, max_depth3): class_list [row[-1] for row in data_set] if len(set(class_list)) 1: return class_list[0] if depth max_depth: return majority_vote(class_list) if len(data_set[0]) 1: return majority_vote(class_list) best_feat_idx, best_val choose_best_split(data_set) if best_feat_idx -1: return majority_vote(class_list) acc_before calc_test_acc(majority_vote(class_list), test_set) acc_after calc_test_acc_with_split( data_set, test_set, best_feat_idx, best_val) if acc_after acc_before: return majority_vote(class_list) left [row for row in data_set if row[best_feat_idx] best_val] right [row for row in data_set if row[best_feat_idx] ! best_val] tree {features[best_feat_idx]: {}} tree[features[best_feat_idx]][是] pre_prune_tree( left, test_set, features, depth 1, max_depth) tree[features[best_feat_idx]][否] pre_prune_tree( right, test_set, features, depth 1, max_depth) return tree参数说明max_depth3是预剪枝的第一道防线先限制树的最大深度防止极端情况。acc_before对应的是当前节点不划分、直接按多数类预测的准确率acc_after对应继续划分后预测的准确率。只要划分后准确率没有严格提升就直接返回主类标签。注意这里用的是而不是即划分后准确率持平也剪枝——这是预剪枝偏保守的常见设定宁可少分一层不要无效分层。calc_test_acc_with_split是辅助函数负责按切分点把测试集分成两组分别交给左右子树预测后汇总准确率。实际运行时会发现预剪枝产出的树往往只有两三层在西瓜数据集2.0 上非常常见。4.3 后剪枝实现从底往上把子树换成一个叶后剪枝的代码结构比预剪枝复杂一些因为它要先建一棵完整树再做二次遍历。CART_剪枝.py 里用一种比较直接的方式实现定义post_prune(tree, test_set)先递归找到所有内部节点然后逐个尝试剪枝。判断逻辑如下def post_prune(tree, test_set, features): if not isinstance(tree, dict): return tree, calc_tree_acc(tree, test_set) key_name list(tree.keys())[0] sub_tree tree[key_name] for branch in sub_tree: sub_tree[branch], _ post_prune(sub_tree[branch], test_set, features) acc_before calc_tree_acc(tree, test_set) majority_label get_majority_label_from_subtree(sub_tree) acc_after calc_single_label_acc(majority_label, test_set) if acc_after acc_before: return majority_label, acc_after return tree, acc_before这段代码的关键在于post_prune的返回值设计它同时返回剪枝后的树结构和对应准确率让上层节点在决定要不要剪掉自己的子树时有数据可参考。calc_tree_acc会递归地把测试集样本一路送到叶子节点比对标签后计算准确率。注意后剪枝的判定条件用的与预剪枝的恰好对称——后剪枝允许准确率持平就剪因为反正已经有一棵完整树兜底局部持平不影响上层判断但能显著简化树结构。在西瓜数据集2.0 上做三棵树的对比实验结果趋势很稳定未剪枝树的训练集准确率接近 100%但测试集准确率波动极大换一组划分方式可能从 70% 掉到 50%预剪枝树测试准确率高于未剪枝但树太矮遇到测试集分布稍有偏移时表现不稳定后剪枝树测试准确率最高且树规模适中。需要说明的是这个对比结果依赖测试集划分方式不要拿某一次运行的数字写死进实验报告正确做法是换 5 个随机种子跑 5 次记录准确率均值。4.4 三种树的比较与记录方法做课程设计或实验报告时建议用一张表记录每次运行的结果结构如下实验组训练集准确率测试集准确率树的叶子数树深度未剪枝高接近100%波动大多深预剪枝中等较稳定少浅后剪枝较高最稳定中中记录时三个要点一是固定随机种子数据集划分用random.seed(42)这类固定值否则每次跑的准确率对不上二是准确率至少保留到小数点后两位用round(acc, 4)输出避免浮点数误差干扰比较三是同一份测试集上比较才有意义——剪枝判断用的测试集和最终评测用的测试集应该是同一个划分不要混合使用。我在跑实验时吃过这个亏训练时用了一份随机测试集做剪枝验证提交前又换了一份测试集评估结果预剪枝的准确率从 0.78 掉到 0.61查了半天才发现是测试集不一致。5. 避坑指南决策树实现中五个改到半夜的问题5.1 CSV 中文列名乱码数据全变 NaN现象pd.read_csv(西瓜数据集2.0.csv)运行后列名变成一堆乱码或者打印出来全是NaN但用 Excel 打开 CSV 文件完全正常。原因数据集文件是 GBK 或 ANSI 编码保存的pandas 默认用 UTF-8 解码遇到中文字段直接失败。解决读取时显式指定编码pd.read_csv(西瓜数据集2.0.csv, encodinggbk)。如果还报错就换encodinggb18030GB18030 是 GBK 的超集兼容性更好。注意转换完 DataFrame 后检查一次df.head()确认色泽根蒂这些列名正常显示再开始取数。这个小动作能省掉后面所有神秘的为什么特征全是一样的值的排查时间。5.2 数据集里混入编号列树长得极深还过拟合现象对 iris 数据集建树时树深度一路涨到 30 层以上而且训练准确率 100%测试准确率却只有 60%。原因iris.csv 或者自己整理的数据集里带了一列样本编号比如 1 到 150 的序号决策树把编号当成特征列参与划分。编号每个取值都不同按编号切分后每个子集只有一条样本信息增益和基尼指数都会认为这是最佳划分。树会被编号优先的结构带偏泛化能力几乎为零。解决建树前先检查特征列范围num_features len(data_set[0]) - 1然后把编号列从特征里删除。更稳妥的写法是读 CSV 时直接df.drop([编号], axis1)。我一般会有个习惯先打印data_set[0]看第一行结构确认列数和数据内容匹配再做特征选择。5.3 递归深度超限RecursionError 还是数据没清理干净现象数据量一上来比如 iris 的 150 条样本运行建树函数直接抛RecursionError: maximum recursion depth exceeded。原因有两类。一类是树的深度真的很大比如没有剪枝、且连续属性反复参与划分时树深可以超过 Python 默认的 1000 层递归限制另一类是数据里有循环引用比如用features列表时没有做切片拷贝导致已删除的特征又出现在下一层的特征列表里树在同一个特征上反复切分。解决先看树结构打印确认是否出现同一特征在两个连续层级都被选中的情况。若是检查sub_features是否用了切片语法而不是原地remove。若树本身合理但深度大在代码开头加import sys; sys.setrecursionlimit(10000)同时加max_depth参数兜底。真正细致的做法是把两者一起做总递归限制调到 10000建树函数里再加深度参数预设为None剪枝时传入具体深度。5.4 matplotlib 画图中文全是方块现象plotTree.py 跑完树结构画出来了但所有中文标签都显示成空心方块。原因matplotlib 默认字体是 DejaVu Sans不支持中文字形。系统里有中文字体比如微软雅黑、SimHei也不会自动切换。解决在画图代码最前面加全局设置import matplotlib.pyplot as plt plt.rcParams[font.sans-serif] [SimHei] plt.rcParams[axes.unicode_minus] Falseaxes.unicode_minus False是防止负号显示成方块画坐标轴时会用到。如果用了 SimHei 还报找不到字体说明系统没装该字体换成[Microsoft YaHei]或[Noto Sans CJK SC]并检查matplotlib.font_manager里是否注册了该字体。更好的做法是脚本启动时动态检测可用字体但课程设计场景下直接指定系统已装字体就够了。5.5 UCI 数据集格式不统一读进来后标签列位置不对现象把 adult-stretch.csv 和 watermelon3_0_En.csv 一起丢进同一个建树函数一个能跑另一个直接报IndexError: list index out of range。原因UCI 数据集和西瓜数据集格式差别很大。adult-stretch.csv 的标签列可能在最后一列但列名没有固定有的数据集类别是字符串如 good/bad有的只有两个属性特征数量不一致。按西瓜书的代码默认最后一列是标签来取数自然会出问题。解决写一个统一的数据加载函数把所有数据集转成同一格式。读取后先打印df.shape和df.columns然后用features list(df.columns[:-1])显式提取特征列名data_set df.values[:, :-1].tolist()提取数据标签单独存。最后用assert len(data_set[0]) len(features) 1做一次断言确保特征数和数据维度一致再进建树函数。这个断言能拦下所有因为列数不对导致的隐蔽错误。6. 把实验做严谨UCI 数据集、交叉验证与显著性检验6.1 四个数据集统一灌进同一套建树代码西瓜书 4.6 要求用 4 个 UCI 数据集对两种算法、三种剪枝方式做实验比较。实操时先把数据集整理成统一结构。我一般会写一个通用的加载函数把 CSV 读进来之后统一处理字符串和标签import pandas as pd def load_uni_form(path, label_colNone): df pd.read_csv(path, encodinggbk) if label_col is None: label_col df.columns[-1] y df[label_col].map(lambda v: 1 if v good else (0 if v bad else v)) features [c for c in df.columns if c ! label_col] X df[features].values.tolist() data_set [row [y.iloc[i]] for i, row in enumerate(X)] return data_set, features参数说明label_col默认取最后一列map函数把字符串标签转成 0/1 数值后续熵和基尼计算只认数值标签。这一步处理的是 adult-stretch 这类数据集的good/bad标签。iris 数据集标签是三个字符串类别转成 0、1、2 即可。注意处理完后打印一次每个类别的样本数确认没有某类样本数量为 0否则建树时某个分支的子集为空会直接抛异常。6.2 交叉验证配对 t 检验别再单次跑完就下结论单次划分训练集和测试集来比较两个算法结论不靠谱。因为划分的随机性足以让准确率差 10 个百分点。正确做法是 k 折交叉验证把同一份数据切 k 份每次留一份做测试其余训练记录两种算法在每一折上的准确率最后用配对 t 检验判断差异是否显著。代码是最后一步import random import numpy as np from scipy import stats def k_fold_evaluate(data_set, k5, seed42): random.seed(seed) indices list(range(len(data_set))) random.shuffle(indices) folds np.array_split(indices, k) acc_id3, acc_cart [], [] for i in range(k): test_idx folds[i] train_idx [j for idx in folds[:i] folds[i 1:] for j in idx] train_set [data_set[j] for j in train_idx] test_set [data_set[j] for j in test_idx] # 分别跑 ID3 和 CART 建树、预测 acc_id3.append(evaluate_tree(build_id3_tree(train_set), test_set)) acc_cart.append(evaluate_tree(build_cart_tree(train_set), test_set)) t_stat, p_value stats.ttest_rel(acc_id3, acc_cart) print(fID3 均值: {np.mean(acc_id3):.4f}, CART 均值: {np.mean(acc_cart):.4f}) print(ft 值: {t_stat:.4f}, p 值: {p_value:.4f}) return t_stat, p_value参数说明np.array_split按索引切分保证每一折样本数基本相等seed42固定随机种子保证结果可复现。evaluate_tree是预测函数把测试集每一行的特征逐层送入树结构到达叶子节点后返回标签。重点解释ttest_rel的用法它要求两组数据长度相等且按相同顺序配对即第一个元素都是第 1 折的结果这里已经满足。p 值小于 0.05 时认为两个算法在统计上有显著差异大于 0.05 则差异不显著——注意 p 值不显著不代表两个算法等价只表示在当前数据量下检不出差异。样本量小时比如西瓜数据集只有 17 条k 折里每折只有 3-4 条测试样本准确率离散化严重t 检验基本不会显著此时更合理的做法是报告各折均值并注明数据规模限制。6.3 我自己的收尾习惯这套代码包我拆过好几遍每次重跑最深的教训是必须先看数据再看代码。拿到任何数据集第一件事是打印df.shape和df.head()确认列数、标签列、中文编码没问题然后固定随机种子让每次实验可复现建树之前先问自己一句这里有连续属性吗有编号列吗把这两类特征处理干净再进递归。从那以后我每次跑决策树实验都强制走一遍这三个动作看数据格式、固定随机种子、画图前确认字体。这套习惯在跑西瓜书第四章的 4.3、4.4、4.6 三道题时帮助我少走了大半弯路希望帮到你。本文还有配套的精品资源点击获取
返回列表