ARTICLE DETAIL

资讯详情

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

机器学习入门必学:决策树原理与Python实战

机器学习入门必学:决策树原理与Python实战 先从一个实际场景说起如果你是一位刚接触 AI 开发的初学者可能已经在网上看到过很多机器学习相关的名词像“监督学习”“分类”“回归”“特征”等等。说得直白一点机器学习就是让计算机从数据中自己总结规律然后用这个规律去预测新数据。而在众多算法中决策树Decision Tree可以说是最直观、也最适合入门的一种。它不需要复杂的数学推导画出来就是一棵“如果……那么……”的判断树和人类做决策的思维方式非常接近。本文会围绕“AI 开发基础快速入门机器学习之决策树”展开先讲清楚决策树的基本概念和核心原理再带领大家完成一个基于 Python scikit-learn 的收入预测分类实战案例。不管是准备机器学习期末复习、做课程实验还是准备面试中的决策树剪枝问题这篇文章都能给你一个比较完整的参考。读完本文你将掌握以下内容决策树是什么能解决什么问题信息熵、信息增益、基尼指数这些核心概念怎么理解决策树的训练、预测、可视化方法如何用 scikit-learn 完成一个完整的决策树分类项目常见报错与调优思路。1. 决策树是什么1.1 从日常决策说起先来思考一个生活中的例子。假设你今天打算出门跑步你会怎么判断“是否适合跑步”你会先看天气如果下雨就不跑如果不下雨再看温度温度太高不跑温度适中再看风力风太大不跑否则就出门跑步。这个过程实际上就是一棵树是否下雨 ├── 下雨 → 不跑 └── 不下雨 ├── 温度 30℃ → 不跑 └── 温度 ≤ 30℃ ├── 风力 5级 → 不跑 └── 风力 ≤ 5级 → 跑步决策树做的事情就是把这套判断逻辑用算法自动学出来。它根据历史数据自动找出“哪个特征最重要”“按什么顺序判断”“阈值设在哪里”最合适。1.2 决策树的正式定义从机器学习角度来说决策树是一种基于树结构进行决策的监督学习算法。它可以处理分类问题也可以处理回归问题。一棵决策树由下面几个部分组成根节点第一个用于划分数据的特征内部节点每一个非叶子节点都对应一个特征判断分支特征判断的不同结果叶节点最终的决策结果也就是类别或数值。因为决策树的结构清晰、可解释性强所以在很多领域都有广泛应用银行信用卡审批根据收入、负债、信用历史判断是否发卡医疗辅助诊断根据症状、检查指标判断患病风险电商用户画像根据浏览行为、消费记录判断用户意向工业异常检测根据设备运行参数判断是否故障。另外还要注意决策树是很多高级算法的基础。比如随机森林Random Forest、XGBoost、LightGBM 等集成学习方法核心基学习器都是决策树。所以想入门 AI 开发、机器学习决策树是绕不开的一环。2. 环境准备与版本说明在动手写代码之前先把开发环境准备好。本文的实战代码基于 Python 和 scikit-learn这也是目前机器学习入门最常用的组合。2.1 环境依赖这里给出本文需要用到的主要库Python 3.8 及以上numpypandasscikit-learnmatplotlibgraphviz可选用于可视化。如果你使用的是 Anaconda安装会非常方便。Anaconda 是一个 Python 发行版自带了大量数据科学库非常适合机器学习入门。2.2 创建虚拟环境建议为机器学习项目单独创建虚拟环境避免不同项目之间依赖版本冲突。打开终端或 Anaconda Prompt执行conda create -n ml_decision_tree python3.9 conda activate ml_decision_tree如果你没有安装 Anaconda也可以使用原生 venvpython -m venv ml_decision_tree # Windows ml_decision_tree\Scripts\activate # macOS / Linux source ml_decision_tree/bin/activate2.3 安装依赖库激活环境后执行下面的命令安装依赖pip install numpy pandas scikit-learn matplotlib如果你还需要把决策树导出成图片可以再安装pip install graphviz需要注意的是graphviz 属于可视化增强工具不安装也不影响决策树模型的训练和预测。版本方面scikit-learn 迭代较快不同版本对部分参数的默认值有调整本文示例以常见稳定版本为准如果你的版本不同请根据实际环境调整参数。2.4 验证安装安装完成后运行下面这段代码import numpy as np import pandas as pd import sklearn print(numpy:, np.__version__) print(pandas:, pd.__version__) print(scikit-learn:, sklearn.__version__)如果能正常输出版本号就说明环境已经准备好了。3. 决策树核心原理拆解3.1 决策树的学习目标决策树的学习过程本质上是一个“递归划分特征空间”的过程。算法从根节点开始每次选择一个最优特征按照某个阈值或类别取值把数据集分成两个或多个子集然后对每个子集重复同样的操作直到满足停止条件为止。关键问题在于怎么判断“最优特征”答案是让划分后的子集尽可能“纯”。也就是说每个子集里的样本尽量属于同一个类别。如果一个节点里全是同一类样本那它就已经“纯”了不需要再分裂。3.2 信息熵与信息增益“纯度”在信息论中可以用“信息熵”来衡量。信息熵表示数据集的不确定性公式如下Entropy(D) - Σ p_i * log2(p_i)其中 p_i 表示第 i 类样本在数据集 D 中占的比例。熵越大说明数据越混乱熵越小说明数据越纯。假设一个二分类数据集有 10 个样本其中 5 个正类、5 个负类那么Entropy(D) - (0.5 * log2(0.5) 0.5 * log2(0.5)) 1.0如果 10 个样本全是正类那么Entropy(D) - (1 * log2(1)) 0可以看到当数据完全纯时熵为 0这是最理想的情况。信息增益则表示“划分前后不确定性下降的程度”Gain(D, a) Entropy(D) - Σ (|D_v| / |D|) * Entropy(D_v)其中 a 是特征D_v 是按特征 a 划分后的第 v 个子集。信息增益越大说明用特征 a 划分后数据纯度提升得越多。ID3 算法就是基于信息增益来选择特征的。3.3 信息增益率与 C4.5信息增益有一个问题它倾向于选择取值较多的特征。比如“身份证号”这个特征每个样本的值都不同按它划分后每个子集只有一个样本纯度极高信息增益很大但这样的划分没有任何泛化意义。C4.5 算法使用了信息增益率来修正这个问题在信息增益的基础上除以特征本身的固有值Intrinsic Value从而惩罚取值过多的特征。不过在 sklearn 的决策树实现中默认使用的不是 ID3 或 C4.5而是 CART。3.4 基尼指数与 CARTCARTClassification And Regression Tree是分类与回归树的简称它使用基尼指数Gini Index来选择划分特征。基尼指数的公式是Gini(D) 1 - Σ p_i^2基尼指数越小表示数据集纯度越高。对于特征 a 划分后的基尼指数可以按子集大小加权求和。CART 树的特点是既可以做分类也可以做回归每次划分只做二叉分裂sklearn 中的 DecisionTreeClassifier 和 DecisionTreeRegressor 都是基于 CART 的。3.5 剪枝防止过拟合的关键如果不加限制决策树会一直生长直到每个叶节点都完全纯。这样的树在训练集上表现几乎完美但在新数据上往往表现很差也就是过拟合。剪枝是决策树防止过拟合的核心手段主要分为两种前剪枝预剪枝在构建过程中如果划分带来的增益不够大或者节点样本数太少就提前停止分裂后剪枝后剪枝先把树完整构建出来再从下往上剪掉对验证集没有贡献的子树。实际使用 sklearn 时我们主要通过限制树的最大深度、叶节点最小样本数等参数来实现前剪枝。这部分会在后面的实战中演示。4. 实战基于决策树的收入预测分类下面进入本文的实战部分。我们将使用一个非常经典的收入预测数据集目标是基于个人的年龄、教育程度、职业、周工时等信息预测个人收入是否超过 50K。这个案例在很多机器学习课程实验中出现过例如“实验3-2 决策树进行收入预测-sklearn版”。4.1 数据集说明收入预测数据集通常可以从公开渠道获取。数据集中每行代表一个人主要字段包括字段含义age年龄workclass工作类型education教育程度education-num受教育年数marital-status婚姻状况occupation职业relationship家庭关系race种族sex性别capital-gain资本收益capital-loss资本损失hours-per-week每周工作小时数native-country国籍income收入标签50K 或 50K在课程实验中这个数据集经常以 CSV 格式提供但没有表头。我们在读取时需要手动指定列名。4.2 创建项目结构为了方便代码管理先创建项目目录decision_tree_project/ ├── data/ │ └── adult.csv ├── main.py └── README.md把数据文件放入data目录然后在项目根目录创建main.py。4.3 读取数据与预处理第一步是读取数据并把缺失值、字符型特征处理成模型可以使用的形式。import pandas as pd # 列名定义 columns [ age, workclass, fnlwgt, education, education-num, marital-status, occupation, relationship, race, sex, capital-gain, capital-loss, hours-per-week, native-country, income ] # 读取数据指定列名 data pd.read_csv(data/adult.csv, headerNone, namescolumns, na_values?) print(data.shape) print(data.head()) print(data.info())代码说明headerNone表示数据文件没有表头namescolumns手动指定列名na_values?把数据中的问号视为缺失值。接下来处理缺失值和类别特征# 去掉包含缺失值的行 data data.dropna() # 分离特征和标签 X data.drop(income, axis1) y data[income] print(特征矩阵大小:, X.shape) print(标签分布:) print(y.value_counts())4.4 处理类别特征收入预测数据集中包含大量字符串类型的特征比如 workclass、education、occupation。机器学习模型无法直接处理字符串所以需要把这些文本特征转换成数值形式。处理方式有很多种这里使用 pandas 的 get_dummies 进行独热编码One-Hot EncodingX pd.get_dummies(X) print(X.shape) print(X.columns[:20])独热编码会把每个类别变成一个 0/1 列比如 workclass 有 Private、Self-emp-not-inc 等取值编码后每一列表示“是否为该类别”。这样做的好处是不会给类别强加顺序关系。如果你希望保持简单的处理方式也可以使用 LabelEncoder 把每个类别编码成整数但在决策树场景下独热编码通常更稳妥。不过要注意独热编码后特征维度会变多这对决策树来说不是问题但如果换成线性模型还需要做特征缩放。4.5 划分训练集和测试集训练模型之前需要把数据划分为训练集和测试集。训练集用来训练模型测试集用来评估模型在未见过的数据上的表现。from sklearn.model_selection import train_test_split X_train, X_test, y_train, y_test train_test_split( X, y, test_size0.3, random_state42, stratifyy ) print(训练集大小:, X_train.shape) print(测试集大小:, X_test.shape)参数说明test_size0.330% 的数据作为测试集random_state42固定随机种子保证结果可复现stratifyy按标签比例分层抽样保持训练集和测试集中正负样本比例一致。4.6 训练决策树分类器接下来使用 scikit-learn 的 DecisionTreeClassifier 训练模型。为了便于初学者理解我们先使用一组较为保守的参数from sklearn.tree import DecisionTreeClassifier from sklearn.metrics import accuracy_score, classification_report # 创建决策树分类器 clf DecisionTreeClassifier( criteriongini, max_depth6, min_samples_leaf5, random_state42 ) # 训练模型 clf.fit(X_train, y_train) # 预测 y_train_pred clf.predict(X_train) y_test_pred clf.predict(X_test) # 评估 train_acc accuracy_score(y_train, y_train_pred) test_acc accuracy_score(y_test, y_test_pred) print(训练集准确率:, train_acc) print(测试集准确率:, test_acc) print(\n测试集分类报告:) print(classification_report(y_test, y_test_pred))参数含义criteriongini使用基尼指数作为划分标准也可以换成entropy使用信息熵max_depth6限制树的最大深度为 6防止树无限生长min_samples_leaf5每个叶节点至少包含 5 个样本random_state42固定随机种子。运行后你应该会得到类似下面的输出训练集准确率: 0.86... 测试集准确率: 0.85...这里的数值会根据数据和参数略有浮动但整体规律是训练集准确率略高于测试集准确率这是正常现象。如果训练集准确率接近 100% 而测试集准确率很低就说明模型过拟合了。4.7 输出 AUC 值准确率是分类问题最直接的指标但在正负样本不均衡时可能不够客观。收入预测数据中收入小于等于 50K 的样本通常多于大于 50K 的样本所以还可以看一下 AUC 值from sklearn.metrics import roc_auc_score y_pred_proba clf.predict_proba(X_test)[:, 1] auc roc_auc_score(y_test, y_pred_proba) print(测试集 AUC:, auc)predict_proba返回每个样本属于各个类别的概率取第二列表示属于正类收入 50K的概率。AUC 值越接近 1说明模型区分正负样本的能力越强。4.8 可视化决策树决策树最吸引人的地方就是可以“画出来”。使用 sklearn 自带的plot_tree可以直接展示树结构import matplotlib.pyplot as plt from sklearn.tree import plot_tree plt.figure(figsize(20, 10)) plot_tree( clf, feature_namesX.columns, class_names[50K, 50K], filledTrue, roundedTrue, fontsize8 ) plt.savefig(decision_tree.png, dpi100, bbox_inchestight) plt.show()如果你安装了 graphviz也可以把树导出为 Graphviz 格式from sklearn.tree import export_graphviz export_graphviz( clf, out_filetree.dot, feature_namesX.columns, class_names[50K, 50K], filledTrue, roundedTrue )生成的图片会非常直观。你可以清楚地看到根节点选择了哪个特征进行划分每个节点的样本数、类别分布、基尼指数等信息都能直接看到。这也是决策树可解释性强的体现。4.9 特征重要性分析训练完成后可以查看模型认为哪些特征最重要importance pd.DataFrame({ feature: X.columns, importance: clf.feature_importances_ }).sort_values(importance, ascendingFalse) print(importance.head(10))输出会列出最重要的前几个特征。在收入预测问题中通常年龄、教育水平、每周工作小时数、资本收益等特征重要性较高。这个结果对业务分析很有参考价值因为我们可以知道模型到底“看中”了什么信息。5. 决策树调参与优化实践5.1 为什么要调参默认参数的决策树通常存在两个问题过拟合树太深训练集准确率很高测试集准确率下降欠拟合树太浅训练集和测试集准确率都不高。调参的目标就是在这两者之间找到平衡点。5.2 常用超参数sklearn 中 DecisionTreeClassifier 常见的可调参数如下参数作用调参方向max_depth控制树的最大深度过大容易过拟合常用值 3~15min_samples_split内部节点再划分所需最小样本数增大可防止过拟合min_samples_leaf叶节点最小样本数增大可防止过拟合max_features每次划分考虑的最大特征数减少可增加随机性criterion划分标准基尼指数或信息熵两者效果通常接近可对比5.3 使用 GridSearchCV 搜索最佳参数手工调参效率较低推荐使用 GridSearchCV 进行网格搜索同时结合交叉验证from sklearn.model_selection import GridSearchCV param_grid { max_depth: [4, 6, 8, 10], min_samples_leaf: [1, 5, 10, 20], criterion: [gini, entropy] } clf DecisionTreeClassifier(random_state42) grid_search GridSearchCV( estimatorclf, 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_) best_model grid_search.best_estimator_ test_acc best_model.score(X_test, y_test) print(测试集准确率:, test_acc)说明cv5表示 5 折交叉验证scoringaccuracy指定评估指标n_jobs-1使用所有 CPU 核心加速计算。网格搜索会遍历参数组合虽然耗时较长但在数据量不大的情况下可以接受。通过这种方式我们可以找到一组相对合适的参数而不用完全靠人工试。5.4 观察不同深度下的过拟合现象为了更直观地理解 max_depth 对模型的影响可以循环训练多个模型并绘制准确率变化曲线import numpy as np train_scores [] test_scores [] depths range(1, 21) for depth in depths: model DecisionTreeClassifier(max_depthdepth, random_state42) model.fit(X_train, y_train) train_scores.append(model.score(X_train, y_train)) test_scores.append(model.score(X_test, y_test)) plt.figure(figsize(10, 5)) plt.plot(depths, train_scores, label训练集准确率, markero) plt.plot(depths, test_scores, label测试集准确率, markers) plt.xlabel(树的最大深度 (max_depth)) plt.ylabel(准确率) plt.title(决策树深度对模型效果的影响) plt.legend() plt.grid(True) plt.show()运行后可以看到随着深度增加训练集准确率持续上升但测试集准确率在某个点之后趋于平缓甚至下降。这个“拐点”附近就是比较合适的 max_depth 取值。这个实验对理解决策树的过拟合非常有帮助建议动手跑一下。6. 常见问题与排查思路实际运行决策树代码时经常会遇到一些报错或效果不理想的情况。下面整理了几类常见问题。6.1 分类特征无法处理问题现象常见原因解决思路ValueError: could not convert string to float数据中有字符串特征而没有编码使用pd.get_dummies()或LabelEncoder进行编码pandas 读取后所有数据变成 object 类型读文件时没有指定分隔符或编码错误检查文件格式指定sep,encoding6.2 模型过拟合问题现象常见原因解决思路训练集准确率接近 1.0测试集准确率较低树过深没有限制生长调小max_depth增大min_samples_leaf测试集准确率波动大数据划分不均使用stratifyy分层抽样多折交叉验证6.3 数据不平衡问题收入预测数据中“50K”的样本通常远多于“50K”的样本。如果类别不平衡严重准确率可能虚高因为模型只要全部预测多数类就能拿到很高的准确率。处理思路使用class_weightbalanced让模型自动调整样本权重关注 AUC、F1-score而不是只看准确率尝试对多数类进行下采样或对少数类进行上采样。6.4 可视化中文乱码在使用plot_tree绘图时如果特征名或类别名包含中文图片中可能显示为方块。这通常是因为 matplotlib 没有正确配置中文字体。解决方法是设置中文字体import matplotlib.pyplot as plt plt.rcParams[font.sans-serif] [SimHei, Microsoft YaHei, Arial Unicode MS] plt.rcParams[axes.unicode_minus] False不同操作系统支持的中文字体不同需要根据实际情况选择。6.5 GridSearchCV 运行缓慢网格搜索的参数组合是指数级增长的。如果在数据量大、参数组合多的情况下运行缓慢可以减少cv的折数缩小参数范围把n_jobs设置为较大的数值先用少量数据测试代码逻辑确认无误后再全量运行。7. 决策树的优缺点总结7.1 优点可解释性强树的结构可以直接可视化业务人员也能看懂不需要特征缩放决策树是基于特征值比较的不受量纲影响因此不需要归一化或标准化能处理数值特征和类别特征对缺失值有一定容忍度训练速度快适合作为 Baseline 模型是随机森林、XGBoost 等集成模型的基础。7.2 缺点容易过拟合尤其是树深度不受限制时对数据中的噪声敏感一点扰动可能导致树结构完全不同单棵树的泛化能力通常不如集成模型当类别特征取值极多时划分会有偏向性树模型是分段常数函数对线性关系较强的数据表达效率不高。正因为单棵决策树有这些不足工业实践中更常用随机森林、梯度提升树等集成方法。但理解单棵决策树仍然是理解集成学习的前提。8. 工程实践建议8.1 先跑 Baseline再调参在实际项目中不建议一开始就追求最佳参数。先用默认参数跑通流程得到一个 Baseline 结果再根据训练集和测试集之间的差距判断是过拟合还是欠拟合以此确定调参方向。8.2 数据预处理要谨慎决策树虽然不需要特征缩放但对数据质量仍然敏感。字符串特征必须编码缺失值必须处理。尤其要注意测试集必须使用与训练集相同的编码规则否则模型预测时会报特征不一致的错误。8.3 使用交叉验证评估模型单次划分训练集和测试集具有一定随机性。为了让评估结果更稳定建议使用交叉验证。尤其是为了写期末作业或课程实验交叉验证的结果会更有说服力。8.4 记录模型参数和结果在做实验时建议把每次运行的模型参数、数据规模、准确率、AUC 等指标记录在表格中。这既是良好的工程习惯也方便自己复盘调参过程。8.5 注意数据泄漏在划分训练集和测试集之前不要对全量数据做预处理比如先对全量数据填充缺失值、做标准化。正确做法是先划分数据集再在训练集上计算填充值或缩放参数然后应用到测试集。对于缺失值简单的 dropna 操作泄漏风险较小但涉及填充均值、中位数等操作时就要注意了。8.6 使用 class_weight 处理不平衡如果你使用的是 sklearn 的 DecisionTreeClassifier可以直接设置class_weightbalanced模型会在计算基尼指数时自动放大少数类的权重。这是处理不平衡问题最简单的方式之一。9. 学习路线与下一步建议如果你现在还是一名机器学习初学者下面这条路线可以作为参考第一步理解决策树的基本原理包括信息熵、信息增益、基尼指数、剪枝策略建议阅读周志华《机器学习》西瓜书或李航《统计学习方法》中的决策树章节第二步手动推导一个例子比如用年龄、性别两个特征构建一棵极小的决策树体会特征选择的过程第三步使用 sklearn 完成一个简单的分类任务比如鸢尾花分类或收入预测熟悉 API 和建模流程第四步尝试可视化决策树观察不同深度下模型的变化加深对过拟合的理解第五步学习随机森林、梯度提升树等集成方法理解为什么多个弱学习器组合后效果更好第六步在 Kaggle 等平台上找一个真实数据集把整套流程走一遍。决策树的内容虽然基础但涉及的概念和思想在后面几乎所有机器学习模型中都会用到。建议不要跳过剪枝和过拟合这两块内容它们是你区别于“只会调用 API”的开发者的重要分界线。最后提一点机器学习入门最忌讳“只看不练”。本文的代码量不大建议你亲手运行一遍改一改参数看看结果有什么变化。如果你是在做课程实验也可以把树的可视化图片和分析结论整理到实验报告中这样一份完整的决策树收入预测项目基本就可以作为一次合格的机器学习实验作业了。如果这篇文章对你有帮助欢迎收藏备用也欢迎在评论区交流你在运行决策树时遇到的问题。
返回列表