Prhub

#29918 [AMD] Gate broken CK block-FP8 GEMM shapes to aiter-triton-GEMM to fix ROCm 7.0 Qwen3.5 accuracy

原始 PR 作者 yctseng0211 合并时间 2026-07-03 09:18 文件变更 1 提交数 2 评论 2 代码增减 +17 / -1

执行摘要

修复 ROCm 7.0 上 FP8 GEMM 的 NaN 问题

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行为正确。

该PR是一个针对特定硬件/软件组合的精准修复,设计思路清晰——按形状和M动态选择后端,而不是全局弃用CK。值得关注其优雅的降级策略,可作为后续同类硬件差异修复的参考模式。如果团队在生产AMD gfx950 + ROCm 7.0,建议立即合并。

讨论亮点

该PR的review讨论很简单,仅包含自动机器人评论和HaiShaw的批准,没有实质性质疑。HaiShaw直接批准了变更。

实现拆解

  1. 定义安全M上限字典:在python/sglang/srt/layers/quantization/fp8_utils.py中新增模块级字典_AITER_GFX95_CK_W8A8_MAX_SAFE_M,记录受影响的( n, k )形状及其确认安全的M上限。
  2. 修改GEMM选择逻辑:在aiter_w8a8_block_fp8_linear函数的elif _use_aiter_gfx95:分支中,当input_2d.shape[0] > _ck_safe_m时,强制use_triton=True,从而回退到数值正确的Triton实现。
  3. 保留小M性能:对于低于安全M上限的场景,仍使用更快的CK路径,避免性能退化。
  4. 不改动其他路径:对ROCm ≥ 7.2、非gfx95硬件、或未在字典中列出的形状,逻辑保持不变。
文件 模块 状态 重要度
python/sglang/srt/layers/quantization/fp8_utils.py 量化层 modified 6.51

关键源码片段

python/sglang/srt/layers/quantization/fp8_utils.py core-logic

核心更改文件:定义安全 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 选择具体算子

评论区精华

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

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

风险与影响

  1. 回归风险低:变更仅作用于gfx95 + ROCm < 7.2 + 特定( n, k )形状 + 大M,且保留了小M的CK路径;其他硬件和ROCm版本不受影响。
  2. 性能风险:大M场景切换到Triton实现,可能比CK慢,但保证了正确性;对于Qwen3.5的prefill阶段,优先级是正确性。
  3. 测试覆盖不足: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路径不受影响。

缺少测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论