预训练中 Loss Spike 可能是什么原因?
简化版
Loss Spike 要先区分单批数据异常、持续数值发散和日志/归约错误。常见根因包括坏数据或分布突变、学习率过大、梯度爆炸、FP16 overflow、优化器状态损坏、并行通信/分片错误以及恢复状态不完整。处理时冻结现场,对齐最早异常 step,检查样本 ID、分域 loss、梯度/激活范数和非有限值;确认根因后从健康 checkpoint 回滚,而不是在坏状态上硬降学习率继续。
详细版
一次尖峰后立即恢复,可能是超长/乱码 batch、数据源切换或记录噪声;连续上升通常更像优化发散;部分 rank 先异常则需查数据分片、GPU 或 collective。建立分层遥测:全局及每 rank loss、有效 Token、学习率、grad norm、AMP scale、NaN/Inf、路由负载、数据源和样本哈希。
排查顺序是先验证指标计算,再定位第一异常 rank 与样本,随后检查数值和系统状态。用故障 batch 在单机小规模复现;若移除样本仍异常,回滚前一个 checkpoint 对照。修复可包括隔离坏数据、延长 warmup/降低 LR、梯度裁剪、改用 BF16/FP32 关键算子、恢复优化器与 RNG、修复通信,并加入 canary 回归。
Spike告警 -> 指标真假 -> 数据还是数值/系统 -> 最早异常rank/step
-> 可复现最小样例 -> 回滚健康点 -> 修复后小步验证
完整版教学
一、先定义什么才是 Spike
训练 loss 本来就有随机波动,尤其 mixture 切换和变长 batch 下。应基于滑动中位数和稳健离差设动态阈值,而不是看到单步变高就中断。还要分别观察 token-weighted loss 与 batch 平均,padding 或空样本会扭曲后者。
例如过去 1000 step 中位数 2.0,正常波动约 0.05,某步跳到 4.8 且梯度范数同时放大十倍,是真异常;若只从 2.0 到 2.15,恰好该批是高难数学,随后恢复,可能只是数据组成变化。告警要携带上下文而非只发一个数字。
二、三类表现对应不同方向
单步尖峰后恢复常由困难/损坏 batch、数据读取或日志计算引起;尖峰后缓慢恢复可能是一次过大更新扰动参数;持续上升并出现 NaN 更像学习率或数值发散。周期性尖峰可能与数据 shard、checkpoint、评测切换或 scheduler 周期有关。
| 现象 | 优先怀疑 | 关键证据 |
|---|---|---|
| 单步高、下一步正常 | 坏 batch/日志 | 样本与有效 Token |
| 所有 rank 同时持续升 | LR/优化器 | grad norm、update norm |
| 单 rank 先异常 | 数据/GPU/分片 | per-rank 遥测 |
| checkpoint 恢复后出现 | 状态缺失 | LR、Adam、RNG、游标 |
| 固定周期出现 | 数据桶/任务 | source、shard、调度日志 |
表型只是排查优先级,不是最终诊断。
三、数据问题怎样制造尖峰
乱码、超长重复、错误编码、空文档、loss mask 错位和异常高权重样本都能抬高 loss。数据 mixture 突然从自然语言切到代码,平均 loss 也会变化;若只看全局聚合,会把分布变化误判成优化不稳。
每个 batch 记录数据源、稳定样本 ID、有效 Token、长度和质量分摘要。触发异常时把 ID 写入隔离队列,用同一模型 checkpoint 重放;若每次在同一 batch 复现,优先修数据或 collator。注意日志不得直接泄露敏感全文,可保存哈希与受控引用。
排障心法:先找“第一个不正常的量”,不要从最终 NaN 倒推;后续数千个错误往往只是最早一次坏更新的结果。
四、学习率和优化器如何发散
学习率过大、warmup 太短或 Batch 改变后未调 LR,会让更新量相对参数过大。Adam 的二阶矩若损坏或恢复时重置,归一化失效,第一步可能突跳。监控参数范数与更新范数比值比只看 grad norm 更直接。
update_ratio = ||Δθ|| / (||θ|| + ε)
若平时约 1e-4,异常 step 到 1e-1,说明更新尺度跃升千倍。梯度裁剪可限制事故幅度,但若根因是错误 loss mask 或状态损坏,裁剪只会掩盖问题,不能替代修复。
五、混合精度有哪些特有故障
FP16 指数范围较窄,激活或梯度可能 overflow;动态 loss scaling 通过缩放避免下溢,但频繁回退说明数值危险。BF16 指数范围接近 FP32,通常更稳,但尾数精度较低,softmax、归一化、loss 聚合等关键算子仍常用 FP32。
要记录各层激活/梯度最大值和首个 NaN/Inf 位置,而不是只在最终 loss 检查。FlashAttention、自定义 kernel 与 fused optimizer 也可能有特定 shape 的数值 bug;切换到参考实现做 A/B,有助于区分模型本身和算子问题。
六、分布式系统也会产生“假数学问题”
某 rank 读到错误分片、all-reduce 缩放不一致、ZeRO 参数错位或通信静默损坏,会让聚合梯度异常。应比较每 rank loss、grad norm 和 checksum,查找最先偏离者。只记录 rank 0 会丢失关键证据。
恢复后 global batch 或梯度累积改变,也会改变有效学习率。TP/PP 拓扑变化时分片映射错误可能仍能运行,却输出错误数值。定期做固定 batch 的多 rank 一致性测试,并为 checkpoint 分片加 checksum。
七、正确的恢复与修复流程
发现持续发散时先停止写新的“latest”,保存诊断快照但不要把它标为可恢复版本。选尖峰前至少一个健康 checkpoint,验证模型、优化器、scheduler 与数据游标。修复后用故障附近数据跑 50~200 个观察 step,比较 loss、梯度和吞吐,再恢复全规模。
如果坏 batch 可安全排除,记录规则和影响范围;若根因不明,仅回滚重跑可能再次发生。事故复盘要写时间线、根因、检测缺口和永久措施,例如 per-rank 监控、样本 quarantine、更新比告警或数值单测。
八、常见误区与追问
- 误区:看到 Spike 就降低学习率继续。 坏优化器状态或 NaN 权重不会因小 LR 自动恢复,应先判断状态是否健康。
- 误区:全局 loss 足以排查。 聚合会隐藏单 rank、单数据源和长度桶异常。
- 误区:梯度裁剪能解决所有不稳定。 它限制幅度但可能掩盖数据、状态或算子根因。
- 追问:如何区分坏数据与模型发散? 从健康 checkpoint 重放同一 batch,并用正常 batch 做交叉对照。
- 追问:为什么 BF16 通常比 FP16 稳? BF16 指数位更多、动态范围更大,较少 overflow。
- 追问:Spike 后回滚多远? 至少回到所有关键指标健康且 checkpoint 校验通过的版本,并留安全余量。
九、加强记忆
Loss Spike 排查可记成“验、分、找、复、回”:先验证指标真假,按数据、优化、数值和系统分层,找到第一异常 step/rank/样本,用固定 batch 复现,最后从健康点回滚并小步验证。关键不是把曲线临时压下去,而是保住现场、证明根因和建立不再复发的监控与隔离机制。