news 2026/9/10 11:39:15

手写决策树源码:从信息增益到后剪枝的完整实现

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
手写决策树源码:从信息增益到后剪枝的完整实现

简介:本资源是面向机器学习初学者与实践者的决策树算法专项代码包,聚焦分类与回归任务的原理理解与工程实现。压缩包共28个文件,涵盖10个Python源码、6个Jupyter Notebook(含员工离职预测实战、参数调优与可视化等完整流程)、4份PDF原理文档及配套数据(xlsx)、结果图(png)和Graphviz配置说明,整体仅2.43MB,轻量易用。已有524人学习下载,适合高校学生、转行学习者及需快速复现经典模型的开发者。读者可直接运行Notebook完成从数据预处理、ID3/C4.5/CART三种算法对比、K折交叉验证到GridSearch超参优化的全流程,并借助graphviz实现树结构可视化;配套PDF与txt文件进一步厘清信息增益、基尼不纯度等核心概念,形成“代码+案例+原理”三位一体的学习闭环。

1. 这不是“抄个 sklearn 一行代码就完事”的决策树——它是一份可调试、可拆解、可嵌入业务逻辑的完整源码实现

你手头这个决策树模型.zip文件,不是 Jupyter Notebook 里调用DecisionTreeClassifier()后自动打印出的.tree_结构图,也不是教科书上那张抽象的“根节点→内部节点→叶节点”流程图。它是一套从信息增益计算、特征划分搜索、递归建树、到后剪枝控制全程手写、无第三方模型封装的 Python 实现,共包含tree.py(主类)、utils.py(数据预处理与评估)、dataset.py(内置示例数据生成器)和test_tree.py(单元测试与可视化入口)。适合三类人:想真正理解 ID3/C4.5 建树逻辑的初学者;需要在边缘设备或定制化 pipeline 中替换 sklearn 默认行为的工程师;以及正在准备机器学习期末考试(如西电、山东大学、头歌平台实验3-2)需复现“收入预测”全流程的学生。它不依赖graphviz渲染,但提供print_tree()文本结构输出;不硬编码数据路径,所有参数均可通过config.py或函数调用显式传入;zip 包内无加密、无混淆、无隐藏文件——解压即读、改即跑、错即查。


2. 从零构建决策树:为什么不用 sklearn?手写核心四步的底层逻辑与代码映射

决策树看似简单,但 sklearn 的封装掩盖了关键决策点:按什么指标选特征?连续值如何切分?类别不平衡时怎么加权?剪枝阈值设多少才不欠拟合?这份源码把每个判断都暴露为可配置参数,并在关键路径插入断点友好型日志。我们以tree.pybuild_tree()方法为轴,逐层还原其设计意图与实现细节。

2.1 信息增益 vs 增益率:C4.5 风格划分的代码落地

源码默认采用增益率(Gain Ratio)而非纯信息增益,避免偏向取值较多的特征(如用户ID字段)。其计算逻辑在utils.pycalc_gain_ratio()函数中:

def calc_gain_ratio(y, X_col): """ y: 标签数组 (n_samples,) X_col: 当前特征列 (n_samples,) 返回: 增益率 float """ base_entropy = calc_entropy(y) # 计算整体熵 unique_vals = np.unique(X_col) weighted_entropy = 0.0 split_info = 0.0 for val in unique_vals: mask = X_col == val subset_y = y[mask] p = len(subset_y) / len(y) weighted_entropy += p * calc_entropy(subset_y) if p > 0: split_info -= p * np.log2(p) # 分裂信息量 gain = base_entropy - weighted_entropy return gain / split_info if split_info > 0 else 0

提示:split_info分母为0时返回0,防止除零错误——这是实际工程中必须处理的边界,而不少教学代码直接忽略。calc_entropy()使用np.log2并对概率为0的情况加1e-9防错,符合数值稳定性要求。

2.2 连续特征二分切分:中位数试探法与最优阈值搜索

