当前位置: 首页 > news >正文

3.一文看懂反向传播:从单个神经元到 PyTorch 自动求导

反向求导,多层次对应一个神经,单个神经元场景

学习这一篇的前提是,已经学会了梯度算法和线性结构算法,不明白的可以去看我之前的文章。

前面看不懂的,直接跳转到 “ 反向传播的流程 ”

底层的数学算法

z 是中间变量 u 的函数,u 是自变量 w(权重)的函数,因此 z 通过 u 间接依赖 w。
公式核心逻辑为:单路相乘

  • 对权重 w 求偏导时,仅存在 z→u→w 一条路径(如果嵌套的跟多,那就继续导),将路径上的偏导数直接相乘即可,无其他分支,无需相加。

注意:

这一篇文章只考虑单个神经元场景,多个场景在下一篇完成。第一张图片是多个神经元的场景,第二章是本文中对应的

底层的运用逻辑

我们知道对一个函数的求导,就是求这个函数的最值。大模型的运用就是面对海量的训练数据,找到一个最贴近输入一个x,输出一个正确的y值。这样我们的大模型在没有对应的y值时,也能通过之前我们训练的函数计算出预算出。所以训练的过程中,我们要不断的求导,直到找到最合适的 w 和 b 。

类比人类的大脑,如图(看不懂的直接看下一张图):

  1. 树突就是我们的影响因数 w 和 b ,但是树突不可能只有一个,是有很多个的。

  2. 轴突,他输出信息的一个过程,在我们的反向传播可以类比成激活函数(他的作用就是放在我们的函数模型太过单一,这里的意思就是嵌套另外一个函数)

上面的类比图只类比了一层,一组 w 和 b(树突)对应 激活函数(轴突)。现在来看多层对应单个神经元的图:

反向传播的流程

好的现在就根据这个来解说反向传播的流程:

  1. 我们先假设值 w1 b1 w2 b2 w3 b3,激活函数是在每一个算完 y = wx + b ,对 f(y) 再进行嵌套的一个函数。

  2. 正向运算,一步一步的计算,得到一个 y_forward 值。

  3. 将 y_forward 与真正的 y_true 进行相减去得到损失值, loss = 1/2*( y_forward - y_true )^2。(为什么要*1/2,这是因为我们方便计算倒数自己规定的,因为导一下就没有常数了,这个不影响,因为反向传播会自己寻找最合适的 w 和 b,进行平方是为了保证非负数)。

经理过这一步了,我们就得到了 w1 b1 w2 b2 w3 b3 y_forward loss这几个值,接下来我们就要进行反向求导,一步一步的得到最佳 w 和 b

  1. 分别对w1 b1 w2 b2 w3 b3进行求导,注意了我们求导的函数是 loss = 1/2*( y_forward - y_true )^2。求导公式如图:

  1. 求导之后,我们得到了w1 b1 w2 b2 w3 b3 他们的导数,然后用跟新公式对他们进行跟新。

  2. 如此反复,我们就能得到一个最佳的 w1 b1 w2 b2 w3 b3值。

代码展示(三层,四层就是多两个参数,思路代码是一样的)
# 这个库里面有自动帮我们计算的求导的方法,就不用自己去手动计算了importtorch# 定义数据x_data=[0.0,1.0,2.0,3.0]y_data=[1.0,3.0,5.0,7.0]# 第1层参数w1=torch.tensor([0.1],requires_grad=True)# 第1层权重b1=torch.tensor([0.0],requires_grad=True)# 第1层偏置# 第2层参数w2=torch.tensor([0.1],requires_grad=True)# 第2层权重b2=torch.tensor([0.0],requires_grad=True)# 第2层偏置# 第3层参数(输出层)w3=torch.tensor([0.1],requires_grad=True)# 第3层权重b3=torch.tensor([0.0],requires_grad=True)# 第3层偏置# 定义学习率lr=0.1# 定义遍历次数epochs=2000# 定义期望函数defforward(x):# 第一层z1=w1*x+b1# a1 = torch.sigmoid(z1) # 激活函数 这里知道了是线性结构就不用激活函数了,如果要的话,就嵌套进去就好了,注意几个传参就是了# 第二层z2=w2*z1+b2# a2 = torch.sigmoid(z2) # 激活函数这里知道了是线性结构就不用激活函数了,如果要的话,就嵌套进去就好了,注意几个传参就是了# 第三层输出y_pred=w3*z2+b3returnz1,z2,y_pred# 定义损失函数defloss_n(y_pred,y_true):return0.5*(y_pred-y_true)**2# 定义变量,方便找到最小的损失值best_loss=float('inf')# 先设成无穷大,方便后面比较best_w1=0best_b1=0best_w2=0best_b2=0best_w3=0best_b3=0# 开始遍历计算forepochinrange(epochs):total_loss=0forx,yinzip(x_data,y_data):# 先将原始数据转化成可以计算的形式# 因为 PyTorch 的运算和自动求导主要针对 tensorx_tensor=torch.tensor([x])y_tensor=torch.tensor([y])# 拿到原始数据z1,z2,y_pred=forward(x_tensor)# 开始计算求导,用链式法则把每一层参数的梯度都求出来loss=loss_n(y_pred,y_tensor)loss.backward()# 梯度跟新公式withtorch.no_grad():w1-=lr*w1.grad b1-=lr*b1.grad w2-=lr*w2.grad b2-=lr*b2.grad w3-=lr*w3.grad b3-=lr*b3.grad# 记录损失,.item()是取里面的值,因为他是tensor对象的数据total_loss+=loss.item()# 清零w1.grad.zero_()b1.grad.zero_()w2.grad.zero_()b2.grad.zero_()w3.grad.zero_()b3.grad.zero_()# 平均损失值total_loss=total_loss/len(x_data)# 找最佳的参数iftotal_loss<best_loss:best_loss=total_loss# 找到最小的参数best_w1=w1.item()best_b1=b1.item()best_w2=w2.item()best_b2=b2.item()best_w3=w3.item()best_b3=b3.item()ifepoch%100==0:# 打印平均损失值print(f'epoch:{epoch}| loss:{total_loss}')print(f'w1:{w1.item()}| b1:{b1.item()}| w2:{w2.item()}| b2:{b2.item()}| w3:{w3.item()}| b3:{b3.item()}')# 打印最终值print("最终参数:==================================================")print("w1 =",w1.item(),"b1 =",b1.item())print("w2 =",w2.item(),"b2 =",b2.item())print("w3 =",w3.item(),"b3 =",b3.item())# 打印最佳值print("最佳参数:==================================================")print("w1 =",best_w1,"b1 =",best_b1)print("w2 =",best_w2,"b2 =",best_b2)print("w3 =",best_w3,"b3 =",best_b3)w_total=w1*w2*w3 b_total=w3*w2*b1+w3*b2+b3print("最终结果============================")print("w =",w_total.item(),"b =",b_total.item())
http://www.jsqmd.com/news/616325/

