
凌晨两点我盯着一块训练看板发呆global step 4721loss 3.21下一秒 loss 直接变成了 nan。4721 步积累的模型状态全部作废几千张卡从上一个 checkpoint 重新来十几个小时白跑。这不是我第一次在 LLM 训练里撞上 NaN但那次之后我下定决心把“算子容错”这件事从玄学变成工程。折腾下来最有效的一套东西就是我标题里写的这个一个从 NaN 到 Safe 的双重降档检测机制。简单说这个机制解决的是 LLM 训练里最让人头疼的问题——模型训练报 NaN。它有两个核心设计第一重检测是低开销的前兆信号识别在数值彻底崩坏之前提前“看到”风险触发第一次精度降档第二重检测是精确的污染源定位在 NaN 已经出现之后用重放和二分定位把肇事算子揪出来触发第二次降档直接切换到 Safe 模式。这套机制做出来之后我这边训练任务的 NaN 崩溃频率降了大约 90%而且每次出事都能在几分钟内定位到具体算子不再是整张计算图推倒重来。这篇文章会把这套机制讲透NaN 到底是怎么来的、双重降档为什么比传统容错方案好用、两级检测的原理和实现、以及落到 PyTorch / 昇腾 CANN 这类框架里要怎么接。适合正在跟大模型训练数值问题较劲的算法工程师、做训练框架底层开发的朋友以及所有被“lossnan”折磨过的人。我不写空话直接讲实操和踩过的坑。1. NaN 从哪里来LLM 训练里的数值风暴1.1 三种最常见的 NaN 来源我先说结论LLM 训练里的 NaN九成以上来自三个地方——前向溢出、反向梯度爆炸、算子内部数值污染。搞清楚这三个来源你才能理解为什么“降档”是有效的。前向溢出是最直观的一种。FP16 的最大表示值只有 65504超过它就直接变 inf很多算子再往后一路算下去inf 减 inf、0 乘 inf最终变成 NaN。大模型里最容易爆的位置是 attention 的 softmax 之前query 和 key 做点积维度一高logits 的绝对值很容易被推到上万。我见过一个典型案例长上下文场景下某个头部的 query 分量分布非常极端点积结果直接越过 65504softmax 分母当场变 NaN。热词里有个说法我很喜欢——token 的三个点key 是“我是谁”、query 是“我在找什么”、value 是“我能提供什么”这三个表示在注意力机制里做高维内积时数值尺度完全可能失配而这种失配正是 NaN 的温床。反向梯度爆炸是第二大类来源。训练越深、batch 越大梯度范数的波动越剧烈。如果 loss scaling 选得太大梯度中的小数直接被 round 成 0参数纹丝不动表面数值没问题实际等于“隐形 NaN”反过来如果梯度极大值突破上限更新量一步就把权重打到爆。更隐蔽的是 AdamW 这类自适应优化器二阶矩估计在训练后期可能趋近某个极小值更新项变成“极小数除以极小数”数值上完全不稳定。第三类是最气人的——算子内部污染。你从外面看张量值都正常但算子内部在某个中间步骤算坏了。最典型的是 attention 的 padding mask把非法位置填充成-inf之后做 softmax如果这一行全部是 maskexp(-inf)等于 0分母也是 00/0 直接 NaN。还有 LayerNorm 的方差在极端情况下趋近 0除以一个接近 0 的数数值瞬间发散。这类问题在 FP16 和 BF16 下表现还不一样排查起来特别费劲。1.2 为什么大模型时代 NaN 问题被放大大模型时代之前很多人训 CV 模型用 FP32 就够了NaN 没有那么频繁。但 LLM 训练几乎是踩着混合精度走的显存不够、算力不够只能用 FP16 或 BF16。这两个格式的动态范围差别很大——FP16 指数位只有 5 位最大值 65504BF16 指数位跟 FP32 一样有 8 位最大值 3.39e38。所以同一个算子用 FP16 可能溢出换成 BF16 就不会。这也是后面讲第一次降档时精度选择的关键依据。还有一个被很多人忽视的因素是算子数量本身的爆炸。现在的 LLM 前向一次要跑几百万次算子调用大量使用算子对硬件性能的挑战是一回事但每个算子都是潜在的风险暴露点。任何一个算子在某个极端输入下吐出一个 NaN整个 autograd 图后面全是 NaN损失函数直接 ruined。张量并行和流水线并行切分之后数值路径变得更长、更复杂跨卡通信的 allreduce 把各卡局部误差累积起来微小的数值偏差被放大误差像滚雪球一样越来越重。叠加起来大模型训练跑出 NaN 几乎是个概率问题只是时间早晚。1.3 一个真实的血泪场景我印象最深的一次事故是在一个 MoEMixture of Experts架构的模型上。某个 expert router 在 FP16 下把 top-k 的 score 转成概率连续两层 softmax 后出现了上溢出router 输出的 one-hot 分布里出现了多个 NaN。更惨的是这个 NaN 是“间歇性”的——有时候跑 200 步没事有时候跑 20 步就炸。因为它是数据相关的跟特定的 token 分布有关你重放同样的 seed 可能复现换一个数据顺序又消失了。这种间歇性 NaN 最折磨人看起来像硬件错误实际上就是数值路径上缺了一个防御点。从那次之后我意识到靠“出了问题再看”是不够的必须有一双眼睛常驻在计算图里盯着数值健康度。2. 双重降档检测机制两层防线两段降档2.1 传统容错方案的三个痛点在讲我的方案之前先把传统思路盘一盘。为什么大家平时不做算子级容错因为现有做法各有各的坑。第一个痛点是全局 loss 检查太被动。训练循环里写一个if loss.isnan(): load_checkpoint()代码上很简单但问题发生时已经晚了你可能已经跑了 50 步、100 步这些步数里的计算全部浪费。而且它给不了任何定位信息——你只知道 loss 是 NaN 了至于哪个算子、哪个张量搞坏的全靠猜。第二个痛点是逐算子检查开销大。给计算图里的每一个算子输出都加一个torch.isnan(tensor).any()检查听起来很稳妥但每个检查都是一次额外的设备端规约操作。我实测过全算子级别检查在 7B 模型训练上会带来 8%~15% 的吞吐损失本来训练成本就高这种开销根本扛不住。而且训练框架里大量算子输出可能不含標量你不知道该在哪个粒度检查。第三个痛点是静态防御手段覆盖不住动态变化。有人会把 loss scaling 调保守、把梯度裁剪阈值收紧这些属于“固定安全区”。但数值分布是动态的——模型某个阶段的 loss 天然波动大梯度范数可能是另一个阶段的 100 倍。固定参数要么误伤正常训练要么根本兜不住异常。2.2 双重降档机制的核心框架所以我的设计原则是八个字廉价检测昂贵定位。把“发现问题”和“定位问题”拆开分别用不同成本的机制去处理。先定义三档安全状态快速档Fast高性能模式FP16 或 BF16 向量化计算所有算子按照默认最快的实现跑。正常训练绝大多数时间在这个档位。稳健档Stable触发第一次降档后进入。关键算子切换到 BF16/FP32 混合累加优化器侧做 loss scaling 调整和梯度裁剪收紧训练还能继续但计算路径更保守。安全档Safe触发第二次降档后进入。算子切换到带补偿求和、动态范围压缩、FP64 归约的 Safe 模式实现。性能最差但几乎不会产生 NaN。双重降档检测机制的“双重”体现在两条线上。第一重是前兆检测流水线边上挂一个极低开销的旁路监听器实时监控张量统计量的健康度捕捉溢出的“前兆信号”一旦发现风险就触发第一次降档快速档→稳健档。第二重是污染源定位如果第一阶段没拦住或者 NaN 已经出现在 loss 里就启动按需重放二分定位精准找到是哪个算子实例吐出了 NaN触发第二次降档稳健档→安全档。这套机制从 NaN 到 Safe 的路径非常清晰前兆出现 → 第一次降档 → 稳健容错 → 若仍异常 → 重放定位 → 第二次降档 → Safe 模式止损 → 训练恢复。每一步都有明确动作不靠玄学。2.3 这套设计的核心优势对比传统方案这套设计的优势很实际。第一是平时的性能开销几乎可以忽略——旁路监听只读张量统计量不是逐元素检查代价控制在 1% 以内而第二重的重放定位是“按需开启”的平时根本不运行。第二是异常发现早——前兆检测的窗口通常比 loss 变 NaN 早几十步提前降档就能把损失控制在几步之内。第三是可解释、可定位——每次异常都能落到具体算子、具体张量、甚至具体通道上不是笼统地“重跑一遍碰运气”。我举个生活化的类比这就像家里热水器第一道防线是温控器——水温异常升高就自动断电前兆检测触发降档第二道防线是泄压阀——温度已经失控时才启动把压力泄掉避免爆炸重放定位后安全模式止损。平时两道防线都不动真正异常的时候各司其职。3. 核心细节解析两级检测与降档实现3.1 第一重检测前兆信号识别第一重检测的思路是NaN 不是突然出现的数值在崩溃之前通常有“前兆”。以 FP16 溢出为例某个张量的绝对值会先逼近 65504再超过它变成 inf。如果我们能监控这个逼近过程就有机会在真正溢出之前踩一脚刹车。我实现前兆检测用的是一套 EMA 基线 z-score 异常评分的组合。具体做法是这样的对重点算子的输出张量维护一个滑动窗口统计——每 N 步计算一次max_abs、mean、var这三个指标用指数移动平均EMA维护一个“健康基线”。检测时把当前监测值和基线做标准化对比得到 z-score超过阈值就判定为溢出风险信号。这里面有一个我之前踩过坑的细节只看单点的 z-score 是不够的。因为数值分布自然波动某些训练阶段张量极值忽高忽低单点 z-score 很容易误报。所以我额外引入了一个二阶差分判定——用类似拉普拉斯算子Laplacian的思路看统计量的“变化率的变化率”。如果max_abs的斜率本身就在急剧陡峭化说明数值正在加速逼近极限这才是真正的危险信号。实际效果单点 z-score 误报率大约 5%加上二阶差分判定之后降到 0.5% 以内。阈值怎么选我试过的经验值z-score 阈值设 4.0 到 6.0 之间窗口长度 64 到 128 步。太灵敏会频繁降档太迟钝来不及反应。这个需要在你的具体模型上调没有万能值。建议先把日志打出来观察一周看正常情况下的 z-score 分布再定阈值。3.2 第二重检测污染源定位如果第一重没拦住loss 已经变成 NaN 了这时候需要的是“精确打击”而不是从头再来。第二重检测做的就是这件事。核心步骤分三步。第一步是状态回滚冻结训练使用最近一个 checkpoint 的模型状态或者更细粒度的 activation 存储作为重放的起点。第二步是确定性重放在内存 buffer 里以固定随机种子、关闭所有非确定性算子比如 cudnn benchmark的环境下重放从 checkpoint 到 NaN 出现位置之间的计算图。第三步是二分定位先重放前半段子图检查输出张量里有没有 NaN如果没有说明问题在后半段把后半段再二分不断缩小范围直到定位到具体的算子实例。这个二分重放法定位到算子的效率很高——一般几秒到几十秒就能把一张几百个算子的计算图缩小到单个算子。定位之后再对该算子的输入张量做通道级扫描找出是哪个通道、哪个归约维度先出现异常。这一步很关键比如一个 LayerNorm 算子你可以进一步定位是均值计算炸了还是方差计算炸了还是 gamma 缩放那一步炸了这样修复方向就明确了。为什么一定要“确定性重放”因为如果重放环境跟原始运行环境不一致比如用了不同的 kernel、不同的 TensorCore 策略NaN 可能复现不出来整个定位就白做了。我之前就在这上面栽过跟头第一次做重放时没有关闭 cudnn benchmark结果 NaN 时有时无定位了三个小时都没结果。所以重放环境这一块必须较真。3.3 降档怎么执行从快速档到安全档检测是“眼睛”降档是“手”。检测到风险之后动作必须跟得上。这里说的“降档”不是简单地把精度调低而是换一条更安全的数值路径。第一次降档快速档→稳健档核心动作是这几条精度切换FP16 算子换成 BF16。原理很简单——BF16 的指数范围跟 FP32 一样大能容纳更大的绝对值解决“溢出型”NaN。如果你的算子原本就在 BF16 上就改成关键归约用 FP32 累加。loss scaling 降档混合精度训练里把 loss scaling 因子从 2^16 降到 2^8 左右降低梯度下溢风险。我实测过这一步对 LLM 训练特别有效因为 LLM 的梯度本身就存在严重的量级不均匀。梯度裁剪收紧阈值临时缩小一半。同时学习率临时乘 0.1给模型一个“冷静期”让参数不要在这个敏感时期剧烈更新。第二次降档稳健档→安全档动作会更重直接把这个算子替换成 Safe 模式实现。我常用的四个手段动态范围压缩检测到算子输入 max 值过大时先把输入除以一个归一化系数比如 65504算完再乘回来。相当于给数值“提前缩放到安全区”。FP64 归约softmax 的分母、LayerNorm 的均值方差、attention 的注意力加权求和这些归约操作全部改用 FP64 累加器。FP64 的动态范围比 FP32 大得多几乎不会溢出。Kahan 补偿求和把每次加法产生的舍入误差单独存进一个补偿项下一次加法时补偿回去。这能有效防止误差在长序列归约里滚雪球。算子替换如果算子本身写死了低精度比如某些融合 kernel直接用等价的安全实现替换。GPU 上可以用 Triton 写个 SafeKernel昇腾场景下替换成 CANN 的高精度算符实现。进入 Safe 模式之后训练会暂时变慢——这是明知的代价。但损失是可控的而且我们通常只在风险算子局部开 Safe 模式其他算子维持稳健档这样整体性能影响能控制在一个可接受的范围。3.4 阈值与参数速查表我把自己调参过程中总结出来的经验值整理成一张表方便你直接参考。注意这些都是起点值具体模型上还是要自己观察调优。参数默认建议作用调参注意EMA 窗口长度128 步决定健康基线更新的平滑程度窗口越短对近期变化越敏感但误报越多z-score 阈值5.0判定前兆信号的门槛调大减少误报调小增加灵敏度二阶差分阈值3.0捕捉统计量的加速陡峭化取 z-score 的 60% 左右比较稳loss scaling 降档倍数1/256第一次降档时的 scaling 缩减注意别降太狠否则梯度直接下溢Safe 模式最长步数5~10 步安全模式最多持续多久后尝试恢复太短恢复不稳定太长模型会偏移恢复窗口64 步二次确认稳定后才恢复快速档必须先看 z-score 连续低于阈值一半这张表的参数是我在 7B/13B 规模模型上反复试出来的风格偏保守。如果你更追求训练吞吐z-score 可以放宽到 6.0 以上如果你被 NaN 折磨得厉害就调严一点宁可多一些降档也不要崩。4. 实操落地把双重降档检测装进训练框架4.1 先规划好检测 Hook 点理论说完了直接上实操。要把这套机制落地第一步是选好检测 Hook 点——你不可能检查所有算子也不该检查所有算子。我的建议是优先在几个高风险位置挂探针注意力算子的输出softmax 前后的张量是最容易出现 NaN 的区域。LayerNorm / RMSNorm 的输出方差趋近 0 是经典雷区。优化器更新后的参数梯度爆炸往往在这一步体现。跨卡通信allreduce之后误差累积的放大位置。PyTorch 里挂探针非常方便用register_forward_hook和register_full_backward_hook就行。前向 hook 抓激活值反向 hook 抓梯度都不用改模型代码。但要注意实际的 LLM 训练经常会用融合算子比如 Flash Attention这类算子在 CUDA graph 里被捕获默认情况下的 hook 可能触发不到中间张量。我踩过这个坑——我以为探针挂了其实 Flash Attention 的中间 softmax 根本不出来。解决方法是给融合 kernel 内部加一个 sideband 输出通道把关键中间统计量带出来。另外说一句推理侧的事。训练侧有 loss 能看推理侧没有 loss但同样会遇到数值问题。这套机制的思路可以沿用到推理服务第一重检测监控生成过程中每层 token 表征的异常第二重检测在出现生成质量大幅劣化时回溯到最近一次正常 token 位置重新解码。我自己在 PPL serving 场景里试过简化版效果不错但这就是另一个话题了今天先聚焦训练。4.2 第一重检测器的实现骨架直接上代码。一个最小的前兆检测器我用 PyTorch 风格写出来import torch import math class PrecursorDetector: 第一重前兆检测器EMA 基线 z-score 二阶差分判定。 挂在算子输出的 hook 上返回 risk_scorefloat。 def __init__(self, window128, z_threshold5.0, diff_threshold3.0): self.window window self.z_threshold z_threshold self.diff_threshold diff_threshold self.ema_max None # max_abs 的 EMA 基线 self.ema_var None # max_abs 的 EMA 方差 self.past_values [] # 最近的 max_abs 序列 self.alpha 2.0 / (window 1) def _update_ema(self, current): if self.ema_max is None: self.ema_max current self.ema_var 0.0 else: delta current - self.ema_max self.ema_max self.alpha * delta self.ema_var (1 - self.alpha) * (self.ema_var self.alpha * delta * delta) return self.ema_max, math.sqrt(self.ema_var) def _second_diff(self): # 二阶差分捕捉统计量的“加速陡峭化”类似拉普拉斯算子的离散形态 if len(self.past_values) 3: return 0.0 d1 self.past_values[-1] - self.past_values[-2] d0 self.past_values[-2] - self.past_values[-3] return abs(d1 - d0) def on_tensor(self, tensor): # 只用设备端规约取统计量避免 .item() 同步阻塞流水线 max_abs tensor.abs().max().detach() current max_abs.item() self.past_values.append(current) if len(self.past_values) self.window: self.past_values.pop(0) base, std self._update_ema(current) if std 1e-12: return 0.0 z_score (current - base) / std second_diff self._second_diff() risk 0.0 if z_score self.z_threshold: risk 0.5 if second_diff self.diff_threshold: risk 0.5 return min(risk, 1.0)细心的人会注意到我在代码注释里特别写了不要用.item()直接同步——这个非常重要。如果你在训练主循环的 hook 里每次都对张量调用.item()它会强制设备端到主机端的同步训练流水线直接被打断吞吐可以掉 30% 以上。正确做法是先用torch.max这类归约在设备端拿到标量再一次性取回。实际工程上我更推荐让检测器跑在独立的 CUDA stream 上跟主计算流并行这样探针的开销还能再压一半。4.3 第二重定位器的实现要点第二重定位器负责在 NaN 已经出现时做精确打击。核心是确定性重放和二分定位。我给一个简化的流程骨架class PollutionLocalizer: 第二重污染源定位器从最近 checkpoint 确定性重放二分定位到肇事算子。 def __init__(self, model, checkpoint_state, seed42): self.model model self.checkpoint_state checkpoint_state self.seed seed def _setup_deterministic(self): # 确定性重放的关键环境配置 torch.manual_seed(self.seed) torch.cuda.manual_seed_all(self.seed) torch.backends.cudnn.benchmark False # 关闭非确定性算法选择 torch.backends.cuda.matmul.allow_tf32 False # 如果用了 CUDA graph需要重新捕获不能复用原始图 def _replay_subgraph(self, op_list): # 以某个算子列表的子图重放返回输出中是否存在 NaN # 实际实现中需要配合 autograd graph 的截断机制 pass def locate(self, full_op_list, nan_step): 二分定位不断缩小算子范围直到定位到单个算子。 full_op_list: 从 checkpoint 到 nan_step 之间的算子执行列表 self._setup_deterministic() low, high 0, len(full_op_list) - 1 while low high: mid (low high) // 2 has_nan self._replay_subgraph(full_op_list[:mid 1]) if has_nan: high mid # NaN 在前半段范围内 else: low mid 1 # NaN 在后半段 bad_op full_op_list[low] # 对当前算子做通道级扫描定位具体张量位置 return bad_op二分定位的时间复杂度是对数级的几百个算子的计算图十次以内重放就能定位到单个算子。这里的_replay_subgraph实现起来有门槛——你要能截断 autograd graph 的中间节点。我的做法是在训练代码里提前把每个算子的输入输出缓存下来只缓存引用不拷贝数据定位时直接用缓存做子图重放速度快很多。缺点是缓存占用显存我一般只缓存最近 10~20 步的激活引用配合 checkpoint 足够用了。有个细节重放时模型一定要切到 eval错必须保持训练模式但关闭梯度计算。因为有些算子比如 Dropout在训练和 eval 模式下行为不同会影响重放的一致性。正确做法是model.train()torch.no_grad()组合——保持所有依赖随机性的算子行为一致同时不产生多余的梯度计算。4.4 降档执行器动态替换算子的实现检测到问题怎么把算子替换掉用 PyTorch 的话最直接的办法是setattr(model, module_name, SafeLinear(...))但这种方法只能在 Module 粒度替换粒度太粗。更精细的做法是给模型的 forward 增加一个“档位路由”class DualDownshiftManager: 降档管理器管理当前档位状态控制算子在普通实现和 Safe 实现之间切换。 MODE_FAST fast # 快速档FP16/BF16 向量化高性能路径 MODE_STABLE stable # 稳健档BF16 FP32 归约 调整优化器参数 MODE_SAFE safe # 安全档补偿求和 FP64 归约 范围压缩 def __init__(self): self.mode self.MODE_FAST self.safe_steps_left 0 def trigger_first_downshift(self): # 第一次降档快速档 - 稳健档 self.mode self.MODE_STABLE # 实际动作切换算子的累积精度、调整 loss scaling、收紧梯度裁剪 # 通过修改 training state 的 config 实现而不是换模型结构 def trigger_second_downshift(self, bad_op_name): # 第二次降档稳健档 - 安全档只针对肇事算子 self.mode self.MODE_SAFE self.safe_steps_left 8 # Safe 模式最多跑 8 步 self._replace_op_with_safe_impl(bad_op_name) def _replace_op_with_safe_impl(self, op_name): # 动态替换算子为安全实现例如把 nn.Linear 替换为 SafeLinear # 升华腾上则是替换成 CANN 的高精度算符 pass def maybe_recover(self): # 稳定后恢复安全档 - 稳健档 - 快速档 if self.safe_steps_left 0: self.safe_steps_left - 1 elif self.mode self.MODE_SAFE: self.mode self.MODE_STABLE # 再观察若干步确认 z-score 连续偏低后升回快速档这里的关键设计是第一次降档不要换模型结构只改配置loss scaling、梯度裁剪、累积精度开关这样开销最小不会打断 CUDA graph 的复用第二次降档才动算子结构而且只替换肇事算子不是全模型替换。我实测在 13B 模型上全模型替换成 Safe 实现吞吐掉 40% 以上只替换肇事算子吞吐只掉 5%~8%完全可接受。4.5 完整触发流程走一遍把整个流程串起来描述一遍你就知道它到底怎么工作了。假设我们在训练一个 13B 的 LLM某一步某个 attention 头部在 softmax 前出现了 logits 溢出前兆Fusion attention 的 sideband 输出带出来一个统计量第一重检测器的 EMA 基线被击穿z-score 达到 5.8二阶差分达到 4.2风险评分拉到 1.0。降档管理器触发第一次降档切换成稳健档该 attention 算子改用 BF16 计算路径loss scaling 下调梯度裁剪收紧。训练继续跑了几步前兆信号消失z-score 回落看起来正常了——这时候机制就算成功训练没有中断只是自动降档规避了风险。但假设这次运气不好降档之后某个相关算子还是崩了loss 变成 NaN。第一重没拦住第二重启动。训练冻结从最近 checkpoint 加载状态关闭 cudnn benchmark设置固定 seed开始确定性重放。二分十几次后定位到一个 Softmax 算子的分母归约——FP32 累加器在超长序列上还是溢出了。第二次降档触发该 Softmax 算子替换成 Safe 模式分母改用 FP64 累加器 Kahan 补偿求和输入做动态范围压缩。训练从 checkpoint 继续Safe 模式跑 8 步z-score 连续低于阈值一半管理器尝试恢复稳健档再观察 64 步确认稳定后回到快速档。整个事故影响控制在几十步之内定位时间不到一分钟。这套流程不是我纸上谈兵是我在真实训练任务上反复跑过的。第一次完整跑通的时候旁边同事看我看板的眼神都不一样了——以前出 NaN 是要拉会讨论的现在只是日志里多几行“downshift triggered”而已。5. 常见问题与排查技巧实录5.1 降档之后还是 NaN怎么回事这是我被问得最多的一个问题——不是说双重降档很厉害吗怎么降完了还是 NaN说实话我第一次自己跑通这套机制的时候也遇到过。排查方向按优先级排列先看降档到底有没有生效。具体地说检查算子是不是真的被替换了。我遇到过在跑 CUDA graph 捕获时替换算子结果新算子根本没进 graph执行的还是老 kernel。这类问题的典型特征是日志显示已触发降档但算子行为没有任何变化。解决办法是降档动作要在下一次 graph 捕获之前完成或者干脆在降档时对相关子图做一次重新捕获。再看归约是不是真的用了 FP64。很多“Safe 实现”看起来写了 FP64实际在 kernel 内部自动被编译器优化回 FP32 了。建议直接在 kernel 里打印归约累加器的数据类型做验证。我之前用 Triton 写 SafeKernel 时就吃过这个亏——Triton 默认会把 FP64 归约在 GPU 上降级需要在 kernel 里显式用tl.float64指定。再看数据侧。有时候 NaN 不是算子问题是数据本身就含 inf 或 NaN比如某些异常样本的 label 是 NaN。前兆检测器对这种数据源头的污染无能为力因为它在算子出口看统计量源头数据已经是坏的。这种情况要加数据管线前置校验。最后看硬件层。某些显存错误、越界写导致的内存破坏也会表现为随机 NaN。这类问题特征更明显NaN 出现的算子是随机的重放也不稳定复现。如果重放定位到的算子每次都不一样优先怀疑硬件。5.2 误报太多训练吞吐掉得快前兆检测器太灵敏动不动就降档虽然不会崩但稳健档跑多了吞吐还是会下降。我遇到过最夸张的情况一个模型上误报率 30%训练速度肉眼可见地变慢。排查和调优经验如下第一步检查 z-score 阈值是不是定得太低了。把阈值从 4.0 放宽到 5.5~6.0误报率通常会下降一个数量级。第二步看二阶差分判定是不是过度敏感。这个判定我引入的时候是为了抓“加速陡峭化”但有些模型阶段比如学习率 warmup 时统计量本来就在快速变化二阶差分天然偏高。这时候应该给二阶差分判定加一个绝对值下限——只有当前 max_abs 已经接近 FP16 上限的 50% 时才允许触发。还有一个实操技巧第一重探针不要所有算子都挂优先挂注意力输出和 LayerNorm 输出。我最初把所有算子的探针都开启了结果日志量爆炸性能也受影响。后来只保留高风险算子的探针误报和性能问题同时缓解。经验数据7B 模型上从全算子探针缩减到 5 个关键探针检测能力几乎不降开销从 1.5% 降到 0.3%。5.3 降档后模型指标掉得厉害降档机制本身是为了保命但如果频繁降档模型训练曲线会受影响。典型症状是跑了一个晚上loss 倒是没有 NaN但最终指标比正常训练差不少。原因之一Safe 模式跑太久。我默认设置是 Safe 模式最多跑 5~10 步就尝试恢复但如果你的恢复判定太保守比如 z-score 阈值设得太高Safe 模式可能会实际持续几十步。Safe 模式里补偿求和、FP64 归约这些操作的计算结果跟快速档不完全一致相当于给模型加了“扰动”跑久了模型自然偏移。原因之二学习率临时缩小的时间窗口太长。我在第一次降档的时候会把学习率乘 0.1如果降档频繁发生这一步的“冷静期”其实会拉低整体训练进度。后来我把逻辑改成了第一次降档不动学习率只调 loss scaling 和梯度裁剪只有第二次降档进入 Safe 模式才动学习率。这样既能止损又不影响正常阶段的推进速度。5.4 硬件平台适配差异这套机制在不同硬件平台上的落地方式差别很大这里把我知道的说一下。GPUNVIDIA上最顺滑PyTorch hook 机制 Triton 写 SafeKernel替换链路最成熟。唯一要注意的是 Flash Attention 这类融合 kernel 的中间张量拿不到需要厂商开放 sideband 接口或者自己改用非融合实现来跑探针。如果训练吞吐吃紧可以只在探针触发后才切换到非融合路径平时用融合 kernel这个策略我实测比较稳。昇腾平台CANN上需要另一套姿势CANN 的算符Operator是有自己注册体系的降档逻辑可以在算子实现的层面做。具体来说CANN 算子优化里本来就支持多种精度实现映射你只要在本地op_type_impl里注册一个“高精度变体”降档时把算子调度切到变体即可。探针的话建议放到 AICore 侧做而不是在 Host 侧反复拉取——Host 侧同步取数在昇腾上开销更大。我之前在昇腾上做的版本检测开销在 0.5% 以内效果跟 GPU 版基本一致。自研 AI 芯片的话这套机制更值得参考——你可以在指令集层面直接把“溢出前兆判定”做进硬件单元比如让向量单元在输出时顺带标记一个“本次计算最大绝对值”寄存器探针只读这个寄存器就行几乎零成本。5.5 问题速查表把上面这些经验浓缩成一张速查表方便你遇到问题时快速对照。症状可能原因处理办法降档后仍 NaN算子没真正替换 / 归约没走 FP64检查 CUDA graph 捕获时机验证 kernel 内累加器类型NaN 位置随机、重放不复现数据源头污染或硬件不稳定前置数据校验跑内存和显存诊断误报频繁、吞吐下降z-score 阈值过低 / 探针过多放宽阈值限制探针数量到高风险算子Safe 模式后指标劣化Safe 模式持续时间过长 / 学习率降档太频繁缩短 Safe 模式最高步数恢复判定放宽融合算子探针抓不到Flash Attention 等 kernel 中间值不暴露用 sideband 输出或切换非融合路径日志量爆炸全算子探针 全 rank 上报只保留关键探针rank0 汇总异步上报5.6 独家避坑清单最后这条避坑清单是我在被 NaN 折磨了无数个通宵之后总结出来的每一条都是真金白银换来的教训。第一探针取数千万不要用.item()直接同步一定要走设备端归约再统一取回。同步一次看似没关系但训练是流水线的同步一次卡一整条流吞吐掉得无声无息。第二降档动作和检测判定要分开先算风险评分再决定降档不要“先降档再做检测确认”。后者会让你在误报的时候白白降档白吃性能损失。这个逻辑顺序我最初写反了改了之后误报率对训练速度的影响小了很多。第三多卡并行DDP/FSDP场景下探针不要每个 rank 都全量上报只让 rank0 做大脑其他 rank 把统计量传给 rank0。否则日志和指标系统会被探针消息直接淹没而且跨卡通信本身也在加剧数值误差累积。第四前兆检测器的基线会被训练阶段变化污染。比如学习率阶段性提升之后loss 曲线和梯度分布会整体变宽EMA 基线需要一段时间重新适应这个阶段误报率会升高。解决办法是在学习率剧变之后的前 30~50 步把检测灵敏度自动降一档。第五恢复策略比降档策略更难写。降档是一锤子买卖恢复是要在“太早导致再次崩”和“太晚导致模型偏移”之间找平衡。我的经验值是安全模式跑 5 步、观察 64 步、分两段恢复先升稳健档确认稳定再升快速档。这个节奏在 7B/13B 级别模型上都很稳定更大的模型可以把观察窗口适当拉长。我个人在实际操作中的体会是这套双重降档检测机制真正解决的不是“NaN 会不会出现”而是“NaN 出现之后你的反应速度有多快、代价有多小”。廉价检测、昂贵定位、分级降档把这三个原则拆开任何训练框架都能接得住。最后再分享一个小技巧把检测器的风险评分、当前的档位状态、loss 曲线这三路信号放到同一个看板上。正常训练时风险评分应该是一条贴着 0 走平的线档位稳定在快速档一旦风险评分开始抖动你就能在 loss 变 NaN 之前看到异常苗头。有了这张看板故障复盘从“三个小时猜原因”变成“一分钟看趋势”。这套机制做到后面你会发现自己对训练数值健康度的掌控感完全不一样了。