news 2026/9/30 7:35:33

混合精度训练崩溃之谜:手写梯度缩放,彻底根治NaN

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
混合精度训练崩溃之谜:手写梯度缩放,彻底根治NaN

训练跑着跑着 loss 变成 NaN,这大概是每个搞深度学习的人都经历过的噩梦。早期我遇到这种情况,第一反应是调小学习率、清理数据、换初始化,结果发现治标不治本。真正让我彻底理解问题根源的,是后来深入研究自动混合精度(AMP)和梯度缩放(Gradient Scaling)的实现细节——原来训练崩溃的元凶往往不是数据或超参,而是梯度本身溢出了 FP16 的动态范围。

这篇文章我不打算写成官方文档的翻译版,而是从实际踩坑和源码实现的角度,把 AMP 里的梯度缩放到底在做什么、为什么非做不可、以及手写实现时有哪些细节容易翻车,一次性讲透。无论你是刚开始尝试混合精度训练,还是已经在用了但被动态缩放搞得一头雾水,这篇都值得收藏。

1. 先搞清楚一个前提:FP16 的动量优势与它的数字短板

要理解梯度缩放,先得知道为什么训练要用 FP16。这个问题的答案其实非常现实:在支持 Tensor Core 的 GPU 上,FP16 的矩阵乘法和卷积算子通常能达到 FP32 的数倍吞吐量,而且显存占用直接减半——这意味着你能塞下更大的 batch size,或者把模型做得更大。在如今动不动几十亿参数的规模下,这已经不是锦上添花,而是能不能跑完训练的区别。

但 FP16 有个硬伤:动态范围太窄了。它只有 5 个指数位和 10 个尾数位,表示的数值范围大约是 65504 的最大值,最小正规格化数约 6.1e-5。对比 FP32 的 1e-38 到 3.4e38,FP16 在极小和极大两个方向上都相当脆弱。

注意一个关键事实:梯度在反向传播中的分布非常不均匀。某些层的梯度可能小到 1e-6,另一些层的梯度可能大到几十甚至上百。FP16 一存,小的直接下溢为 0,大的直接上溢为 inf。下溢还不那么致命,顶多是某些参数更新不了,让训练收敛变慢;上溢则是灾难性的,一个 inf 传进更新公式,loss 立刻变成 NaN,整个训练报废。

在我实际遇到的案例里,有一种特别隐蔽的情况是:早期训练挺正常,跑到第几百步后突然 loss 跳动然后崩掉。检查数据没问题、学习率没问题,最后发现是某些层的梯度在某些 batch 上出现了比较大的异常值,用 FP16 存储时发生了溢出。这个场景解释了一个核心问题:为什么混合精度不能简单地把模型和梯度全切成 FP16,而必须保留 FP32 的 master weight,并对梯度做缩放处理。

2. 混合精度的设计逻辑:哪部分用 FP16,哪部分必须留在 FP32

真正落地的时候,AMP 的流程是这样的:前向传播和反向传播的矩阵运算用 FP16 跑,得到速度提升;但模型的权重保留一份 FP32 的 master copy,用于参数更新,保证数值稳定性。每个训练步骤里,FP16 权重负责前向和反向,计算出 FP16 梯度后,再转回 FP32 并对 master weight 做更新。

这里有一个很容易忽略的细节:为什么不能直接更新 FP16 权重?因为权重更新公式是 (w = w - \eta \cdot g),学习率 (\eta) 通常很小,比如 1e-3,梯度也很小,比如 1e-4,那么更新量就是 1e-7。FP16 的最小可表示间隔大约在 1e-8 到 1e-7 这个量级(取决于数值本身的大小),更新量会被直接吞掉。简单说,FP16 存不住"微小但关键"的权重变化,模型就无法精细收敛。所以 master weight 必须是 FP32,计算出的 FP16 梯度必须转回 FP32 再做更新。

但问题来了:前向用 FP16 权重,反向算出的梯度也是 FP16。如果某个梯度值是 1e-5(完全在 FP16 的动态范围内),它虽然没下溢到 0,但精度已经非常差——FP16 在接近 1e-5 这个量级时的尾数分辨率不够,梯度近似误差会被放大,训练质量下降。更糟的是,如果梯度值超过 65504,直接变成 inf。

所以单靠"FP16 前向、FP32 更新"还不够,必须引入梯度缩放来人为地把梯度从很小的量级抬到 FP16 能准确表示的范围。这就是整个 AMP 技术里最核心、也最常被忽视的机制。

3. 梯度缩放的工作机制:为什么放大 loss 等于放大梯度