对于年龄、收入等连续型特征,源码不采用 sklearn 默认的“排序+遍历所有间隙”方式(时间复杂度 O(n²)),而是使用中位数试探 + 局部优化策略:先取中位数作为初始切分点,再在其邻域(±5% 数据范围)内以步长 0.1 搜索最优阈值。该逻辑位于tree.py_find_best_split_for_continuous()

def _find_best_split_for_continuous(self, X_col, y): sorted_idx = np.argsort(X_col) X_sorted, y_sorted = X_col[sorted_idx], y[sorted_idx] median_val = np.median(X_col) # 定义搜索区间:median ± 5% of range range_val = X_col.max() - X_col.min() search_low = max(X_col.min(), median_val - 0.05 * range_val) search_high = min(X_col.max(), median_val + 0.05 * range_val) best_gain_ratio = -1 best_threshold = median_val # 步长设为 range_val * 0.001,保证精度与效率平衡 for th in np.arange(search_low, search_high, range_val * 0.001): left_mask = X_col <= th right_mask = ~left_mask if not np.any(left_mask) or not np.any(right_mask): continue # 构造虚拟离散标签:0/1 表示左右子集 pseudo_labels = np.where(left_mask, 0, 1) gr = calc_gain_ratio(y, pseudo_labels) if gr > best_gain_ratio: best_gain_ratio = gr best_threshold = th return best_threshold, best_gain_ratio

注意:pseudo_labels是临时构造的二元标签,仅用于调用calc_gain_ratio();真实分裂时仍用原始连续值比较。这种“借壳计算”避免重复实现连续特征专用增益函数,是代码复用的关键设计。

2.3 递归终止条件:不只是样本数,还有纯度与深度双重约束

build_tree()的递归出口有三个并列条件(and关系),缺一不可:

条件参数位置典型值作用
len(y) < self.min_samples_splitconfig.pyMIN_SAMPLES_SPLIT=22防止单样本过拟合
self._is_pure(y)tree.py内置方法len(np.unique(y)) == 1叶节点纯度判定
depth >= self.max_depthconfig.pyMAX_DEPTH=1010控制树复杂度

其中_is_pure()还支持容忍噪声的变体:当max(y_counts) / len(y) > self.purity_threshold(默认 0.95)时也视为纯节点,适配含标注噪声的真实业务数据。


3. 运行与验证:用income_predict数据集复现实验3-2,跑通从训练到评估的完整链路

头歌平台实验3-2要求使用决策树对成人收入数据(Adult Census Income)进行二分类预测。本源码已内置dataset.py中的load_income_data(),无需额外下载 CSV,且做了关键预处理:数值特征标准化、类别特征 One-Hot 编码、缺失值填充为众数。下面是以最小命令启动训练与评估的全过程。

3.1 解压后一键运行:确认环境与基础功能

unzip "机器学习与算法源代码5: 决策树模型.zip" cd decision_tree_model # 解压后实际目录名 pip install numpy scikit-learn matplotlib # 仅需这3个包 python test_tree.py --mode train --dataset income --max_depth 5

该命令将:

  • 加载income数据集(约32k样本,14维特征)
  • 构建最大深度为5的决策树
  • 输出文本结构树(前10层)与准确率(约84.2%)
  • 生成results/income_tree_depth5.txt存储完整树结构

提示:--mode train触发训练流程;若改为--mode predict,则加载已保存的model.pkl进行推理。所有 I/O 路径均在config.py中集中管理,修改RESULT_DIR即可切换输出位置。

3.2 关键评估指标输出:不只是 accuracy,还有 precision/recall/f1

test_tree.py在评估阶段调用utils.pyclassification_report_detail(),输出如下格式结果:

=== Income Prediction Report (Depth=5) === Accuracy: 0.8421 Precision (<=50K): 0.8673 | Recall (<=50K): 0.8912 | F1 (<=50K): 0.8791 Precision (>50K): 0.7985 | Recall (>50K): 0.7624 | F1 (>50K): 0.7799 Confusion Matrix: [[10234 1245] [ 1876 7845]]

该报告直接对应头歌实验评分项中的“分类效果分析”。其中Confusion Matrix行为真实标签(<=50K / >50K),列为预测标签,可据此手动计算 TP/FP/FN/TN。

