导读:本期聚焦于越南程序员创作的《混合精度训练出现NaN怎么办?Loss Scaling原理与溢出检查方法详解》,敬请观看详情。训练到一半突然出现NaN,loss变成inf,这几乎是每个尝试混合精度训练的人都会踩的坑。问题的根源在于FP16的数值范围有限,权重更新时的小梯度乘以学习率后很容易下溢成零,而某些中间计算又会超出上限溢出。本文从FP16的数值表示入手,分析动态范围缩小带来的上溢与下溢两类风险,详细讲解Loss Scaling如何通过放大loss来保住小梯度,并对比静态缩放与动态缩放的实现差异。同时介绍PyTorch的GradScaler与TensorFlow的LossScaleOptimizer的用法,以及如何通过分布式训练中的AllReduce检查、逐层排查等手段定位溢出发生的具体位置,帮助读者建立一套完整的溢出预防和诊断方案。

深度学习模型越做越大,显存压力也越来越明显,混合精度训练因此成了主流选择:用FP16或BF16保存权重和激活值,用FP32维护主权重,可以节省接近一半的显存并加速计算。但混合精度并不是免费午餐,最常见的问题就是训练过程中loss突然变成NaN,或者梯度中出现inf。这背后的罪魁祸首通常是FP16的数值范围问题,而Loss Scaling正是为了解决它而生的。这篇文章会从原理到实践,把溢出的成因、Loss Scaling的机制以及溢出检查的手段讲清楚。

混合精度训练出现NaN怎么办?Loss Scaling原理与溢出检查方法详解

为什么FP16会产生溢出

要理解溢出,得先看FP16的数值表示。FP16用1位符号位、5位指数位、10位尾数位表示一个数,指数位决定了动态范围。FP16能表示的最大正数约为65504,最小的正规数约为6.1e-5,小于这个值的数会进入非规格化区间,再小就直接变成零。

对比FP32,它的指数位有8位,最大数约3.4e38,最小正规数约1.2e-38,动态范围大得多。也就是说,从FP32切换到FP16,动态范围从大约76个数量级压缩到了不到9个数量级。这个差距带来了两类风险。

第一类是上溢。训练中某些中间结果,比如某些分类网络未做数值稳定处理的softmax、loss本身过大,或者梯度的累加,都可能超过65504。一旦超过,FP16直接返回inf,后续计算全部污染成NaN。第二类是下溢,这个问题更隐蔽也更普遍。深度网络的很多梯度其实非常小,尤其在训练后期梯度趋近于零,FP16中1e-8这样的数根本无法表示,会被直接截断为0。梯度变零意味着这部分参数得不到更新,模型精度悄悄下降,甚至loss不降反升。

需要特别指出的是,上溢容易被发现,因为NaN会让训练立刻崩溃;下溢却往往无声无息,训练看起来正常,但收敛质量变差。这也是为什么混合精度训练必须配合合理的溢出检测机制,而不能只靠肉眼观察loss曲线。

Loss Scaling的工作原理

Loss Scaling的思路很直接:既然小梯度会被FP16截断成零,那就在反向传播开始前,把loss放大一个倍数。根据链式法则,所有梯度都会等比例放大,原本处于下溢边界的梯度被抬升到FP16可表示的范围内。等梯度计算完成后,在更新权重前再除以这个倍数,恢复真实梯度数值。

整个过程可以用伪代码描述:

# 前向计算
output = model(data)
loss = criterion(output, target)

# 放大loss,假设scale factor为S
loss = loss * S

# 反向传播,所有梯度被放大S倍
loss.backward()

# 反常缩放,恢复真实梯度
for param in model.parameters():
    param.grad = param.grad / S

# 用FP32主权重执行优化器更新
optimizer.step()

关键点在于梯度除以S的操作是在FP32下进行的。放大是为了让梯度在FP16的反向传播中存活下来,而缩放回去并更新参数时使用FP32,就不会再引入精度问题。这就是混合精度训练的标准套路:前向和反向用FP16,主权重维护和参数更新用FP32。

Scaling因子的选择分两种方案。静态缩放在整个训练过程中使用固定倍数,比如常见的1024或2048,实现简单但不够灵活:倍数太小压不住下溢,太大又会把本不会溢出的梯度推过65504的上限,反而制造上溢。动态缩放则从一个较大的初始值开始(如2的16次方),每次反向传播后检查梯度中是否出现inf或NaN,如果出现就跳过这次更新,并把缩放因子减半;如果连续若干步(比如2000步)都没有溢出,就把缩放因子翻倍,逐步逼近可用范围的上限。动态缩放在几乎不损失训练效果的前提下,把溢出管理完全自动化了。

框架中的实现:GradScaler与LossScaleOptimizer

PyTorch提供了torch.cuda.amp模块,核心是GradScaler类。典型用法如下:

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler(init_scale=65536.0, growth_interval=2000)

