ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

深入理解PyTorch自动求导:从计算图到实战避坑指南

2026/9/9 17:02:11 拓冰建站 浏览量
深入理解PyTorch自动求导:从计算图到实战避坑指南 1. 自动求导到底自动了什么从手推公式到框架一行代码1.1 从损失函数到参数更新梯度从哪来刚开始学深度学习的时候我做过一件蠢事照着别人的代码把模型跑通然后看着那一行loss.backward()发呆脑子里全是问号——它到底是怎么知道每个参数该往哪个方向调、调多少的后来翻《动手学深度学习》的自动求导那一节才意识到整个训练循环的根基就藏在这里。所谓训练本质上是在解一个最优化问题找到一个参数组合让损失函数的值尽可能小。而绝大多数优化算法SGD、Adam 这些都依赖一个信息——损失函数对每个参数的梯度。梯度就是损失函数在某个点上对参数的偏导数向量它指向损失增长最快的方向取反就是下降方向。问题来了一个真正的神经网络参数少则几万多则几亿。你不可能像高中数学课那样用笔和纸一条一条地推偏导公式。哪怕真的花一个星期推出来了换一个网络结构又得全部重推。所以框架必须有一种通用、高效、自动的办法把求梯度这件事自动化这就是自动求导Automatic Differentiation常简称 autograd存在的意义。我记得第一次拿 PyTorch 写下面这段代码时心里的石头落地的感觉import torch x torch.arange(4.0, requires_gradTrue) y 2 * torch.dot(x, x) y.backward() print(x.grad) # tensor([ 0., 4., 8., 12.])y 2 * x · x对每个x_i的偏导是4 * x_i所以 x [0,1,2,3] 时梯度是 [0,4,8,12]。框架在backward()之后把结果塞进了x.grad整个过程没有手推任何公式。这篇文章就是想把这一行代码背后发生了什么彻底讲透顺便把我踩过的坑和排查思路都交代出来。1.2 数值求导、符号求导与自动求导为什么只能选第三条路在理解自动求导之前有必要把另外两种求导方式和它做个对比不然你会搞不清楚框架到底聪明在哪里。第一种是数值求导Numerical Differentiation原理就是导数的定义f(x) ≈ (f(x h) - f(x)) / h给一个很小的 h比如 1e-6暴力算差商。这个方法实现起来确实简单写十行代码就能对任意函数求导。但两个致命问题一是慢每求一个参数的导数都要额外跑一次前向计算几百万参数就得多跑几百万次二是不准h 太小会触发浮点舍入误差h 太大会引入截断误差精度两头受气。平时拿它验证手推的梯度对不对可以实际训练根本不能用。第二种是符号求导Symbolic Differentiation就是像 Mathematica 那样对表达式做符号层面的变换把sin(x^2)自动变成2x * cos(x^2)这种数学公式。它的问题是表达式爆炸一个稍微复杂的网络中间变量的符号表达式会指数级膨胀内存和时间都扛不住。而且深度学习代码里有大量控制流循环、分支符号求导对这类动态结构非常不友好。自动求导的思路完全不同它不追求一个完整的符号表达式而是把整个计算过程拆成一堆最基本的运算加减乘除、幂、指数、对数、三角函数这些然后沿着计算顺序逐步运用链式法则把导数算出来而不是推出来。每一步都是数值计算精度和直接计算浮点数一样高每一轮只需要一次前向加一次反向传播所有参数的梯度就都齐了。这条路同时解决了精度、效率和通用性的问题所以现在 PyTorch、TensorFlow、MindSpore 这些主流框架清一色走的是自动求导路线。求导方式精度计算效率实现难度适用场景数值求导受 h 影响不高极慢低验证梯度正确性符号求导高表达式易爆炸高数学公式推导自动求导高高一次前向一次反向中深度学习训练2. 计算图与链式法则反向传播为什么天生适合深度学习2.1 一个简单函数的前向计算图自动求导的地基是计算图Computational Graph。你在框架里做任何一次张量运算框架都会把这个运算记录成图里的一个节点节点之间的边代表数据的流向。这个图在学术上叫有向无环图DAG因为数据只能往前流不能形成循环。举个例子假设有个函数y x1 * x2 sin(x2)把它拆成基本运算前向过程是这样的import torch x1 torch.tensor([2.0], requires_gradTrue) x2 torch.tensor([3.0], requires_gradTrue) a x1 * x2 # 第一步乘法 b torch.sin(x2) # 第二步正弦 y a b # 第三步加法 y.backward() print(x1.grad) # tensor([3.0]) print(x2.grad) # tensor([2.0, cos(3.0)]) 约等于 tensor([1.0100])手推验证一下∂y/∂x1 x2 3∂y/∂x2 x1 cos(x2) 2 cos(3) ≈ 1.01完全一致。在这背后框架悄悄记录了y - a - x1和y - a - x2、y - b - x2这几条依赖关系。backward()调用后它从 y 出发沿着这些依赖反向走在每个节点上应用链式法则一层层把梯度传回去。这里有个关键点计算图是在前向过程中动态构建的。你执行a x1 * x2时图里才出现节点 a你执行y a b时图里才出现节点 y。这也是 PyTorch 被称为动态图框架的原因——图跟着代码走代码里有if、for循环都没关系每跑一遍代码就建一张新图。对比之下TensorFlow 1.x 时代的静态图需要先定义完整的图再执行调试起来痛苦得多这也是 PyTorch 能迅速占领研究社区的一个核心原因。2.2 反向模式为什么胜出参数数量和中间变量的博弈自动求导有两种模式正向模式Forward Mode和反向模式Reverse Mode。正向模式沿着前向计算的方向同时把导数值和函数值一起算出来每跑一次只能得到一个输入变量比如 x1对所有中间变量的导数。如果网络有 100 万个参数就得跑 100 万次前向显然不现实。反向模式则是先做一次标准的前向传播把每个节点的值存下来然后从最终输出损失值开始逆向地应用链式法则把梯度逐层向前传播。它的神奇之处在于只跑一次反向传播就能拿到损失函数对所有参数的梯度。为什么反向模式这么划算核心在于深度学习任务的特殊性——输出端通常是一个标量损失值输入端是海量参数。设输入维度是 n输出维度是 m正向模式的计算量正比于 n反向模式的计算量正比于 m。在训练场景里 n 是千万级别、m 是 1反向模式的效率优势是指数级的。如果把计算图想象成一条河损失是入海口参数是上游千万条支流反向模式相当于从入海口逆流而上一次性勘探完所有支流而正向模式是每条支流单独探一次谁划算一目了然。这也解释了为什么所有深度学习框架都把backward()设计成从损失张量出发——它天然就是反向模式的实现。理解这一点对你后面理解梯度计算顺序、理解为什么可以多次backward()、理解显存里到底存了什么都有直接帮助。3. 动手跑通 autogradPyTorch 求导的最小闭环3.1 标量损失与非标量输出backward 的参数到底怎么传先看最标准的用法。训练时损失一定是个标量所以最常见的求导代码长这样import torch x torch.arange(4.0, requires_gradTrue) y 2 * torch.dot(x, x) # 点积结果是标量 y.backward() print(x.grad)注意torch.dot(x, x)返回的是标量backward()不需要额外参数。但如果你对一个非标量张量直接调backward()框架会报错提示你必须传入一个与它 shape 相同的gradient参数。这个参数叫上游梯度upstream gradient表示从更外层传入的梯度。很多人第一次碰到这个报错会懵我解释一下它的语义。假设x torch.arange(4.0, requires_gradTrue) y x * x # y 是向量 [0, 1, 4, 9] y.backward() # 报错grad can be implicitly created only for scalar outputs为什么不行因为向量的梯度本身是个不完整信息——你要的是 y 里每个分量对 x 的偏导吗还是 y 的某种加权和框架为了避免歧义强制你显式传入一个权重向量 v实际计算的是v^T * J其中 J 是雅可比矩阵y 对 x 的偏导矩阵。最常见的做法是传全 1 向量等价于对 y 求和后再求导y.sum().backward() # 等价写法先求和再反向传播 print(x.grad) # tensor([0., 2., 4., 6.])从这个细节能看出 autograd 的设计哲学它把如何把非标量损失聚合成标量的决策权交给你而不是替你瞎猜。这也解释了为什么 PyTorch 里大家写损失函数时总习惯加.sum()或.mean()——不只是为了数值大小合适更是为了让反向传播的语义清晰。3.2 叶子节点与非叶子节点为什么 y.grad 是 None第二个常见的困惑明明 y 也参与了计算为什么y.grad打印出来是 Nonex torch.arange(4.0, requires_gradTrue) y x * x # y 不是叶子节点 z y.sum() z.backward() print(x.grad) # tensor([0., 2., 4., 6.]) print(y.grad) # None原因在于 PyTorch 默认只给叶子节点分配.grad内存。叶子节点是指直接由用户创建、requires_gradTrue的张量它是计算图的边界。而 y 是中间结果框架默认认为你只关心叶子节点的梯度所以不给你存。这既是省内存的设计也是常见 bug 的来源——你想检查中间层的梯度结果看到 None以为是梯度消失了。解决办法有两个。第一用retain_grad()显式让非叶子节点保留梯度x torch.arange(4.0, requires_gradTrue) y x * x y.retain_grad() z y.sum() z.backward() print(y.grad) # tensor([0., 1., 2., 3.])第二用torch.autograd.grad()这个函数接口它直接返回指定张量的梯度而不往.grad里写x torch.arange(4.0, requires_gradTrue) y x * x z y.sum() grad_y torch.autograd.grad(z, y) # (tensor([0., 1., 2., 3.]),)调完backward()之后计算图默认会被释放这是为了省内存。如果你要对同一个计算图连续调两次backward()比如某些元学习算法里一个计算图要贡献两个损失就必须在第一次调用时传retain_graphTrue否则第二次会报Trying to backward through the graph a second time的错误。x torch.arange(4.0, requires_gradTrue) y x * x z1 y.sum() z2 (y * 2).sum() z1.backward(retain_graphTrue) z2.backward() print(x.grad) # tensor([0., 3., 6., 9.])两次梯度累加需要注意的是这里两次梯度是累加的这正是下一个部分要展开的大坑。4. 梯度相关的坑与排查链路zero_grad、原地操作和 detach4.1 梯度累积导致结果漂移为什么每个 step 都要 zero_grad我在带新人入门时几乎每周都能遇到一次loss 不降反升或训练曲线像锯齿的问题最后十有八九是忘了optimizer.zero_grad()。这个错误如此常见是因为PyTorch 的.grad默认是累加的。为什么这样设计因为 autograd 的底层语义是梯度累积。在某些场景下比如用梯度做 mimo-batch 采样、或者梯度压缩通信你可能希望把多个小 batch 的梯度加起来当一个大 batch 用。所以框架选择了把新旧梯度叠加而不是覆盖。这是刻意的设计但代价就是如果你每个 step 不手动清零上一轮的梯度会和这一轮叠加成了错误的平均数参数更新方向就乱了。正确的训练循环长这样for epoch in range(num_epochs): for batch in dataloader: optimizer.zero_grad() # 清零梯度 loss model(batch) # 前向 loss.backward() # 反向传播梯度写入 .grad optimizer.step() # 参数更新optimizer.zero_grad()本质上是遍历所有参数把.grad置零。这里有个更隐蔽的坑如果某一个参数在前向过程中被detach了或者根本没有参与本次计算但它的.grad里还残留着上一次的梯度那么optimizer.step()依然会用这个过期梯度更新它。所以排查某个参数更新得不对时除了看backward()有没有执行还要看zero_grad()是否覆盖到了这个参数。经验之谈调试梯度问题时第一件事就是打印model里每个参数的.grad看哪些是 None、哪些是 0、哪些值异常大。这一步能过滤掉八成的问题。4.2 原地操作in-place引发的求导错误PyTorch 的报错信息里有段话你可能见过a leaf Variable that requires grad is being used in an in-place operation.这句话很多新手看不懂但它背后的机制其实不复杂。原地操作是指x.add_(1)、x.mul_(2)这类带下划线的方法它们直接修改原张量的值而不是创建新张量。问题在于autograd 在反向传播时需要用到前向过程中保存的中间值。比如计算y x * x时反向需要知道 x 的值如果你在backward()之前用原地操作把 x 改了那反向时拿到的 x 已经不是前向时的 x 了算出来的梯度就是错的。可以这样复现一下报错x torch.tensor([1.0], requires_gradTrue) y x * x x.add_(1.0) # 原地修改 x y.backward() # RuntimeError: a leaf Variable that requires grad...PyTorch 的应对策略是宁可报错也不给错结果。它在前向传播时给每个张量打标记如果检测到某个叶子节点在反向前被原地修改就直接抛异常。明白了这个逻辑你就知道该怎么做不要在requires_gradTrue的张量上做原地操作改用x x 1这样创建新张量的方式如果必须原地操作先x.data脱离 autograd 跟踪再改但这样做风险自负一般只有经验丰富的人才会在极个别场景比如某些权重裁剪算法里这么干。4.3 detach 的真实使用场景梯度截断和特征冻结detach()是一个高频且容易误用的方法。它的作用是返回一个新张量这个新张量和原张量共享同一块数据但requires_gradFalse而且完全脱离计算图。对新张量做的任何操作都不会把梯度传回原来的分支。最常见的场景有三个。第一是固定特征提取器你用一个预训练模型比如 ImageNet 上训好的 ResNet提取特征只想训练后面的分类头那就把前面所有层requires_grad_(False)或者在前向时把特征detach()出来这样backward()就不会更新前面的网络。features backbone(inputs) # backbone 参数 requires_gradFalse logits classifier(features.detach())第二是避免梯度跨分支传播。GAN 训练里判别器要更新、生成器要不更新或者强化学习里某个目标值不希望梯度回流这些场景都要用 detach 把计算图切断。第三是对比学习里的 stop-gradient 技巧SimSiam、BYOL 这些方法里一个分支的梯度会被显式截断这是整个算法成立的关键设计而不是偷懒——刚开始学的时候别盲目删掉这个操作删掉模型可能直接崩。使用 detach 时要注意detach 出来的张量虽然不求导了但它和原张量共享数据内存。如果你对 detach 后的张量做原地修改原张量也会变这同样会破坏 autograd 的记录。判断一个张量是否在计算图里可以用x.requires_grad和x.grad_fn两个属性。grad_fn为 None 且requires_gradTrue的是叶子节点grad_fn不为 None 的是中间节点它记录了产生这个张量的运算类型。4.4 梯度为 None 的排查链路训练时最让人头大的错误之一是某个参数的梯度是 None参数直接不更新。我把排查链路整理成一套固定的流程遇到问题按顺序走确认这个参数是否参与过前向计算。如果它压根没被用就没有梯度路径.grad 自然是 None。确认前向过程中没有对它执行.detach()或者被传入了torch.no_grad()上下文。torch.no_grad()里面做的任何计算都不建图梯度也不会留在外边。推理模式下用model.eval()with torch.no_grad():是标准做法但如果你在训练循环里误把某段前向包在no_grad里梯度就断了。确认损失是否经过了该参数。比如你写了一个自定义损失但实际只用了logits的某一部分另外一部分参数可能就没接到梯度。检查有没有在backward()之前修改过计算图结构。这套流程我写进过团队的文档里新同学照着排查基本都能在一刻钟内定位问题。梯度是 None和梯度异常大是两个极端前者通常是图断了后者通常是数值稳定性问题比如梯度爆炸需要做梯度裁剪torch.nn.utils.clip_grad_norm_。5. 自动求导的进阶用法二阶导、自定义算子和显存权衡5.1 二阶梯度create_graph 到底创建了什么图自动求导不仅能求一阶梯度还能继续对梯度求导得到二阶梯度。这在一些优化算法如 Newton 法、部分元学习算法 MAML里会用到。PyTorch 里求二阶导的关键是create_graphTruex torch.tensor([2.0], requires_gradTrue) y x ** 3 # 一阶导同时保留计算图 first_grad, torch.autograd.grad(y, x, create_graphTrue) print(first_grad) # tensor([12.]), 3*x^2 12 # 对一阶导再求导 second_grad, torch.autograd.grad(first_grad, x) print(second_grad) # tensor([12.]), 6*x 12create_graphTrue的含义是在计算 first_grad 的过程中把求梯度这个过程本身的操作也记录成一张计算图。这样 first_grad 就不再是独立的数值而是一个可以继续求导的节点。代价是显存占用显著增加因为框架不仅要保存前向的中间值还要保存反向过程的中间值。平时训练完全不需要二阶导只有做特定研究时才用所以默认是 False。5.2 自定义 autograd.Function当你需要定义自己的前向和反向大多数时候用框架现成的算子就够了但有三种情况需要你自定义求导逻辑一是写了一个 PyTorch 没有的运算二是某个算子 PyTorch 的反向实现效率太低你想手写一个更快的三是你在实现某个对数值稳定性要求极高的算法想手动控制反向公式。标准写法是继承torch.autograd.Function实现静态方法forward和backwardimport torch class MyReLU(torch.autograd.Function): staticmethod def forward(ctx, input): # ctx 是上下文对象用来保存反向时需要的中间值 ctx.save_for_backward(input) return input.clamp(min0) staticmethod def backward(ctx, grad_output): input, ctx.saved_tensors grad_input grad_output.clone() grad_input[input 0] 0 return grad_input # 使用 x torch.tensor([-1.0, 2.0, -3.0], requires_gradTrue) y MyReLU.apply(x) y.sum().backward() print(x.grad) # tensor([0., 1., 0.])有几个细节值得注意。backward的返回值数量必须和forward的入参数量一致ctx.save_for_backward保存的是原始张量后面从ctx.saved_tensors取出来grad_output是上游传下来的梯度你需要把它继续向后传如果某个输入不需要梯度backward 里对应位置返回 None 即可。我用这个方法实现过一些自定义算子比如带掩码的稀疏注意力。最大的教训是一定要写梯度检查来验证自定义 backward 的正确性。直接用数值求导和 autograd 求导对比from torch.autograd import gradcheck x torch.randn(3, requires_gradTrue, dtypetorch.float64) assert gradcheck(MyReLU.apply, (x,), eps1e-6, atol1e-4)gradcheck会用数值方法近似梯度然后跟你的自定义 backward 对比偏差在容差内才算通过。数值求导虽然训练时不能用但作为校验工具价值很大。5.3 从自动求导反观显存占用中间结果都存哪儿了理解了 autograd 之后很多工程问题也有了答案。最典型的一个为什么训练时的显存占用远大于推理时的显存占用原因就是反向传播需要重放前向的运算而重放时需要前向的中间结果。以最简单的y x * x为例反向时需要知道前向时的 x 是多少所以框架自动保存了它。一个 50 层的网络每一层的激活值都得存下来供反向传播使用。Batch size 越大、网络越深要存的中间值越多显存就爆了。针对这个瓶颈目前主流方案是梯度检查点Gradient Checkpointing。它的思想很直接不存所有中间结果只存一部分关键检查点反向传播时遇到没存的中间值就现场重新前向计算一遍。这就是用时间换空间训练速度会慢一些但显存占用能大幅下降。PyTorch 里用torch.utils.checkpoint包一下前向函数即可import torch.utils.checkpoint as cp def forward_with_checkpoint(module, x): return cp.checkpoint(module, x)另外把模型切到推理模式时用torch.no_grad()本质就是告诉 autograd不要建图、不要存中间值所以推理时显存占用小得多。明白了这些机制你在显存不够时就知道该往哪个方向调了。最后再多说一句和自动求导直接相关的实践心得想真正理解这个机制别只看书上的结论一定要自己动手把torch.autograd.grad、backward(retain_graphTrue)、create_graphTrue、自定义Function这四件事各写一遍并且用前面提到的gradcheck验证。我对这章的理解就是在反复复现这些代码、亲手踩过那些报错之后才建立的。这套知识在之后接触更复杂的模型结构、分布式训练、混合精度的时候会反复用到。