· 模型原理
前向传播与反向传播概述
微调场景下的前向产物、反向传播机制,以及冻结参数与 LoRA 等不同策略对反向计算图的实际影响。
核心结论:前向传播真正的产物不是最终那个 loss,而是计算图和各层缓存的中间激活值。反向传播以标量 loss 为唯一种子,沿计算图逆序执行 VJP,消费「上游梯度 + 前向缓存」,产出「下游梯度 + 本层参数梯度」。微调策略(LoRA、冻结、gradient checkpointing)的差异,全部落在两件事上:哪些路径参与反向计算,以及激活值是保留还是重算。
一、前向传播的产物
一次前向传播结束后,除了标量 loss,还会留下两项反向传播必须依赖的东西。
为什么 loss 必须是标量
这不是实现约定,而是 VJP 机制的硬性要求。反向传播的起点需要一个种子梯度;标量 loss 天然提供 ∂L/∂L = 1,一次逆序遍历即可得到全部参数的梯度。若 loss 是 n 维向量,就必须为每个分量提供种子并各跑一次扫描,等价于把反向计算量乘以 n。这也是 CrossEntropyLoss 默认 reduction='mean' 的原因。
缓存的内容
被缓存的并非前向全部中间结果,而是「局部导数足以算出」的最低限度:
| 算子 | 缓存内容 | 不可省略的原因 |
|---|---|---|
| MatMul / Linear | 输入 x(含形状) | dW = gᵀ·x,x 是自变量 |
| Softmax | 输出 y 本身 | 梯度公式直接含 y |
| LayerNorm / RMSNorm | μ、σ(或 1/σ) | 归一化导数需要该批的统计量 |
| Dropout | 0/1 mask + 缩放系数 | 随机性必须复现,否则梯度错位 |
| SiLU / GELU | 输入 x | 导数本身含 x,如 SiLU′ = σ(x) + xσ(x)(1−σ(x)) |
| CrossEntropy | 概率 p(无需保留 logits) | 化简后 dL/dlogits = p − onehot |
| Attention | K、V(或整块 attn 权重) | FlashAttention 的分歧点:不缓存,反向重算 |
激活缓存的逐项拆解
以 7B 模型(32 层 / hidden 4096 / 32 head / 词表 151936)、batch 1、序列长度 512 为例:
| 缓存位置 | 张量形状 | bf16 体积 | 用途 |
|---|---|---|---|
| 每层 Attention 的 Q / K / V | 3 × (1, 32, 512, 128) | 12.6 MB / 层 | 计算 Wqkv 的梯度 |
| 每层 Attention 输出 | (1, 512, 4096) | 4.2 MB / 层 | 残差与下一层输入 |
| 每层 Dropout mask | (1, 512, 4096) | 4.2 MB / 层 | 复现前向的随机性 |
| 每层 FFN 中间激活 | (1, 512, 16384) | 16.8 MB / 层 | SiLU 的导数依赖它 |
| 每层 LayerNorm 的 μ、σ | 2 × (1, 512, 1) | 约 2 KB / 层 | 归一化反向所需 |
| LM Head 输出的 logits | (1, 512, 151936) | 155 MB | dlogits 的计算起点 |
| 合计(32 层 + 词表头) | 约 1.5 GB | 与权重(14 GB)同量级 |
该表定位了长上下文场景的瓶颈。序列长度从 512 增至 4096 时,Q/K/V、mask、FFN 中间激活按长度线性增长 8 倍;Attention 分数矩阵按长度平方增长 64 倍(FlashAttention 因此选择不缓存、反向重算)。同一模型的激活缓存会由 1.5 GB 增至 10 GB 以上。
对应的估算公式:激活显存 ≈ batch × seq × hidden × layers × 常数。这是长上下文的瓶颈落在激活而非权重上的原因。
二、反向传播的机制
反向传播本身不引入新的数学,它是链式法则在计算图上的逆序编排。每个算子被调用时的动作固定:
- 确定种子。
loss.backward()等价于loss.backward(torch.ones_like(loss))。对非标量张量直接调用 backward 会报grad can be implicitly created only for scalar outputs。 - 逆拓扑序遍历。按计算图依赖的逆序处理节点,保证节点被访问时,其所有下游分支的梯度均已到达并可合并。顺序错误会导致重复计算或结果不完整。
- 只做 VJP,不构造完整 Jacobian。对
y = Wx,完整 Jacobian 是out × in个元素;实际只需针对当前的上游梯度 g 计算gᵀ · x,一次矩阵乘法。计算量因此从参数量级降到激活量级——这是 autograd 可行的前提。 - 梯度累加而非覆盖。同一张量被多个算子消费(残差分支)或权重共享(tied embedding、共享专家)时,梯度按加法合并。
.grad是累加缓冲,因此每个 step 必须显式清零。
(一次算 dX,一次算 dW)
由标量 loss 保证
对激活显存的压缩效果
(2× 前向 → 3× 前向)
三、微调策略的差异
LoRA 冻结底模权重,仅训练旁路矩阵 A、B。反向传播的覆盖范围因此收缩到旁路支路——但这不改变激活缓存的需求。
LoRA 的前向为 h = Wx + (α/r)·B·A·x,求导后两个待训练矩阵的梯度分别为:
dL/dB = (α/r)·g·(Ax)ᵀ:需要前向缓存的 Ax(A 的输出)。dL/dA = (α/r)·Bᵀg·xᵀ:需要前向缓存的 x(该层输入)。
两个梯度都依赖前向缓存。LoRA 缩小的是求导范围,并未减少激活的保留量。
把这一机制折算到显存(7B 量级估算):
三种策略对反向传播的覆盖范围与代价:
| 策略 | 反向传播覆盖范围 | 显存收益 | 无收益 |
|---|---|---|---|
| 全量 FT | 计算图全路径 | — | — |
| 冻结 + LoRA | 仅 LoRA 旁路,至冻结层终止 | 参数梯度 + Adam 状态 + fp32 主权重 | 激活值完全不变 |
| Gradient checkpointing | 全路径,中途重算前向以重建激活 | 激活从 O(L) 降至 O(√L) | 反向计算耗时 |
误解一:LoRA 的显存收益来自激活
并非如此。激活缓存量不变,前向计算量也不变。LoRA 的收益集中在优化器状态:Adam 为每个参与训练的参数维护两份 fp32 状态(一阶矩 m、二阶矩 v)。7B 模型全量微调时该项为 参数量 × 4 字节 × 2 = 56 GB,加上 fp32 主权重 28 GB,合计约 84 GB,即图 4 中占比最大的橙色段。LoRA 仅对 A、B 维护优化器状态,合计不足 0.5 GB。
结论:LoRA 使训练能够在小显存设备上启动,但不解决上下文长度增长带来的激活压力。前者由优化器状态决定,后者只能依靠 gradient checkpointing 或长上下文并行策略。
误解二:gradient checkpointing 是零成本优化
不是。选择不保留中间激活,就必须在反向传播到达该段时重新执行一次前向以重建它们。反向传播本身约为 2 倍前向 FLOPs,重算再叠加 1 倍前向,总开销由 2 倍升至 3 倍,即约 +50% 训练时间。这是用计算换显存,实践中通常只对部分层启用(use_reentrant=False + 选择性 checkpoint)。
四、实现要点
.grad是累加缓冲而非返回值。漏掉zero_grad(set_to_none=True)会导致梯度跨 step 累积,表现为 loss 突然发散。- 前向若被
@torch.no_grad()或inference_mode()包裹,计算图不会构建,.backward()直接抛出element 0 of tensors does not require grad。 - AMP 下
loss.backward()得到的是被缩放过的梯度,必须先scaler.unscale_(optimizer)再裁剪,否则clip_grad_norm_的阈值失去意义。 - 中间张量默认
requires_grad=True,但其.grad为None(除非调用retain_grad())。排查梯度连通性应检查grad_fn链是否中断。 - 冻结层设为
requires_grad=False后,若其输入仍需梯度(要传递到更靠前的 LoRA 层),该通路必须保持可微。对整个模块调用detach()会静默切断梯度且不报错,属于最难排查的一类问题。
附:三个算子的 VJP 推导
Linear(y = xWᵀ + b,g 为上游梯度)
dL/dx = g · W
dL/dW = gᵀ · x # 依赖前向缓存的输入 x
dL/db = Σ_batch g # 在 batch 维上求和Softmax + CrossEntropy——完整 Jacobian 为 V × V,与交叉熵复合后大幅化简:
dL/dz_i = p_i - y_i # p 为前向输出的概率,y 为 onehot
# 因此只需缓存 p,无需保留 logitsLayerNorm(ŷ = (x − μ)/σ):
dL/dx = (1/σ) · ( g − mean(g) − ŷ · mean(g · ŷ) )
# 依赖前向缓存的 μ、σ附:训练循环骨架(冻结、AMP 与梯度裁剪)
# 1. 冻结底模:前向正常执行,反向不产生梯度
for p in model.base_model.parameters():
p.requires_grad = False
# 2. loss 必须标量化,反向传播才具备种子
optimizer.zero_grad(set_to_none=True) # 比 zero_grad() 少一次 memset
loss = criterion(model(**batch).logits, batch["labels"])
assert loss.dim() == 0
# 3. AMP:先还原梯度尺度,再裁剪
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(trainable_params, 1.0)
scaler.step(optimizer)
scaler.update()
# 4. 校验:冻结参数不应出现 .grad
for n, p in model.base_model.named_parameters():
assert p.grad is None, n第 4 步的断言建议长期保留。它能在一次迭代内发现「漏冻结」「requires_grad 作用域错误」这类不产生任何报错的静默故障。
grad_fn 指针构成计算图;反向传播以标量 loss 为种子,沿图逆序执行 VJP,逐层完成链式法则,最终产出写入 .grad 的参数梯度。微调策略的差异,本质上是决定哪些路径参与反向计算,以及激活值是保留还是重算。