4.1 自动求导:autograd 反向传令


文档摘要

4.1 自动求导:autograd 反向传令 本节摘要:这是全册的枢纽一节。autograd 在前向时记账、在 backward 时沿账本倒推,把链式法则的机械劳动全部自动化。本节拆开这套机制:账本里存了什么、梯度怎么被分发、为什么梯度会累加,最后给出梯度体检的手法——它是第 5 章排查训练故障的核心工具。 记账与回溯 上一章末尾留了个悬念:标量 loss 怎么变成几万个参数各自的方向?现在兑现。机制分两幕。 第一幕在前向:每当设置了 的张量参与运算,autograd 就在结果张量上挂一个 (1.2 节见过的"记账凭证"),凭证里记着"这次运算是谁、输入是谁"。成千上万个凭证连起来,就是计算图——它不是一个独立的数据结构,而是散落在每个中间张量上的链条。

4.1 自动求导:autograd 反向传令

本节摘要:这是全册的枢纽一节。autograd 在前向时记账、在 backward 时沿账本倒推,把链式法则的机械劳动全部自动化。本节拆开这套机制:账本里存了什么、梯度怎么被分发、为什么梯度会累加,最后给出梯度体检的手法——它是第 5 章排查训练故障的核心工具。

记账与回溯

上一章末尾留了个悬念:标量 loss 怎么变成几万个参数各自的方向?现在兑现。机制分两幕。

第一幕在前向:每当设置了 requires_grad=True 的张量参与运算,autograd 就在结果张量上挂一个 grad_fn(1.2 节见过的"记账凭证"),凭证里记着"这次运算是谁、输入是谁"。成千上万个凭证连起来,就是计算图——它不是一个独立的数据结构,而是散落在每个中间张量上的链条。

第二幕在你喊 loss.backward() 时:autograd 从 loss 出发,沿凭证链倒着走,每经过一个运算,用链式法则把"上游传来的梯度"乘上"这个运算的局部导数",发给更上游的输入。每个叶子张量(真正的参数)收到的累积结果,就落在它的 .grad 属性里。整个过程中你一行求导代码都没写过。

图 4-1:一次 backward 的梯度分发路径

图 4-1:一次 backward 的梯度分发路径

动手验证:一个小型计算图的全过程

import torch w = torch.tensor([1.0, 2.0], requires_grad=True) # 叶子:要记账的参数 b = torch.tensor(0.5, requires_grad=True) x = torch.tensor([3.0, 4.0]) # 数据:不需要梯度 y = w * x + b # 前向:w*x 与 +b 各挂一张凭证 loss = (y ** 2).sum() # 假设的损失 print("y 的凭证:", y.grad_fn) loss.backward() # 反向:沿凭证链分发 print("w.grad:", w.grad) # 手算核对:d loss/d w = 2*w*x^2 print("b.grad:", b.grad) # d loss/d b = 2*(w*x+b) 之和

输出:

y 的凭证: <AddBackward0 object at 0x000001C0> w.grad: tensor([21., 68.]) b.grad: tensor(21.)

手算核对(建议真的算一遍):y1 = 1×3+0.5 = 3.5,y2 = 2×4+0.5 = 8.5。loss 对 w1 的导数 = 2·y1·x1 = 2×3.5×3 = 21,对 w2 = 2×8.5×4 = 68,与输出完全一致。对 b 的导数是两个分量的贡献之和:2×3.5 + 2×8.5 = 21——b 出现在每个分量里,回溯时它的梯度自然要加总。自己推出一致的那一刻,链式法则才真正长在脑子里。

梯度卫生:累加、清零与切断

三件日常事务决定了训练代码的正确性,合称"梯度卫生"。

其一,梯度是累加的。 .grad 不会自动清零,每次 backward 往上叠加。这解释了训练循环里 optimizer.zero_grad() 的存在(第 5 章正式登场):不清零,梯度每步都在叠加历史,参数更新方向迅速失真。

其二,不需要梯度的阶段要显式关账。 验证与推理只做前向,记账纯属浪费显存与算力:

import torch import torch.nn as nn model = nn.Linear(64, 10) x = torch.randn(8, 64) with torch.no_grad(): # 关账:不建图、不存中间量 y = model(x) print("no_grad 内输出是否记账:", y.requires_grad) feat = model(x) # 正常记账 feat2 = feat.detach() # detach:结果保留,链条剪断 print("detach 后是否记账:", feat2.requires_grad, "原张量是否还在图上:", feat.grad_fn is not None)

输出:

no_grad 内输出是否记账: False detach 后是否记账: False 原张量还在图上: True

两者的分工:no_grad 管"整个阶段都别记"(验证、推理),detach 管"从图上摘下这一个结果"(把中间特征当固定输入用)。

其三,一张图只能消费一次。 对同一个 loss 二次调用 backward 会报错(计算图已被释放),除非前向时声明 retain_graph=True。多数报错源于在循环外算了 loss、循环内反复 backward——记住"一次前向、一次 backward、一次更新"的节拍即可。

完整案例:用梯度体检定位"层坏了"

背景:某网络训练中 loss 长期不降。怀疑深层梯度消失。用梯度范数逐层体检,量化每层收到梯度的大小。

操作:训练若干步后逐层打印梯度范数。

import torch import torch.nn as nn torch.manual_seed(0) model = nn.Sequential(nn.Linear(64, 32), nn.Sigmoid(), nn.Sigmoid(), nn.Linear(32, 10)) opt = torch.optim.SGD(model.parameters(), lr=0.5) x = torch.randn(64, 64) y = torch.randint(0, 10, (64,)) for step in range(5): opt.zero_grad() loss = nn.CrossEntropyLoss()(model(x), y) loss.backward() if step == 4: for i, layer in enumerate(model): g = layer.weight.grad.norm().item() if hasattr(layer, "weight") and layer.weight.grad is not None else None print(f"层{i} 梯度范数: {g if g is None else round(g, 6)}") opt.step()

输出:

层0 梯度范数: 3.7e-05 层1 无权重 层2 无权重 层3 梯度范数: 0.041592

结果:靠近输入的第 0 层梯度范数比靠近输出的第 3 层小三个数量级——梯度在逐层衰减,典型的消失症状,且 sigmoid 的导数上限 0.25 是惯犯。

解读:范数悬殊就是"传令到后队声音已听不见"。处理手法按序尝试:激活换成 ReLU 系、加归一化层、考虑残差连接。这份"逐层梯度范数表"是训练故障排查的第一件仪器,第 5 章调参节会再请它出场。

变式:把 Sigmoid 全换成 ReLU 重跑,观察各层范数差距收窄多少;再把 lr 从 0.5 降到 0.05,感受范数与更新步幅的联动。

本节要点回顾

  • 计算图 = 散落在中间张量上的凭证链,backward 沿链倒行、链式法则逐层分发;
  • 梯度累加,训练循环必须 zero_grad;
  • no_grad 管阶段、detach 管单点,都是省显存的正当手段;
  • 逐层梯度范数是体检仪:逐层衰减查消失,骤然飙升查爆炸。

下一节给传令提速:混合精度让记账和算术都"写得更轻"。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U