决策树入门实战:用scikit-learn实现分类、剪枝与调参全流程
2026/9/7 17:25:14 网站建设 项目流程

今天直接进入主题:AI开发入门,第一个必须吃透的机器学习算法,我建议选决策树(Decision Tree)。

原因很简单。决策树不依赖复杂数学推导,训练结果能直接画成一棵树,每一步判断都看得见、解释得了。它既能做分类,也能做回归,还是随机森林、XGBoost、LightGBM 这些工业级模型的底层构件。换句话说,决策树不是“玩具算法”,而是后续所有树模型家族的地基。

这篇文章会按本地实战路线走:先讲决策树到底在算什么,然后直接在 Python 环境里用 scikit-learn 完成两次完整建模,一次是经典的鸢尾花分类,一次是收入预测任务。过程中会覆盖特征选择、树的可视化、剪枝、网格搜索调参、模型评估这些必踩环节。所有代码都是可复制的,环境要求也非常低,不需要独立显卡,一台普通 CPU 电脑就能跑完。

读完你会有两方面的收获:一是彻底搞懂决策树的原理,不靠死记硬背;二是手里多一套可以直接复用的 sklearn 决策树建模模板,后续换成自己的表格数据也能快速上手。

1. 核心能力速览

先给一张能力速览表,把决策树在 AI 开发基础阶段的位置说清楚。

能力项说明
算法类型监督学习,支持分类和回归
输入数据结构化表格数据,特征可以是数值型,也可以是离散型
常用库scikit-learn(sklearn)、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 数据集说明

本文使用两个数据集:

  • 鸢尾花数据集(Iris):sklearn 内置,包含 150 条样本,4 个特征,3 个类别,适合跑通第一个分类模型;
  • 收入预测数据集:使用模拟生成的工资收入数据,包含年龄、教育年限、工作时长等特征,目标变量为“收入是否超过 5 万/年”。

收入预测部分我用代码现场构造数据,不依赖外部下载,这样在任何网络环境下都能复现。

4. 决策树核心原理与代码实现

4.1 决策树在算什么

决策树做的事情可以概括为一句大白话:通过一系列“是/否”判断,把数据一步步分到不同的类别里。

比如判断一个人收入是否超过 5 万:

  • 第一步:教育年限是否大于 12 年?
    • 是:进入下一步;
    • 否:大概率预测为“不超过 5 万”。
  • 第二步:工作时长是否大于 40 小时?
    • 是:预测为“超过 5 万”;
    • 否:预测为“不超过 5 万”。

这个“下一步分到哪个特征、以什么值作为分割点”就是决策树训练阶段要解决的核心问题。

4.2 特征分裂的数学依据

训练时,算法会遍历所有特征以及所有可能的分裂点,挑选一个“让分裂后的数据更纯”的特征作为当前节点。

“更纯”在数学上有两种常见度量:

度量方式公式思路特点
信息熵系统混乱程度,熵越低越纯ID3、C4.5 算法使用
基尼指数随机抽取两个样本类别不一致的概率CART 算法使用,sklearn 默认

sklearn 的DecisionTreeClassifier默认使用criterion='gini',也就是基尼指数。你可以把基尼指数理解为“不纯度”:如果某个节点的样本全是同一类别,基尼指数就是 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_size=0.3, random_state=42 ) # 创建决策树模型 clf = DecisionTreeClassifier(random_state=42) # 训练 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, columns=iris.feature_names) y = pd.Series(iris.target, name="target") # 2. 划分数据集 X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.3, random_state=42, stratify=y ) # 3. 初始化模型 clf = DecisionTreeClassifier( criterion="gini", max_depth=3, min_samples_leaf=2, random_state=42 ) # 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_depth=3min_samples_leaf=2,属于预剪枝操作,目的是先把树限制在较小的规模,观察基本效果。

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", ascending=False) print(importance_df)

输出结果中 importance 值越大,说明该特征越早参与分裂,对最终预测的影响越大。这是决策树可解释性的重要体现。

6. 决策树可视化与结果解释

6.1 可视化方法概述

sklearn 提供了内置的plot_tree函数,可以直接把树结构画成图,不需要额外安装 Graphviz:

import matplotlib.pyplot as plt from sklearn.tree import plot_tree plt.figure(figsize=(16, 10)) plot_tree( clf, feature_names=iris.feature_names, class_names=iris.target_names, filled=True, rounded=True ) 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_names=iris.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 到 10
min_samples_leaf叶子节点最少样本数2 到 5
max_features每次分裂最多考虑的特征数训练速度慢时可使用

修改上面的模型参数,对比不同深度下的测试集准确率:

