梯度下降(Gradient Descent)原理以及Python代码

2024-05-26 08:48

本文主要是介绍梯度下降(Gradient Descent)原理以及Python代码,希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

给定一个函数f(x),我们想知道当x是值是多少的时候使这个函数达到最小值。为了实现这个目标,我们可以使用梯度下降(Gradient Descent)进行近似求解。

梯度下降是一个迭代算法,具体地,下一次迭代令

x_{n+1} = x_{n} - \eta {f}'(x_{n})

{f}'(x)是梯度,其中\eta是学习率(learning rate),代表这一轮迭代使用多少负梯度进行更新。梯度下降非常简单有效,但是其中的原理是怎么样呢?

原理

为什么每次使用负梯度进行更新呢?这要从泰勒公式(Taylor's formula)说起:

f(x) = f(x_{0}) + \frac{​{f}'(x_{0})}{1!}(x-x_{0}) + \frac{f{}''(x_{0})}{2!}(x-x_{0}) + ...

泰勒公式的目的是使用x-x_{0}的多项式去逼近函数f(x),这里可以理解泰勒公式在x-x_{0}的展开是原函数的一个近似函数。

那泰勒公式跟梯度下降有什么关系呢?

我们的目标是使f(x_{n+1})\leq f(x_{n}),我们对f(x_{n+1})x_{n}处进行一阶泰勒展开:

f(x_{n+1}) \approx f(x_{n}) + {f}'(x_{n})(x_{n+1}-x_{n})

由此可知,我们只需令x_{n+1}-x_{n} = -{f}'(x_{n}),就会使f(x_{n+1})\leq f(x_{n})

所以迭代公式可以为x_{n+1}= x_{n} -\eta {f}'(x_{n})

案例

下面我们看具体例子,假设我们有以下函数

f(x) = \frac{1}{2}\left \| Ax-b \right \|^2

矩阵和A向量b已知,我们想知道当x取值为多少的时候,函数f(x)的值最小。

根据梯度下降法,我们只需计算出负梯度,然给定一个初始值x_{0},不断迭代就能找到一个近似解了。负梯度计算如下:

{f}'(x) =A ^{T}(Ax-b)=A ^{T}Ax-A ^{T}b

接下来让我写一段代码解决这个问题

定义梯度下降函数

首先,定义cal_gradient函数用来计算梯度,然后使用gradient_decent进行迭代,其中learning_rate就是公式中的\eta,这个值需要合理设置,过大的话会导致震荡,过下的话又会导致迭代时间过长。step代表迭代的次数,理想情况下找到满意的解就停止。

我们会在代码中调整这两个参数查看它们对求解过程的影响。

import numpy as np
import time#calculate gradient
def cal_gradient(A, b, x):left = np.dot(np.dot(A.T, A), x)right = np.dot(A.T, b)gradient = left - rightreturn gradient# iteration
def gradient_decent(x, A, b, learning_rate, step):start = time.time()for i in range(step):gradient = cal_gradient(A, b, x)delta = learning_rate * gradientx = x - deltaend = time.time()time_cost = round(end - start, 4)print('done! x = {a}, time cost = {b}s'.format(a=x, b=time_cost))

求解过程

我们给了矩阵和A向量b的值以及标准答案 [29, 16, 3],然后我们随机初始化一个x_{0},让学习率\eta =0.01,迭代次数step=1000000

A = np.array([[1.0, -2.0, 1.0], [0.0, 2.0, -8.0], [-4.0, 5.0, 9.0]])
b = np.array([0.0, 8.0, -9.0])
# Giveb A and b,the solution x is [29, 16, 3]x0 = np.array([1.0, 1.0, 1.0])
learning_rate = 0.01
step = 1000000gradient_decent(x0, A, b, learning_rate, step)

结果

以下为结果,可以看出求得的近似解和标准答案 [29, 16, 3]还是非常接近的。

done! x = [28.98272933 15.99042465  2.99763054], time cost = 4.6037s

调整学习率

其他参数都一样,我们让学习率变小,运行相同的步数,从以下结果看到求得的近似解跟标准答案还有一定差距。这意味着小的过小学习率需要学习更久的时间。

learning_rate = 0.001# result
# done! x = [15.8048349   8.68422815  1.18968306], time cost = 4.5997s

调整初始值

我们只调整初始值,学习相同的步数,发现求得的近似解尽管与标准答案相似,但是不如第一个方法求得解。这说明梯度下降方法也会受到初始值得影响。

x0 = np.array([1000, 1000, 1000])# result
# done! x = [29.78036839 16.43265826  3.10706301], time cost = 4.5528s

总结

梯度下降方法是一种非常有效的优化方法,它的效果会受到初始值、学习率、步数的影响。如果要说缺点的话,就是它容易找到局部最优解,有时候会发生震荡现象。

 

参考

https://sm1les.com/2019/03/01/gradient-descent-and-newton-method/

这篇关于梯度下降(Gradient Descent)原理以及Python代码的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/1003996

相关文章

Python中的filter() 函数的工作原理及应用技巧

《Python中的filter()函数的工作原理及应用技巧》Python的filter()函数用于筛选序列元素,返回迭代器,适合函数式编程,相比列表推导式,内存更优,尤其适用于大数据集,结合lamb... 目录前言一、基本概念基本语法二、使用方式1. 使用 lambda 函数2. 使用普通函数3. 使用 N

Python如何实现高效的文件/目录比较

《Python如何实现高效的文件/目录比较》在系统维护、数据同步或版本控制场景中,我们经常需要比较两个目录的差异,本文将分享一下如何用Python实现高效的文件/目录比较,并灵活处理排除规则,希望对大... 目录案例一:基础目录比较与排除实现案例二:高性能大文件比较案例三:跨平台路径处理案例四:可视化差异报

python之uv使用详解

《python之uv使用详解》文章介绍uv在Ubuntu上用于Python项目管理,涵盖安装、初始化、依赖管理、运行调试及Docker应用,强调CI中使用--locked确保依赖一致性... 目录安装与更新standalonepip 安装创建php以及初始化项目依赖管理uv run直接在命令行运行pytho

Python中yield的用法和实际应用示例

《Python中yield的用法和实际应用示例》在Python中,yield关键字主要用于生成器函数(generatorfunctions)中,其目的是使函数能够像迭代器一样工作,即可以被遍历,但不会... 目录python中yield的用法详解一、引言二、yield的基本用法1、yield与生成器2、yi

深度解析Python yfinance的核心功能和高级用法

《深度解析Pythonyfinance的核心功能和高级用法》yfinance是一个功能强大且易于使用的Python库,用于从YahooFinance获取金融数据,本教程将深入探讨yfinance的核... 目录yfinance 深度解析教程 (python)1. 简介与安装1.1 什么是 yfinance?

Python脚本轻松实现检测麦克风功能

《Python脚本轻松实现检测麦克风功能》在进行音频处理或开发需要使用麦克风的应用程序时,确保麦克风功能正常是非常重要的,本文将介绍一个简单的Python脚本,能够帮助我们检测本地麦克风的功能,需要的... 目录轻松检测麦克风功能脚本介绍一、python环境准备二、代码解析三、使用方法四、知识扩展轻松检测麦

Python多线程应用中的卡死问题优化方案指南

《Python多线程应用中的卡死问题优化方案指南》在利用Python语言开发某查询软件时,遇到了点击搜索按钮后软件卡死的问题,本文将简单分析一下出现的原因以及对应的优化方案,希望对大家有所帮助... 目录问题描述优化方案1. 网络请求优化2. 多线程架构优化3. 全局异常处理4. 配置管理优化优化效果1.

IDEA与MyEclipse代码量统计方式

《IDEA与MyEclipse代码量统计方式》文章介绍在项目中不安装第三方工具统计代码行数的方法,分别说明MyEclipse通过正则搜索(排除空行和注释)及IDEA使用Statistic插件或调整搜索... 目录项目场景MyEclipse代码量统计IDEA代码量统计总结项目场景在项目中,有时候我们需要统计

MyBatis-Plus 与 Spring Boot 集成原理实战示例

《MyBatis-Plus与SpringBoot集成原理实战示例》MyBatis-Plus通过自动配置与核心组件集成SpringBoot实现零配置,提供分页、逻辑删除等插件化功能,增强MyBa... 目录 一、MyBATis-Plus 简介 二、集成方式(Spring Boot)1. 引入依赖 三、核心机制

MySQL设置密码复杂度策略的完整步骤(附代码示例)

《MySQL设置密码复杂度策略的完整步骤(附代码示例)》MySQL密码策略还可能包括密码复杂度的检查,如是否要求密码包含大写字母、小写字母、数字和特殊字符等,:本文主要介绍MySQL设置密码复杂度... 目录前言1. 使用 validate_password 插件1.1 启用 validate_passwo