基于Python的机器学习系列(18):梯度提升分类(Gradient Boosting Classification)

news/2024/11/13 9:02:13/

简介

        梯度提升(Gradient Boosting)是一种集成学习方法,通过逐步添加新的预测器来改进模型。在回归问题中,我们使用梯度来最小化残差。在分类问题中,我们可以利用梯度提升来进行二分类或多分类任务。与回归不同,分类问题需要使用如softmax这样的概率模型来处理类别标签。

梯度提升分类的工作原理

        梯度提升分类的基本步骤与回归类似,但在分类任务中,我们使用概率模型来处理预测结果:

  1. 初始化模型:选择一个初始预测器,这里使用DummyClassifier来作为第一个模型。
  2. 计算梯度:计算每个样本的梯度,梯度是当前预测值与真实标签之间的差异。
  3. 训练新预测器:用计算得到的梯度作为目标,训练一个新的分类器。
  4. 更新模型:将新预测器的结果加到现有模型中。
  5. 重复步骤:重复上述步骤,逐步添加更多的预测器以改进模型的分类能力。

分类示例

        在二分类任务中,梯度提升分类器的工作流程如下:

  1. 预测概率:通过softmax将预测值转换为概率。
  2. 更新模型:利用当前的梯度来训练下一个分类器。

代码示例

        下面的代码示例展示了如何实现一个梯度提升分类器,包括支持二分类和多分类任务:

python">from sklearn.tree import DecisionTreeRegressor
from sklearn.dummy import DummyRegressor, DummyClassifier
from sklearn.model_selection import train_test_split
from sklearn.datasets import load_digits, load_breast_cancer
import numpy as npclass GradientBoosting:def __init__(self, S=5, learning_rate=1, max_depth=1, min_samples_split=2, regression=True, tol=1e-4):self.S = Sself.learning_rate = learning_rateself.max_depth = max_depthself.min_samples_split = min_samples_splitself.regression = regression# 初始化回归树tree_params = {'max_depth': self.max_depth, 'min_samples_split': self.min_samples_split}self.models = [DecisionTreeRegressor(**tree_params) for _ in range(S)]if regression:# 回归模型的初始模型self.models.insert(0, DummyRegressor(strategy='mean'))else:# 分类模型的初始模型self.models.insert(0, DummyClassifier(strategy='most_frequent'))def grad(self, y, h):return y - hdef fit(self, X, y):# 训练第一个模型self.models[0].fit(X, y)for i in range(self.S):# 预测yhat = self.predict(X, self.models[:i+1], with_argmax=False)# 计算梯度gradient = self.grad(y, yhat)# 训练下一个模型self.models[i+1].fit(X, gradient)def predict(self, X, models=None, with_argmax=True):if models is None:models = self.modelsh0 = models[0].predict(X)boosting = sum(self.learning_rate * model.predict(X) for model in models[1:])yhat = h0 + boostingif not self.regression:# 使用softmax转换为概率yhat = np.exp(yhat) / np.sum(np.exp(yhat), axis=1, keepdims=True)if with_argmax:yhat = np.argmax(yhat, axis=1)return yhat# 示例:使用乳腺癌数据集进行二分类
X, y = load_breast_cancer(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.3, random_state=42)# 创建和训练梯度提升分类器
gb = GradientBoosting(S=50, learning_rate=0.1, regression=False)
gb.fit(X_train, y_train)# 预测并计算准确率
y_pred = gb.predict(X_test)
from sklearn.metrics import accuracy_score
print(f'Accuracy: {accuracy_score(y_test, y_pred)}')

总结

        梯度提升分类器通过逐步减少分类错误来提高模型的性能。这种方法在处理分类任务时,能够有效提高预测准确率。与回归任务类似,分类任务中的梯度提升也能通过逐步添加预测器来优化模型。通过调整学习率和模型参数,我们可以进一步提高模型的表现。

如果你觉得这篇博文对你有帮助,请点赞、收藏、关注我,并且可以打赏支持我!

欢迎关注我的后续博文,我将分享更多关于人工智能、自然语言处理和计算机视觉的精彩内容。

谢谢大家的支持!


http://www.ppmy.cn/news/1518900.html

相关文章

2024.8.31 Python,合并区间,用sort通过列表第一个元素给列表排序,三数之和,跳跃游戏

1.合并区间 以数组 intervals 表示若干个区间的集合,其中单个区间为 intervals[i] [starti, endi] 。请你合并所有重叠的区间,并返回 一个不重叠的区间数组,该数组需恰好覆盖输入中的所有区间 。 示例 1: 输入:inter…

ARCGIS 纸质小班XY坐标转电子要素面(2)

本章用于说明未知坐标系情况下如何正确将XY转要素面 背景说明 现有资料:清除大概位置,纸质小班图,图上有横纵坐标,并已知小班XY拐点坐标,但未知坐标系。需要上图 具体操作 大部分操作同这边文章ARCGIS 纸质小班XY…

Java | Leetcode Java题解之第387题字符串中的第一个唯一字符

题目&#xff1a; 题解&#xff1a; class Solution {public int firstUniqChar(String s) {Map<Character, Integer> position new HashMap<Character, Integer>();Queue<Pair> queue new LinkedList<Pair>();int n s.length();for (int i 0; i …

模型 错位竞争(战略规划)

系列文章 分享 模型&#xff0c;了解更多&#x1f449; 模型_思维模型目录。与其更好&#xff0c;不如不同。 1 错位竞争的应用 1.1 美团的错位竞争策略 美团&#xff0c;作为中国领先的电子商务平台&#xff0c;面临着阿里巴巴等电商巨头的竞争压力。为了在市场中获得独特的…

MATLAB虫害检测预警系统

一、课题介绍 本课题是基于MATLAB颜色的植物虫害检测识别&#xff0c;可以辨析植物叶子属于是轻度虫害&#xff0c;中度虫害&#xff0c;严重虫害&#xff0c;正常等四个级别。算法流程&#xff1a;每种等级叶子分别放在同一个文件夹&#xff0c;训练得到每个文件夹每个叶…

续:MySQL的gtid模式

为什么要启用gtid? master端和slave端有延迟 ##设置gtid master slave1 slave2 [rootmysql1 ~]# vim /etc/my.cnf [rootmysql1 ~]# cat /etc/my.cnf [mysqld] datadir/data/mysql socket/data/mysql/mysql.sock symbolic-links0 log-binmysql-bin server-id1 slow_query_lo…

AI学习指南深度学习篇-门控循环单元中的门控机制

AI学习指南深度学习篇-门控循环单元中的门控机制 引言 深度学习是当前人工智能领域的一个重要方向&#xff0c;而循环神经网络&#xff08;RNN&#xff09;在处理序列数据方面展现出了强大的能力。然而&#xff0c;标准的RNN在处理长序列时存在长期依赖问题&#xff0c;容易导…

jenkins安装k8s插件发布服务

1、安装k8s插件 登录 Jenkins&#xff0c;系统管理→ 插件管理 → 搜索 kubernetes&#xff0c;选择第二个 Kubernetes&#xff0c;点击 安装&#xff0c;安装完成后重启 Jenkins 。 2、对接k8s集群、申请k8s凭据 因为 Jenkins 服务器在 kubernetes 集群之外&#xff0c;所以…