for depth in [2, 3, 5, 10, None]: clf_temp = DecisionTreeClassifier( random_state=42, max_depth=depth ) 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(f"max_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_state=42) 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_state=42) # 网格搜索 grid_search = GridSearchCV( estimator=base_clf, param_grid=param_grid, cv=5, scoring="accuracy", n_jobs=-1, verbose=1 ) 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_size=0.3, random_state=42, stratify=y )

9.3 训练并评估收入预测模型

clf_income = DecisionTreeClassifier( criterion="gini", max_depth=4, min_samples_leaf=5, random_state=42 ) 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", ascending=False) 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_state=42, max_depth=4)) ]) 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_weight='balanced'
网格搜索速度过慢参数空间大且数据量大查看控制台任务状态缩小参数范围,使用 n_jobs=-1,或改用 RandomizedSearchCV

10.1 类别不平衡的处理

在收入预测这类问题中,如果高收入样本占比很低,模型容易把所有样本都预测为“低收入”,准确率看起来还不错,实际毫无意义。

判断方法:

print(data["high_income"].value_counts(normalize=True))

如果某一类占比超过 80%,就需要在模型中加入类别权重参数:

clf_balance = DecisionTreeClassifier( random_state=42, class_weight="balanced", max_depth=4 )

10.2 随机数种子问题

决策树训练本身有一定随机性,尤其是特征较多时。如果不固定random_state,每次运行结果可能不同。建模时建议都加上random_state=42,方便复现实验结果。

11. 决策树在 AI 开发全局中的定位与最佳实践

11.1 和现在热门 AI 开发方向的关系

当前热词里有大量“AI Agent 开发”“机器学习应用流程”“AI 应用开发”的内容。很多人以为 AI 开发只等于大模型和提示词工程,实际上,Agent 的很多底层能力依赖经典机器学习算法:

  • 意图分类可以用决策树快速建立基线;
  • 用户画像分层适合用树模型做可解释规则;
  • 表格数据的自动化分析任务里,决策树和随机森林依然是主流底座;
  • 大模型生成的伪代码或自动化建模工具,很大一部分也是在自动调 sklearn 这类经典库。

所以,先掌握决策树,不是为了停留在“调库跑准确率”层面,而是为了理解监督学习的完整流程:数据处理、特征工程、模型训练、参数搜索、评估上线。这套流程在大模型时代同样适用,只是工具形态变了。

11.2 最佳实践建议

结合整个实战过程,给出几条工程化建议:

  1. 第一次建模不要一上来就追求高准确率,先跑通全流程,再逐步加参数;
  2. 训练集、验证集、测试集划分比例建议 6:2:2,分类任务要使用stratify保持类别分布一致;
  3. 模型文件用joblib.dump保存,预测时重新加载,避免每次重训:
import joblib # 保存模型 joblib.dump(clf_income, "income_tree_model.pkl") # 加载模型 loaded_clf = joblib.load("income_tree_model.pkl")
  1. 每个实验固定random_state,方便对比不同模型的效果;
  2. 树模型不需要对特征做标准化和归一化,这会减少一部分预处理工作量;
  3. 数据量达到数万条以上、特征维度较多时,优先考虑随机森林或梯度提升树,单棵树的稳定性会不足;
  4. 输出规则时,建议把提取出来的 if-then 规则同步给业务方确认,让模型结论落地到业务中。

11.3 下一步学习方向

如果你已经能独立完成本文的前两个实战,下一步建议按这个顺序展开:

阶段学习内容目标
算法扩展信息增益率、C4.5、CART 回归树理解决策树的理论全貌
集成学习随机森林、Bagging、AdaBoost理解“多个弱模型组合成强模型”
梯度提升XGBoost、LightGBM、CatBoost掌握工业界表格数据主力模型
工程化Pipeline、交叉验证、模型持久化具备上线能力
自动机器学习GridSearchCV、Optuna学会用自动化方式调参和选模型

12. 总结与下一步

决策树是整个机器学习体系里最适合当作第一个落地产物的算法。你不需要 GPU,不需要分布式环境,只需要一个 Python 环境和几百条结构化数据,就能完成从数据预处理、模型训练、可视化解释到参数搜索的完整流程。

建议你拿到代码后,按这个顺序动手验证:

  1. 先跑通鸢尾花分类,观察plot_tree输出的树结构;
  2. max_depth从 2 到 10 做一组实验,直观理解过拟合;
  3. 跑收入预测流程,把 Pipeline 固化下来;
  4. 尝试把DataFrame换成自己的 Excel 数据,把流程改造成自己的建模小工具。

最容易踩的坑有两个:一是忽略数据中的字符串特征和缺失值,导致训练直接报错;二是不做任何剪枝,模型训练集准确率虚高,上线后测试集效果崩盘。这两个问题在本地很容易复现,排查方法在这里也写清楚了,遇到时不用慌。

决策树之后,随机森林和梯度提升树都是非常自然的下一站。它们底层还是“树”,只不过用了不同的组合方式。当你理解了这个演进过程,再看 AI Agent 自动建模、AutoML 自动调参这些新工具时,就会发现核心逻辑都是一样的:给定数据、给定候选模型、自动搜索最优配置。把今天这套流程练熟,后面的路会顺很多。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询