· 模型原理

前向传播与反向传播概述

微调场景下的前向产物、反向传播机制,以及冻结参数与 LoRA 等不同策略对反向计算图的实际影响。

核心结论:前向传播真正的产物不是最终那个 loss,而是计算图和各层缓存的中间激活值。反向传播以标量 loss 为唯一种子,沿计算图逆序执行 VJP,消费「上游梯度 + 前向缓存」,产出「下游梯度 + 本层参数梯度」。微调策略(LoRA、冻结、gradient checkpointing)的差异,全部落在两件事上:哪些路径参与反向计算,以及激活值是保留还是重算。

一、前向传播的产物

一次前向传播结束后,除了标量 loss,还会留下两项反向传播必须依赖的东西。

前向传播的三个产物前向传播除输出标量 loss 外,还缓存每层激活值与 autograd 计算图,后两者供反向传播使用。前向传播forward pass① 标量 loss(0-dim)唯一种子:∂L/∂L = 1.0② 各层缓存的中间激活值x、attn 权重、LN 的 σ、dropout mask③ autograd 计算图op 的 grad_fn 指针串成的 DAG
图 1 · 前向传播的三个产物,只有 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/σ)归一化导数需要该批的统计量
Dropout0/1 mask + 缩放系数随机性必须复现,否则梯度错位
SiLU / GELU输入 x导数本身含 x,如 SiLU′ = σ(x) + xσ(x)(1−σ(x))
CrossEntropy概率 p(无需保留 logits)化简后 dL/dlogits = p − onehot
AttentionK、V(或整块 attn 权重)FlashAttention 的分歧点:不缓存,反向重算

激活缓存的逐项拆解

以 7B 模型(32 层 / hidden 4096 / 32 head / 词表 151936)、batch 1、序列长度 512 为例:

缓存位置张量形状bf16 体积用途
每层 Attention 的 Q / K / V3 × (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 MBdlogits 的计算起点
合计(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 × 常数。这是长上下文的瓶颈落在激活而非权重上的原因。

二、反向传播的机制

反向传播本身不引入新的数学,它是链式法则在计算图上的逆序编排。每个算子被调用时的动作固定:

单个算子的反向传播数据流算子的 backward 接收上游梯度,结合前向缓存的输入,产出一份传给下游的输入梯度和一份累加到参数的梯度。上游梯度 g = ∂L/∂yop.backward(g) — 一次 VJP读取前向缓存:输入 x,部分算子还需 mask / σ / pg_in = J^T · g继续向上一层传递W.grad += g^T · x累加进参数梯度缓冲
图 2 · 单个算子的固定动作:接收一份上游梯度,产出两份结果
  1. 确定种子。loss.backward() 等价于 loss.backward(torch.ones_like(loss))。对非标量张量直接调用 backward 会报 grad can be implicitly created only for scalar outputs
  2. 逆拓扑序遍历。按计算图依赖的逆序处理节点,保证节点被访问时,其所有下游分支的梯度均已到达并可合并。顺序错误会导致重复计算或结果不完整。
  3. 只做 VJP,不构造完整 Jacobian。y = Wx,完整 Jacobian 是 out × in 个元素;实际只需针对当前的上游梯度 g 计算 gᵀ · x,一次矩阵乘法。计算量因此从参数量级降到激活量级——这是 autograd 可行的前提。
  4. 梯度累加而非覆盖。同一张量被多个算子消费(残差分支)或权重共享(tied embedding、共享专家)时,梯度按加法合并。.grad 是累加缓冲,因此每个 step 必须显式清零。
反向传播相对前向的 FLOPs
(一次算 dX,一次算 dW)
1 次
反向扫描次数
由标量 loss 保证
O(L) → O(√L)
gradient checkpointing
对激活显存的压缩效果
+50%
checkpointing 的时间代价
(2× 前向 → 3× 前向)

三、微调策略的差异

LoRA 冻结底模权重,仅训练旁路矩阵 A、B。反向传播的覆盖范围因此收缩到旁路支路——但这不改变激活缓存的需求。

LoRA 微调的前反向路径与缓存依赖冻结的 W 不参与梯度计算,LoRA 支路上的 A 和 B 可训练,其梯度依赖前向缓存的 x 与 Ax。xW · 冻结A · 可训练B · 可训练+hW 冻结 —— 不算 dW、不写 .grad,反向传到 W 即止;但前向激活照常缓存A / B 可训练 —— 梯度仍依赖前向结果:dL/dA 用 x,dL/dB 用 Ax
图 3 · LoRA 的两条支路:底模冻结不产生梯度,激活缓存不受影响

LoRA 的前向为 h = Wx + (α/r)·B·A·x,求导后两个待训练矩阵的梯度分别为:

  • dL/dB = (α/r)·g·(Ax)ᵀ:需要前向缓存的 Ax(A 的输出)。
  • dL/dA = (α/r)·Bᵀg·xᵀ:需要前向缓存的 x(该层输入)。

两个梯度都依赖前向缓存。LoRA 缩小的是求导范围,并未减少激活的保留量。

把这一机制折算到显存(7B 量级估算):

三种微调策略的显存构成估算全量微调的优化器状态占主导,LoRA 把优化器状态压到可忽略,但激活值完全不变。全量微调 Full FT124 GBLoRA(r=16,冻结底模)26 GBLoRA + gradient checkpointing18 GBLoRA 的梯度与优化器状态合计不足 0.5 GB,为可见性以最小宽度绘制。权重 bf16梯度优化器状态 Adam fp32激活值量级估算:7B / bf16 权重 / Adam 状态 fp32 未分片 / seq 4096 / bs 1。激活随批大小线性增长。
图 4 · 三种策略的显存构成(7B 量级估算):LoRA 压缩优化器状态,激活不变

三种策略对反向传播的覆盖范围与代价:

策略反向传播覆盖范围显存收益无收益
全量 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,但其 .gradNone(除非调用 retain_grad())。排查梯度连通性应检查 grad_fn 链是否中断。
  • 冻结层设为 requires_grad=False 后,若其输入仍需梯度(要传递到更靠前的 LoRA 层),该通路必须保持可微。对整个模块调用 detach() 会静默切断梯度且不报错,属于最难排查的一类问题。
附:三个算子的 VJP 推导

Lineary = 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,无需保留 logits

LayerNormŷ = (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 的参数梯度。微调策略的差异,本质上是决定哪些路径参与反向计算,以及激活值是保留还是重算
自包含文档:全部图示为内联 SVG,样式为内联 CSS,无外部资源依赖,离线可回看。 显存与体积数值为量级估算,口径见图 4 标注与正文表格。
分享:
返回文章列表

相关文章

全部文章 »