拓冰建站拓冰建站
首页 / 资讯中心 / 正文

混合精度训练崩溃之谜:手写梯度缩放,彻底根治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.4e38FP16 在极小和极大两个方向上都相当脆弱。注意一个关键事实梯度在反向传播中的分布非常不均匀。某些层的梯度可能小到 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.01FP16 能表示得很精确原本是 0.001 的梯度变成 1.0存储精度也没有损失。算完梯度之后我只在最后一步——更新权重之前——把梯度除以 (s)恢复到真实梯度值。这个乘上去再除回来的过程就是整个缩放机制的全部秘密。有人会问那这和自己手动把梯度放大了一下有什么区别没有区别本质上就是一样的。区别只在于它是自动化、动态调整的由库来管理缩放系数程序员不需要去分析每层的梯度谱来手工挑系数。3.1 缩放系数的动态调整策略缩放系数 (s) 不是拍脑袋定的。太小了小的梯度仍然下溢太大了放大后的梯度会在 FP16 里上溢。所以框架必须动态监测。PyTorch 和 NVIDIA 的实现逻辑是梯度监测驱动的初始设置一个缩放系数通常是 65536。这个数的来源很有趣因为 FP16 最大值为 65504如果初始系数是 65536相当于让大部分原梯度在被放大后刚好处于 FP16 的表示上限附近——理论上偏向给梯度最大的放大空间。每个训练步骤反向传播时检查这一步的梯度中是否出现了 inf 或 NaN。如果出现了说明缩放太激进了把系数减半实际是乘上某个衰减系数当前的优化器更新直接跳过这一轮白跑。如果连续很多步都没有出现 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)内部做的事比想象中多遍历optimizer的参数检查梯度中有没有 inf/NaN。如果发现溢出直接跳过optimizer.step()不更新任何参数这一点非常重要否则这次迭代的权重被污染。如果没有溢出把所有梯度除以 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(), lr1e-3)这样 optimizer 操作的还是 FP16 的参数更新量先变成 FP16这违背了混合精度的初衷。真正可靠的做法是optimizer torch.optim.Adam(model.parameters(), lr1e-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 在内部做临时转换梯度在反向传播后是 FP16GradScaler 负责梯度缩放和保护。这里真正的权衡点是模型保持在 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, biasFalse).half().cuda() optimizer torch.optim.SGD(model.parameters(), lr0.1) # 制造一个极小梯度场景: 输入极小, 目标也是极小 x torch.tensor([[1e-4, 1e-4]], devicecuda).half() y torch.tensor([[1e-5]], devicecuda).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, biasFalse).half().cuda() model2.load_state_dict(model.state_dict()) optimizer2 torch.optim.SGD(model2.parameters(), lr0.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 组合过程中记录每一步的梯度统计。这套诊断流程比盲目调整学习率和清洗数据有效得多。梯度缩放不是银弹它是数值保护里的一块拼图但把这块拼图补上之后你会发现很大一部分难以解释的训练崩溃突然就都有了答案。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门