【Transformer 与注意力机制】36|训练稳定性:损失尖峰、混合精度与梯度爆炸

💡 原文中文,约15500字,阅读约需37分钟。
📝

内容提要

大模型训练不稳定是常见问题,表现为loss尖峰、发散或NaN。关键预警信号是梯度范数异常,而非loss本身。常用修复手段各有针对:warmup解决Adam早期方差,Pre-LN保持恒等映射,BF16扩大数值范围,梯度裁剪仅作保险丝。loss spike根因包括数值、优化和数据问题。小模型稳定超参放大后失效,因学习率安全窗口随规模收窄,可用μP等方法解决。

🔎

延伸解读

监控指标比loss更早预警

训练不稳定时,loss本身往往不是最早的信号。OPT-175B和GLM-130B的实践表明,梯度范数、激活范数和动态loss scale的异常通常领先于loss发散几步出现。例如,GLM-130B发现embedding层梯度范数的尖峰可提前预示崩溃。因此,监控面板应持续记录这些伴随指标,而不仅盯着loss曲线,才能争取宝贵的反应时间。

修复手段各有针对性

warmup、Pre-LN、BF16和梯度裁剪并非万能,各自修复特定的假设。warmup解决Adam早期二阶矩方差过大;Pre-LN保持主路径近似恒等映射;BF16扩大数值动态范围;梯度裁剪只是保险丝,防不住根因。理解这些工具的适用边界,才能避免误用,例如BF16不能解决embedding层梯度异常,需配合EGS等定向手段。

小模型超参放大失效的机制

小模型上稳定的超参放大后失效,并非模型变娇气,而是学习率安全窗口随规模(尤其深度)收窄。Wortsman等人的研究表明,不稳定性可在小模型上用高学习率复现,且学习率敏感度随深度增加。μP参数化可稳定最优学习率随宽度的迁移,但无法解决数值溢出,需与qk-layernorm等架构改动配合。

排查loss spike的严谨流程

遇到loss spike,不应直觉归咎于数据。PaLM的对照实验显示,单独重放异常batch并不复现spike,说明是数据与参数状态的组合问题。正确流程包括:对齐时间线、单独重放、检查batch统计特征、排除系统因素,最后才考虑数据处理。这能避免掩盖真正的数值或优化问题。

Q&A

大模型训练不稳定的早期预警信号是什么?

梯度范数异常是比loss更早的预警信号,loss spike往往滞后于梯度范数尖峰几步。此外,loss scale下降和激活范数飙升也是重要信号。

warmup为什么能提高训练稳定性?

warmup通过早期使用较小学习率,压低了Adam优化器在训练初期二阶矩估计方差过大导致的自适应学习率方差,等统计量积累足够后再提高学习率,从而避免早期更新过大。

Pre-LN相比Post-LN为什么更稳定?

Pre-LN将残差主路径保持为近似恒等映射,每个子层只贡献可控增量,避免了梯度反复穿过LayerNorm的Jacobian,从而在深层网络中更容易训练。

BF16相比FP16为什么更稳定?

BF16的指数位有8位,与FP32相同,动态范围远大于FP16,能避免溢出和下溢,因此很多在FP16下需要loss scaling的问题在BF16下自然规避。

梯度裁剪能解决所有训练不稳定问题吗?

不能。梯度裁剪只是保险丝,防止单次异常梯度把参数推得太远,但无法修复梯度异常的根源。例如PaLM在梯度裁剪开启时仍出现约20次loss spike。

loss spike的常见根因有哪些?

loss spike的根因分为数值类(如FP16溢出、attention logits过大)、优化类(如学习率过高、warmup太短)和数据类(如异常batch),但数据类往往不是独立原因,而是特定batch与参数状态的组合。

为什么小模型上稳定的超参放大后会失效?

因为学习率安全窗口随模型规模(尤其是深度)增大而收窄,小模型上稳定的绝对学习率在更大模型上可能落在窗口之外。μP等方法可以解决最优学习率随规模漂移的问题。

如何排查异常batch导致的loss spike?

先对齐时间线,定位触发异常的step;然后单独重放该batch,从更早检查点重新训练,若不复现则说明是数据与参数状态的组合问题;再检查该batch的统计特征;排除系统性因素;最后才考虑数据处理。

🏷️

标签

➡️

继续阅读