尧图网络 高端网站定制 · 原创设计
免费咨询热线
400-888-6620
免费获取方案
10_自动微分模块与计算图[pytorch框架与神经网络基础]
8.自动微分模块训练神经网络时最常用的算法就是反向传播。在该算法中参数(模型权重)会根据损失函数关于对应参数的梯度进行调整。为了计算这些梯度PyTorch内置了名为 torch.autograd 的微分引擎(自动微分模块)。一般来说,pytorch中模型的计算流程为前向计算 → 构建计算图 → 计算损失 → 反向传播 → 自动微分 → 优化器更新参数 自动微分负责“怎么算梯度” 反向传播是计算梯度的一种高效算法利用链式法则从后往前算 优化器拿到梯度后负责“怎么更新参数”8.1 前向传播与计算图的构建前向传播把输入数据xxx沿着神经网络层层递进,经过所有参数和激活函数的运算,最终得到输出结果的过程。计算图是一个有向无环图(DAG),由两部分组成节点(Nodes)张量(数据),参数(w和b)或操作(如加、减、乘、卷积、激活函数)。pytorch中,直接创建的张量被称为叶子节点(Leaf Tensor)边(Edges)代表数据的流向(即张量作为输入,经过运算,产生输出)。每执行一次前向传播,就是遍历一次模型代码,同时pytorch.autograd会遍历其中的所有数学运算,构建出计算图,在调用backward()时提供精确的“导数路线”。requires_grad和梯度累计你可以在张量初始化时指定requires_gradTrue或调用.retain_grad()来声明此张量数据是可训练的参数,对每个requires_gradTrue的叶子节点(Leaf Tensor),都会关联一个AccumulateGrad节点。同时,对requires_gradTrue具有传播性与其有关的张量都具备requires_gradTrue。每个requires_gradTrue的叶子节点(Leaf Tensor)其梯度都会被保存在内存中。AccumulateGrad结点是pytorch中自动微分引擎中一个特殊的、底层的 Function节点它专门负责接收反向传播来的梯度并将其累加到对应叶子张量的.grad属性中,直到调用zero_grad()才会将梯度清零。下图展示了一个简单的计算图,其数学表达式如下lossMSE(y,z),zw∗xb loss MSE(y,z) \quad ,z w * x blossMSE(y,z),zw∗xb在这个例子中,www和bbb的梯度计算方式为∂l∂w∂l∂z⋅∂z∂w2(w∗xb−y)∗x∂l∂b∂l∂z⋅∂z∂b2(w∗xb−y)∗1 \begin{align} \frac{\partial l}{\partial w} \frac{\partial l}{\partial z} \cdot \frac{\partial z}{\partial w} 2(w * x b - y) * x \\ \\ \frac{\partial l}{\partial b} \frac{\partial l}{\partial z} \cdot \frac{\partial z}{\partial b} 2(w * x b - y) * 1 \end{align}​∂w∂l​∂z∂l​⋅∂w∂z​2(w∗xb−y)∗x∂b∂l​∂z∂l​⋅∂b∂z​2(w∗xb−y)∗1​​其中,∂l∂z2(z−y)\frac{\partial l}{\partial z} 2(z - y)∂z∂l​2(z−y),被称为顶层误差信号。并在后续计算当中被传递到∂l∂w\frac{\partial l}{\partial w}∂w∂l​和∂l∂b\frac{\partial l}{\partial b}∂b∂l​之中。如定义初始常量x5,y0,w1,b3x5,y0,w1,b3x5,y0,w1,b3,可以得到其计算图的自顶向下的拓扑结构loss (grad_fnMseLossBackward0) ← 均方误差结点 (关联了z和y以及MSE的求导规则) └── z (grad_fnAddBackward0) ← 加法节点 (关联了b和w*x以及加法求导规则) ├── (左) MulBackward0 ← 乘法节点 (关联了w和x以及乘法求导规则) │ ├── w (叶子, AccumulateGrad) │ └── x (常量, 不追踪) └── (右) b (叶子, AccumulateGrad)其程序及pytorch的计算过程如下# 当X为标量时梯度的计算defscaler_grad_compute():xtorch.tensor(5)# 目标值: labelytorch.tensor(0.)# 设置要更新的权重和偏置的初始值wtorch.tensor(1,requires_gradTrue,dtypetorch.float32)btorch.tensor(3,requires_gradTrue,dtypetorch.float32)# 设置网络的输出值zw*xb# 设置损失函数,并进行损失的计算losstorch.nn.MSELoss()lossloss(z,y)# 自动微分loss.backward()# 打印 w,b 变量的梯度# backward 函数计算的梯度值会存储在张量的 grad 变量中print(fw-{w.grad})print(fb-{b.grad})scaler_grad_compute()计算顶层误差:∂l∂z2(z−y)2∗(wxb−y)16\frac{\partial l}{\partial z} 2(z - y) 2 * (wx b - y) 16∂z∂l​2(z−y)2∗(wxb−y)16,并传递到AddBackward0(加法)节点中,此时loss、y、MSE被“移出”计算图。AddBackward0把误差信号复制成两份因为加法节点的梯度分流(左) 乘法结点根据乘法求导法则计算出www的梯度∂l∂w∂l∂z⋅∂z∂w16∗580\frac{\partial l}{\partial w} \frac{\partial l}{\partial z} \cdot \frac{\partial z}{\partial w} 16 * 5 80∂w∂l​∂z∂l​⋅∂w∂z​16∗580,并传递给其叶子节点(右) 计算bbb的梯度∂l∂b∂l∂z⋅∂z∂b16∗116\frac{\partial l}{\partial b} \frac{\partial l}{\partial z} \cdot \frac{\partial z}{\partial b} 16 * 1 16∂b∂l​∂z∂l​⋅∂b∂z​16∗116,由于其requires_gradTrue且为叶子节点(Leaf Tensor),其梯度都会被保存在内存中 b.grad16。此时、*、x(requires_gradFalse)被“移出”计算图,而w节点接收上层的梯度并保存 w.grad 80。ad16。此时、*、x(requires_gradFalse)被“移出”计算图,而w节点接收上层的梯度并保存 w.grad 80。这种通过计算图拓扑的方法简化了计算梯度的复杂度。这就是反向传播算法(Back Propagation)的核心思想。
RELATED