for data, target in train_loader:
    optimizer.zero_grad()
    with autocast():  # 前向过程自动选择低精度算子
        output = model(data)
        loss = criterion(output, target)
    scaler.scale(loss).backward()   # 缩放loss并反向传播
    scaler.step(optimizer)          # 检查梯度,正常则更新参数
    scaler.update()                 # 动态调整缩放因子

这里的scaler.step内部做了两件事:先调用unscale_把梯度恢复到真实数值,然后检查是否包含inf或NaN,有问题就跳过该步的optimizer.step。scaler.update则根据检查结果调整缩放因子。整个过程对训练逻辑侵入很小,只需要把loss.backward换成scaler.scale(loss).backward,把optimizer.step换成scaler.step即可。

TensorFlow的对应机制是LossScaleOptimizer,它包裹原始优化器并提供动态损失缩放:

import tensorflow as tf

optimizer = tf.keras.optimizers.SGD(learning_rate=0.01)
optimizer = tf.keras.mixed_precision.LossScaleOptimizer(optimizer)

with tf.GradientTape() as tape:
    output = model(data, training=True)
    loss = loss_fn(target, output)
    scaled_loss = optimizer.get_scaled_loss(loss)

scaled_grads = tape.gradient(scaled_loss, model.trainable_variables)
grads = optimizer.get_unscaled_gradients(scaled_grads)
optimizer.apply_gradients(zip(grads, model.trainable_variables))

两个框架的思路完全一致,只是API形态不同。值得注意的是,如果使用BF16(Brain Float 16),情况会有所不同。BF16用8位指数位,动态范围与FP32相同,因此基本不会发生上溢或下溢,无需Loss Scaling。这也是新版硬件纷纷推荐BF16的原因。判断是否需要缩放的简单原则:FP16必须配合Loss Scaling,BF16可以省略。

溢出检查与定位方法

即便使用了动态缩放,训练仍可能偶发NaN,这时需要一套排查手段确定溢出发生在哪里。最基础的检查是在每个反向传播后扫描梯度:

import torch

def check_gradients(model):
    for name, param in model.named_parameters():
        if param.grad is None:
            continue
        if torch.isinf(param.grad).any() or torch.isnan(param.grad).any():
            print(f"溢出出现在参数: {name}")
            return name
    return None

如果逐个参数都正常,但loss仍然是NaN,说明问题出在前向过程。可以在关键节点上做数值断言,检查每一层输出的最大值和最小值,找到第一个出现inf或NaN的层。常见的嫌疑对象包括:没有做数值稳定的softmax(应使用log-softmax配合NLL损失)、exp运算直接对大数求值、除法中分母可能为零、LSTM中的循环累加导致激活值爆炸等。

在分布式训练中,排查逻辑还要考虑梯度聚合环节。数据并行训练会对各卡的梯度做AllReduce求平均,任何一张卡上的溢出都会扩散到全局。PyTorch的分布式场景下需要配合定制的通信钩子做梯度检查,或者在梯度同步前先做局部过滤。如果某一卡持续溢出,往往是该卡的数据中存在异常样本(比如脏数据、极端离群值),需要对数据管道单独排查。

另一个实用技巧是梯度裁剪的顺序。正确的做法是先unscale再裁剪,也就是在梯度恢复真实数值之后再按阈值裁剪。如果在缩放状态下直接裁剪,阈值本身也要乘以缩放因子,非常容易出错,而且裁剪会破坏溢出检测的正确性。PyTorch中使用scaler.unscale_(optimizer)之后再调用nn.utils.clip_grad_norm_,既保证裁剪阈值有意义,也不会干扰后续的溢出判断。

实践建议与总结

总结一下避免混合精度溢出的完整方案。首选确认硬件和框架支持BF16,能切换就切换,可以从根源上消除动态范围问题。如果必须使用FP16,务必启用动态Loss Scaling,不要自己手写静态缩放。模型层面优先选择数值稳定的算子:用log-softmax代替softmax,用稳定的注意力实现(内部对相似度矩阵减去最大值),loss函数尽量选择框架自带的数值稳定版本。

排查问题时遵循固定流程:先确认溢出发生在前向还是反向,再逐层定位具体位置,最后检查数据管道和分布式聚合环节。开启PyTorch的TORCH_SHOW_CPP_STACKTRACES或设置异常检测上下文,也能帮助快速定位NaN产生的源头。只要把数值范围管理好,混合精度带来的加速和显存收益是完全值得的,NaN问题并不是放弃它的理由。

混合精度训练Loss Scaling梯度溢出修改时间:2026-09-04 20:01:17

免责声明:已尽一切努力确保本网站所含信息的准确性。网站作品多为原创整理与精心创作,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们进行处理Email:chomcom@qq.com。
引用或转载本作品时,请注明当前出处:https://www.ipipp.com/html/20260904/50449.html,基于非商业用途的前提下,欢迎转载或二创本作品。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。