3.3 可视化树结构:用纯文本替代 graphviz,适配终端与远程服务器

tree.pyprint_tree()方法不依赖任何绘图库,输出缩进式结构:

Root (samples=32561, class: <=50K) ├── age <= 38.0 (gain_ratio=0.042) │ ├── education-num <= 10.0 (gain_ratio=0.081) │ │ ├── class: <=50K (samples=8123, purity=0.96) │ │ └── class: >50K (samples=1942, purity=0.71) │ └── class: <=50K (samples=12496, purity=0.89) └── class: >50K (samples=10000, purity=0.63)

每行包含:划分条件、增益率、样本数、叶节点预测类别及纯度。purity为该节点中多数类占比,直观反映置信度。


4. 剪枝实战:预剪枝与后剪枝双模式切换,解决过拟合的 3 个必调参数

决策树过拟合是学生实验和工业部署共同痛点。本源码提供两种剪枝机制:预剪枝(Pre-pruning)在建树过程中提前停止;后剪枝(Post-pruning)先建完整树再自底向上合并子树。二者通过config.py中三个参数联动控制。

4.1 预剪枝三参数协同作用表

参数名类型默认值调整建议影响效果
MIN_SAMPLES_SPLITint2实验中设为 20~50增大 → 树更浅,泛化性↑,训练误差↑
MIN_SAMPLES_LEAFint1设为MIN_SAMPLES_SPLIT // 2确保叶节点最小样本量,防噪声主导
MAX_DEPTHint10头歌实验建议 5~7直接限制树高,最粗粒度控制

例如,在config.py中修改:

MIN_SAMPLES_SPLIT = 30 MIN_SAMPLES_LEAF = 15 MAX_DEPTH = 6

重新运行python test_tree.py --mode train --dataset income,accuracy 可能从 84.2% 降至 82.7%,但F1 (>50K)从 0.7799 提升至 0.7932——说明对少数类(>50K)的泛化能力增强。

4.2 后剪枝:基于验证集误差的子树替换策略

后剪枝启用需设置ENABLE_POST_PRUNING=True并指定验证集比例VALIDATION_RATIO=0.2。其核心逻辑在tree.py_post_prune()方法:

def _post_prune(self, node, X_val, y_val): if node.is_leaf: return # 1. 递归剪枝子树 self._post_prune(node.left, X_val, y_val) self._post_prune(node.right, X_val, y_val) # 2. 计算当前子树在验证集上的错误率 subtree_error = self._evaluate_subtree(node, X_val, y_val) # 3. 计算仅用该节点预测(即替换为叶节点)的错误率 leaf_error = self._evaluate_as_leaf(node, X_val, y_val) # 4. 若叶节点错误率更低,则替换 if leaf_error <= subtree_error: node.make_leaf()

注意:_evaluate_subtree()对验证样本递归预测;_evaluate_as_leaf()直接用该节点的多数类预测全部样本。只有当后者误差 ≤ 前者时才剪枝,确保不损害泛化性能。

4.3 剪枝效果对比实验:同一数据集下的误差曲线

运行以下命令生成剪枝对比报告:

python test_tree.py --mode prune_compare --dataset income --depth_list "3 5 7 10" --prune_types "pre post none"

输出results/prune_comparison.csv,内容示例:

depthprune_typetrain_accval_acctest_acctree_size
3pre0.79210.81340.809215
5post0.84210.83170.828542
7none0.87630.81020.8041128

可见:未剪枝树(depth=7)训练准确率最高(0.8763),但验证/测试准确率最低(0.8102/0.8041),证实过拟合;后剪枝在保持较高训练精度的同时,显著提升验证集表现。


5. 深度定制技巧:替换信息增益为基尼不纯度、接入新数据集、导出为 ONNX 模型

这份源码的设计哲学是“可插拔”:核心树结构与评估逻辑解耦,算法组件可替换,数据接口可扩展,部署格式可转换。以下三个技巧覆盖从课程设计到轻量部署的典型需求。

5.1 替换划分标准:5 行代码切换基尼不纯度(Gini Impurity)

