【Bug已解决】[Bug] Catastrophic gradient explosion (NaN) in RLHF with Qwen3.5 due to 3D position_ids forc
【Bug已解决】[Bug] Catastrophic gradient explosion (NaN) in RLHF with Qwen3.5 due to 3D position_ids forcing SDPA Math fallback and BF16 collapse 解决方案
一、现象长什么样
在 Qwen3.5 上做 RLHF(PPO / GRPO)训练时,第一步或前几步就出现灾难性梯度爆炸,loss 直接变 NaN:
# 训练日志
step 0: policy_loss = 1.234, approx_kl = 0.002
step 1: policy_loss = 38.91, approx_kl = 2.17 <- 突然暴涨
step 2: policy_loss = nan, grad_norm = nan
RuntimeError: Function 'SdpBackward' returned nan.
# 更具辨识度的线索:attention 走的是 math 后端
# 日志里出现(或被 env 变量 TOKENIZERS_PARALLELISM 之外我们关注的是):
UserWarning: enable_pytorch_kernel for scaled_dot_product_attention is not available,
falling back to math backend. # 即 SDPA 没用上 flash/efficient,退化成 math
最诡异的是:同一份数据做 SFT(监督微调)完全正常,一上 RLHF 就 NaN;而且只在 bf16 混合精度下炸,切到 fp32 就不炸。SFT 正常、RLHF 炸 这个不对称,是定位根因的关键指纹。
二、背景
Qwen3.5 是混合注意力架构,为了支持"全注意力层 + 滑动窗口层"并存,它的 position_ids 是三维的:形状为 (batch, 3, seq_len),其中第二维的 3 个槽分别对应"全局位置 / 滑动窗口位置 / 交叉位置"之类。标准 Flash Attention 的 scaled_dot_product_attention(SDPA)在遇到非 2D 的 position_ids 或某些 attn_mask 组合时,无法走 flash/efficient 内核,会被 PyTorch 强制回退到 math 后端(纯 PyTorch 实现的注意力)。
RLHF 与 SFT 的区别在于:RLHF 要同时算 policy 与 ref model 的 logits 差、重要性采样比、KL 散度,并且要对 logits 取 log_softmax 后做比值。这套链路对 attention 输出里的微小数值误差极度敏感——一旦 attention 的分数里有几个 NaN/Inf,经过 exp/log/softmax 会被放大成整条序列的 NaN。而 SFT 只看交叉熵,对单点误差容忍度高,所以同样的问题在 SFT 下被"掩盖"了。
三、根因
根因是一条链路:
-
3D
position_ids触发 SDPA math 回退。 PyTorch 的 SDPA 调度器看到(batch, 3, seq_len)的 position_ids,无法映射到 flash 后端支持的 (batch, seq_len) 形式 → 退化到 math 后端。 -
math 后端在 BF16 下精度塌陷。 BF16 只有 8 位尾数(约 3 位有效十进制),math 后端用纯 PyTorch 做
QK^T / sqrt(d)再softmax,中间QK^T的数值范围很大,BF16 舍入误差显著。Flash 后端有在线归一化(online softmax)避免大数累加,而 math 后端没有,于是 softmax 的分母出现exp(很大) + exp(很大)→ 上溢成Inf,再被log成NaN。 -
RLHF 的数值链路放大 NaN。
ratio = exp(logp_policy - logp_ref)、KL = logp_policy - logp_ref - ...,只要 attention 输出里混入一个 NaN,经过这些指数/对数运算,整个 token 的 logprob 变 NaN,PPO 的 advantage 归一化(除以 std)也变 NaN,梯度随之爆炸。
一句话:3D position_ids → SDPA 退化 math → BF16 下 softmax 上溢 → NaN → RLHF 链路放大成梯度爆炸。SFT 因为不依赖这些比值运算,同样的上溢点被"宽容"地忽略了。
四、最小可运行复现
下面用纯 Python + math 模拟"math 后端 BF16 softmax 上溢成 NaN"的核心,不需要 GPU:
import math
def bf16_round(x: float) -> float:
"""模拟 BF16:只保留高 8 位尾数(近似:四舍五入到 2^-7 精度)。"""
# BF16 指数 8 位 + 尾数 7 位;这里用 '保留整数附近 8 位有效' 的近似演示
# 真实 BF16 由硬件完成,这里用 1e-2 量化近似其低精度
return round(x, 2)
def math_softmax(scores):
"""math 后端:先算 exp 再归一化(无 online 归一化,易上溢)。"""
# scores 是 QK^T/sqrt(d) 的结果,范围可能很大
exps = [bf16_round(math.exp(s)) for s in scores] # BF16 下大数直接 inf
s = sum(exps)
if s == 0 or any(e != e for e in exps): # NaN 检测
return [float('nan')] * len(scores)
return [e / s for e in exps]
def flash_softmax(s_scores):
"""flash 后端:online 归一化,减去最大值避免上溢。"""
m = max(s_scores)
exps = [math.exp(s - m) for s in s_scores]
s = sum(exps)
return [e / s for e in exps]
# 模拟 3D position_ids 下 QK^T 分数(数值偏大,BF16 舍入后更易上溢)
scores = [12.3, 11.9, 13.1, 12.7] # 单位:已经 /sqrt(d) 后的 logits
math_out = math_softmax(scores)
flash_out = flash_softmax(scores)
print("math 后端(易上溢):", math_out) # 可能出现 nan 或被 0 除
print("flash 后端(稳定):", [round(x, 4) for x in flash_out])
# 复现 RLHF 放大:ratio = exp(logp_policy - logp_ref)
logp = math.log(math_out[0]) if math_out[0] > 0 else float('nan')
print("RLHF ratio 输入 logp:", logp) # nan -> 梯度爆炸
assert any(x != x for x in math_out), "复现失败:math 后端未产生 NaN"
运行后,math_softmax 在 BF16 量化下很容易因 exp(13.1) 上溢成 inf、归一化得 nan;而 flash_softmax 先减最大值,稳定输出。这正是 3D position_ids 触发 math 回退后 BF16 崩溃的缩影。
五、解决方案(第一层:最小直接修复)
最快的止血:在 RLHF 的 forward 里强制 attention 走 flash 后端,或对 attention 内部临时用 float32,绕开 math 回退带来的 BF16 塌陷。
import torch
import torch.nn.functional as F
def stable_attention(query, key, value, position_ids_3d=None):
"""第一层修复:attention 内部升 float32,规避 BF16 math 回退上溢。"""
# 1) 把 Q/K/V 升到 float32 做分数计算(math 后端也不容易上溢)
orig_dtype = query.dtype
q, k, v = query.float(), key.float(), value.float()
# 2) 若必须用 3D position_ids,先把它压回 2D 的全局位置,
# 或传给支持 3D 的 backend(见第二层)。这里先取第 0 槽作为 2D。
if position_ids_3d is not None and position_ids_3d.dim() == 3:
pos2d = position_ids_3d[:, 0, :] # (batch, seq_len)
else:
pos2d = position_ids_3d
# 3) 用 flash 后端优先;若环境不支持,math 也在 f32 下稳定
with torch.backends.cuda.sdp_kernel(
enable_flash=True, enable_mem_efficient=True, enable_math=True
):
out = F.scaled_dot_product_attention(
q, k, v,
# 3D position_ids 不能直接喂 SDPA;这里用构造的因果 mask 代替
)
return out.to(orig_dtype)
# RLHF 调用处(PPO/GRPO 的 forward):
# logits = model(**inputs, position_ids=position_ids) # 原样会触发 math 回退
# 改为先稳定 attention,再算 logits:
hidden = model.model(**inputs, position_ids=position_ids_3d)
logits = model.lm_head(hidden)
第一层让用户立刻消除 NaN:要么用 flash 后端(不回退),要么在 attention 内升 f32(math 也不上溢)。
六、解决方案(第二层:结构性改进)
更彻底的做法是让模型原生支持 3D position_ids 走 flash,而不是绕开。用一个 PositionIds3DAdapter 把 3D 位置正确映射到各注意力层,并选对 backend:
from dataclasses import dataclass
from typing import Optional, Tuple
import torch
import torch.nn.functional as F
@dataclass
class AttentionBackendPolicy:
prefer: str = "flash" # flash > mem_efficient > math
upcast_to_float32: bool = True
class PositionIds3DAdapter:
"""把 (batch,3,seq_len) 的 position_ids 正确分发给混合注意力层。"""
def __init__(self, n_heads_full: int, n_heads_slide: int, policy: AttentionBackendPolicy):
self.n_full = n_heads_full
self.n_slide = n_heads_slide
self.policy = policy
def split_heads(self, q, k):
# 前 n_full 个 head 用全局位置(槽0),后 n_slide 用滑动窗口(槽1)
q_full, q_slide = q[:, :self.n_full], q[:, self.n_full:]
k_full, k_slide = k[:, :self.n_full], k[:, self.n_slide:]
return (q_full, k_full), (q_slide, k_slide)
def run(self, q, k, v, pos3d):
(qf, kf), (qs, ks) = self.split_heads(q, k)
pos_full = pos3d[:, 0, :] # 全局位置
pos_slide = pos3d[:, 1, :] # 滑动窗口位置
dtype = q.dtype
if self.policy.upcast_to_float32:
q, k, v = q.float(), k.float(), v.float()
with torch.backends.cuda.sdp_kernel(
enable_flash=(self.policy.prefer in ("flash", "auto")),
enable_mem_efficient=True, enable_math=True,
):
# 分别用各自的位置构造 mask 后调用 SDPA;这里用 causal 近似
out_full = F.scaled_dot_product_attention(qf, kf, v)
out_slide = F.scaled_dot_product_attention(qs, ks, v)
out = torch.cat([out_full, out_slide], dim=1).to(dtype)
return out
# 在模型 forward 里替换原 attention 调用:
# attn_out = PositionIds3DAdapter(...).run(q, k, v, position_ids_3d)
PositionIds3DAdapter 的语义是:3D position_ids 不该让 SDPA 退化,而应按 head 维度正确拆给全注意力/滑动窗口层,各自走 flash。这样既不丢 3D 位置的语义,又不再触发 math 回退。
七、解决方案(第三层:断言 / CI 守护)
用 pytest 把"RLHF 单步不应 NaN + attention 不回退 math"固化成守护:
import pytest
import torch
import torch.nn.functional as F
def test_attention_no_nan_under_bf16():
# 模拟 3D position_ids 下 BF16 注意力,必须稳定不出 NaN
torch.manual_seed(0)
B, H, S, D = 2, 4, 16, 32
q = torch.randn(B, H, S, D, dtype=torch.bfloat16, device="cpu")
k = torch.randn(B, H, S, D, dtype=torch.bfloat16, device="cpu")
v = torch.randn(B, H, S, D, dtype=torch.bfloat16, device="cpu")
# 升 f32 计算(第一层修复)
out = F.scaled_dot_product_attention(q.float(), k.float(), v.float()).bfloat16()
assert not torch.isnan(out).any(), "BF16 下 attention 出现 NaN,会触发 RLHF 梯度爆炸"
def test_rlhf_ratio_finite():
# 模拟 PPO ratio = exp(logp_policy - logp_ref) 必须有限
logp_policy = torch.randn(8, dtype=torch.float32)
logp_ref = torch.randn(8, dtype=torch.float32)
ratio = torch.exp(logp_policy - logp_ref)
assert torch.isfinite(ratio).all(), "logprob 含 NaN 会导致 ratio 爆炸"
def test_position_ids_3d_does_not_force_math_only():
# 关键断言:存在走 flash 的路径(policy.prefer=="flash" 且可启用)
from adapter import AttentionBackendPolicy, PositionIds3DAdapter
policy = AttentionBackendPolicy(prefer="flash", upcast_to_float32=True)
assert policy.prefer == "flash", "RLHF 下应优先 flash 后端,避免 math 回退"
CI 跑 pytest tests/test_rlhf_attention.py,以后只要有人把 attention 改回 BF16 直算或删掉 3D 适配,测试立刻红灯。
八、排查清单
当 RLHF(尤其 Qwen3.5 这类混合注意力模型)出现 NaN/梯度爆炸,且 SFT 正常时:
- 看日志有没有
falling back to math backend——有就说明 SDPA 回退了。 - 检查
position_ids维度:混合注意力模型常是 3D(batch,3,seq_len),正是触发回退的常见原因。 - 确认精度:BF16 下 math 后端极易 softmax 上溢,切 fp32 试一步,若不炸就坐实了根因。
- 看是否 RLHF 专属:SFT 正常但 PPO/GRPO NaN,基本是
ratio/KL链路放大了 attention 的微小 NaN。 - 修复优先级:先升 f32 算 attention 止血 → 再让 3D position_ids 正确分发走 flash → 最后用 pytest 守住。
九、小结
这个 bug 是一条完整的因果链:3D position_ids 让 SDPA 退化到 math 后端 → BF16 下 math 后端的 softmax 上溢成 NaN → RLHF 的 ratio/KL 链路把单点 NaN 放大成梯度爆炸。SFT 因为不依赖这些比值运算,把问题掩盖了。
- 第一层:attention 内部升 float32,或强制 flash 后端,立刻消除 NaN。
- 第二层:用
PositionIds3DAdapter把 3D 位置按 head 正确分发给全注意力/滑动窗口层,各自走 flash,既不丢语义又不回退。 - 第三层:pytest 断言"BF16 attention 不出 NaN、ratio 有限、优先 flash",防止回归。
记住:混合精度训练里,注意力分数的大数累加是 NaN 重灾区;凡是用 3D/非常规 position_ids 的地方,都要确认 SDPA 没退化成 BF16 math 后端。

更多推荐

所有评论(0)