在训练类似Transformer的深度模型时,我们常常会遇到一个棘手的问题:训练初期一切正常,但经过若干个epoch后,Loss突然变为NaN,导致整个训练过程崩溃。尤其是在ImageNet这样的数据集上训练时,经过几个epoch(30-40个),Loss就变成了NaN”。

一、 问题根源:混合精度的“双刃剑”特性

混合精度训练(Mixed Precision Training)通过在计算中同时使用16位(FP16/BF16)和32位(FP32)浮点数,旨在加速训练并减少显存占用。然而,FP16的动态范围远小于FP32(FP16的有效范围约为 6e-5 到 6e4),这使其在处理Transformer这类模型时尤为脆弱。

Transformer模型的核心——自注意力机制和层归一化(LayerNorm)——会产生数值范围极广的中间结果(激活值和梯度)。在FP16的有限表示能力下,这些值极易发生两种灾难性错误:

  1. **上溢出 **(Overflow):当某个值的绝对值超过FP16能表示的最大值(约65504)时,它会被表示为正无穷(inf)或负无穷(-inf)。
  2. **下溢出 **(Underflow):当某个值的绝对值小于FP16能表示的最小正数(约 6e-5)时,它会被直接舍入为0。

一旦计算图中出现 inf 或 0(在错误的位置),后续的运算(如softmax、除法、对数等)就会产生 NaN(Not a Number),并迅速污染整个网络的梯度,最终导致Loss变为NaN。

二、 核心机制:损失缩放(Loss Scaling)

为了解决下溢出问题,现代深度学习框架(如PyTorch)引入了自动混合精度(AMP)和梯度缩放器(GradScaler)。

其工作原理如下:

  1. 前向传播:在 autocast 上下文中,大部分计算以FP16进行。
  2. 损失缩放:在反向传播前,将Loss乘以一个很大的缩放因子(scale factor,通常初始值为 2^24)。
  3. 反向传播:由于Loss被放大,反向传播计算出的梯度也会相应放大,从而避免了因数值过小而下溢为0。
  4. 梯度反缩放与更新:在优化器更新参数前,GradScaler 会先将放大的梯度反缩放(unscale)回原始大小,然后用FP32的主权重(master weights)进行更新,以保证数值稳定性。

三、 为何Loss仍会变为NaN?

尽管AMP机制设计精巧,但在实践中仍可能失败,主要原因如下:

  1. 缩放因子失效:如果模型产生的梯度本身就包含 inf(上溢出),那么无论缩放因子多大,inf * scale 仍然是 inf。GradScaler 在反缩放时检测到 inf,会跳过本次参数更新,并自动减小缩放因子。但如果上溢出持续发生,缩放因子会不断减小直至失效,最终导致梯度中出现 NaN。
  2. 不恰当的学习率:过高的学习率会直接导致权重更新步长过大,产生数值不稳定的激活值和梯度,极易触发上溢出。这与您观察到的“学习率也要适合”完全吻合。混合精度训练通常需要比FP32训练更低的学习率。
  3. 模型架构的内在不稳定性:某些自定义的激活函数、归一化层或注意力实现可能存在数值不稳定的隐患,在FP32下被掩盖,但在FP16下被放大。

四、 系统性解决方案

  1. 首要检查:学习率。尝试将当前学习率降低一个数量级(例如,从 1e-4 降到 1e-5),这是最简单有效的第一步。
  2. 正确使用AMP。确保您的训练循环严格按照PyTorch AMP的最佳实践编写:
    from torch.cuda.amp import autocast, GradScaler
    
    scaler = GradScaler()
    for data, target in dataloader:
        optimizer.zero_grad()
        with autocast(): # 开启自动混合精度上下文
            output = model(data)
            loss = loss_fn(output, target)
        # 缩放loss并反向传播
        scaler.scale(loss).backward()
        # 用缩放后的梯度更新优化器
        scaler.step(optimizer)
        # 更新缩放因子
        scaler.update()
    
  3. 监控缩放因子。在训练日志中打印 scaler.get_scale() 的值。如果它在持续、快速地下降,说明模型存在严重的梯度上溢出问题,需要从学习率或模型架构入手解决。
  4. 考虑使用BF16。如果您的硬件支持(如Ampere架构及以后的NVIDIA GPU),可以尝试使用 bfloat16 (BF16)。BF16的指数位与FP32相同,因此其动态范围与FP32几乎一致,能从根本上避免上/下溢出问题,虽然精度略低于FP16,但对于大多数模型来说足够了。
  5. 终极方案:回退到FP32。如果以上方法均无效,或者您的任务对数值稳定性要求极高,那么暂时放弃混合精度,使用纯FP32训练是保证训练成功的可靠选择,正如您所验证的那样。

总结

我是放弃混合精度训练,直接用FP32才解决的。不过显存一下子上去了,没办法。

更多推荐