当前位置:

首页 > 业界资讯 > XGBoost机器学习模型的决策过程

XGBoost机器学习模型的决策过程

使用XGBoost的算法在Kaggle和其它数据科学竞赛中经常可以获得好成绩,因此受到了人们的欢迎。本文用一个具体的数据集分析了XGBoost机器学习模型的预测过程,通过使用可视化手段展示结果,我们可以更好地理解模型的预测过程。随着机器学习的产业应用不断发展,理解、解释和定义机器学习模型的工作原理似乎已成日益明显的趋势。对于非深度学习类型的机器学习分类问题,XGBoost是最流行的库。由于XGBoost可以很好地扩展到大型数据集中,并支持多种语言,它在商业化环境中特别有用。例如,使用X

XGBoost机器学习模型的决策过程

使用 XGBoost 的算法在 Kaggle 和其它数据科学竞赛中经常可以获得好成绩,因此受到了人们的欢迎。本文用一个具体的数据集分析了 XGBoost 机器学习模型的预测过程,通过使用可视化手段展示结果,我们可以更好地理解模型的预测过程。

随着机器学习的产业应用不断发展,理解、解释和定义机器学习模型的工作原理似乎已成日益明显的趋势。对于非深度学习类型的机器学习分类问题,XGBoost 是最流行的库。由于 XGBoost 可以很好地扩展到大型数据集中,并支持多种语言,它在商业化环境中特别有用。例如,使用 XGBoost 可以很容易地在 Python 中训练模型,并把模型部署到 Java 产品环境中。

虽然 XGBoost 可以达到很高的准确率,但对于 XGBoost 如何进行决策而达到如此高的准确率的过程,还是不够透明。当直接将结果移交给客户的时候,这种不透明可能是很严重的缺陷。理解事情发生的原因是很有用的。那些转向应用机器学习理解数据的公司,同样需要理解来自模型的预测。这一点变得越来越重要。例如,谁也不希望信贷机构使用机器学习模型预测用户的信誉,却无法解释做出这些预测的过程。

另一个例子是,如果我们的机器学习模型说,一个婚姻档案和一个出生档案是和同一个人相关的(档案关联任务),但档案上的日期暗示这桩婚姻的双方分别是一个很老的人和一个很年轻的人,我们可能会质疑为什么模型会将它们关联起来。在诸如这样的例子中,理解模型做出这样的预测的原因是非常有价值的。其结果可能是模型考虑了名字和位置的独特性,并做出了正确的预测。但也可能是模型的特征并没有正确考虑档案上的年龄差距。在这个案例中,对模型预测的理解可以帮助我们寻找提升模型性能的方法。

在这篇文章中,我们将介绍一些技术以更好地理解 XGBoost 的预测过程。这允许我们在利用 gradient boosting 的威力的同时,仍然能理解模型的决策过程。

为了解释这些技术,我们将使用 Titanic 数据集。该数据集有每个泰坦尼克号乘客的信息(包括乘客是否生还)。我们的目标是预测一个乘客是否生还,并且理解做出该预测的过程。即使是使用这些数据,我们也能看到理解模型决策的重要性。想象一下,假如我们有一个关于最近发生的船难的乘客数据集。建立这样的预测模型的目的实际上并不在于预测结果本身,但理解预测过程可以帮助我们学习如何最大化意外中的生还者。

import pandas as pd
from xgboost import XGBClassifier
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score
import operator
import matplotlib.pyplot as plt
import seaborn as sns
import lime.lime_tabular
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import Imputer
import numpy as np
from sklearn.grid_search import GridSearchCV
%matplotlib inline

我们要做的首件事是观察我们的数据,你可以在 Kaggle 上找到(https://www.kaggle.com/c/titanic/data)这个数据集。拿到数据集之后,我们会对数据进行简单的清理。即:

  • 清除名字和乘客 ID
  • 把分类变量转化为虚拟变量
  • 用中位数填充和去除数据

这些清洗技巧非常简单,本文的目标不是讨论数据清洗,而是解释 XGBoost,因此这些都是快速、合理的清洗以使模型获得训练。

data = pd.read_csv("./data/titantic/train.csv")
y = data.Survived
X = data.drop(["Survived", "Name", "PassengerId"], 1)
X = pd.get_dummies(X)

现在让我们将数据集分为训练集和测试集。

X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.33, random_state=42)

