执行摘要
- 一句话:修复ROCm 7.0上FP8 GEMM的NaN问题
- 推荐动作:该PR是一个针对特定硬件/软件组合的精准修复,设计思路清晰——按形状和M动态选择后端,而不是全局弃用CK。值得关注其优雅的降级策略,可作为后续同类硬件差异修复的参考模式。如果团队在生产AMD gfx950 + ROCm 7.0,建议立即合并。
功能与动机
Qwen3.5-397B-A17B-FP8在AMD gfx950 + ROCm 7.0上准确率从正常0.95跌至0.24,且约67%输出无效。PR body中详细分析了根因:由于ROCm 7.0 hipcc编译问题,bpreshuffle CK GEMM被禁用,fallback至非bpreshuffle的ck_gemm_a8w8_blockscale,该实现在特定形状和大M值下返回NaN。已在PR body附上形状与NaN阈值的对照表,并在ROCm 7.2上验证CK行为正确。
实现拆解
- 定义安全M上限字典:在
python/sglang/srt/layers/quantization/fp8_utils.py中新增模块级字典_AITER_GFX95_CK_W8A8_MAX_SAFE_M,记录受影响的( n, k )形状及其确认安全的M上限。
- 修改GEMM选择逻辑:在
aiter_w8a8_block_fp8_linear函数的elif _use_aiter_gfx95:分支中,当input_2d.shape[0] > _ck_safe_m时,强制use_triton=True,从而回退到数值正确的Triton实现。
- 保留小M性能:对于低于安全M上限的场景,仍使用更快的CK路径,避免性能退化。
- 不改动其他路径:对ROCm ≥ 7.2、非gfx95硬件、或未在字典中列出的形状,逻辑保持不变。
关键文件:
python/sglang/srt/layers/quantization/fp8_utils.py(模块 量化层;类别 source;类型 core-logic): 核心更改文件:定义安全M字典并修改 GEMM 选择逻辑,直接修复精度问题。
关键符号:未识别
关键源码片段
python/sglang/srt/layers/quantization/fp8_utils.py
核心更改文件:定义安全M字典并修改 GEMM 选择逻辑,直接修复精度问题。
# 位于 _use_aiter_bpreshuffle_gfx95 定义之后
# gfx95 + ROCm < 7.2: bpreshuffle CK 被禁用(如上),非 bpreshuffle 的
# ck_gemm_a8w8_blockscale 后备路径在某些形状下对大 M 返回 NaN
# (实测 NaN 起始点 : (2560,4096)@M>=4096, (4096,1024)@M>=8192),
# 这会导致 prefill 阶段被污染。
# 这里将每个受影响的 (n, k) 映射到确认安全的 CK 最大 M 值
# (保守取最后一个安全 M)。在不超过该 M 时保留较快的 CK 路径,
# 超过时回退到数值正确的 Triton FP8 GEMM。该问题在 ROCm 7.2 中已修复。
_AITER_GFX95_CK_W8A8_MAX_SAFE_M = {
(2560, 4096): 2048,
(4096, 1024): 4096,
}
# ... 省略其他代码 ...
def aiter_w8a8_block_fp8_linear(
input: torch.Tensor,
weight: torch.Tensor,
block_size: List[int],
weight_scale: torch.Tensor,
input_scale: Optional[torch.Tensor] = None,
bias: Optional[torch.Tensor] = None,
) -> torch.Tensor:
input_2d = input.view(-1, input.shape[-1])
output_shape = [*input.shape[:-1], weight.shape[0]]
n, k = weight.shape
# 原有分支 : 当 bpreshuffle 可用则优先使用 tuned gfx950
if _use_aiter_bpreshuffle_gfx95:
use_triton = use_aiter_triton_gemm_w8a8_tuned_gfx950(n, k)
elif _use_aiter_gfx95:
# gfx95 + ROCm < 7.2: 在安全 M 以下保留较快 CK 路径;高于阈值时
# ck_gemm_a8w8_blockscale 返回 NaN,故改用 Triton。未列出的形状保持原有决策。
_ck_safe_m = _AITER_GFX95_CK_W8A8_MAX_SAFE_M.get((n, k))
use_triton = use_aiter_triton_gemm_w8a8_tuned_gfx950(n, k) or (
_ck_safe_m is not None and input_2d.shape[0] > _ck_safe_m
)
else:
use_triton = True
# 后续根据 use_triton 选择具体算子
评论区精华
该PR的review讨论很简单,仅包含自动机器人评论和HaiShaw的批准,没有实质性质疑。HaiShaw直接批准了变更。
风险与影响
- 风险:
- 回归风险低:变更仅作用于gfx95 + ROCm < 7.2 + 特定( n, k )形状 + 大M,且保留了小M的CK路径;其他硬件和ROCm版本不受影响。
- 性能风险:大M场景切换到Triton实现,可能比CK慢,但保证了正确性;对于Qwen3.5的prefill阶段,优先级是正确性。
- 测试覆盖不足:PR描述提到已通过本地和CI测试验证精度恢复,但未包含直接的单元测试文件,手动构造的fp8 GEMM测试未纳入回归套件。
- 影响:直接影响:修复Qwen3.5-397B-A17B-FP8在AMD gfx950 + ROCm 7.0上的推理精度,GSM8K从0.244恢复到0.950,达到ROCm 7.2参考水平。影响范围:仅限使用AMD gfx950且ROCm版本低于7.2,同时运行时触发受影响的GEMM形状((2560,4096)和(4096,1024))且M值大于2048或4096的场景。性能:大M场景下Triton代替CK,预计prefill用时略有增加,但decode等小M路径不受影响。
- 风险标记:缺少测试覆盖
关联脉络
- PR #23319 Unknown (mentioned as PR#23319): PR body中提及的ROCm 7.0 hipcc miscompile问题,是导致bpreshuffle被禁用的根本原因,本PR在此基础上进一步修复了fallback路径的NaN bug。
参与讨论