ARTICLE DETAIL

资讯详情

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

决策树入门实战:用scikit-learn实现分类、剪枝与调参全流程

决策树入门实战:用scikit-learn实现分类、剪枝与调参全流程 今天直接进入主题AI开发入门第一个必须吃透的机器学习算法我建议选决策树Decision Tree。原因很简单。决策树不依赖复杂数学推导训练结果能直接画成一棵树每一步判断都看得见、解释得了。它既能做分类也能做回归还是随机森林、XGBoost、LightGBM 这些工业级模型的底层构件。换句话说决策树不是“玩具算法”而是后续所有树模型家族的地基。这篇文章会按本地实战路线走先讲决策树到底在算什么然后直接在 Python 环境里用 scikit-learn 完成两次完整建模一次是经典的鸢尾花分类一次是收入预测任务。过程中会覆盖特征选择、树的可视化、剪枝、网格搜索调参、模型评估这些必踩环节。所有代码都是可复制的环境要求也非常低不需要独立显卡一台普通 CPU 电脑就能跑完。读完你会有两方面的收获一是彻底搞懂决策树的原理不靠死记硬背二是手里多一套可以直接复用的 sklearn 决策树建模模板后续换成自己的表格数据也能快速上手。1. 核心能力速览先给一张能力速览表把决策树在 AI 开发基础阶段的位置说清楚。能力项说明算法类型监督学习支持分类和回归输入数据结构化表格数据特征可以是数值型也可以是离散型常用库scikit-learnsklearn、pandas、matplotlib核心优势可解释性强训练速度快对数据分布没有强假设主要缺点单棵树容易过拟合对噪声敏感决策边界是轴平行分割是否支持批量调参支持配合 GridSearchCV 可自动搜索最优参数组合是否有 API 服务算法本身不提供 API但可导出为模型文件后用 FastAPI 等封装硬件要求CPU 即可运行无需 GPU典型任务鸢尾花分类、收入预测、客户流失预测、信用风险评估适合读者AI 开发初学者、准备机器学习面试的开发者、想快速建立建模流程的工程师从能力速览可以看出决策树的准入门槛在主流机器学习算法里属于最低的一档。你不需要 8G 显存、不需要 CUDA、不需要配置大模型推理环境只要有一个能跑 Python 的笔记本就能把整个流程走通。2. 决策树的适用场景与使用边界2.1 适合什么场景决策树最擅长的场景是“表格数据 明确的预测目标”。比如根据用户的年龄、收入、历史消费次数预测是否愿意续费根据天气、温度、风力预测是否适合出行根据病人的各项体检指标辅助判断风险等级根据订单金额、退货率、发货时长判断是否存在异常交易。这类任务的共同点是数据是结构化的特征含义明确业务方还经常要求“你得告诉我为什么这么判断”。决策树的天然优势就在这里它能输出一条完整的 if-then 规则链比如“如果年龄大于 30 且收入大于 1 万则预测为续费用户”。2.2 不适合什么场景决策树不适合以下场景提前说明可以帮你绕过坑图像分类、语音识别等非结构化数据任务应该交给深度学习模型超高维稀疏数据比如用户行为序列直接展开成上万维特征单棵决策树效果通常不理想特征之间高度非线性且关系极其复杂的任务单棵树的表达能力有限数据量极大时单棵树不一定打不过集成模型。2.3 使用边界与合规提醒决策树模型本身是通用算法但使用时要注意数据边界训练数据涉及用户个人信息时要进行脱敏处理明确授权范围企业内部的业务数据不要随意上传到公网 Notebook 平台模型输出的是统计预测结果不能直接作为医疗诊断、信贷审批等高风险场景的唯一依据如果后续要把模型封装成 API 服务要注意接口鉴权和访问控制避免数据泄露。3. 环境准备与前置条件3.1 确认 Python 环境建议使用 Python 3.9 及以上版本。打开终端执行以下命令确认版本python --version如果还没有安装 Python可以从官网下载安装包安装时勾选“Add Python to PATH”。3.2 安装依赖库本文的实战环节只需要以下四个库pip install scikit-learn pandas matplotlib安装完成后可以用下面的命令验证核心库能否正常导入python -c import sklearn, pandas, matplotlib; print(deps ok)如果看到deps ok说明环境已经就绪。3.3 数据集说明本文使用两个数据集鸢尾花数据集Irissklearn 内置包含 150 条样本4 个特征3 个类别适合跑通第一个分类模型收入预测数据集使用模拟生成的工资收入数据包含年龄、教育年限、工作时长等特征目标变量为“收入是否超过 5 万/年”。收入预测部分我用代码现场构造数据不依赖外部下载这样在任何网络环境下都能复现。4. 决策树核心原理与代码实现4.1 决策树在算什么决策树做的事情可以概括为一句大白话通过一系列“是/否”判断把数据一步步分到不同的类别里。比如判断一个人收入是否超过 5 万第一步教育年限是否大于 12 年是进入下一步否大概率预测为“不超过 5 万”。第二步工作时长是否大于 40 小时是预测为“超过 5 万”否预测为“不超过 5 万”。这个“下一步分到哪个特征、以什么值作为分割点”就是决策树训练阶段要解决的核心问题。4.2 特征分裂的数学依据训练时算法会遍历所有特征以及所有可能的分裂点挑选一个“让分裂后的数据更纯”的特征作为当前节点。“更纯”在数学上有两种常见度量度量方式公式思路特点信息熵系统混乱程度熵越低越纯ID3、C4.5 算法使用基尼指数随机抽取两个样本类别不一致的概率CART 算法使用sklearn 默认sklearn 的DecisionTreeClassifier默认使用criteriongini也就是基尼指数。你可以把基尼指数理解为“不纯度”如果某个节点的样本全是同一类别基尼指数就是 0说明节点非常纯。4.3 最简单的决策树训练代码先写一个最小可运行版本from sklearn.datasets import load_iris from sklearn.tree import DecisionTreeClassifier from sklearn.model_selection import train_test_split # 加载数据 iris load_iris() X iris.data y iris.target # 划分训练集和测试集 X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.3, random_state42 ) # 创建决策树模型 clf DecisionTreeClassifier(random_state42) # 训练 clf.fit(X_train, y_train) # 预测 y_pred clf.predict(X_test) # 评估 print(准确率:, (y_pred y_test).mean())这段代码就是整个决策树建模的骨架。你不需要手动实现熵的计算和特征选择sklearn 已经把细节封装好了。5. 鸢尾花分类实战模型训练与预测5.1 完整建模流程在最小代码的基础上补全评估指标和预测演示import pandas as pd from sklearn.datasets import load_iris from sklearn.tree import DecisionTreeClassifier from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score, classification_report, confusion_matrix # 1. 加载数据 iris load_iris() X pd.DataFrame(iris.data, columnsiris.feature_names) y pd.Series(iris.target, nametarget) # 2. 划分数据集 X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.3, random_state42, stratifyy ) # 3. 初始化模型 clf DecisionTreeClassifier( criteriongini, max_depth3, min_samples_leaf2, random_state42 ) # 4. 训练 clf.fit(X_train, y_train) # 5. 预测 y_train_pred clf.predict(X_train) y_test_pred clf.predict(X_test) # 6. 评估 print(训练集准确率:, accuracy_score(y_train, y_train_pred)) print(测试集准确率:, accuracy_score(y_test, y_test_pred)) print(\n混淆矩阵:\n, confusion_matrix(y_test, y_test_pred)) print(\n分类报告:\n, classification_report(y_test, y_test_pred))这里先设置了max_depth3和min_samples_leaf2属于预剪枝操作目的是先把树限制在较小的规模观察基本效果。5.2 判断模型是否成功判断标准如下训练集准确率通常接近 0.95 以上测试集准确率保持在 0.9 左右训练集和测试集准确率差距不大说明没有明显过拟合。如果发现训练集准确率接近 1.0测试集只有 0.8说明树太深记住了太多噪声需要降低max_depth或提高min_samples_leaf。5.3 特征重要性分析决策树训练完成后可以通过feature_importances_查看每个特征对预测的贡献importance_df pd.DataFrame({ feature: iris.feature_names, importance: clf.feature_importances_ }).sort_values(importance, ascendingFalse) print(importance_df)输出结果中 importance 值越大说明该特征越早参与分裂对最终预测的影响越大。这是决策树可解释性的重要体现。6. 决策树可视化与结果解释6.1 可视化方法概述sklearn 提供了内置的plot_tree函数可以直接把树结构画成图不需要额外安装 Graphvizimport matplotlib.pyplot as plt from sklearn.tree import plot_tree plt.figure(figsize(16, 10)) plot_tree( clf, feature_namesiris.feature_names, class_namesiris.target_names, filledTrue, roundedTrue ) plt.title(Decision Tree Visualization) plt.show()可视化结果中每个节点包含以下信息分裂条件比如petal length (cm) 2.45节点的基尼指数当前节点的样本数量每个类别的样本分布该节点被判定为哪个类别。6.2 如何阅读决策树从根节点开始按照“满足条件走左分支不满足走右分支”的规则一路走到叶子节点。叶子节点就是最终预测结果。对于鸢尾花数据集你很可能看到第一个分裂特征就是petal length说明花瓣长度是最关键的区分特征。这个结论和领域知识相符也验证了模型确实学到了有意义的规律。6.3 输出树结构的文本规则如果不想画图也可以直接打印决策树的规则文本from sklearn.tree import export_text text_representation export_text( clf, feature_namesiris.feature_names.tolist() ) print(text_representation)输出格式类似|--- petal length (cm) 2.45 | |--- class: 0 |--- petal length (cm) 2.45 ...这里不会贴全部输出你在本地运行时可以直接复制这段规则到项目文档中作为业务解释材料。7. 决策树剪枝与过拟合处理7.1 为什么需要剪枝如果不限制树的深度决策树会一直分裂到所有叶子节点都“纯”为止。结果是训练集准确率接近 100%但测试集表现很差。这就是典型的过拟合。从材料中的热词“决策树剪枝面试题”也能看出来剪枝是面试和实际建模都绕不开的点。7.2 预剪枝参数预剪枝在训练过程中直接限制树的生长。常用参数如下参数含义建议max_depth树的最大深度从 3 到 5 开始尝试min_samples_split内部节点再分裂所需的最小样本数5 到 10min_samples_leaf叶子节点最少样本数2 到 5max_features每次分裂最多考虑的特征数训练速度慢时可使用修改上面的模型参数对比不同深度下的测试集准确率for depth in [2, 3, 5, 10, None]: clf_temp DecisionTreeClassifier( random_state42, max_depthdepth ) clf_temp.fit(X_train, y_train) train_acc accuracy_score(y_train, clf_temp.predict(X_train)) test_acc accuracy_score(y_test, clf_temp.predict(X_test)) print(fmax_depth{depth}, 训练集{train_acc:.4f}, 测试集{test_acc:.4f})这段实验可以直观看到depth 从小到大时训练集准确率持续上升但测试集准确率可能在某个深度后开始下降或不再提升。选择测试集准确率最高的那个深度即可。7.3 后剪枝成本复杂度剪枝sklearn 还支持基于ccp_alpha的成本复杂度剪枝。思路是先生成一棵完整的大树然后通过参数ccp_alpha去掉那些对整体精度贡献不大的子树。实际操作时先获取可用的 alpha 候选值再遍历选择效果最好的import numpy as np # 获取成本复杂度剪枝路径 clf_prune DecisionTreeClassifier(random_state42) path clf_prune.cost_complexity_pruning_path(X_train, y_train) ccp_alphas path.ccp_alphas print(候选 alpha 数量:, len(ccp_alphas))然后针对每一组 alpha 训练模型并记录测试集表现。不要直接使用最后一个极端值通常选择测试集准确率最高的中等 alpha 值。8. 网格搜索与批量调参8.1 手动循环的局限上一节我们手动写了一个for循环来测试不同深度。当参数组合增加到 3 个、4 个时手动嵌套循环会变得很难维护。这时候需要用GridSearchCV完成批量搜索。8.2 GridSearchCV 批量调参示例from sklearn.model_selection import GridSearchCV # 参数空间 param_grid { max_depth: [3, 5, 7, 10], min_samples_split: [2, 5, 10], min_samples_leaf: [1, 2, 4], criterion: [gini, entropy] } # 基础模型 base_clf DecisionTreeClassifier(random_state42) # 网格搜索 grid_search GridSearchCV( estimatorbase_clf, param_gridparam_grid, cv5, scoringaccuracy, n_jobs-1, verbose1 ) grid_search.fit(X_train, y_train) print(最优参数:, grid_search.best_params_) print(最优交叉验证准确率:, grid_search.best_score_) print(测试集准确率:, accuracy_score(y_test, grid_search.best_estimator_.predict(X_test)))n_jobs-1表示使用所有 CPU 核心并行搜索。这里再次强调决策树训练非常快不需要 GPU普通笔记本就能完成网格搜索。8.3 批量调参的实际意义从工程角度讲网格搜索本质上就是“批量任务”给定一组候选参数组合系统自动依次训练、评估、比较最终返回最优结果。这个思路和后续使用 AI Agent 自动调参是一致的先学会 GridSearchCV以后理解 AutoML 工具会轻松很多。9. 收入预测实战从特征工程到模型评估9.1 构造模拟收入数据收入预测是机器学习教材和高频面试题中的经典场景正好对应“决策树进行收入预测-sklearn版”。这里构造一个结构化表格数据集import numpy as np from sklearn.model_selection import train_test_split from sklearn.tree import DecisionTreeClassifier from sklearn.preprocessing import LabelEncoder from sklearn.metrics import accuracy_score, classification_report import pandas as pd # 设置随机种子保证可复现 np.random.seed(42) n_samples 2000 # 生成特征 age np.random.randint(18, 65, n_samples) education_years np.random.randint(6, 22, n_samples) hours_per_week np.random.randint(20, 80, n_samples) occupation np.random.choice( [engineer, teacher, sales, admin, manager], n_samples ) # 构造收入标签教育年限和工作时长的权重更高 income_score ( education_years * 0.15 hours_per_week * 0.01 age * 0.005 np.random.normal(0, 0.3, n_samples) ) # 转换为二分类标签是否超过 5 万 年收入 income_label (income_score np.median(income_score)).astype(int) # 组合成 DataFrame data pd.DataFrame({ age: age, education_years: education_years, hours_per_week: hours_per_week, occupation: occupation, high_income: income_label }) print(data.head()) print(data[high_income].value_counts())这是一个演示数据集目的是跑通流程不是说“学历、工时决定收入”。真实项目中需要结合业务背景采集数据并判断特征合法性。9.2 类别特征编码决策树虽然能处理部分类别特征但 sklearn 的实现要求输入全部是数值型。这里对occupation进行标签编码encoder LabelEncoder() data[occupation_encoded] encoder.fit_transform(data[occupation]) X data[[age, education_years, hours_per_week, occupation_encoded]] y data[high_income] X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.3, random_state42, stratifyy )9.3 训练并评估收入预测模型clf_income DecisionTreeClassifier( criteriongini, max_depth4, min_samples_leaf5, random_state42 ) clf_income.fit(X_train, y_train) y_pred clf_income.predict(X_test) print(收入预测测试集准确率:, accuracy_score(y_test, y_pred)) print(\n分类报告:\n, classification_report(y_test, y_pred))9.4 查看特征重要性feature_importance pd.DataFrame({ feature: X.columns, importance: clf_income.feature_importances_ }).sort_values(importance, ascendingFalse) print(feature_importance)在这个模拟数据集中education_years大概率是重要性最高的特征。这说明模型确实捕捉到了构造数据时设定的规律。换成真实业务数据时这个输出可以帮我们快速锁定哪些字段对预测影响最大从而减少无效特征的采集。9.5 用 Pipeline 固化流程工程化建模时推荐把编码和模型封装到 Pipeline 中避免每次预测重复写转换逻辑from sklearn.pipeline import Pipeline from sklearn.compose import ColumnTransformer from sklearn.preprocessing import OneHotEncoder # 定义列处理器年龄、教育年限等数值列直接使用职业列做 One-Hot 编码 preprocessor ColumnTransformer( transformers[ (num, passthrough, [age, education_years, hours_per_week]), (cat, OneHotEncoder(), [occupation]) ] ) pipeline Pipeline(steps[ (preprocessor, preprocessor), (classifier, DecisionTreeClassifier(random_state42, max_depth4)) ]) pipeline.fit(X_train, y_train) print(Pipeline 测试集准确率:, accuracy_score(y_test, pipeline.predict(X_test)))Pipeline 的好处是训练时自动处理特征预测时用同一套规则处理新数据不会出现训练和上线特征不一致的问题。10. 常见问题与排查方法下面是实际学习过程中最常遇到的问题整理成排查表。问题现象可能原因排查方式解决方案pip install scikit-learn失败网络源慢或 Python 版本不兼容检查 pip 版本和 Python 版本使用pip install scikit-learn -i https://pypi.tuna.tsinghua.edu.cn/simple数据包含字符串特征训练报错sklearn 默认要求数值输入查看 DataFrame 的 dtypes使用 LabelEncoder 或 OneHotEncoder 编码存在缺失值训练报错决策树无法处理 NaN使用df.isnull().sum()统计缺失填充均值/中位数或删除缺失行训练集准确率 1.0测试集很低树过深过拟合对比 train/test 准确率降低 max_depth提高 min_samples_leaf测试集准确率始终上不去特征太弱或数据本身难以区分查看特征重要性增加特征、构造新特征或更换算法可视化图形中文乱码matplotlib 字体问题查看控制台报错使用英文标签或配置中文字体模型预测结果偏向某一类类别不平衡打印类别分布使用 class_weightbalanced网格搜索速度过慢参数空间大且数据量大查看控制台任务状态缩小参数范围使用 n_jobs-1或改用 RandomizedSearchCV10.1 类别不平衡的处理在收入预测这类问题中如果高收入样本占比很低模型容易把所有样本都预测为“低收入”准确率看起来还不错实际毫无意义。判断方法print(data[high_income].value_counts(normalizeTrue))如果某一类占比超过 80%就需要在模型中加入类别权重参数clf_balance DecisionTreeClassifier( random_state42, class_weightbalanced, max_depth4 )10.2 随机数种子问题决策树训练本身有一定随机性尤其是特征较多时。如果不固定random_state每次运行结果可能不同。建模时建议都加上random_state42方便复现实验结果。11. 决策树在 AI 开发全局中的定位与最佳实践11.1 和现在热门 AI 开发方向的关系当前热词里有大量“AI Agent 开发”“机器学习应用流程”“AI 应用开发”的内容。很多人以为 AI 开发只等于大模型和提示词工程实际上Agent 的很多底层能力依赖经典机器学习算法意图分类可以用决策树快速建立基线用户画像分层适合用树模型做可解释规则表格数据的自动化分析任务里决策树和随机森林依然是主流底座大模型生成的伪代码或自动化建模工具很大一部分也是在自动调 sklearn 这类经典库。所以先掌握决策树不是为了停留在“调库跑准确率”层面而是为了理解监督学习的完整流程数据处理、特征工程、模型训练、参数搜索、评估上线。这套流程在大模型时代同样适用只是工具形态变了。11.2 最佳实践建议结合整个实战过程给出几条工程化建议第一次建模不要一上来就追求高准确率先跑通全流程再逐步加参数训练集、验证集、测试集划分比例建议 6:2:2分类任务要使用stratify保持类别分布一致模型文件用joblib.dump保存预测时重新加载避免每次重训import joblib # 保存模型 joblib.dump(clf_income, income_tree_model.pkl) # 加载模型 loaded_clf joblib.load(income_tree_model.pkl)每个实验固定random_state方便对比不同模型的效果树模型不需要对特征做标准化和归一化这会减少一部分预处理工作量数据量达到数万条以上、特征维度较多时优先考虑随机森林或梯度提升树单棵树的稳定性会不足输出规则时建议把提取出来的 if-then 规则同步给业务方确认让模型结论落地到业务中。11.3 下一步学习方向如果你已经能独立完成本文的前两个实战下一步建议按这个顺序展开阶段学习内容目标算法扩展信息增益率、C4.5、CART 回归树理解决策树的理论全貌集成学习随机森林、Bagging、AdaBoost理解“多个弱模型组合成强模型”梯度提升XGBoost、LightGBM、CatBoost掌握工业界表格数据主力模型工程化Pipeline、交叉验证、模型持久化具备上线能力自动机器学习GridSearchCV、Optuna学会用自动化方式调参和选模型12. 总结与下一步决策树是整个机器学习体系里最适合当作第一个落地产物的算法。你不需要 GPU不需要分布式环境只需要一个 Python 环境和几百条结构化数据就能完成从数据预处理、模型训练、可视化解释到参数搜索的完整流程。建议你拿到代码后按这个顺序动手验证先跑通鸢尾花分类观察plot_tree输出的树结构用max_depth从 2 到 10 做一组实验直观理解过拟合跑收入预测流程把 Pipeline 固化下来尝试把DataFrame换成自己的 Excel 数据把流程改造成自己的建模小工具。最容易踩的坑有两个一是忽略数据中的字符串特征和缺失值导致训练直接报错二是不做任何剪枝模型训练集准确率虚高上线后测试集效果崩盘。这两个问题在本地很容易复现排查方法在这里也写清楚了遇到时不用慌。决策树之后随机森林和梯度提升树都是非常自然的下一站。它们底层还是“树”只不过用了不同的组合方式。当你理解了这个演进过程再看 AI Agent 自动建模、AutoML 自动调参这些新工具时就会发现核心逻辑都是一样的给定数据、给定候选模型、自动搜索最优配置。把今天这套流程练熟后面的路会顺很多。
返回列表