梯度缩放的基本想法非常反直觉:训练时把 loss 乘上一个大于 1 的系数,再反向传播。

反直觉的点在于:我们平时训练都希望 loss 越小越好,为什么还要主动放大 loss?答案在链式法则里。反向传播的梯度计算是逐层累积的,根据链式法则:

[ \frac{\partial L}{\partial w} = \frac{\partial L}{\partial \text{output}} \cdot \frac{\partial \text{output}}{\partial w} ]

如果 (L) 被放大为 (s \cdot L),那么每层梯度都会等比例乘以 (s),甚至因为逐层链式相乘,整体梯度会被放大 (s) 倍。这就是"调大 loss 数值能等比放大梯度数值"的原理。

放大之后,原本是 1e-5 的梯度变成 0.01,FP16 能表示得很精确;原本是 0.001 的梯度变成 1.0,存储精度也没有损失。算完梯度之后,我只在最后一步——更新权重之前——把梯度除以 (s),恢复到真实梯度值。这个"乘上去再除回来"的过程,就是整个缩放机制的全部秘密。

有人会问:那这和自己手动把梯度放大了一下有什么区别?没有区别,本质上就是一样的。区别只在于它是自动化、动态调整的,由库来管理缩放系数,程序员不需要去分析每层的梯度谱来手工挑系数。

3.1 缩放系数的动态调整策略

缩放系数 (s) 不是拍脑袋定的。太小了,小的梯度仍然下溢;太大了,放大后的梯度会在 FP16 里上溢。所以框架必须动态监测。

PyTorch 和 NVIDIA 的实现逻辑是梯度监测驱动的:

  1. 初始设置一个缩放系数,通常是 65536。这个数的来源很有趣:因为 FP16 最大值为 65504,如果初始系数是 65536,相当于让大部分原梯度在被放大后刚好处于 FP16 的表示上限附近——理论上偏向"给梯度最大的放大空间"。
  2. 每个训练步骤反向传播时,检查这一步的梯度中是否出现了 inf 或 NaN。
  3. 如果出现了,说明缩放太激进了,把系数减半(实际是乘上某个衰减系数),当前的优化器更新直接跳过,这一轮白跑。
  4. 如果连续很多步都没有出现 inf/NaN,说明还有放大空间,把系数适度调大(通常是乘 2 的幂次倍数)。

这个逻辑我写过自己的最小实现,核心代码大概长这样:

# 简化版动态缩放逻辑 scale = 65536.0 growth_factor = 2.0 backoff_factor = 0.5 steps_since_last_overflow = 0 growth_interval = 2000 for step in range(total_steps): optimizer.zero_grad() # 放大loss loss = compute_loss(model, batch) * scale loss.backward() # 检查梯度是否有溢出 grad_has_overflow = False for p in model.parameters(): if p.grad is not None: if not torch.isfinite(p.grad).all(): grad_has_overflow = True break if grad_has_overflow: # 溢出处理: 跳过更新, 缩小scale scale *= backoff_factor steps_since_last_overflow = 0 continue # 梯度反缩放 + 更新 with torch.no_grad(): for p in model.parameters(): if p.grad is not None: p.grad /= scale # 恢复到真实梯度 optimizer.step() steps_since_last_overflow += 1 if steps_since_last_overflow >= growth_interval: scale *= growth_factor steps_since_last_overflow = 0

注意一个核心操作顺序:溢出检查是在缩放后的梯度上做的,反缩放是在检查通过之后才做的。如果先反缩放再检查,小的梯度变成极小值,你根本检查不出 FP16 上溢的危险。这个顺序很多手写实现会搞反,导致检查形同虚设。

4. PyTorch AMP 里的梯度缩放:从 GradScaler 到 autocast 的配合

你不需要手写上面的逻辑,因为 PyTorch 的torch.cuda.amp.GradScaler已经帮你封装好了。但理解了原理之后,用起来才会明白每一步的语义。

标准用法是这样的:

scaler = torch.cuda.amp.GradScaler() for epoch in range(epochs): for batch in dataloader: optimizer.zero_grad() with torch.cuda.amp.autocast(): loss = model(batch) loss = criterion(loss, target) # 关键: 这里的loss会被内部放大 scaler.scale(loss).backward() # 关键: 反缩放 + 裁剪 + 更新 scaler.step(optimizer) # 更新scale系数 scaler.update()

