简介:本资源是面向机器学习初学者与实践者的决策树算法专项代码包,聚焦分类与回归任务的原理理解与工程实现。压缩包共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.py中build_tree()方法为轴,逐层还原其设计意图与实现细节。
2.1 信息增益 vs 增益率:C4.5 风格划分的代码落地
源码默认采用增益率(Gain Ratio)而非纯信息增益,避免偏向取值较多的特征(如用户ID字段)。其计算逻辑在utils.py的calc_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_split | config.py中MIN_SAMPLES_SPLIT=2 | 2 | 防止单样本过拟合 |
self._is_pure(y) | tree.py内置方法 | len(np.unique(y)) == 1 | 叶节点纯度判定 |
depth >= self.max_depth | config.py中MAX_DEPTH=10 | 10 | 控制树复杂度 |
其中_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.py的classification_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.py的print_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_SPLIT | int | 2 | 实验中设为 20~50 | 增大 → 树更浅,泛化性↑,训练误差↑ |
MIN_SAMPLES_LEAF | int | 1 | 设为MIN_SAMPLES_SPLIT // 2 | 确保叶节点最小样本量,防噪声主导 |
MAX_DEPTH | int | 10 | 头歌实验建议 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,内容示例:
| depth | prune_type | train_acc | val_acc | test_acc | tree_size |
|---|---|---|---|---|---|
| 3 | pre | 0.7921 | 0.8134 | 0.8092 | 15 |
| 5 | post | 0.8421 | 0.8317 | 0.8285 | 42 |
| 7 | none | 0.8763 | 0.8102 | 0.8041 | 128 |
可见:未剪枝树(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.py的DATASET_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 设备固件。
本文还有配套的精品资源,点击获取