news 2026/8/25 13:43:22

刘二大人深度学习实践笔记--梯度下降算法

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
刘二大人深度学习实践笔记--梯度下降算法

目录

梯度下降算法(GD)

代码

随机梯度下降算法 (SGD)

代码

批量梯度下降算法(MBGD)


在寻找loss最低的权重时,之前使用穷举法,但是如果模型中有多个需要确定的参数,会导致需要列举的点过多

我们把找到合适的求函数最小值的方法叫做优化问题,下面介绍梯度下降算法

梯度下降算法(GD)

梯度:目标函数对权重求导,梯度的正方向是上升的,负方向是下降的

更新方向:取梯度的负方向,乘以学习率a(一般a要取得比较小,否则每次调整的过多,会无法收敛)

应用了贪心的算法,对于非凸函数,不一定能得到全局最优解,但是能得到局部的最优解

深度学习的目标函数中,不一定有很多局部最优点,所以经常使用梯度下降算法,但是会存在鞍点

鞍点:这个点的梯度值为0,如果陷入鞍点,就无法移动无法迭代了

代码

import numpy as np import matplotlib.pyplot as plt #对数据集保存,x为输入,y为输出,对应位置为一组 x_data = [1.0, 2.0, 3.0] y_data = [2.0, 4.0, 6.0] # 初始猜测权重 w = 1.0 def forward(x): return x * w # 计算mse def cost(xs,ys): cost = 0 for x, y in zip(xs, ys): y_pred = forward(x) cost += (y_pred - y)**2 return cost/len(xs) # 计算梯度 def gradient(xs, ys): grad = 0 for x, y in zip(xs, ys): grad += 2*x*(x*w - y) return grad / len(xs) w = 1.0 w_list = [] cost_list = [] print('Predict before training', 4, forward(4)) for epoch in range(100): cost_val = cost(x_data, y_data) grad_val = gradient(x_data, y_data) #计算梯度 cost_list.append(cost_val) w -= 0.01 * grad_val #更新 print('Epoch:',epoch, 'w=', w, 'loss=', cost_val) print('Predict after training', 4, forward(4)) plt.plot(np.arange(1, 101, 1), cost_list) plt.ylabel('cost') plt.xlabel('epoch') plt.show()

在训练集上训练时,以epoch为横坐标,mse为纵坐标,正常的函数图像应该是逐渐趋于0收敛的,类似于

如果最后反而上升了,说明训练失败了,结果发散了,可能是因为学习率a取得太大了

随机梯度下降算法 (SGD)

梯度下降算法是将所有样本的损失取平均值作为权重更新的依据

随机梯度下降从N个数据中随机选一个的损失作为权重更新的依据,引入了一个随机噪声,这样即使陷入了鞍点,也很可能脱离鞍点

代码

import numpy as np import matplotlib.pyplot as plt #对数据集保存,x为输入,y为输出,对应位置为一组 x_data = [1.0, 2.0, 3.0] y_data = [2.0, 4.0, 6.0] # 初始猜测权重 w = 1.0 def forward(x): return x * w # 计算mse,修改为只计算单个样本,而不是平均值 def loss(x,y): y_pred = forward(x) return (y_pred - y)**2 # 计算梯度,修改为只计算单个样本,而不是平均值 def gradient(x, y): return 2*x*(x*w - y) w = 1.0 cost_list = [] print('Predict before training', 4, forward(4)) # 修改为在每个样本更新权重,而不是每个轮次更新 for epoch in range(100): epoch_loss = 0 for x, y in zip(x_data, y_data): grad = gradient(x, y) w = w - 0.01 * grad epoch_loss += loss(x, y) print('Epoch:',epoch, 'w=', w, 'epoch_loss=', epoch_loss) cost_list.append(epoch_loss) print('Predict after training', 4, forward(4)) plt.plot(np.arange(1, 101, 1), cost_list) plt.ylabel('loss') plt.xlabel('epoch') plt.show()

注意:

梯度下降算法(GD)中,两个样本之间是可以并行的,不相互影响,因为w每个轮次才改变一次,计算效率高,但是可能陷入鞍点

随机梯度下降算法(SGD),两个样本之间互相影响,因为后面的样本使用的w受前面的改变,计算效率低,但是得到的最优点更好

因此选取折中方案,批量梯度下降算法(MBGD)

批量梯度下降算法(MBGD)

每个batch更新一次权重

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/25 13:38:35

Agent系列

开篇:行业不景气,一个普通程序员决定把会的东西全写出来(附 42 篇系列目录)CSDN 承接版 2026-08-23 同步发抖音/小红书图文(卡片版);本篇是 CSDN 专栏的第一篇 四个免费专栏的开篇说明先交代…

作者头像 李华
网站建设 2026/8/25 13:33:36

跨设备剪贴板共享工程化笔记:可读代码与可复现页面并行推进

手机剪贴板与设备粘贴演示:从输入到历史记录的完整交互 写在前面 很多人第一次看到“跨设备剪贴板”这样的页面标题时,会自然联想到系统级剪贴板、设备发现、账号同步和真正的跨设备传输。不过,判断一个页面能做什么,不能只看标题…

作者头像 李华
网站建设 2026/8/25 13:24:57

模板小程序、SaaS商城和定制开发有什么区别?费用和交付边界对比

模板小程序、SaaS商城和定制开发有什么区别?费用和交付边界对比模板小程序、SaaS商城和定制开发的区别,主要在交付方式和维护责任。模板偏页面和快速搭建,SaaS商城偏成熟后台和年费制使用,定制开发偏复杂功能、源码、接口和私有化…

作者头像 李华
网站建设 2026/8/25 13:24:47

模板建站哪个平台更适合企业?页面效果、后台权限和后期修改对比

模板建站哪个平台更适合企业?页面效果、后台权限和后期修改对比模板建站平台适不适合企业,不能只看模板截图。企业官网至少要经得起三件事:页面效果能不能体现品牌,后台权限能不能分工,后期产品、案例、新闻和表单能不…

作者头像 李华
网站建设 2026/8/25 13:24:32

PLC S7-1200 1214C电源电路故障维修

STD10P6F6(10P6F)∶P-MOS(60V/10A)TSM680P06∶P-MOS(60V/18A)56 1EWX :TPS1H200ASTD10P6F6∶P-MOS(60V/10A)2N7002K(7KU)∶N-MOS(60V/0.15A)SJ 5A:二极管LTV-208SA:双光耦LTV-063L6:高速双光耦(10M) 2颗TSM680P06∶P-MOS(60V/18A) 2颗STD10P6F6∶P-MOS(…

作者头像 李华