【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 散度,并且要对 logitslog_softmax 后做比值。这套链路对 attention 输出里的微小数值误差极度敏感——一旦 attention 的分数里有几个 NaN/Inf,经过 exp/log/softmax 会被放大成整条序列的 NaN。而 SFT 只看交叉熵,对单点误差容忍度高,所以同样的问题在 SFT 下被"掩盖"了。

三、根因

根因是一条链路:

  1. 3D position_ids 触发 SDPA math 回退。 PyTorch 的 SDPA 调度器看到 (batch, 3, seq_len) 的 position_ids,无法映射到 flash 后端支持的 (batch, seq_len) 形式 → 退化到 math 后端。

  2. math 后端在 BF16 下精度塌陷。 BF16 只有 8 位尾数(约 3 位有效十进制),math 后端用纯 PyTorch 做 QK^T / sqrt(d)softmax,中间 QK^T 的数值范围很大,BF16 舍入误差显著。Flash 后端有在线归一化(online softmax)避免大数累加,而 math 后端没有,于是 softmax 的分母出现 exp(很大) + exp(很大) → 上溢成 Inf,再被 logNaN

  3. RLHF 的数值链路放大 NaNratio = 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 正常时:

  1. 看日志有没有 falling back to math backend——有就说明 SDPA 回退了。
  2. 检查 position_ids 维度:混合注意力模型常是 3D (batch,3,seq_len),正是触发回退的常见原因。
  3. 确认精度:BF16 下 math 后端极易 softmax 上溢,切 fp32 试一步,若不炸就坐实了根因。
  4. 看是否 RLHF 专属:SFT 正常但 PPO/GRPO NaN,基本是 ratio/KL 链路放大了 attention 的微小 NaN。
  5. 修复优先级:先升 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 后端。

更多推荐