scaler.scale(loss).backward()内部做的是我上面那段代码的封装:loss 乘以当前 scale,然后反向传播。scaler.step(optimizer)内部做的事比想象中多:

  1. 遍历optimizer的参数,检查梯度中有没有 inf/NaN。
  2. 如果发现溢出,直接跳过optimizer.step(),不更新任何参数(这一点非常重要,否则这次迭代的权重被污染)。
  3. 如果没有溢出,把所有梯度除以 scale,然后调用原来的optimizer.step()。

这里有一个极其容易踩的坑:如果用了梯度裁剪(gradient clipping),顺序不能搞错。官方推荐的顺序是:

scaler.scale(loss).backward() # 先反缩放,再裁剪,最后更新 scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm) scaler.step(optimizer) scaler.update()

为什么不直接clip_grad_norm_在缩放后的梯度上?因为梯度裁剪的阈值是针对真实梯度设定的。缩放后的梯度整体大了 65536 倍,直接裁剪会把所有梯度的模长压到一个离谱的范围,导致更新量严重失真。所以必须先scaler.unscale_(optimizer)把梯度还原,再裁剪,再更新。

4.1 autocast 的作用边界

再深入说一句autocast到底自动了什么。很多人误以为 autocast 自动做了混合精度的一切,其实它只负责前向计算中的算子精度选择——在支持 FP16 的算子(如 matmul、conv)上自动用 FP16,在不支持的算子(如某些归一化层、softmax)上保持 FP32。

而反向传播的梯度也是通过前向保存的中间激活和 FP16 权重计算出来的,所以梯度天然就是 FP16 的。这时梯度缩放就登场了,它不依赖 autocast,而是独立运作的。你甚至在完全没有 autocast 的纯 FP16 训练脚本里也能用 GradScaler 做梯度管理。

这两个机制一个是精度分配策略,一个是数值保护策略,缺一不可。理解了这层关系,就不会再问"为什么有了 autocast 还要显式写 GradScaler"这种问题了。

5. 一个容易被忽略的角落:优化器内部状态与 master weight 的关系

用 AMP 的时候,我不建议直接把 optimizer 挂在 FP16 模型参数上。原因前面提过:FP16 权重更新时,微小更新量会被舍入吞掉。PyTorch 中一种常规做法是:

model = model.half() # 模型FP16 optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)

这样 optimizer 操作的还是 FP16 的参数,更新量先变成 FP16,这违背了混合精度的初衷。真正可靠的做法是:

optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) # AMP内部维护一份model的FP32副本作为master weight # 或者你自己手动维护: model_fp32 = copy.deepcopy(model).float() # 每次更新后将model_fp32同步给model(转为FP16)

但绝大多数情况下,你用 torch.cuda.amp 时是不需要手动维护 master weight 的。PyTorch 1.6+ 的 AMP 训练流程中,模型参数本身保留 FP32,只是在前向传入 autocast 时转换运算。你可以把模型保持 FP32,让 autocast 在内部做临时转换,梯度在反向传播后是 FP16,GradScaler 负责梯度缩放和保护。

这里真正的权衡点是:模型保持在 FP32,前向时自动转 FP16 运算,那么模型自身的内存节省效果就没有了。想要真正的内存减半,还是得手动model.half()把权重存成 FP16,同时维护 FP32 master weight。当下很多大模型训练框架选择的是:权重 FP16 存储 + FP32 优化器状态 + 动态节点管理,这样才能真正吃到显存红利。这也是我之前做大规模模型训练时反复对比过的方案,细节差异很大。

6. 实战避坑:动态范围问题、NaN 排除和其他放大注意点

纸上得来终觉浅,我把实际训练中遇到的和梯度缩放直接相关的一系列坑整理出来,每个都是环境变量级别的教训。

6.1 溢出后跳步造成的"幽灵卡顿"

GradScaler 在检测到溢出时会跳过这一轮优化器更新。如果训练配置较大、batch 很大,一次溢出跳过的算力成本不小。我遇到过一种情况:某个模型频繁出现溢出,导致有效训练步数只有理论的一半,loss 曲线像锯齿一样上下摆动,进展缓慢。

排查下来,问题出在某个网络层的输入包含着可能出现较大方差的特征。解法有两种:一是降低初始缩放系数(比如从 65536 改到 128),牺牲一部分对小梯度的精度保护,换取更少的溢出跳步;二是找到溢出的源头,比如疑似瓶颈层,改用更稳定的激活函数或归一化。

我个人更倾向后者,因为降低初始系数是一种"向下兼容"的妥协——它会让小梯度重新靠近 FP16 的下溢边界,尤其在小 batch 和长训练后期,这个副作用会被放大。

6.2 等梯度流经多个算子时,放大系数的累积效应

