观文听傑

返回

上一篇把学习率绑定到成功消费的 token;现在每次更新何时发生已经清楚。新的瓶颈是数值格式:Transformer 的大矩阵乘法用 FP32 往往没有充分利用低精度硬件,而粗暴地对模型调用 .half() 又可能让小梯度归零、大激活溢出。

本篇聚焦一个核心问题:Automatic Mixed Precision(自动混合精度,AMP)怎样选择运算精度,以及 FP16 训练为什么需要 Gradient Scaling(梯度缩放)。

01 “少一半位宽”究竟少了什么?#

浮点数可抽象为

x=(1)s×m×2ex=(-1)^s\times m\times 2^e

ss 是符号,ee 的位数决定动态范围,mm 的位数决定相邻可表示数的精细程度。

格式总位数指数位尾数位核心取舍
FP3232823范围和精度都较好
FP1616510精度较细,但范围窄
BF161687接近 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]
mermaid

03 小梯度如何在 FP16 中消失?#

考虑参数 ww 的真实梯度 g=230g=2^{-30}。若反向路径要把它存为 FP16,这个值可能低于可表示范围并舍入为 0。于是

wwη0w\leftarrow w-\eta\cdot0

该参数看似“没有梯度”。把 loss 乘尺度 S=216S=2^{16} 后,链式法则让梯度变成

g=Sg=216230=214g'=Sg=2^{16}\cdot2^{-30}=2^{-14}

它更容易被 FP16 表示。优化前再除以 SS,恢复 g=g/Sg=g'/S。缩放不会改变理想数学更新,只是把反向中间量暂时搬进可表示区间。

04 为什么尺度不能无限大?#

若另一处梯度为 g=2g=2,同样乘 2162^{16} 得 131072,超过 FP16 最大有限值,成为 inf。因此动态 scaler 在连续若干次梯度有限时增大 SS,发现 inf/NaN 时跳过本次更新并减小 SS

scale=65536 -> overflow -> skip step -> scale=32768
scale=32768 -> finite   -> update
...连续稳定若干步...
scale=65536
text

05 当前 PyTorch 的最小正确循环#

PyTorch 2.14 推荐统一的 torch.autocasttorch.amp.GradScaler;旧的 torch.cuda.amp.* 入口已弃用。Autocast 只包住 forward 和 loss,backward 放在上下文外。

输入 token 仍是 int64;autocast 只影响符合条件的浮点运算,不会把类别索引变成浮点数。不要在进入 autocast 前对模型或输入手工调用 .half()

06 stepupdate 与学习率日程怎样配合?#

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()
python

unscale_ 每个 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 多,范围却窄,通常要动态缩放。选择流程应是:

  1. 确认目标 GPU/加速器对哪种低精度有原生高吞吐;
  2. 跑 FP32 小基线,保存 loss 与梯度范数;
  3. 优先测试 BF16 autocast(硬件支持时);
  4. 使用 FP16 时启用 GradScaler;
  5. 比较吞吐、峰值显存、验证指标与非有限更新率。

“没有 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)
python

Softmax 前手工减最大值、使用 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 输出 NaNautocast 不知道其稳定 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 今天真正需要记住什么?#

  1. Autocast 按算子选择低精度或 FP32;不要把模型和输入粗暴地全部 .half()
  2. FP16 的窄范围会让小梯度下溢、大值上溢;GradScaler 通过 scale、检查、跳步和回退保护更新。
  3. 裁剪前必须 unscale;梯度累积窗口内必须保持同一 scale。
  4. BF16 范围更宽但尾数更短,是否更快、更稳取决于硬件和模型,必须用 FP32 基线验证。

16 思考题与小练习#

  1. 对梯度 [2^-30, 2^-20, 2],分别用 scale 2^102^16 计算缩放值,判断哪个更可能下溢或上溢。
  2. 给累积 3 个不同有效 token 数 micro-batch 的循环加入精确 token 归一化、unscale 与 gradient clipping,并标出每一步张量/标量 dtype。
  3. 设计一个定位首个非有限激活的 hook;要求只保存层名、dtype、shape、最大绝对值与有限值比例,避免复制完整张量拖慢训练。

相关工作#

  1. Micikevicius et al., Mixed Precision Training,系统化提出 FP16 主干、FP32 主权重与 loss scaling。
  2. Kalamkar et al., A Study of BFLOAT16 for Deep Learning Training,分析 BF16 的范围、精度与训练表现。
  3. Micikevicius et al., FP8 Formats for Deep Learning,把混合精度设计推进到 FP8 格式与缩放策略。

17 下一篇预告#

混合精度减少了计算和激活成本,但大模型仍可能放不进单卡。下一篇将研究 activation checkpointing 如何用重算换显存,以及它与训练 checkpoint 文件为何只是同名、不是同一件事。

半精度为何会让梯度变成 0 或 NaN?Autocast、FP16/BF16 与 GradScaler
https://zwjcode.cn/blog/automatic-mixed-precision-fp16-bf16-gradscaler
作者
发布于 2026年9月14日
版权协议 CC BY-NC-SA 4.0
评论加载似乎遇到了问题,请尝试刷新页面。