Prhub

#25733 [Bug] Fix V4-Pro NaN on Blackwell by converting fp8_einsum input scale to ue8m0

原始 PR 作者 yhyang201 合并时间 2026-05-19 14:48 文件变更 1 提交数 1 评论 10 代码增减 +1 / -0

执行摘要

修复 Blackwell 上 DeepSeek-V4-Pro NaN 问题

DeepSeek-V4-Pro 在 Blackwell GPU 上做自回归解码时产生乱码(NaN),根本原因是 deep_gemm 的 transpose_and_pack_fp32_into_ue8m0 CUDA 内核在打包 fp32 指数时没有屏蔽尾数位,导致非 2 的幂的 scale 值被破坏,进而在 batch=1(单 token 解码)时 fp8_einsum 产生 NaN。PR body 和 commit message 均明确指出了根本原因和修复思路。

值得精读:这是一个经典的低级别数值精度 Bug,展示了尾数位泄露如何导致 NaN,以及如何通过一个简单的 scale 校正函数来解决。对于运行 DeepSeek-V4-Pro 在 Blackwell 上的团队,此修复是必须集成的。

讨论亮点

review 过程中无争议讨论,ch-wan 直接 approved。PR 评论中 nai-kon 询问该修复是否包含在 v0.5.13 中,Fridge003 确认已包含。

实现拆解

  1. python/sglang/srt/models/deepseek_v4.pyforward 方法中,定位到 _FP8_WO_A_GEMM 分支下的量化后 scale 变量 o_s 的使用位置。
  2. 在第 626 行插入对 deep_gemm.ceil_to_ue8m0(o_s) 的调用,将 quantization scale 转换为 ue8m0 格式,确保尾数位被清零。
  3. 改动仅一行,不涉及其他逻辑变更;由于 scale tensor 很小(如 shape (2,32)),ceil_to_ue8m0 的额外开销可忽略。
  4. 无配置文件、测试或部署改动。
文件 模块 状态 重要度
python/sglang/srt/models/deepseek_v4.py 模型层 modified 5.65

关键符号

forward

关键源码片段

python/sglang/srt/models/deepseek_v4.py data-contract

核心修复文件,在 forward 方法中的 FP8 GEMM 路径前插入 ceil_to_ue8m0 调用,消除 NaN 根源。

# 文件 : python/sglang/srt/models/deepseek_v4.py
# 关键改动在第 626 行:在调用 fp8_einsum 前对量化 scale 执行 ue8m0 格式化
# 防止 deep_gemm 的内部打包内核因尾数位泄露导致 NaNif _FP8_WO_A_GEMM:
    import deep_gemm
​
    T, G, D = o.shape
    R = self.o_lora_rank
    o_fp8, o_s = sglang_per_token_group_quant_fp8(
        o.reshape(T * G, D).contiguous(),
        group_size=128,
    )
    # [ 关键修复 ] 将 scale 转换为 ue8m0 格式,清零尾数位
    # deep_gemm 的 ue8m0 打包内核在非 2 的幂 scale 时会泄露尾数位,
    # 导致 fp8_einsum 在 batch=1 时产生 NaN
    o_s = deep_gemm.ceil_to_ue8m0(o_s)
    output = torch.empty(T, G, R, device=o.device, dtype=torch.bfloat16)
    deep_gemm.fp8_einsum(
        "bhr,hdr->bhd",
        (o_fp8.view(T, G, D), o_s.view(T, G, -1)),
        (self.wo_a.weight.view(G, R, D), self.wo_a.weight_scale_inv.data),
        output,
        recipe=(1, 1, 128),
    )
    o = output
else:
    # 非 FP8 路径保持不变
    wo_a = self.wo_a.weight.view(self.n_local_groups, self.o_lora_rank, -1)
    o = torch.einsum("tgd,grd->tgr", o, wo_a)

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险极低:改动只增加了一行对 deep_gemm.ceil_to_ue8m0 的调用,不影响其他分支逻辑。如果 deep_gemm 后续修复了自身打包 bug,这一行可以作为冗余保护,不会引发新问题。但当前修复依赖于 deep_gemm 库的新函数 ceil_to_ue8m0,若运行时 deep_gemm 版本不含该函数,会引发错误。不过 PR 测试已通过 B200 环境验证。

修复仅限于 Blackwell GPU(B300/B200)上运行 DeepSeek-V4-Pro 模型、且使用了 FP8 partial offline activation GEMM 路径的场景。对非 Blackwell 硬件、不使用 FP8 路径或其他模型无影响。由于修复锁定在核心解码路径,且测试已通过,对受影响用户来说是一个关键的 Bug 修复。

依赖外部库接口

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论