这条经验是在实现自定义算子时领悟的。梯度缩放虽然全套包装在框架里,可一旦你写了自定义的 autograd.Function,情况就变了。如果你在自定义反向函数内部用到了需要精确 FP16 的中间结果,就必须清楚当前作用的 scale 是多少,否则自定义反向传播里再手动乘除一下,数值就乱套了。

更微妙的是:梯度在跨过多层反向传播时是链式相乘的,scale 只作用于最外层的 loss,但随着链式规则逐层往后,scale 的量级会在每一层等比例传递。也就是说,如果你在某个中间层手动插入了一次除法,那不是把最外层的 scale 效应去除,而是把一个庞大的因子从后续所有反传路径上砍掉,梯度直接乱掉。所以永远不要在自定义反向里尝试"还原梯度原来的样子",除非你明确知道整体的链式缩放关系。要处理梯度,就在优化器更新前统一 unscale_。

6.3 检查 FP16 下的 inf 要区分上溢和下溢

前面我说过,inf 通常是上溢。但有一种情况是极小概率下的"数据本身不合法"——比如你的标签里有 NaN,或者数据预处理在某条样本上产生了 inf,反向传播时梯度沿着这条路径变 NaN。这种情况下的"溢出"不是 FP16 导致的,而是源数据污染。区分两者的方法很简单:把梯度转回 FP32 后重新检查一遍是否仍然 NaN。如果 FP32 下也是 NaN,说明是数据问题;如果 FP32 下正常、FP16 下才溢出,那才是动态范围的锅。

很多初学者把这两类问题混为一谈,要么疯狂调 scale 却始终无效,要么拼命清洗数据却还是崩。我建议在训练脚本里做一次双轨检查:先关闭 autocast 和 GradScaler,在纯 FP32 下跑 50 步观察是否出 NaN。FP32 干净的话,问题百分百出在混合精度的数值处理上;FP32 本身都炸,那就先修数据和计算图。

6.4 嵌入式场景的延伸:低精度推理时的近似缩放思路

顺便提一句,我在 RK3506 这类嵌入式 SoC 上做推理优化时,也遇到过和梯度缩放本质类似的"精度困境":中间特征图的动态范围很宽,但低精度整型表示的步长有限。解决方案思想上非常接近——对特征图做逐通道的缩放(per-channel scale),把动态范围压缩到可表示空间内,推理结束再还原。虽然训练阶段梯度缩放解决的是反向传播问题,推理阶段解决的是前向数值分布问题,但"动态范围不够,用缩放系数来凑"这个思路是一脉相承的。你如果在嵌入式端部署模型时对动态范围控制有困惑,可以把训练阶段对梯度的缩放心态搬到推理阶段的特征归一化上去,排查路径是类似的。

6.5 缩放系数增长策略的调整

前面代码里的growth_interval和growth_factor是经典配置,但实际训练中我根据任务调整过几轮。比如:使用重梯度噪声的任务(如强化学习),梯度方差极大,频繁溢出,需要更保守的策略:更低初始系数、更大的增长间隔。而视觉分类这类梯度相对稳定的任务,可以更激进地增大系数去保护极小梯度。

关于增长策略,有一个来自实践的小技巧:把 GradScaler 的_growth_interval暴露到日志里监控,小幅调整到 1000~4000 区间,观察溢出步数和 loss 收敛速度的权衡曲线。你会发现 2000 不是神圣不可动的数字,而是各种标准任务折中的产物。我自己在训练一批分割模型时,把 interval 从 2000 调到 5000,溢出次数没增加多少,但后期小梯度的精度保护明显更到位,最终精度有小幅提升。

7. 手写一个最小可运行的梯度缩放验证实验

为了验证前面所有解释,我建议你自己动手做一个 20 行以内的实验。核心思路是构造一个梯度值极小的简单模型,对比两种情况下的更新效果。

import torch import torch.nn as nn torch.manual_seed(42) model = nn.Linear(2, 1, bias=False).half().cuda() optimizer = torch.optim.SGD(model.parameters(), lr=0.1) # 制造一个极小梯度场景: 输入极小, 目标也是极小 x = torch.tensor([[1e-4, 1e-4]], device='cuda').half() y = torch.tensor([[1e-5]], device='cuda').half() # 情况1: 无缩放的FP16训练 optimizer.zero_grad() loss = (model(x) - y).pow(2) print("初始loss:", loss.item()) loss.backward() print("梯度:", model.weight.grad.item()) optimizer.step() print("无缩放更新后权重:", model.weight.data.item()) # 情况2: 缩放后的更新 model2 = nn.Linear(2, 1, bias=False).half().cuda() model2.load_state_dict(model.state_dict()) optimizer2 = torch.optim.SGD(model2.parameters(), lr=0.1) scale = 65536.0 optimizer2.zero_grad() loss2 = (model2(x) - y).pow(2) * scale loss2.backward() print("缩放后梯度:", model2.weight.grad.item()) # 手动反缩放 model2.weight.grad.data /= scale optimizer2.step() print("缩放更新后权重:", model2.weight.data.item())