并通过少量的超参数测试构建一个训练管道。

pipeline = Pipeline(
[('imputer', Imputer(strategy='median')),
('model', XGBClassifier())])
parameters = dict(model__max_depth=[3, 5, 7],
model__learning_rate=[.01, .1],
model__n_estimators=[100, 500])
cv = GridSearchCV(pipeline, param_grid=parameters)
cv.fit(X_train, y_train)

接着查看测试结果。为简单起见,我们将会使用与 Kaggle 相同的指标:准确率。

test_predictions = cv.predict(X_test)
print("Test Accuracy: {}".format(
accuracy_score(y_test, test_predictions)))

Test Accuracy: 0.8101694915254237

至此我们得到了一个还不错的准确率,在 Kaggle 的大约 9000 个竞争者中排到了前 500 名。因此我们还有进一步提升的空间,但在此将作为留给读者的练习。

我们继续关于理解模型学习到什么的讨论。常用的方法是使用 XGBoost 提供的特征重要性(feature importance)。特征重要性的级别越高,表示该特征对改善模型预测的贡献越大。接下来我们将使用重要性参数对特征进行分级,并比较相对重要性。

fi = list(zip(X.columns, cv.best_estimator_.named_steps['model'].feature_importances_))
fi.sort(key = operator.itemgetter(1), reverse=True)
top_10 = fi[:10]
x = [x[0] for x in top_10]
y = [x[1] for x in top_10]
top_10_chart = sns.barplot(x, y)
plt.setp(top_10_chart.get_xticklabels(), rotation=90)

XGBoost机器学习模型的决策过程

从上图可以看出,票价和年龄是很重要的特征。我们可以进一步查看生还/遇难与票价的相关分布:

sns.barplot(y_train, X_train['Fare'])

XGBoost机器学习模型的决策过程

我们可以很清楚地看到,那些生还者相比遇难者的平均票价要高得多,因此把票价当成重要特征可能是合理的。

特征重要性可能是理解一般的特征重要性的不错方法。假如出现了这样的特例,即模型预测一个高票价的乘客无法获得生还,则我们可以得出高票价并不必然导致生还,接下来我们将分析可能导致模型得出该乘客无法生还的其它特征。

这种个体层次上的分析对于生产式机器学习系统可能非常有用。考虑其它例子,使用模型预测是否可以某人一项贷款。我们知道信用评分将是模型的一个很重要的特征,但是却出现了一个拥有高信用评分却被模型拒绝的客户,这时我们将如何向客户做出解释?又该如何向管理者解释?

