执行摘要
- 一句话:修复 Blackwell 上 DeepSeek-V4-Pro NaN 问题
- 推荐动作:值得精读:这是一个经典的低级别数值精度 Bug,展示了尾数位泄露如何导致 NaN,以及如何通过一个简单的 scale 校正函数来解决。对于运行 DeepSeek-V4-Pro 在 Blackwell 上的团队,此修复是必须集成的。
功能与动机
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 均明确指出了根本原因和修复思路。
实现拆解
- 在
python/sglang/srt/models/deepseek_v4.py 的 forward 方法中,定位到 _FP8_WO_A_GEMM 分支下的量化后 scale 变量 o_s 的使用位置。
- 在第 626 行插入对
deep_gemm.ceil_to_ue8m0(o_s) 的调用,将 quantization scale 转换为 ue8m0 格式,确保尾数位被清零。
- 改动仅一行,不涉及其他逻辑变更;由于 scale tensor 很小(如 shape (2,32)),ceil_to_ue8m0 的额外开销可忽略。
- 无配置文件、测试或部署改动。
关键文件:
python/sglang/srt/models/deepseek_v4.py(模块 模型层;类别 source;类型 data-contract;符号 forward): 核心修复文件,在 forward 方法中的 FP8 GEMM 路径前插入 ceil_to_ue8m0 调用,消除 NaN 根源。
关键符号:forward
关键源码片段
python/sglang/srt/models/deepseek_v4.py
核心修复文件,在 forward 方法中的 FP8 GEMM 路径前插入 ceil_to_ue8m0 调用,消除 NaN 根源。
# 文件 : python/sglang/srt/models/deepseek_v4.py
# 关键改动在第 626 行:在调用 fp8_einsum 前对量化 scale 执行 ue8m0 格式化
# 防止 deep_gemm 的内部打包内核因尾数位泄露导致 NaN
if _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)
评论区精华
review 过程中无争议讨论,ch-wan 直接 approved。PR 评论中 nai-kon 询问该修复是否包含在 v0.5.13 中,Fridge003 确认已包含。
风险与影响
- 风险:风险极低:改动只增加了一行对
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 修复。
- 风险标记:依赖外部库接口
关联脉络
- PR #27896 [Perf] Skip per-call mat_a/scales_a padding in cutlass FP8 blockwise GEMM: 同属 FP8 量化路径下的性能/正确性优化链,修改了 fp8_utils.py 中的量化逻辑,与本 PR 的 FP8 scale 处理有间接关联。
参与讨论