如果你在真机上跑这个实验,会看到两个结果之间的差异:无缩放情况下,梯度可能非常小,FP16 存储后更新量几乎为 0,权重几乎不变;而有缩放的情况,更新能正常作用。这就在最简层面证明了缩放的价值:不是把数值人为变大好看,而是把有效信息从 FP16 的精度盲区里捞出来。建议你把 scale 换成 1.0、16.0、65536.0 各跑一次,观察权重的变化曲线,数值上的差异会非常直观。

说回经验和总结的话:我的核心体会是,AMP 这套东西写起来不算复杂,五个函数调用就能跑通,但真正有价值的不是流式调用,而是理解每一个内部数值行为背后的动机。NaN 从哪来、为什么要调 loss、为什么先 unscale 再裁剪、为什么溢出要跳步,这些点连起来之后,你才能自如地调整混合精度配置,去适配那些"标准平台、标准模型"之外的个性化训练任务。

如果你现在正处于被 NaN 折磨的阶段,我建议的排查顺序是:先纯 FP32 跑通,再用原生 FP16 + 无缩放跑,观察差值,再加入 autocast + GradScaler 组合,过程中记录每一步的梯度统计。这套诊断流程比盲目调整学习率和清洗数据有效得多。梯度缩放不是银弹,它是数值保护里的一块拼图,但把这块拼图补上之后,你会发现很大一部分难以解释的训练崩溃,突然就都有了答案。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/30 7:35:18

HR 避坑:2026 人才测评工具 5 大常见使用误区拆解

过去一年里,人才测评工具的市场规模持续走高,全球员工绩效评估平台预计在2026年达到54亿美元。但工具多了,踩坑的人也多了。不少HR把测评买回来才发现:候选人不配合、报告看不懂、数据躺在后台没人用。问题不在于工具本身&#xf…

作者头像 李华
网站建设 2026/9/30 7:35:05

同程出行服务客服咨询AI流量赋能,同程出行服务科技重塑智能体验新标杆

近期,由湖南改变生物科技有限公司主办、本因内酵未徕品牌协办的“生物科技健康论坛暨AI赋能大健康产业启动会”在长沙市步步高福鹏喜来登酒店隆重举行。活动以“AI流量赋能实体破局——中小企业增长峰会”为主题,汇聚全国大健康行业专家、中小企业负责人、机构代表及…

作者头像 李华
网站建设 2026/9/30 7:34:26

DDNS攻击手法与防御体系全面解析:从DNS重绑定到域名劫持

1. DDNS攻击目标画像:为什么攻击者死盯动态域名1.1 DDNS到底是怎么工作的:三分钟搞懂核心机制DDNS的设计初衷很朴素:你家里或小公司的公网出口IP是动态的,宽带运营商隔一段时间就重新分配一次地址,但你的NAS、摄像头、…

作者头像 李华
网站建设 2026/9/30 7:34:26

深度学习大模型全链路实战:从环境搭建到ONNX部署的避坑指南

简介:这份资源是一套面向深度学习研发人员、数据科学家及技术爱好者的全链路实战指南,聚焦大模型从构建到部署的完整流程,帮助具备一定理论基础的学习者打通环境搭建、数据处理、模型选择与训练、评估优化到最终部署的关键环节。资源包内含1个…

作者头像 李华
网站建设 2026/9/30 7:33:39

DAMON实战:异构内存冷热数据自动迁移,驱动内存数据库性能飙升

先把场景摆出来。一台双路服务器,装了 256GB DRAM 加 512GB 持久内存,跑内存数据库。压测一上量,p99 延迟比纯 DRAM 机器差了快三倍。原因不复杂:异构内存访问效率的核心就是热数据别放慢速层,可传统机制根本没法精细地…

作者头像 李华
网站建设 2026/9/30 7:33:04

生成对抗网络训练逻辑详解:从损失函数到交替更新

这次我们来看生成对抗网络(GAN)中最核心的一个问题:训练逻辑到底是什么。很多人第一次接触 GAN 时,会看到一张生成器和判别器相互博弈的示意图,但这张图距离真正理解训练流程还差很远。真正困惑人的地方在于&#xff1…

作者头像 李华