半精度为何会让梯度变成 0 或 NaN?Autocast、FP16/BF16 与 GradScaler
从浮点数指数与尾数预算出发,手算梯度下溢和 loss scaling,拆解 PyTorch 2.14 autocast 与 GradScaler 数据流,并给出裁剪、累积、调试和失败边界。
上一篇把学习率绑定到成功消费的 token;现在每次更新何时发生已经清楚。新的瓶颈是数值格式:Transformer 的大矩阵乘法用 FP32 往往没有充分利用低精度硬件,而粗暴地对模型调用 .half() 又可能让小梯度归零、大激活溢出。
本篇聚焦一个核心问题:Automatic Mixed Precision(自动混合精度,AMP)怎样选择运算精度,以及 FP16 训练为什么需要 Gradient Scaling(梯度缩放)。
01 “少一半位宽”究竟少了什么?#
浮点数可抽象为
是符号, 的位数决定动态范围, 的位数决定相邻可表示数的精细程度。
| 格式 | 总位数 | 指数位 | 尾数位 | 核心取舍 |
|---|---|---|---|---|
| FP32 | 32 | 8 | 23 | 范围和精度都较好 |
| FP16 | 16 | 5 | 10 | 精度较细,但范围窄 |
| BF16 | 16 | 8 | 7 | 接近 FP32 范围,精度粗 |
FP16 最大有限值约为 65504;BF16 保留 8 位指数,因此更不容易因范围不足而 overflow(上溢),但它不是“更准确”,因为尾数更短。
02 为什么整个模型 .half() 很危险?#
神经网络不同运算需要不同数值性质:矩阵乘法通常能从低精度 Tensor Core 获益;softmax、归一化、指数和大规模 reduction(归约)更需要 FP32 的范围或累加精度。若所有参数、输入和运算一刀切成 FP16,模型失去高精度主权重与稳定运算的保护。
Autocast(自动类型转换)按算子策略选 dtype,而不是把整张图永久转换:
flowchart LR
A[FP32 参数与输入] --> B{autocast 算子策略}
B -->|matmul/linear/conv| C[FP16 或 BF16]
B -->|loss/reduction 等| D[FP32]
C --> E[FP32 loss]
D --> E
E --> F[scaled backward]
F --> G[unscale gradients]
G --> H{梯度有限?}
H -->|是| I[clip + optimizer.step]
H -->|否| J[跳过更新并减小 scale]mermaid03 小梯度如何在 FP16 中消失?#
考虑参数 的真实梯度 。若反向路径要把它存为 FP16,这个值可能低于可表示范围并舍入为 0。于是
该参数看似“没有梯度”。把 loss 乘尺度 后,链式法则让梯度变成
它更容易被 FP16 表示。优化前再除以 ,恢复 。缩放不会改变理想数学更新,只是把反向中间量暂时搬进可表示区间。
04 为什么尺度不能无限大?#
若另一处梯度为 ,同样乘 得 131072,超过 FP16 最大有限值,成为 inf。因此动态 scaler 在连续若干次梯度有限时增大 ,发现 inf/NaN 时跳过本次更新并减小 。
scale=65536 -> overflow -> skip step -> scale=32768
scale=32768 -> finite -> update
...连续稳定若干步...
scale=65536text05 当前 PyTorch 的最小正确循环#
PyTorch 2.14 推荐统一的 torch.autocast 与 torch.amp.GradScaler;旧的 torch.cuda.amp.* 入口已弃用。Autocast 只包住 forward 和 loss,backward 放在上下文外。
import torch
import torch.nn.functional as F
model = MyModel().cuda() # 参数保持 FP32
optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4)
scaler = torch.amp.GradScaler("cuda")
for input_ids, labels in loader:
input_ids = input_ids.cuda(non_blocking=True) # [B,L]
labels = labels.cuda(non_blocking=True) # [B,L]
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type="cuda", dtype=torch.float16):
logits = model(input_ids) # [B,L,V]
loss = F.cross_entropy(
logits.transpose(1, 2), labels,
ignore_index=-100,
) # scalar, normally FP32
scaler.scale(loss).backward()
scaler.step(optimizer) # 先检查非有限梯度;必要时跳过
scaler.update() # 调整下一步 scalepython输入 token 仍是 int64;autocast 只影响符合条件的浮点运算,不会把类别索引变成浮点数。不要在进入 autocast 前对模型或输入手工调用 .half()。
06 step、update 与学习率日程怎样配合?#
scaler.step(optimizer) 会先 unscale 并在梯度非有限时跳过 optimizer.step();scaler.update() 根据本轮结果调整 scale。若学习率日程按成功更新计时,就必须判断参数是否真的更新。
一个可检查的方法是比较 update 前后的 scale:发生 overflow 时新 scale 通常下降,且 optimizer step 被跳过。训练框架最好显式返回 update_succeeded,再决定是否推进 token 时钟与 EMA;不要无条件调用 scheduler。
old_scale = scaler.get_scale()
scaler.step(optimizer)
scaler.update()
new_scale = scaler.get_scale()
update_succeeded = new_scale >= old_scale
if update_succeeded:
token_schedule.step(global_valid_tokens)text这依赖动态缩放的默认回退行为;封装层若改变 growth/backoff 策略,应使用其明确的 skipped-step 信号。
07 梯度裁剪为什么必须先 unscale?#
若真实梯度范数是 2,而 scale 是 65536,直接裁剪看到的是 131072,会错误地把本来正常的梯度压小。顺序应是:
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
grad_norm = torch.nn.utils.clip_grad_norm_(
model.parameters(), max_norm=1.0,
)
scaler.step(optimizer)
scaler.update()pythonunscale_ 每个 optimizer 每步只能调用一次。多个 optimizer 时分别 unscale、检查和 step,并明确某一方 overflow 时是否允许另一方单独更新。排查问题时可临时给 clip_grad_norm_ 设置 error_if_nonfinite=True,让首个坏窗口立刻失败;常规动态缩放则应把跳步交给 scaler。
08 与梯度累积组合时,scale 何时更新?#
同一个有效 batch 的所有 micro-batch 必须使用同一 scale;只在完整 accumulation window 结束时 unscale、step 和 update。
optimizer.zero_grad(set_to_none=True)
for micro in micro_batches:
with torch.autocast("cuda", dtype=torch.float16):
loss = loss_fn(model(micro.x), micro.y) / len(micro_batches)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()python若每个 micro-batch 都 update(),同一组累积梯度混入不同尺度,最后无法用一次除法恢复。上一篇讨论的按有效 token 精确归一化仍然适用:可以累积 loss_sum,在 unscale 后再按全局 token 分母缩放 .grad。
09 FP16 与 BF16 应该怎样选?#
BF16 的指数范围与 FP32 相近,通常不需要 GradScaler;代价是有效数字更少,而且硬件必须高效支持。FP16 尾数比 BF16 多,范围却窄,通常要动态缩放。选择流程应是:
- 确认目标 GPU/加速器对哪种低精度有原生高吞吐;
- 跑 FP32 小基线,保存 loss 与梯度范数;
- 优先测试 BF16 autocast(硬件支持时);
- 使用 FP16 时启用 GradScaler;
- 比较吞吐、峰值显存、验证指标与非有限更新率。
“没有 NaN”不等于数值等价。要对固定 batch 比较 logits、loss、梯度方向和短程收敛,而不是要求逐位相同。
10 哪些运算需要特别留意?#
PyTorch autocast 有按设备维护的 op eligibility(算子资格)列表:某些算子转低精度,某些强制 FP32,另一些提升到最宽输入类型。自定义 CUDA op 或 autograd.Function 不会自动获得正确策略。
若某段在低精度不稳定,可嵌套禁用:
with torch.autocast("cuda", dtype=torch.float16):
hidden = encoder(x) # [B,L,D], maybe FP16
with torch.autocast("cuda", enabled=False):
stable = fragile_reduction(hidden.float()) # force FP32
logits = head(stable)pythonSoftmax 前手工减最大值、使用 cross_entropy 而非先 softmax 再 log、归一化统计量用稳定实现,仍然重要。AMP 不能修复数学上不稳定的自定义公式。
11 性能和显存为何不一定正好翻倍?#
低精度减小部分激活与临时张量,并加速合适尺寸的矩阵乘法;但 FP32 主参数、optimizer states、部分 FP32 运算和非浮点张量仍然存在。小模型可能受 Python、DataLoader 或 kernel launch 限制,AMP 转换开销反而盖过收益。
| 指标 | 说明 |
|---|---|
| 有效 tokens/s | 端到端学习吞吐 |
| 峰值 allocated/reserved 显存 | 区分真实张量与缓存池 |
| scaler scale | 是否持续回退 |
| skipped updates | 数值失败频率 |
| FP32 对照 loss | 精度漂移基线 |
| 验证指标 | 最终目标是否受损 |
预热若干步后再计时,并同步设备;否则异步 CUDA 会让 wall-clock 结果失真。
12 常见错误与最短调试路径#
| 症状 | 常见原因 | 最短检查 |
|---|---|---|
| 梯度裁剪后几乎为 0 | 对 scaled gradient 直接裁剪 | clip 前调用 unscale_ |
| 学习率日程偶尔抢跑 | overflow 跳步仍推进 scheduler | 同时记录 scale、lr、参数校验和 |
| BF16 模型转 FP16 后常溢出 | FP16 动态范围不足 | 改 BF16/FP32,检查激活最大值 |
| loss 正常但参数不再变化 | 小梯度下溢或连续 skip | 统计零梯度比例与 skipped updates |
| 自定义 op 输出 NaN | autocast 不知道其稳定 dtype | 局部禁用并强制 FP32 |
| AMP 没有提速 | 瓶颈不在低精度矩阵乘法 | profiler + FP32/AMP 端到端对照 |
| 恢复后行为改变 | 未保存 scaler state | 比较 scaler.state_dict() |
定位 NaN 时先固定同一 batch:依次运行 FP32、BF16 autocast、FP16 autocast 无 scaler、FP16 + scaler;逐层 hook 只记录 isfinite、绝对值最大值和 dtype,找到第一个异常算子,而不是等最终 loss 报错。
13 Checkpoint 还要多保存什么?#
除了 model、optimizer 和上一篇的 token schedule,还要保存 scaler:
torch.save({
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"schedule": schedule.state_dict(),
"scaler": scaler.state_dict(),
"global_tokens": global_tokens,
}, path)python若丢掉 scaler state,恢复后 scale 回到初始值,可能先连续 overflow,改变成功 update 的序列。保存只是第一步;恢复测试应从同一 checkpoint 分叉跑 3–5 步,比较 batch id、lr、scale、skip 标记和 loss。
14 失败场景与相近方法#
AMP 是训练数值格式策略,不等于 Quantization(量化):INT8/INT4 量化通常需要 scale/zero-point、校准或量化感知训练,目标常是推理压缩。TF32 只改变支持硬件上的 FP32 矩阵乘法内部精度,也不等同于把张量存成 FP16。
模型若包含极端指数、病态线性代数、自定义低精度 kernel,AMP 仍可能失败。应允许局部 FP32 或整体回退,并优先修正异常初始化、错误 loss、未归一化输入和爆炸梯度。
15 今天真正需要记住什么?#
- Autocast 按算子选择低精度或 FP32;不要把模型和输入粗暴地全部
.half()。 - FP16 的窄范围会让小梯度下溢、大值上溢;GradScaler 通过 scale、检查、跳步和回退保护更新。
- 裁剪前必须 unscale;梯度累积窗口内必须保持同一 scale。
- BF16 范围更宽但尾数更短,是否更快、更稳取决于硬件和模型,必须用 FP32 基线验证。
16 思考题与小练习#
- 对梯度
[2^-30, 2^-20, 2],分别用 scale2^10和2^16计算缩放值,判断哪个更可能下溢或上溢。 - 给累积 3 个不同有效 token 数 micro-batch 的循环加入精确 token 归一化、unscale 与 gradient clipping,并标出每一步张量/标量 dtype。
- 设计一个定位首个非有限激活的 hook;要求只保存层名、dtype、shape、最大绝对值与有限值比例,避免复制完整张量拖慢训练。
相关工作#
- Micikevicius et al., Mixed Precision Training ↗,系统化提出 FP16 主干、FP32 主权重与 loss scaling。
- Kalamkar et al., A Study of BFLOAT16 for Deep Learning Training ↗,分析 BF16 的范围、精度与训练表现。
- Micikevicius et al., FP8 Formats for Deep Learning ↗,把混合精度设计推进到 FP8 格式与缩放策略。
17 下一篇预告#
混合精度减少了计算和激活成本,但大模型仍可能放不进单卡。下一篇将研究 activation checkpointing 如何用重算换显存,以及它与训练 checkpoint 文件为何只是同名、不是同一件事。