返回专辑
·Johan·4 分钟阅读

混合精度:溢出与缩放

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 文档

python
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 决策依据。排查非数时应同时对照缩放曲线与学习率日程。 训练日志除缩放曲线外,还应记录峰值梯度范数与非数发生步,便于判断是数值问题还是学习率问题。卡型迁移时重新做全精度对照,不要假设上一代的混合精度配置可直接沿用。

← 全部文章

johan's blog