Skip to content

混合精度与显存优化 ​

标签
AI/infra/显存管理
AI/infra/精度
字数
2112 字
阅读时间
9 分钟

前面几篇讲的都是「切到多卡」。这一篇讲单卡内部还能省什么 —— 三类与并行维度正交、几乎所有训练都会用到的技术:混合精度、梯度累积、激活重计算。

「正交」的意思是:它们和 DP/TP/PP 叠加使用,互不冲突。任何一次真实训练都是并行策略 × 这三类技术的组合。

混合精度 ​

为什么能省显存又提速 ​

低位宽 → 显存减半 + Tensor Core 算力翻倍。 这两条同时成立,所以混合精度是少见的「没有代价的优化」。

类型位宽指数位尾数位数值范围场景
FP3232823大master weight、优化器
FP1616510小,易溢出旧方案
BF161687与 FP32 相同大模型训练首选
FP8 (E4M3)843很小H100+ 的前向

判别要点:FP16 与 BF16 同为 16 位,差别全在指数位。完整的机制、溢出/下溢/舍入的分析见 01-数值计算与精度。

BF16 混合精度流程 ​

前向 / 反向:BF16 计算(走 Tensor Core)
参数更新:   FP32 Master Weight

为什么要留一份 FP32 master:低精度权重上过小的更新会被直接舍掉 —— 若 |δ| 远小于该数值附近的可表示间距,fl(x+δ)=x,这一步的更新等于没做。

这份 FP32 副本正是 Adam 的 12Ψ 里的一项(见 02-分布式训练总论与显存账本)—— 所以它既是精度保障,也是显存大头,才成为 ZeRO 第一刀的目标。

BF16 混合精度里每一步用什么精度(以及 FP32 master 为什么必须在):

  ┌─────────────── 前向 / 反向 ────────────────┐
  │  BF16 权重副本(工作副本)                  │  走 Tensor Core ⇒ 算力翻倍
  │        │                                   │  显存是 FP32 的一半
  │        └──▶ BF16 计算(矩阵乘 / 激活)      │
  └────────────────────┬───────────────────────┘
                       │ 产出
                   BF16 梯度
                       │
  ┌────────────────────▼───────────────────────┐
  │  FP32 Master Weight + Adam 状态             │  参数更新在这里做
  │  (这份 FP32 副本正是 Adam $12\Psi$ 的一项) │
  └────────────────────┬───────────────────────┘
                       │ 更新完再转回 BF16
                       ▼
                  BF16 权重副本

  ⇒ 为什么必须留 FP32 master:低精度权重上过小的更新会被直接舍掉
    ($|\delta|$ 远小于该处可表示间距 ⇒ $\operatorname{fl}(x+\delta) = x$,这一步等于没做)

Loss Scaling:FP16 必须、BF16 通常不需要 ​

反向里的 FP16 小梯度可能下溢为 0。做法是把 loss 放大 S 倍(∇L′=S∇L),反向完再除回来(∇L=∇L′/S)。

动态 Loss Scaling 的四步:

  1. 用当前 scale 算反向
  2. 检查梯度是否含 inf / NaN
  3. 若溢出 → 跳过更新并减小 scale
  4. 连续稳定若干步 → 尝试增大 scale

BF16 因为动态范围与 FP32 相同,通常不需要 Loss Scaling —— 这是它取代 FP16 的核心原因,省掉一整套需要调参的机制。

FP8 训练 ​

H100 的 Transformer Engine 走得更远:

部分精度
前向FP8(E4M3)+ per-tensor 动态缩放
反向梯度BF16

配套的机制是逐层动态判断:数值稳定的层用 FP8,数值敏感的层自动降回 BF16(见 02-NVIDIA GPU 架构演进:Volta 到 Blackwell)。风险在精度 —— FP8 只有 3 位尾数,需要依赖 Transformer Engine 的动态缩放兜底。

梯度累积 ​

目的:在有限显存下模拟更大的 Effective Batch Size。

Effective Batch Size=单卡 batch×world_size×accumulation_steps

机制:连续 K 步前向+反向累积梯度,每 K 步做一次参数更新。

两个容易错的点 ​

一、Loss 必须除以 K。

python
loss = criterion(output, target) / accumulation_steps
loss.backward()                   # 梯度是累加的,所以要提前平均
if step % accumulation_steps == 0:
    optimizer.step(); optimizer.zero_grad()

