混合精度:溢出与缩放
AMP fp16/bf16 加速训练,但 gradient underflow 会 NaN——GradScaler 动态缩放是标配,bf16 在 A100+ 上 often 更省心。

1. fp16 省显存,underflow 会 NaN
ResNet50 ImageNet fp32 占 11GB;开 torch.autocast + GradScaler 后同 batch 约 7GB,吞吐涨一截。但有一次 loss 在 step 3000 突然 NaN——log 里 loss_scale 已从 65536 掉到 1,gradient underflow 后 optimizer 在噪声里走,某步又 overflow。fp16 指数范围窄,小梯度会 underflow 成零;GradScaler 动态放大 loss 再 backward,是 fp16 训练的标配。
A100 上 bf16 exponent 与 fp32 同宽,常更省心——多数模型可直接 dtype=torch.bfloat16 少调 scale。
2. 标准 AMP 写法
PyTorch AMP 文档:
scaler = torch.cuda.amp.GradScaler()
for x, y in loader:
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type="cuda", dtype=torch.float16):
loss = criterion(model(x), y)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()autocast 只包 forward 和 loss;backward 在 scaled loss 上。unscale_ 后 clip 再 step——顺序错则 clip 无效。scaler.update() 根据是否 inf/nan 调整 scale。
3. fp16 vs bf16
fp16 尾数多、精度高,但指数窄易 overflow/underflow,通常需要 GradScaler。bf16 尾数少但指数与 fp32 同宽,A100/H100 上常可省略 GradScaler,代码更简单。V100 等老卡 fp16 仍是主流,必须 scaler 加监控 scale 曲线。
4. 失败模式
scale 持续下降 → 梯度经常 inf/nan,查 LR、loss 设计、数据。scale 长期不变在极大值 → 可能 underflow 被掩盖,看 weight 是否更新。某些 op 在 autocast 下数值不稳定——可 autocast(enabled=False) 包该 op 强制 fp32。
LayerNorm、Softmax 在 fp16 有时需 fp32 内部计算;PyTorch autocast 对多数 op 已处理,自定义 op 需自查。Adam eps=1e-8 在 fp16 下有时改 1e-6。
5. 案例:scale 掉到 1 后 NaN
某 Transformer 训练 scale 从 65536 单调降到 1,随后 NaN——根因是某层梯度长期 underflow,权重几乎不更新后在 bad batch 上 overflow。换 bf16 或增大初始 scale 并查 LR 后稳定。
6. 验收
- fp16 路径含 GradScaler,log 含 scale 曲线。
- unscale → clip → step 顺序正确。
- bf16 在支持硬件上对比 fp32 baseline,acc 应 parity。
- NaN 时查 scale 历史与 LR,不盲目关 AMP。
- 显存与吞吐相对 fp32 有 measurable 改善。
7. 落地与衔接
训练配置显式声明半精度与脑浮点路径:半精度必须配梯度缩放并记录缩放曲线;脑浮点在支持硬件上可省略缩放但仍监控损失。持续集成用单步反向传播测反缩放裁剪更新顺序,防止重构打乱混合精度契约。出现非数报警时保存缩放历史与当前批次,便于判断下溢还是学习率问题。上线导出前用全精度基线对比关键层数值一致性,再决定是否全链路混合精度推理。文档写清何种硬件用何种 dtype。
8. 案例复盘
某 Transformer 半精度训练缩放从六万五千三百三十六单调降到一后非数,日志显示长期下溢后权重几乎不更新。换脑浮点并略降学习率后缩放稳定、精度与全精度持平。团队现在在实验模板里固定记录缩放曲线与峰值梯度范数,混合精度相关合并必须附全精度对照。省下的显存若未换更大批次或更深模型,要在文档里说明收益去向,避免开了混合精度却无任何吞吐提升。
混合精度是算力换数值风险。GradScaler 管理 fp16 的 underflow;bf16 在 modern GPU 上往往是更省心的默认。无论哪种,监控 scale 和 loss 曲线是必选项。把 AMP 选型与 scale 曲线写进发版材料,后续排查 NaN 才有时间轴可对齐。推理侧是否沿用训练 dtype 应单独评估精度与延迟。发版材料须附缩放曲线与全精度对照结论。训练任务卡型变更时重新评估半精度与脑浮点选型,须在实验日志里记录硬件代际与 dtype 决策依据。排查非数时应同时对照缩放曲线与学习率日程。 训练日志除缩放曲线外,还应记录峰值梯度范数与非数发生步,便于判断是数值问题还是学习率问题。卡型迁移时重新做全精度对照,不要假设上一代的混合精度配置可直接沿用。
相关
也可以看看
- ·18 分钟阅读
混合精度下的梯度累积:loss 缩放、裁剪与 step 顺序
沿 autocast、GradScaler、unscale_、裁剪与 optimizer.step 的数据流说明静默错误;区分 micro-batch 平均与求和,并用等价性测试验收。
- ·3 分钟阅读
grad_norm spike 与 misclass grid
只看 train loss 不够——要 log lr、grad norm、val slice、GPU 利用率和样例可视化,异常通常是 loss 曲线先以外的信号。
- ·5 分钟阅读
beta=0.999 与 eval 权重
训练时对参数做指数滑动平均,验证和部署常用 EMA 权重——曲线更顺,有时泛化更好,但要与 optimizer 步和 checkpoint 策略一致。
johan's blog