混合精度与显存优化
前面几篇讲的都是「切到多卡」。这一篇讲单卡内部还能省什么 —— 三类与并行维度正交、几乎所有训练都会用到的技术:混合精度、梯度累积、激活重计算。
「正交」的意思是:它们和 DP/TP/PP 叠加使用,互不冲突。任何一次真实训练都是并行策略 × 这三类技术的组合。
混合精度
为什么能省显存又提速
低位宽 → 显存减半 + Tensor Core 算力翻倍。 这两条同时成立,所以混合精度是少见的「没有代价的优化」。
| 类型 | 位宽 | 指数位 | 尾数位 | 数值范围 | 场景 |
|---|---|---|---|---|---|
| FP32 | 32 | 8 | 23 | 大 | master weight、优化器 |
| FP16 | 16 | 5 | 10 | 小,易溢出 | 旧方案 |
| BF16 | 16 | 8 | 7 | 与 FP32 相同 | 大模型训练首选 |
| FP8 (E4M3) | 8 | 4 | 3 | 很小 | H100+ 的前向 |
判别要点:FP16 与 BF16 同为 16 位,差别全在指数位。完整的机制、溢出/下溢/舍入的分析见 01-数值计算与精度。
BF16 混合精度流程
前向 / 反向:BF16 计算(走 Tensor Core)
参数更新: FP32 Master Weight为什么要留一份 FP32 master:低精度权重上过小的更新会被直接舍掉 —— 若
这份 FP32 副本正是 Adam 的
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 放大
动态 Loss Scaling 的四步:
- 用当前 scale 算反向
- 检查梯度是否含
inf/NaN - 若溢出 → 跳过更新并减小 scale
- 连续稳定若干步 → 尝试增大 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。
机制:连续
两个容易错的点
一、Loss 必须除以
loss = criterion(output, target) / accumulation_steps
loss.backward() # 梯度是累加的,所以要提前平均
if step % accumulation_steps == 0:
optimizer.step(); optimizer.zero_grad()不除的话,累积 01-PyTorch 框架与训练循环)。
二、DDP 下要用 no_sync() 跳过中间步的 AllReduce。
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(),累积
no_sync()是一个「不改数学正确性、纯省通信」的优化 —— 这类优化在分布式训练里很值得找。
激活重计算(Activation Checkpointing)
问题:前向保存的激活随层数
核心思想:只保存部分层的激活(checkpoint),其余层反向时重新前向计算。
| 方案 | 只保存什么 | 收益 | 代价 |
|---|---|---|---|
| Full Checkpointing | 每个 Transformer Block 的输入,反向时重算整个 Block | 激活显存 | 约 +33% 前向计算量 |
| Selective Checkpointing | 只重算计算量小但激活大的操作(如 Attention 的 softmax 输出) | 同样省显存,但计算代价更低 | 需要人工挑算子 |
Selective 的判据是「计算量与激活量的比值」 —— 重算代价大的算子保留激活,重算便宜的丢掉激活。这与 06-GEMM 性能优化 里「用带宽而非延迟做判据」是同一套思维方式。
与并行策略的配合:PP 天然每个 Stage 只需存自己的激活;TP 下 SP 已减少激活冗余(见 08-张量并行与序列并行)。
显存优化的全景
四类手段,作用对象不同:
| 类别 | 手段 | 作用在显存的哪部分 |
|---|---|---|
| 切分类(多卡) | ZeRO/FSDP 切状态、TP 切权重、PP 切层、SP/CP 切激活 | 参数 / 梯度 / 优化器状态 / 激活 |
| 降精度类 | 混合精度(BF16 / FP8) | 全部四部分 |
| 时间换空间类 | 梯度累积(省激活)、激活重计算 | 激活 |
| 卸载类 | ZeRO-Offload / Infinity | 优化器状态 / 参数 → CPU / NVMe |
这张表回扣了 02-分布式训练总论与显存账本 的四项账本:参数
、梯度 、优化器状态 、激活( )。 显存优化的问题因此可以问得更精确:「我装不下的到底是这四项里的哪一个?」—— 不同答案对应完全不同的手段:
装不下的是 该上 优化器状态 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 │ ○ │ ○ │ ● │ · │
└──────────────────┴────────┴────────┴──────────────┴──────────────┘
● 主要作用 ○ 部分作用 · 不直接作用
⇒ 显存问题因此能问得更精确:「我装不下的到底是这四项里的哪一个?」相关
- 01-数值计算与精度 —— 精度格式、动态范围、Loss Scaling 的完整机制
- 02-分布式训练总论与显存账本 —— 四项显存账本与五大并行策略
- 07-ZeRO 显存优化系列 —— 切分优化器状态与梯度
01-PyTorch 框架与训练循环—— 混合精度的代码写法与梯度累加- 02-NVIDIA GPU 架构演进:Volta 到 Blackwell —— FP8 与 Transformer Engine 的硬件基础
参考
- https://caomaolufei.github.io/AIInfraGuide/guides/模块三-分布式训练/第8章-混合精度与显存优化
- https://arxiv.org/abs/1710.03740
- https://arxiv.org/abs/2205.05198
- https://docs.nvidia.com/deeplearning/transformer-engine/user-guide/index.html
- https://pytorch.org/docs/stable/amp.html
YJ