相关推荐

StoryDiffusion:长序列故事图像生成工具,让多帧漫画中的角色保持一致

StoryDiffusion:长序列故事图像生成工具,让多帧漫画中的角色保持一致

StoryDiffusion:长序列故事图像生成工具,让多帧漫画中的角色保持一致 【免费下载链接】StoryDiffusion Accepted as [NeurIPS 2024] Spotlight Presentation Paper 项目地址: https://gitcode.com/GitHub_Trending/st/StoryDiffusion StoryDiffus…

📅 2026/9/24 14:25:00
IronClaw 能力调用膜(CapabilityHost):特权效果的单一授权入口架构解析

IronClaw 能力调用膜(CapabilityHost):特权效果的单一授权入口架构解析

人工智能AI 应用交互助手AI Agent 【免费下载链接】ironclaw IronClaw is an Agent OS focused on privacy, security and extensibility 项目地址: https://gitcode.com/gh_mirrors/iro/ironclaw 点击查看 免费下载 导读 本文以 crates/kernel/ironclaw_capabili…

📅 2026/9/24 14:25:00
深入解析 OpenKruise:Kubernetes 增强工作负载与原地升级实战指南

深入解析 OpenKruise:Kubernetes 增强工作负载与原地升级实战指南

深入解析 OpenKruise:Kubernetes 增强工作负载与原地升级实战指南 【免费下载链接】kubernetes-handbook Kubernetes 架构与生态:从云原生到 AI 原生基础设施的构建指南 项目地址: https://gitcode.com/gh_mirrors/ku/kubernetes-handbook OpenKr…

📅 2026/9/24 14:25:00
MORE NEWS

更多资讯

📰

【Dify】YouTube全自动内容生成与多平台分发应用

自媒体视频内容的生产和分发,已成为内容创业与个人品牌塑造的重要途径。高效的视频自动化处理工具,能够显著提升内容制作与运营的效率。 本文聚焦于YouTube及多平台自媒体场景,介绍一个覆盖从素材导入、音频转写、语义分析、文案生成、分段整理到成品分发的全流程工作流。通…

📰

WeChatMsg:免费开源,五分钟把微信聊天记录导出成文档,顺带生成年度报告

WeChatMsg:免费开源,五分钟把微信聊天记录导出成文档,顺带生成年度报告 【免费下载链接】WeChatMsg 提取微信聊天记录,将其导出成HTML、Word、CSV文档永久保存,对聊天记录进行分析生成年度聊天报告 项目地址: https:…

📰

RT-Thread 在先楫 HPM6P00EVK(RISC-V 双核)上的 BSP 移植与快速上手实战指南

操作系统嵌入式物联网嵌入式OSRTOS 【免费下载链接】rt-thread RT-Thread is an open source IoT Real-Time Operating System (RTOS). https://rt-thread.github.io/rt-thread/ 项目地址: https://gitcode.com/gh_mirrors/rt/rt-thread 点击查看 免费下载 本篇技术…

📰

【Dify】火车票发票自动识别与智能比对应用

自动化发票识别与核查已成为数字化财务管理的关键一环,特别是在海量火车票数据处理中,传统手工方式难以兼顾效率与准确性。 本文介绍基于深度学习模型与智能节点协作的火车票发票比对工作流,涵盖从图片导入、图像优化到信息提取和比对输出的全流程,适合实际企业和个人票据…

📰

【Dify】智能体驱动的高效智能客服应用

智能客服在互联网服务与健康管理领域快速普及,自动化、个性化的用户响应成为应用升级的核心需求。依托多智能体协作和大模型赋能,智能客服已具备文本、图片等多模态信息处理能力,并支持细致化问题分类与定制化回复。 本文聚焦于一套以Dify为核心的智能客服工作流,展示其节…

📰

templ ADR 0001 解析:如何在文本行内书写 `@Component()` 组件表达式

开发工具代码生成后端 【免费下载链接】templ A language for writing HTML user interfaces in Go. 项目地址: https://gitcode.com/gh_mirrors/te/templ 点击查看 免费下载 符号在 templ 模板语言中用于调用组件表达式(templ element expression&…

TODAY

今日更新

THIS WEEK

本周精选

THIS MONTH

本月热门

读完文章,想聊聊您的网站?

告诉我们您的行业与需求,资深顾问一对一梳理方案与报价,全程免费。

📞 💬