← 返回题目列表

Transformer 训练中的激活显存由什么组成?

高频 困难 第 22 / 25 题 更新于 2026/09/17
Transformer激活显存训练优化FlashAttention

简化版

训练显存不只有参数,还要保存反向传播所需激活。Transformer 激活大致随 batch × sequence × hidden × layers 线性增长,朴素注意力还可能保存 batch × heads × sequence² 的概率矩阵;FFN 中间维通常是隐藏维的 3~4 倍,也会形成大激活。估算时要结合精度、张量生命周期、并行切分和内核实现,不能只用参数量推显存。

详细版

每层常见激活包括残差输入、Q/K/V、注意力输出、softmax/Dropout 中间量、Norm 统计和 FFN 上投影/门控。反向所需保存项取决于框架:FlashAttention 不物化完整 S×S 矩阵,checkpointing 则少存一部分并在反向重算。

bytes ≈ dtype_bytes × Σ(saved tensor elements)
主线性项 ~ O(B·S·H·L)
朴素 attention score ~ O(B·A·S²·L)

混合精度下参数、梯度、优化器和激活精度可能不同;数据并行也不切激活,序列/张量并行才可能分摊部分维度。实际排查用框架 memory snapshot 和逐层峰值,而不是把所有张量尺寸简单相加,因为生命周期会重叠或释放。

完整版教学

一、训练显存的四本账

训练通常包含参数、梯度、优化器状态和激活。前三项主要随参数量增长,激活则强烈依赖 batch、序列长度、隐藏维、层数和具体 kernel。

推理时不做反向,可释放大多数训练激活;自回归推理的主要动态内存变成 KV Cache。把训练 OOM 与推理 OOM 混为一谈会用错优化手段。

记忆钩子:参数显存问“模型多大”,激活显存问“一次让多少 token 穿过多少层且要记住什么”。

二、为什么反向需要保存激活

线性层反向计算权重梯度需要输入激活,非线性反向也需要前向值。自动微分会保留计算图中必要张量,直到对应反向完成。

若把全部前向中间量释放,反向就要重新从较早检查点计算。Activation Checkpointing 正是用额外计算换减少保存。

训练模式下 Dropout、Norm 等也会保存 mask 或统计量,具体数量由融合实现决定。

三、线性项如何估算

一个 [B,S,H] 张量元素数为 B×S×H。例如 B=8、S=4096、H=4096、bf16 两字节:

8 × 4096 × 4096 × 2 ≈ 256 MiB

一层若同时存数个这种张量,几十层很快达到数十 GB。不能把这 256 MiB 当整层精确值,但它提供数量级。

梯度累积若采用多个 micro-batch 顺序反向,通常不同时保存所有 micro-batch 激活;峰值由单个 micro-batch 决定。

四、注意力的平方项

朴素注意力 logits/probability 形状常为 [B,A,S,S]。B=2、A=32、S=4096、每元素 2 字节时,仅一个矩阵约:

2 × 32 × 4096² × 2 ≈ 2 GiB

每层都物化会无法承受。FlashAttention 分块在线计算 softmax,只保存小统计与输入输出,将注意力中间显存从平方级降为近线性级,但计算语义仍是全注意力。

五、FFN 中间激活常被低估

传统 FFN 先从 H 扩到 rH 再降回 H,r 常见约 4;SwiGLU 有 gate 和 value 两个分支,参数匹配时扩展维选择也不同。

激活典型形状规模特征
ResidualB×S×H每层基础项
Q/K/VB×S×H 合计若干倍与投影实现相关
Attention scoreB×A×S×S朴素实现平方项
FFN intermediateB×S×rHr 倍线性项

长序列下注意力受关注,较短序列大 batch 时 FFN 中间量也可能主导。

六、融合算子改变保存集合

融合 LayerNorm、bias、activation 或 Dropout 可减少中间张量写回显存,并让反向重算部分轻量值。两个数学等价实现,峰值显存可能不同。

内存估算必须标注是否使用 FlashAttention、fused MLP、in-place 与 checkpointing。论文公式与当前 PyTorch kernel 不一致时,以实际 snapshot 为准。

融合也会影响数值和可用 dtype,需要回归正确性,不只是看显存。

七、并行策略如何影响激活

数据并行每张卡处理不同 batch,模型完整复制,激活不会因卡数自动切小;减小每卡 micro-batch 才会降峰值。张量并行可切 hidden/head 维,但某些操作需要 all-gather;序列并行切 token 维。

Pipeline 并行每阶段只放部分层,却可能同时保留多个 micro-batch 激活以填充流水线。内存取决于调度策略和 in-flight 数量。

ZeRO 主要切参数、梯度和优化器状态,对激活的直接帮助有限。

八、如何实测和优化

先分别记录 allocated、reserved 与峰值,使用 memory snapshot 找到大张量及生命周期。按序列长度和 micro-batch 做二分,确认是线性还是平方增长。

优化顺序可为:启用高效注意力、减小 micro-batch 并梯度累积、Activation Checkpointing、序列/张量并行、缩短或分桶序列。CPU offload 可省显存但增加传输。

每项都记录 tokens/s、峰值、数值与收敛,避免只解决 OOM 却让训练慢数倍。

九、常见误区与追问

  • 误区:模型权重能放下就一定能训练。 梯度、优化器和激活往往更大。
  • 误区:显存按所有中间张量相加即可。 生命周期、释放和融合决定峰值。
  • 误区:混合精度所有状态都是 2 字节。 优化器和主权重常保留 fp32。
  • 误区:ZeRO 会显著切分激活。 它主要处理模型状态,激活需其他技术。
  • 误区:FlashAttention 是稀疏注意力。 它计算精确全注意力,只改变 IO 与物化方式。
  • 追问:序列翻倍激活怎样变? 线性项翻倍,朴素注意力矩阵约四倍。
  • 追问:梯度累积是否增加峰值激活? 顺序前后向通常不按累积步数倍增,但吞吐和状态需实测。
  • 追问:为什么 reserved 大于 allocated? 缓存分配器保留内存以复用,不等于当前张量占用。

十、加强记忆

Transformer 激活显存记住“B、S、H、L 四个线性维,朴素注意力再带 S²”:残差、QKV 和 FFN 形成 B·S·H·L 主项,完整注意力矩阵形成 B·A·S²·L,FlashAttention 避免其物化。融合、checkpointing 和并行会改变保存集合,所以公式只做数量级,最终用逐层 snapshot、峰值和 tokens/s 一起验证。