混合精度训练中数值稳定性的核心问题在于FP16的窄动态范围,主要靠损失缩放、主权重副本和梯度裁剪这三招兜底,其中梯度下溢是最隐蔽的坑。
混合精度训练和单精度训练区别在哪
混合精度训练不是把模型从FP32整体换成FP16,它让权重和激活值用FP16存储和计算,同时保留一份FP32主权重用于参数更新,单精度训练则是一路FP32走到底,这个区别直接决定了数值稳定性的难度。
FP16和FP32的底层差异,重点看三处:
- 指数位宽度:FP32用8位存指数,FP16只有5位。
- 动态范围:FP32覆盖约1e-38到3e38,FP16只能在约1e-5到6e4之间表现正常。
- 最小正规数:FP16的最小正规数是2^-14,约6e-5,再小的正数已经不再是标准浮点。
| 指标 | FP16 | FP32 |
|---|---|---|
| 指数位 | 5位 | 8位 |
| 动态范围 | 约1e-5到6e4 | 约1e-38到3e38 |
| 最小正规数 | 约6e-5 | 约1.2e-38 |
| 尾数有效精度 | 约3位十进制 | 约7位十进制 |
动态范围窄意味着什么?训练中如果某个梯度小于FP16的最小正规数,它在传播途中就被吞成了0,注意力机制里softmax之后的反向梯度,就经常落在极小的数值区间,业内专家指出,这类梯度在FP16下直接消失,模型前几层卷积核就会退化。
另一个常见选择是BF16,它同样占16位,但指数位有8位,动态范围和FP32几乎一样宽,代价是尾数更短,在英伟达A100、H100这类算力集群上做分布式训练,BF16比FP16不容易下溢,但舍入误差更大,取舍要看模型对精度的敏感程度。
混合精度训练梯度下溢怎么办
梯度下溢是混合精度训练最常遇到、也最容易误判的现象,loss曲线突然变成一条直线,或者训练了好几个小时验证集指标纹丝不动,先别急着调学习率,检查一下梯度值。

下溢是怎么发生的
反向传播计算出的梯度,在FP32下可能小到1e-8甚至更低,FP16能表示的最小正数约为5.96e-8,小于这个值的正数会被直接抹掉,底层梯度的信息一旦丢光,网络前几层就学不到任何东西。
用日志确认是否在下溢
在PyTorch里手动监控梯度范数,做法是每次反向传播后遍历模型参数,记录param.grad.abs().max(),如果大多数层的最大梯度绝对值长期小于6e-5,说明模型正处在下溢边缘。
动态损失缩放是标准解法
损失缩放的核心思路:前向传播时把loss乘上一个缩放因子,让梯度整体变大,完成反向传播后再把梯度除回来,用于参数更新,这样底层的小梯度有机会落入FP16的可表示范围。
- 手动固定缩放:适合临时调试,但缩放太小仍然下溢,太大会导致梯度溢出为inf。
- 动态缩放:PyTorch的
torch.cuda.amp.GradScaler会自动调整因子,梯度溢出时减小缩放,连续若干步无溢出就增大。
在老的Apex里,amp.initialize也能实现类似能力,但动态缩放逻辑需要手动搭得更细致一些。
混合精度训练损失缩放设置:从固定到动态
损失缩放里的缩放值设置,没有万能答案,但有一些经验法则可以遵循。
缩放因子怎么选
初始缩放因子,行业共识一般从1024或32768开始,动态缩放会自动调整,但前提是给它留够空间。
- 在
torch.cuda.amp里,GradScaler(init_scale=214, growth_factor=2.0, backoff_factor=0.5)是常见的起点。 - 最大缩放因子建议设到2的16次方以上,防止动态调整时增长受限。
- 如果训练中频繁出现inf或nan,先看缩放因子是否过大,再看学习率是否过高引发梯度爆炸,两者会同时触发溢出。

损失缩放和batch size的关系
Batch size越大,平均梯度越平滑,通常可以用更小的缩放因子,相反,batch size小、梯度噪声大时,动态缩放会频繁调整,建议把growth_interval调大,比如每2000步增长一次,而不是默认的200步。
更新时机不能搞错
用GradScaler时,optimizer.step()必须放在scaler.step(optimizer)里,scaler.update()在每个batch末尾调用,顺序一旦颠倒,缩放逻辑直接失效,这也是不少新手混合精度训练不稳定的根源。
只开AMP还不够:主权重和梯度裁剪
很多人以为调好损失缩放就万事大吉,实际上主权重副本和梯度裁剪同样关键。
主权重副本的作用
FP16的有效精度只有约3位十进制数字,而权重更新量通常比权重本身小好几个数量级,如果直接在FP16权重上做更新,更新量会因为舍入误差被丢得干干净净,因此训练时必须保留一份FP32主权重,每次更新在FP32上执行,再转回FP16用于前向传播,PyTorch AMP和Apex默认做了这件事,但如果你写自定义优化器,需要确认这个逻辑没被破坏。
梯度裁剪的数值稳定技巧
梯度裁剪在混合精度下有个容易被忽略的点:如果先裁剪再缩放,裁剪阈值会被缩放因子改变,结果跟预期不一致,正确的顺序是先用scaler.unscale_把梯度还原,再做torch.nn.utils.clip_grad_norm_裁剪,最后让优化器按正常逻辑更新。
训练中的数值稳定监控清单
这份清单是混合精度训练期间应该定期检查的项目,顺序很重要:
- 看loss曲线拐点:如果loss在某个step突然跳高又回落,先怀疑缩放因子调整,而不是单纯的学习率问题。
- 记录梯度范数

:每N步打印一次梯度的L2范数,与单精度训练的基线对比,差太多说明数值链路有异常。
- 检查激活值范围:在前向传播时插入hook,看中间层激活值是否落在FP16合理范围,大致1e-3到1e4之间。
- 验证权重更新量:计算
weight_fp32 - weight_prev的最大绝对值,如果这个值持续小于FP16最小正规数,优化器更新逻辑可能有问题。 - 定期跑一次完整评估:混合精度训练通常不会明显掉点,如果发现验证集精度比单精度训练低不少,回头看是否在
autocast之外手写了强制FP16计算的算子。
混合精度训练数值稳定性Q&A
混合精度训练会不会降低模型最终精度?
正确配置时,混合精度训练的最终精度与单精度训练基本持平,主流框架使用动态损失缩放和主权重副本后,精度差距通常控制在可接受范围内,如果出现显著掉点,优先排查损失缩放和梯度裁剪配置,而不是直接归咎于FP16。
混合精度训练显存占用能省多少?
显存节省主要来自激活值和权重缓存改用FP16存储,在视觉模型和Transformer结构上,整体显存占用比单精度训练省掉接近一半,具体数值取决于模型结构和batch size,省显存的代价是代码里需要多一层数值监控,否则inf和nan会悄悄出现。
为什么我的混合精度训练loss变成NaN?
最常见的原因有两个:一是梯度溢出,缩放因子过大导致梯度在FP16中变成inf;二是学习率过高,参数更新后权重直接溢出,解决办法是开启动态缩放、适当降低学习率,并在scaler.update()之前检查scaler.scale的值是否异常。
归根结底,混合精度训练是用理解换速度,摸清FP16的脾气,把损失缩放和主权重这两件事做扎实,数值稳定性问题就能被控制住。