幸运的是,近期出现了华盛顿大学关于解释任意分类器的预测过程的研究。他们的方法称为 LIME,已经在 GitHub 上开源(https://github.com/marcotcr/lime)。本文不打算对此展开讨论,可以参见论文(https://arxiv.org/pdf/1602.04938.pdf)

接下来我们尝试在模型中应用 LIME。基本上,首先需要定义一个处理训练数据的解释器(我们需要确保传递给解释器的估算训练数据集正是将要训练的数据集):

X_train_imputed = cv.best_estimator_.named_steps['imputer'].transform(X_train)
explainer = lime.lime_tabular.LimeTabularExplainer(X_train_imputed,
feature_names=X_train.columns.tolist(),
class_names=["Not Survived", "Survived"],
discretize_continuous=True)

随后你必须定义一个函数,它以特征数组为变量,并返回一个数组和每个类的概率:

model = cv.best_estimator_.named_steps['model']
def xgb_prediction(X_array_in):
if len(X_array_in.shape) < 2:
X_array_in = np.expand_dims(X_array_in, 0)
return model.predict_proba(X_array_in)

最后,我们传递一个示例,让解释器使用你的函数输出特征数和标签:

X_test_imputed = cv.best_estimator_.named_steps['imputer'].transform(X_test)
exp = explainer.explain_instance(
X_test_imputed[1],
xgb_prediction,
num_features=5,
top_labels=1)
exp.show_in_notebook(show_table=True,
show_all=False)

XGBoost机器学习模型的决策过程

在这里我们有一个示例,76% 的可能性是不存活的。我们还想看看哪个特征对于哪个类贡献最大,重要性又如何。例如,在 Sex = Female 时,生存几率更大。让我们看看柱状图:

sns.barplot(X_train['Sex_female'], y_train)

XGBoost机器学习模型的决策过程

所以这看起来很有道理。如果你是女性,这就大大提高了你在训练数据中存活的几率。所以为什么预测结果是「未存活」?看起来 Pclass =2.0 大大降低了存活率。让我们看看:

sns.barplot(X_train['Pclass'], y_train)

XGBoost机器学习模型的决策过程

看起来 Pclass 等于 2 的存活率还是比较低的,所以我们对于自己的预测结果有了更多的理解。看看 LIME 上展示的 top5 特征,看起来这个人似乎仍然能活下来,让我们看看它的标签:

y_test.values[0]>>>1

这个人确实活下来了,所以我们的模型有错!感谢 LIME,我们可以对问题原因有一些认识:看起来 Pclass 可能需要被抛弃。这种方式可以帮助我们,希望能够找到一些改进模型的方法。

本文为读者提供了一个简单有效理解 XGBoost 的方法。希望这些方法可以帮助你合理利用 XGBoost,让你的模型能够做出更好的推断。

本文内容来源于互联网,如有侵权请联系删除。
作者最新文章
业界资讯 机器学习
相关文章 更多
机器学习之父MichaelJordan:AGl是炒作,AI的下一个战场是经济学
机器学习之父MichaelJordan:AGl是炒作,AI的下一个战场是经济学

机器学习奠基人MichaelJordan认为AGI是公关炒作,AI的真正挑战在于让系统学会协调。他主张机器学习必须与经济学结合,关注不确定性、激励机制和数据流动等社会系统性问题,而非单纯追求模型能力。

手写线程池,对照学习ThreadPoolExecutor线程池实现原理!
手写线程池,对照学习ThreadPoolExecutor线程池实现原理!

持续坚持原创输出,点击蓝字关注我吧 ❝ 沉淀、分享、成长,让自己和他人都能有所收获! ❞ 目录 一、前言二、面试题三、线程池讲解1. 先看个例子2. 手写一个线程池3. 线程池源码分析四、总结五、系列推荐一、前言人看手机,机器学习!正好是2020年,看到这张图还是蛮有意思的。以前小时候总会看到一些科

deepseek-智能助手入口
deepseek-智能助手入口

当前数字时代,deepseek入口智能助手凭借前沿技术架构为用户带来革新体验 在如今这个数字浪潮翻涌的时代,deepseek入口智能助手凭借其前沿的技术架构,为用户带来了真正意义上的革新体验。这个平台巧妙融合了自然语言处理与知识图谱技术,能够精准解析那些看似复杂的语义需求,已然成为提升工作效率的智能

用相关性分析消除“冗余特征”
用相关性分析消除“冗余特征”

刚开始学机器学习时,我也像大部分人 一样,有个很朴素的想法:特征越多,模型能学到的信息就越多,效果自然越好。 所以每次拿到数据,我就拼命往模型里塞特征——能算出来的统计量全加上,能衍生出来的比率全拼上。 一个原本 20 列的数据集,经常被我搞到七八十列。 然后我就发现事情不太对劲了: 训练时间明显变

用方差阈值过滤掉“惰性特征”
用方差阈值过滤掉“惰性特征”

机器学习实战中,大家往往把精力花在调参和选模型上,却忽略了一个更基础的问题:喂给模型的数据里,有多少特征是真正有用的? 刚开始学机器学习那会儿,很多人特别喜欢堆特征——不管有用没用先一股脑全塞进去,总觉得特征越多模型越聪明。直到有一次跑练习数据集,X_train.shape 直接干到了 (50000

超细空气颗粒物每年致约199万人早亡
超细空气颗粒物每年致约199万人早亡

空气中直径小于100纳米的超细颗粒物每年导致约199万人早逝,其中近半因心血管疾病。该颗粒物主要来自化石燃料,可通过呼吸进入血液,引发氧化应激和内皮功能紊乱,是未被充分认识的心血管危险因素。若年均浓度控制在每立方厘米5000个以下,全球超额死亡率可降低约45%。

当我们用错误的尺子量对了东西,会发生什么?
当我们用错误的尺子量对了东西,会发生什么?

这项研究来自Ezgi Korkmaz,论文以预印本形式发布于2026年7月8日,编号为arXiv:2607.07769,发表于cs.LG(机器学习)领域,并将出版于2026年人工智能促进协会(AAAI)会议论文集。有兴趣深入了解的读者可以通过arXiv编号2607.07769查询完整论文。 深度强化

学少儿编程还是机器人编程
学少儿编程还是机器人编程

机器人编程侧重硬件组装与程序驱动,编程受限于特定机器人;少儿编程系统化教授编程知识,从Scratch启蒙延伸至Python、C++等高级语言。两者各有侧重,选择取决于希望孩子提升的能力方向。

少儿编程分类全解析:软件编程 VS 硬件编程,孩子该怎么选?
少儿编程分类全解析:软件编程 VS 硬件编程,孩子该怎么选?

少儿编程分为软件与硬件两大方向。软件编程侧重逻辑与代码能力,含图形化、Python、C++,适配学业竞赛;硬件编程侧重动手与工程实践,含乐高搭建与机器人。选择应基于孩子兴趣与年龄。

Python 基础(一):入门必备知识
Python 基础(一):入门必备知识

学习Python,从基础语法开始是最稳妥的路径。这篇文章整理了Python入门阶段必须掌握的核心概念,包括标识符、关键字、引号、编码、输入输出、缩进、多行、注释、数据类型以及运算符等,后面还整理了基础进阶、爬虫、自动化、数据分析、小游戏、趣味项目以及自学路线等系列内容,方便你按需查阅。 目录 1 标

查看更多
精品专题 更多
装机必备
装机必备

正软商城装机必备专区,精选办公、浏览器、安全防护、影音播放、压缩解压、设计创作和系统工具等电脑常用正版软件,帮助用户快速完成新电脑软件配置。

Windows
Windows

正软商城Windows软件专区,汇集适用于Windows电脑的办公、设计、安全防护、影音播放、开发工具和系统优化软件,提供软件介绍、系统要求、正版授权及购买下载服务。

macOS软件
macOS软件

正软商城macOS软件专区,精选适用于Mac电脑的办公、设计、影音、效率、开发和系统工具,提供软件功能介绍、macOS兼容版本、正版授权及购买下载服务。

Mac软件 更多
灵活计算器
灵活计算器
macOS/iOS/Android

灵活计算器是一款笔记式算数应用,支持实时计算、动态关联和云端同步功能。记录、整理和输出之间的过渡会更自然,适合长期写作、做笔记或持续沉淀个人内容。

赤友清理大师
赤友清理大师
macOS

赤友清理大师是一款为 Mac 设计的智能清理优化工具,可精准扫描垃圾、大文件、重复文件等,释放磁盘空间。做扫描整理、文字提取和表格转换时,它能把识别后的处理步骤接得更顺,资料录入这类场景会省下不少时间。

极度公式
极度公式
Windows/macOS/Linux

极度公式是一款跨平台专业LaTeX公式识别编辑软件,支持OCR公式识别和多平台编辑。和使用说明,避免使用,享受完整功能与稳定支持。做扫描整理、文字提取和表格转换时,它能把识别后的处理步骤接得更顺,资料录入这类场景会省下不少时间。

WINDOWS 更多
Windows 10
Windows 10
Windows

Windows 10 是一款微软推出的经典操作系统,拥有硬件兼容性与多任务处理能力。它更偏向把系统状态查看和常用调节动作放在一起,适合需要持续观察和微调设备状态的场景。

极度公式
极度公式
Windows/macOS/Linux

极度公式是一款跨平台专业LaTeX公式识别编辑软件,支持OCR公式识别和多平台编辑。和使用说明,避免使用,享受完整功能与稳定支持。做扫描整理、文字提取和表格转换时,它能把识别后的处理步骤接得更顺,资料录入这类场景会省下不少时间。

密码键盘
密码键盘
Windows/macOS/iOS/Android

密码键盘是一款兼具安全性与便捷性的高效密码管理器。日常使用里的持续防护和信息管理会更突出,适合把安全控制放进长期使用流程中的场景。