不除的话,累积 K 步后的梯度是单步的 K 倍 —— 等效于把学习率放大了 K 倍。梯度累加这个设计本身依赖 PyTorch「梯度不自动清零」的行为(见 01-PyTorch 框架与训练循环)。

二、DDP 下要用 no_sync() 跳过中间步的 AllReduce。

python
for i, batch in enumerate(loader):
    ctx = model.no_sync() if i % K != 0 else nullcontext()
    with ctx:
        loss = ...; loss.backward()
    if i % K == 0:
        optimizer.step(); optimizer.zero_grad()

没有 no_sync(),累积 K 步就白搭了 K−1 次 AllReduce 的通信 —— 因为中间步骤的梯度反正还要继续累加,此时同步毫无意义。

no_sync() 是一个「不改数学正确性、纯省通信」的优化 —— 这类优化在分布式训练里很值得找。

激活重计算(Activation Checkpointing) ​

问题:前向保存的激活随层数 L 与 batch 线性增长,大模型训练里激活显存可能超过参数显存。

核心思想:只保存部分层的激活(checkpoint),其余层反向时重新前向计算。

方案只保存什么收益代价
Full Checkpointing每个 Transformer Block 的输入,反向时重算整个 Block激活显存 O(L)→O(L)约 +33% 前向计算量
Selective Checkpointing只重算计算量小但激活大的操作(如 Attention 的 softmax 输出)同样省显存,但计算代价更低需要人工挑算子

O(L) 怎么来的:每 L 层存一个 checkpoint,其余层重算。这是「分段 + 段内重算」的经典组合 —— 与 01-数值计算与精度 里 Welford 的「分块可合并」是同一类思路:把状态量从「每层一份」压到「每段一份」。

Selective 的判据是「计算量与激活量的比值」 —— 重算代价大的算子保留激活,重算便宜的丢掉激活。这与 06-GEMM 性能优化 里「用带宽而非延迟做判据」是同一套思维方式。

与并行策略的配合:PP 天然每个 Stage 只需存自己的激活;TP 下 SP 已减少激活冗余(见 08-张量并行与序列并行)。

显存优化的全景 ​

四类手段,作用对象不同:

类别手段作用在显存的哪部分
切分类(多卡)ZeRO/FSDP 切状态、TP 切权重、PP 切层、SP/CP 切激活参数 / 梯度 / 优化器状态 / 激活
降精度类混合精度(BF16 / FP8)全部四部分
时间换空间类梯度累积(省激活)、激活重计算激活
卸载类ZeRO-Offload / Infinity优化器状态 / 参数 → CPU / NVMe

这张表回扣了 02-分布式训练总论与显存账本 的四项账本:参数 2Ψ、梯度 2Ψ、优化器状态 12Ψ、激活(∝b⋅s⋅h⋅L)。

显存优化的问题因此可以问得更精确:「我装不下的到底是这四项里的哪一个?」—— 不同答案对应完全不同的手段:

装不下的是该上
优化器状态ZeRO-1
梯度ZeRO-2
参数ZeRO-3 / TP
激活梯度累积 / 激活重计算 / SP / CP
全部组合,或 Offload

四类显存手段各自作用在账本的哪一项:

  ┌──────────────────┬────────┬────────┬──────────────┬──────────────┐
  │ 手段              │ 参数2Ψ │ 梯度2Ψ │ 优化器 12Ψ    │ 激活 ∝ b·s·h·L│
  ├──────────────────┼────────┼────────┼──────────────┼──────────────┤
  │ ZeRO / FSDP      │   ○    │   ○    │      ●       │      ·       │
  │ TP               │   ●    │   ●    │      ○       │      ○       │
  │ PP               │   ·    │   ·    │      ·       │ ●(每个 Stage │
  │                  │        │        │              │   只存自己的)│
  │ SP / CP          │   ·    │   ·    │      ·       │      ●       │
  │ 混合精度 / FP8    │   ●    │   ●    │ ○(master 仍 │      ●       │
  │                  │        │        │   是 FP32)   │              │
  │ 激活重计算        │   ·    │   ·    │      ·       │      ●       │
  │ 梯度累积          │   ·    │   ·    │      ·       │      ●       │
  │ Offload          │   ○    │   ○    │      ●       │      ·       │
  └──────────────────┴────────┴────────┴──────────────┴──────────────┘

  ● 主要作用    ○ 部分作用    · 不直接作用

  ⇒ 显存问题因此能问得更精确:「我装不下的到底是这四项里的哪一个?」

相关 ​

参考 ​

贡献者 ​

文件历史 ​