相关文章:

  • 像个小丑,我花了一周把 APP 做上线,又花两小时改回本地
  • 2026ukey集中管理技术解析:u盾集中管理/网银密钥集中/网银盾安全集中/网银盾集中/ukey安全/选择指南 - 优质品牌商家
  • 加入csdn 5周年
  • C# OnnxRuntime 部署 RMBG-2.0 实现高精度背景去除
  • 2026年热门的层叠式过滤器/过滤器/海宁层叠式过滤器厂家推荐与选型指南 - 行业平台推荐
  • 动态规划——01背包问题、完全背包(python、一维DP)
  • vue数据更新了,但视图不刷新的诡异问题 this.$set
  • 当Windows 10的OneDrive无法彻底卸载时,这个批处理脚本是你的终极解决方案
  • 2026年全网视频去水印实测:6款消除字幕工具上手,哪款更适合你
  • OpenClaw个人知识库:Qwen3-4B驱动文档智能检索
  • 深圳游戏主板品牌怎么选:2026年华硕、七彩虹、技嘉、微星产品线全解析
  • Arduino RTCtime库:标准time.h兼容的DS1307/DS3231驱动
  • 鸿蒙_ArkUI组件同时支持双击和单击事件
  • FastAPI实战:WebSocket vs Socket.IO,这回真给我整明白了!侣
  • 问题解决策略基础算法实现训练1
  • 2026年浙江六甲基二硅氮烷口碑产品推荐分析,耐高温、高硬度、疏水疏油涂层/聚硅氮烷陶瓷先驱体,六甲基二硅氮烷厂家推荐 - 品牌推荐师
  • lvgl-micropython、lv_micropython和lv_binding_micropython到底啥关系?一文读懂永
  • 2026年白发养护品牌盘点:白养黑/禾亚美养发馆/禾亚美加盟/禾亚美效果/禾亚美毛发管理中心/禾亚美白发养护/禾亚美门店/选择指南 - 优质品牌商家
  • 2026年往复式提升机采购指南:液压升降平台、液压升降机、液压货梯、高速提升机、往复式提升机、液压升降台、升降机选择指南 - 优质品牌商家
  • 【国家级数字农业项目技术白皮书节选】:PHP轻量化时序数据处理框架如何扛住每秒8700+传感器上报?
  • 【OpenClaw】通过 Nanobot 源码学习架构---()总体患
  • 2026浙江岗亭企业盘点:台州岗亭、吸烟亭、嘉兴岗亭、宁波岗亭、浙江岗亭、湖州岗亭、移动卫生间、移动厕所、移动垃圾分类房选择指南 - 优质品牌商家
  • GLM-. 全面支持与 Gemini CLI 集成:HagiCode 的多模型进化之路估
  • 2026大板专用瓷砖胶技术解析:德高和亿固瓷砖胶/瓷砖胶十大名牌/瓷砖胶十大品牌/瓷砖胶口碑排行/选择指南 - 优质品牌商家
  • OpenClaw+Qwen3.5-9B组合优势:3个不可替代的使用场景
  • 2026户外监控必选无电无网款!68%人选它,Q1销量榜单给你参考;格行模式值得行业借鉴;AOV低功耗+黑光夜视,解决无电无网痛点
  • 2026年行业内质量好的锻件企业选哪家,压力容器法兰/船用法兰/高温合金法兰/锻件/不锈钢法兰/法兰,锻件厂商找哪家 - 品牌推荐师
  • 2026年口碑好的玻璃钢化粪池/陕西化粪池横向对比厂家推荐 - 行业平台推荐
  • 10分钟搞懂 RAG:大模型如何边检索边生成答案
  • eVTOL 研制必读 | 厘清研制保证与设计保证的边界