LLM算子NaN双重降档机制:从检测到恢复的工程实践

发布时间:2026/10/4 22:49:01
LLM算子NaN双重降档机制:从检测到恢复的工程实践 处理过LLM训练和在线推理性价比的人几乎都跟 NaNNot a Number打过照面。模型跑得好好的某一步 loss 突然变成 nan后面所有 checkpoint 都像“废纸”在线服务里更尴尬一个算子输出踩到 NaN生成出来的 token 直接变成乱码不说KV Cache 里还会残留污染后面的请求跟着遭殃。我这两年一直在做算子层的稳定性保障团队里陆续经历了“模型训练报nan”的上千次案例沉淀下来一套通用方案——从 NaN 到 Safe 的双重降档检测机制。简单说就是给算子的执行状态定义成几档正常档、观察档、安全档、兜底档用一套检测层实时盯着算子输出、梯度和 loss一旦发现异常征兆就自动把算子降到更稳的实现上跑等状态恢复之后再慢慢升回来。这篇文章我会把机制的设计逻辑、检测层的工程实现、Safe 算子库的改造细节以及一次完整的故障注入验证过程都拆开讲正被训练不收敛或者推理偶发崩坏折磨的同学应该能直接用上。1. LLM算子里的NaN源头、危害与那句“模型训练报nan”1.1 一条NaN会怎样毁掉整条计算链我先说个最反常识的点NaN 不是“崩掉”的是“传染”的。GPU 上任何一次浮点运算只要某条计算路径出现 NaN它参与的所有后续运算——不管是最大、最小、求和还是乘法——结果都会变成 NaN。在 LLM 这种层数深、算子多、计算图很大的场景里意味着你往往在 loss 里看到 NaN 的时候它已经从源头跑了十几个算子、污染了十几层激活、或许还写进了 KV Cache。重启没用因为下一轮同样的激活路径还会在同一个位置炸出来直接加载上一个好 checkpoint 也没用因为如果根因数值还在训练恢复后大概率还会复现。这正是“模型训练报nan”最折磨人的地方你根本不知道它是哪一步炸的也不知道要回滚到哪里才算安全。那为什么说“传染”很要命举个具体场景attention 计算里 QK^T 以后会过一个 Softmax如果 QK 矩阵里混入一个 Inf正无穷Softmax 分母上的累加也会是 Inf比例变成 0/0 之后整行概率分布的梯度在反向传播里就是 NaN。这时候哪怕其他 99.99% 的激活都正常这个 NaN 也会沿着梯度链往回传把前面的 weight 梯度一并污染。等 optimizer step 做完模型权重可能就永久“漂移”了。1.2 三类最容易爆雷的算子根据我自己的统计LLM 训练和推理中 NaN 的高发区非常集中主要就这几类Softmax / Attention输入是长尾 logit 分布时exp(大数)直接打到 Inf。fp16 下最大值只有 65504logit 只要大于 11exp 就爆了。别笑长上下文里 QK^T 的方差会随着长度增加单词条 logit 数值很容易冲到几十甚至上百。LayerNorm / RMSNorm方差项在精度损失以后可能变成 0 或者接近 0这时候“除以标准差”就成了除以一个几乎为 0 的极小值结果直接炸成 NaN反向传播里分母还会再平方一次更危险。MatMul大批量矩阵乘法的中间累加结果在低精度累加器上很容易溢出尤其是 bf16/fp16 下跑大 K 维的 GEMM几个大数累加之后突破 65504 是常事。另外还有一类容易被忽略的GELU 等激活函数对输入做近似数学运算时某些特殊输入比如极端负值下 sqrt 内部的近似参数偏了也会偶发 NaN/Inf。这类问题最难定位因为它的爆点是随输入分布的不是固定的某一层。1.3 fp16/bf16精度下问题被放大的具体原因其实很多人会疑惑为什么换成 fp32 就没 NaN 了这不是算法变了而是数值表示范围变了。fp16 的指数位只有 5 位最大值 65504bf16 的指数位有 8 位范围和 fp32 一样大但它的小数位只有 7 位相对精度只有约 2^-8。两个低精度格式在 LLM 里的分工是fp16 常用于激活和权重传输范围有限但精度相对好一点bf16 常用于梯度更新范围大不怕溢出但精度差。混合精度训练里模型参数用 fp32 保存但 forward 和 backward 都在低精度下跑于是计算路径中任何一个中间结果越界就立刻产生 NaN/Inf。所以“精度升档到 fp32”是最直接有效的兜底手段这也是我们后面“第二重降档”里的核心动作之一。我记得有一次线上推理服务反馈长上下文场景下连续几条生成结果全部乱码。我们把一个爆了的 token 对应的输入日志捞出来发现第一个出现 NaN 的算子根本不是 attention而是前几层的 RMSNorm那条样本的 embedding 经过 MLP 激活后某个 channel 的方差在 bf16 下被抹成 0除以一个极小值后直接拉爆。这种问题只看 loss 是看不出来的因为 loss 可能过了好几步才变而且当 batch 里其他样本正常时单个样本的 NaN 会先污染该样本的 attention row再通过反向传播污染对应梯度。2. 双重降档机制的设计不是一句“检测到NaN就恢复”那么简单2.1 从状态机视角重新理解“降档”一提到容错很多人第一反应是“检测到 NaN 就回滚上一个 checkpoint”。这种做法的毛病很明显对于长上下文在线推理回滚的代价可能是一个完整请求重算对于训练回滚意味着丢失最近几万步的有效计算。我们需要的不是在灾难发生后修复而是在灾难发生过程中让算子自动选择更安全的执行路径。我把这个过程建模成一个状态机。每个算子或者一组相关算子都有自己的执行档位状态NORMAL默认状态。融合算子、目标精度fp16/bf16、性能调优全部打开。OBSERVE检测层捕获到疑似 NaN 信号但还没确认。此时算子不切换实现只是把输出扫描频率从采样模式提升到全量模式。SAFE_STAGE_1第一重降档。算子切换到 Safe 实现但精度档位仍保持目标精度。说白了换算法不换精度。SAFE_STAGE_2第二重降档。算子进一步反融合中间计算提升到 fp32输入做保护性裁剪KV Cache 做净化。VERIFY算子已经连续 N 步无 NaN且激活和梯度的分布与正常基线对齐开始验证是否可以升档。RECOVER逐步升档先试回 SAFE_STAGE_1确认稳定后再试回 NORMAL。这里的关键是“降档动作不是一次性到底”。我刻意把它拆成两级因为每一级都有完全不同的成本收益曲线后面会细说。2.2 第一重降档从性能模式切换到安全算子第一重降档解决的是“算法本身在边界条件下不稳定”的问题。典型做法是把算子替换成数值上更稳的等价实现。举个例子标准 Softmax 是先对 logits 求 exp再除以总和Safe Softmax 则是先减去该行的最大值再 exp分母用减完 max 之后的和。这两个在数学上完全等价但数值稳定性天差地别。第二类替换比如 LayerNorm 的方差计算标准实现是直接算 E(x^2) - E(x)^2在低精度下这个“差”会严重丢失有效位改用 Welford 公式逐样本更新或者至少把均值/方差的计算放在 fp32就能显著降低“除以 0”和微小方差被抹平的风险。第一重降档的最大优点是代价低。我们实测下来Safe 算子相比融合算子通常只慢 5%~15%在端到端时间里占比很小。因此这一重降档的触发条件可以设得宽松一点宁可偶尔误触发也不要漏掉真实故障。2.3 第二重降档精度与算法路径的彻底兜底如果第一重降档做了算子还在报 NaN那说明问题不只是“算法边界不稳定”更大概率是低精度的表示范围根本承载不了当前输入分布。此时第二重降档会执行三件事中间计算全部提升到 fp32。Softmax 的 exp、LayerNorm 的方差、MatMul 的累加器全部在 fp32 下完成只在算子边界转回低精度。算子反融合。把一个大的融合 kernel 拆成若干小的子步骤步骤之间插入数值检查点和裁剪点。比如把 Attention 拆成 QK^T、缩放、Softmax、与 V 相乘几个子步骤每个步骤结束后先看输出是否越界再决定下一步。输入保护。对极端值做饱和裁剪clip 到某个经验上限尤其是 attention logits、embedding 输入、以及位置编码相加之后的输入。第二重降档的代价比第一重大不少算子耗时可能增加 20%~30%但因为它是“只要不炸就行”的兜底档触发条件要严格得多一般是第一重降档之后连续超过 2 步还有 NaN才进入。2.4 为什么必须“双重”而不是一步到位我见过一些团队直接用“高精度兜底”作为唯一降档策略检测到 NaN 就把所有算子全切成 fp32。这是最粗暴的做法它忽视了第一重降档的价值大部分真实 NaN 场景问题出在“算法实现不稳”而不是“精度不够”。如果所有算子都直接切 fp32性能损耗从一开始就是满的。再好的量化方案、融合 kernel 设计在这一刻全部形同虚设。双重的意义在于分层逼近先付出 5% 的代价解决 80% 的问题剩下 20% 再付出多一点代价去兜底。另外还有一个实际原因降档容易升档难。一级一级地降意味着恢复时也可以一级一级地升每一步都验证过以后才继续。如果你一步到位降到底恢复的时候也一步切回大概率会出现“切换瞬间又触发 NaN 抖动”的尴尬循环。3. 检测层的工程实现在不伤性能的前提下抓到那只NaN3.1 算子输出张量扫描isnan/isinf计数与分级阈值检测是整个机制的“眼睛”。我最开始做这层的时候第一直觉是每个 step 把每个算子的输出都过一遍 torch.isnan结果性能立刻掉了将近 20%训练从 3 天变成 4 天根本没法用。后来我把扫描策略改成分级NORMAL 状态每 K 步抽一次样只扫描最容易爆的算子attention、layer_norm、matmul其他算子靠统计采样。OBSERVE 状态每步全量扫描所有关注算子的输出但只统计 NaN/Inf 数量不打 log。SAFE_STAGE_1/2 状态全量扫描 记录算子输入端和输出端的 min/max方便之后定位根因。阈值这一块要说清楚不要一看到 NaN 就立刻降档。因为小概率的异步 kernel 竞态或者 L2 cache 抖动可能产生一个孤立 NaN但下一步就消失了。我们的实践是同一算子在连续 2~3 步内反复检测到 NaN或者单次检测中 NaN 占比超过该张量元素数的十万分之一才触发降档。按张量 shape 归一化很重要否则 batch size 或 seq len 放大后同样的噪声强度会被误判成故障。3.2 训练流程级的早期信号梯度与loss的联动监控算子张量扫描能发现“已经发生的 NaN”但很多时候我们想在“即将发生之前”预警。训练流程里有三个信号是完全免费的grad_normoptimizer step 之前对全部梯度求范数。如果 grad_norm 在正常的下降趋势中突然跳变到 Inf/NaN基本可以断定某个算子的反向输出炸了。spike 检测对 loss 做指数滑动平均如果当前 loss 与 EMA 的偏差超过历史标准差的 8~10 倍哪怕还没到 NaN也值得进入 OBSERVE。梯度的 NaN 占比逐个参数张量统计 NaN 比例如果某个特定层连续多次出现可以精确定位到层。这三个信号结合算子张量扫描能覆盖“将炸未炸”的时间窗。我们的检测层会把它们汇总成一个轻量级的故障信号位图每个 bit 代表一个信号源进入 OBSERVE 只需要 2 个信号源同时报警。3.3 在线检测的隐形开销采样频率与被检查算子的选择我不建议对全部算子一律全量检测因为很多算子比如 reshape、transpose本身不产生新数值检查它们完全没有意义。真正需要定期检查的是那些可能产生新数值的算子所有归一化算子LayerNorm/RMSNorm所有点积类算子Attention、MatMul所有概率类算子Softmax、log_softmax、sigmoid所有激活函数GELU、SiLU 等optimizer step 前的梯度把检查列表从“全部算子”收敛到“数值产生点”之后检测层的性能开销能压到 5% 以内。另外一个经验是扫描动作放在 host 侧异步进行不阻塞主计算流。正确做法是把 isnan 检查 kernel 插入在独立的 CUDA stream 里用 event 和主 stream 做松耦合同步让扫描本身可以和下一个计算 kernel 并行。3.4 采样检测与全量检测的配合具体采样策略NORMAL 状态在训练中每 50 步扫一次推理服务中每处理 50 个请求随机抽 1 个请求扫一次OBSERVE 状态每步/每请求全量扫。为什么不是每步都全量扫因为 scan kernel 的开销在任何加速卡上都是真金白银50 步扫一次已经足够覆盖 95% 的持续性故障。而真正的持续性故障一旦上报就会进入 OBSERVE此时全量扫反而是必须的因为我们要确认降档是否真的生效。此外还有误报与漏报的取舍。误报会导致系统频繁降档训练变慢漏报更致命因为 NaN 最后会在 loss 里出现然后污染权重。我们调参时宁可让误报多一点也不要漏报。上面说连续 2~3 步触发其实已经算比较保守的如果系统整体对性能波动不敏感甚至可以把阈值降到“单步出现 1 个 NaN 就进入 OBSERVE连续 2 步就降档”。4. Safe算子库的落地从Softmax、LayerNorm到Attention的改造实录4.1 别在低精度下重新实现“数学等价”这里有个大坑很多人听说 Safe Softmax就拿着公式一改完工。他们做完了才发现根本没有效果因为在 autocast 模式下你写的中间操作会被自动转回 fp16/bf16等于白改。Safe 算子的前提是“中间计算过程的精度必须钉死”。最直接的方式是把关键数值路径包在明确的 fp32 cast 中或者干脆用 fp32 副本去算最后再转回目标精度。比如def safe_softmax(logits, dim-1): logits_f32 logits.float() # 强制提升到fp32 m logits_f32.max(dimdim, keepdimTrue).values e (logits_f32 - m).exp() s e.sum(dimdim, keepdimTrue) return (e / s).to(logits.dtype) # 输出转回目标精度就这么几行能解决 attention 中 95% 由 Softmax 导致的 NaN。因为先减掉 maxexp 的输入上限接近 0不会再出现 Inffp32 的指数范围又足够大累加和不会溢出。4.2 Safe LayerNorm与RMSNorm把方差计算留在fp32LayerNorm 的经典实现是 y (x - mean) / sqrt(var eps)。问题出在 var E(x^2) - E(x)^2 这一步bf16 下 x^2 本身就有比较大的相对误差再做一个减法有效位几乎丢光算出来的 var 可能变成 0 甚至负数。负数再被 sqrt直接 NaN。我们的 safe 实现有三层保险均值、方差、std 全部在 fp32 下计算方差计算用两遍扫描先算均值再算差的平方平均不要用 E(x^2) - E(x)^2eps 从一个固定值升级为一套动态下限把当前张量能表示的最小 normal 值做一个 scale保证分母永远大于某个正阈值。很多框架其实已经内置了 T5 风格的 RMSNorm 实现但要注意它默认的 dtype。RMSNorm 没有 mean只有 rms sqrt(E(x^2) eps)如果 E(x^2) 在低精度下算成 0一样会爆。4.3 Attention模块的净化和检测点Attention 是 LLM 里最容易出现 NaN 的地方也是 KV Cache 污染的重灾区。针对它我们的 Safe 实现做了三件事QK^T 之后的缩放因子固定用 fp32 计算不参与低精度自动缩放Softmax 走上面说的 safe 路径对 KV Cache 的写入做净化写入前检测新算出的 K/V 张量中 NaN/Inf 的位置把异常位置替换成该张量的均值或者前一步的缓存值。我为什么要强调最后一步因为在线推理时一旦某个 token 的 K/V 是 NaN它对后面所有 token 的 attention 都产生影响。KV Cache 里的 NaN 不会因为后续几个 token 正常就消失它是持久性的污染源。净化 KV Cache 之后可以保证即使当前这个 token 的激活有问题也不会影响下一个 token 的计算。4.4 降档切换的原子性与状态隔离这里说一个翻车点如果你在一个融合 kernel 执行的过程中动态切换算子实现可能会遇到“图中某些子步骤已经跑了新版另一部分还在跑旧版”的中间态。我们的解法是在计算图的调度层面设置一个“降档屏障”。检查到需要降档时先把当前异步流上的所有算子排空也就是 kernel 执行到同步点再以新的算子配置重新生成后续计算图。这会有一次几十到几百毫秒的卡顿但无论如何好过让一批不一致的算子在低精度和 fp32 之间反复横跳。同时为了避免一个算子降档导致整个模型状态被扰乱降档上下文是按张量维度隔离的。同一个 layer 里前面几个 token 可能还在正常档后面新的 token 已经进入 Safe 档靠的是不同请求/token 的上下文区分。这在高并发在线推理场景下尤其重要降档不能是全局性的否则一个坏请求会把所有正常请求都拖慢。5. 一次完整的故障注入与恢复验证从NaN到Safe走过的路5.1 实验设计人为注入极大logit为了验证这套机制不是纸上谈兵我们在测试环境里模拟过很多次故障注入。最经典的是在 attention 的 logits 上加一个极值把某个 token 位置对应的 logits 整体加上 10.0然后用 fp16 精度跑 Softmax。fp16 下 exp(11) 已经是几万直接把两个这样的 logit 加在一起Inf 就出现了。这种注入方式能稳定复现 NaN而且不会影响其他算子非常适合做单变量实验。5.2 第一重降档的响应曲线与loss恢复故障注入后的前两步检测层会捕获到逃逸的 NaN梯度里出现loss 还没怎么变。到第 3 步grad_norm 和 loss 同时报警触发第一重降档。注意此时我们并没有把全部算子切成 fp32只切了 Softmax 和 LayerNorm 两个算子到 safe 路径。响应时间大约几十毫秒。第 4 步开始注入的 NaN 源头被 safe 路径把 exp 限制住了loss 马上恢复下降趋势。整体表现像是一次轻微抖动而不是一次“训练失败事故”。我截过几条曲线第一条是 loss第二条是 grad_norm第三条是检测层的 NaN 计数。NaN 计数从第 1 步到第 2 步是 1、3第 3 步警报到降档第 4 步清零。曲线非常干净。5.3 第二重降档的触发边界也有一个用例需要第二重降档把注入的 logits 极值加大到 1e4 级别。此时 Safe Softmax 虽然因为减 max 避免了 exp 爆发但后续的 e/s 比例计算在 fp32 下虽然不炸在输出转回 fp16 时却会直接饱和成 Inf。这个场景单靠第一重降档救不回来。检测层会在第一重降档生效后继续看到 NaN随即触发第二重降档attention 内部的计算全部留在 fp32、KV Cache 净化打开、logits 入口加饱和裁剪。经过这轮操作即使是 1e4 的极端 logit 也能稳定输出到 fp32 边界再转回 fp16 时只会得到 Inf 而不是 NaN——二者在后续处理策略里完全不同Inf 可以被 clipNaN 只会继续传染。这里值得单独强调一下 Inf 和 NaN 的区别。Inf 是“太大”NaN 是“无意义”。遇到 Inf 可以用饱和裁剪解决遇到 NaN 没有“往下救赎”的空间。Safe 实现的目标之一就是尽量避免产生 NaN哪怕产生 Inf后面也可以裁剪回来。这个小细节在很多工程实现里反而是决定成败的关键。5.4 恢复后的状态复核与性能损耗测试降档不是终点关键是恢复。VERIFY 阶段我们会检查三件事连续 N 步通常 N20无任何 NaN/Inf 计数激活分布的关键统计量均值、方差、最大绝对值与未注入故障时对齐梯度范数的趋势已经回到历史区间。都满足后算子先回到 SAFE_STAGE_1再跑若干步确认没问题后再回到 NORMAL。整个恢复过程会在几分钟内逐步完成。性能方面第一重降档造成的端到端耗时增加约 5%第二重降档约 20%。相比“发现 NaN 后整机重启 加载 checkpoint”动辄几分钟到几十分钟的代价这套机制节省的算力是实打实的。6. 我在实战中踩过的坑和一些留给你的建议6.1 阈值按模型规模动态缩放别用“零容忍”一刀切第一个坑刚开始我们把阈值设成“检测到 1 个 NaN 就进 OBSERVE”在 7B 模型上稳如老狗一上 70B 模型就疯狂降档。原因是超大规模模型张量本身就大激活里偶尔会有那么一两个孤立的数值毛刺比如分布式通信时某些 pipeline stage 的边界值。这种毛刺不是持续性故障你降档、恢复然后再降档、再恢复纯粹在消耗算力。后来我们把阈值改成“按元素数量动态缩放”连续 2 步以上检测到、或者单次占比超过 1e-6才触发降档。这个经验在长上下文场景也同样适用context 长度翻倍时单步张量元素数量会翻倍固定阈值就会失去意义。6.2 恢复后别急着一刀切回高性能模式回环抖动问题第二个坑也是我最想提醒的一句恢复比降档更容易翻车。理论上连续 20 步无 NaN 就可以切回 NORMAL但很多情况下故障源是一个持续存在的异常数据分布你只是暂时压制了它的“显性症状”。直接切回高性能模式等于是把同一个高压输入再跑一遍大概率又踩中之前的雷马上又降档。这种“降档-恢复-再降档”的回环抖动会让系统整体一直处于半降低状态性能损耗反而比一直降档更大。我的建议是恢复过程也要分级先从 SAFE_STAGE_2 升回 SAFE_STAGE_1跑一轮更长的验证比如 50 步再看激活分布、梯度分布是否和基线一致最后才升回 NORMAL。同时记录故障点的指纹比如“input logits 超过阈值的具体是哪个 channel”下次遇到同样的输入形状时可以提前预警。6.3 降档期间的日志策略静默计数优先别一上来就拉堆栈第三个坑早期版本一检测到 NaN 就打全量堆栈把算子的输入输出都 dump 出来。结果故障本身没影响多少日志反倒把 I/O 打爆训练卡死。后来改成“静默计数优先”的策略OBSERVE 阶段只累计计数不输出任何日志真正触发降档后才在故障上下文里记录一次 min/max 和 NaN 位置等到二次复现才 dump 堆栈做离线分析。这样既保住了可观测性又不至于被日志淹没。6.4 这套机制的边界哪些NaN救不了最后说点实在的这套双重降档机制不是万能的。以下三类情况它救不了数据本身的 NaN。比如训练样本读取时反序列化出错送入模型的特征就带着 NaN这时候再安全的算子也拦不住必须在数据 loader 入口就做清洗。硬件故障引起的 NaN。比如显存发生单粒子翻转或者通信链路在传输中出现 bit flip这属于物理层错误单纯切换算子实现毫无意义。检测层如果发现“所有算子都换到哪里都还出现 NaN”应该直接拉起硬件自检。分布式训练中跨卡梯度同步AllReduce阶段出现的 NaN。如果某一卡内部已经出现异常AllReduce 会把 NaN 广播到所有卡。这个需要在通信算子层面加“NaN 屏蔽”把异常值替换成 0 或发送方平均值而不是靠单机内的算子降档解决。当你发现第二重降档都生效了、NaN 还在往外冒别再继续降了。再降下去就是纯 Python 逐算子回退性能没法看不如直接停机、加载最近好的 checkpoint再用降档配置重放最近几步看能不能复现。我个人的体会是容错机制最大的价值不是让系统永远不会坏而是让系统的“坏”变得可预期、可恢复、可诊断。把故障从“训练报 nan 我就完蛋”变成一次可观测的状态跳转整个基础设施的稳定性会上一个台阶。如果你也被这类问题折磨建议先别急着堆监控面板把状态机、降档层级和恢复策略这三件事想清楚大概率能少走不少弯路。

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询