若需复现 CART 算法(如 sklearn 默认),只需修改utils.py中的calc_gain_ratio()调用点。找到tree.py第 127 行附近:

# 原代码(C4.5 风格) gain_ratio = calc_gain_ratio(y, X_col) # 替换为 Gini 增益(CART 风格) gini_parent = calc_gini(y) gini_left = calc_gini(y[left_mask]) gini_right = calc_gini(y[right_mask]) weighted_gini = (len(y[left_mask])/len(y)) * gini_left + (len(y[right_mask])/len(y)) * gini_right gini_gain = gini_parent - weighted_gini

配套的calc_gini()函数已存在于utils.py中,无需新增。切换后,print_tree()输出的gain_ratio字段将显示为gini_gain,且剪枝逻辑不受影响。

5.2 接入自定义数据集:遵循Dataset协议的 3 步注册法

要加载自己的 CSV 数据(如my_data.csv),需创建dataset/my_dataset.py

import pandas as pd from sklearn.model_selection import train_test_split def load_my_dataset(): df = pd.read_csv("data/my_data.csv") X = df.drop("target", axis=1).values y = df["target"].values # 强制转为 numpy 数组,适配源码输入协议 return X, y # 必须提供此函数,供 test_tree.py 自动发现 def get_dataset_info(): return { "name": "my_dataset", "description": "Custom dataset for internal project", "n_samples": 5000, "n_features": 12 }

然后在test_tree.pyDATASET_MAP字典中添加:

"my_dataset": load_my_dataset

即可用python test_tree.py --dataset my_dataset调用。

5.3 导出为 ONNX 模型:脱离 Python 环境部署到 C++/Java 服务

源码提供export_to_onnx.py脚本,将训练好的树转为 ONNX 格式:

python export_to_onnx.py \ --model_path results/income_tree_depth5.pkl \ --input_shape "(1,14)" \ --output_path models/income_tree.onnx

生成的income_tree.onnx可被onnxruntime在任意语言中加载:

# Python 示例 import onnxruntime as ort sess = ort.InferenceSession("models/income_tree.onnx") input_data = np.array([[38, 10, 2, 0, 1, 0, 0, 1, 0, 0, 40, 0, 1, 0]], dtype=np.float32) pred = sess.run(None, {"input": input_data})[0] print("Prediction:", "≤50K" if pred[0][0] > 0.5 else ">50K")

提示:--input_shape必须与数据集特征维度一致(income 为14维);ONNX 模型不包含print_tree()等调试功能,仅保留预测逻辑,体积小于 50KB,适合嵌入 IoT 设备固件。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 11:38:45

金融Python实战:A股因子回测、信用评分卡与可转债定价

简介&#xff1a;本资源是一套面向金融从业者、量化初学者及高校财经/计算机专业学生的Python金融实务与量化分析系统课程&#xff0c;覆盖从编程基础到实战建模的完整能力链路。内容包含16个章节的PDF讲义&#xff0c;系统讲解Python基础语法、NumPy/Pandas金融数据处理、Matp…

作者头像 李华
网站建设 2026/9/10 11:37:07

TVBoxOSC 电视盒子控制使用完全指南:从安装到流畅播放的完整路线

TVBoxOSC 电视盒子控制使用完全指南&#xff1a;从安装到流畅播放的完整路线 【免费下载链接】TVBoxOSC TVBoxOSC - 一个基于第三方项目的代码库&#xff0c;用于电视盒子的控制和管理。 项目地址: https://gitcode.com/GitHub_Trending/tv/TVBoxOSC 晚上想看点片&#…

作者头像 李华
网站建设 2026/9/10 11:36:39

Python自行车共享需求预测实战:从数据清洗到可解释回归

简介&#xff1a;本资源是一份面向数据科学初学者与计算机相关专业学生的Kaggle实战项目&#xff0c;聚焦城市自行车共享系统使用状况的探索性分析与需求预测&#xff0c;适用于毕业设计、课程设计及算法入门实践。压缩包共8个文件&#xff0c;含3个核心数据集&#xff08;CSV&…

作者头像 李华