Prhub

#52293 [ROCm][Perf] Enable fused KDA decode on gfx942 (MI325X)

原始 PR 作者 mpashkovskii 合并时间 2026-08-19 02:01 文件变更 3 提交数 8 评论 7 代码增减 +15 / -14

执行摘要

解锁 MI325X 上 Kimi-K3 fused KDA 解码

PR body 明确指出问题来源:PR #50654 added the fused Kimi-K3 KDA decode kernel ... but gated it to gfx950 (MI355X) in two places, so it never activates on gfx942 (MI325X)。kernel 只使用 CDNA 通用原语(wave64、DPP row/quad permutes、__builtin_nontemporal_load/store、bf16 intrinsic),无 MFMA/FP8/gfx950-only 路径,LDS 约 2.8 KiB、寄存器约 130 VGPR 均满足 CDNA3 限制,因此无需改动 kernel 源码,只需放宽构建与运行时两处 gate。作者还说明当前端到端收益有限的原因:Kimi-K3 decode 步骤由 MXFP4 MoE(int4)与 MLA attention 加 TP8 collectives 主导,KDA 融合节省的约 4.6 ms/token 固定成本占比尚小,待 #50817(MXFP4 MoE Triton a16w4 for gfx942)落地后会更为显著。

值得精读。虽然 diff 极小(+15/-14),但 PR body 的 profiling 分析展示了 launch overhead(每 kernel 约 4 us)在未融合路径中的累积成本,以及"端到端收益受主导算子遮蔽"的判断方法——这是评估 kernel fusion 真实价值的常见陷阱。代码层面,"构建 gate、运行时 gate、测试 gate 三处同步修改"保证了单一事实源,避免 op 已编译但运行时不可达、或测试在目标机器上被误跳过的不一致状态。建议后续平台适配(如 gfx942r1、gfx12xx)复用该模式。

讨论亮点

reviews 中几乎没有实质技术交锋:claude[bot] 指出该 PR 来自 fork,自动 review 被禁用,需维护者用 @claude review 触发一次性评审;hongxiayang 直接 APPROVED,未留评论。因此最有价值的论证集中在 PR body 里作者自己的设计说明:

  • kernel 跨架构复用的依据:kernel 仅依赖 CDNA 通用原语(wave64、DPP row/quad permutes、__builtin_nontemporal_load/store、bf16 intrinsic),不含 MFMA/FP8/gfx950 专属路径,故无需改动 kernel 源码即可解锁 gfx942。
  • 端到端收益的理性预期:作者明确承认当前 -0.77% 的 TPOT 提升很小,因为 decode 成本由 MXFP4 MoE(int4)与 MLA attention 加 TP8 collectives 主导;KDA 融合节省的约 4.6 ms/token 固定成本要待 #50817 压缩 MoE 后才会显著。
  • AI 辅助声明:作者声明该变更借助 Claude 起草,要求 human 逐行复核并实机验证后才可合并;最终合并行为可视为该复核流程已完成。

实现拆解

本 PR 是典型的三层 gate 同步解禁:构建、运行时、测试三处 gate 缺一不可,改动虽小但互相依赖。

  1. 构建 gate——CMakeLists.txt:将 FUSED_KDA_DECODE_HIP_ARCHS 的过滤正则从 gfx950 改为 gfx942|gfx950。这是最底层的开关:gfx942-only 构建此前完全不编译 csrc/libtorch_stable/kimi_k3/fused_kda_decode_kernel_rocm.cu,导致 torch.ops._C.fused_kda_decode op 根本不存在,运行时 gate 永远拿不到 kernel;放宽后 MI325X 构建会生成该 op。注释同步说明 multi-arch 构建下每个匹配 arch 都会输出该源,由运行时 gate 在其余设备上关闭路径。

  2. 运行时 gate——vllm/models/kimi_k3/amd/ops/kda_decode.pyis_fused_kda_decode_supported() 的返回值从 on_gfx950() 变为 on_gfx950() or on_gfx942()。函数开头的 shape、dtype、kernel 可用性 guard(如 num_heads 是否落在 SUPPORTED_NUM_HEADS 内、hasattr(torch.ops._C, "fused_kda_decode"))原样保留,确保不支持的配置仍回退到三次 launch 的 Triton 路径。原 TODO 注释("只测过 gfx950")被替换为对 CDNA3 共享原语的说明。

  3. 测试 gate——tests/models/kimi_k3/test_amd_kda_decode.py:辅助函数从 _on_gfx950() 重命名为 _on_supported_arch()pytestmark 的 skipif 条件扩展为 on_gfx942()on_gfx950(),使 fused kernel 与 Triton 回退链的对照测试在 MI325X 上真正执行而非整模块跳过。

  4. 验证配套:作者在 MI325X 上完成三层验证——op 注册检查(hasattr(torch.ops._C, "fused_kda_decode") 为 True)、kernel 正确性测试(12/12 通过,覆盖 heads {12, 24, 96} × tokens {1, 7, 128} 及 null-slot、no-output-norm、padding 边界)、gsm8k 精度 parity(fused 与 Triton 路径 exact match 均为 1.00)。端到端基准显示 mean TPOT 从 425.95 ms 降至 422.66 ms(-0.77%)。

文件 模块 状态 重要度
vllm/models/kimi_k3/amd/ops/kda_decode.py KDA 解码 modified 4.7
CMakeLists.txt 构建配置 modified 3.26
tests/models/kimi_k3/test_amd_kda_decode.py KDA 测试 modified 5.3

关键符号

is_fused_kda_decode_supported _on_supported_arch

关键源码片段

vllm/models/kimi_k3/amd/ops/kda_decode.py core-logic

运行时 gate 所在文件,`is_fused_kda_decode_supported()` 从仅 gfx950 放宽为 gfx950 或 gfx942,是决定 MI325X 是否走 fused 路径的核心开关;shape/dtype/op 可用性 guard 原样保留。

def is_fused_kda_decode_supported(
    num_heads: int,
    num_tokens: int,
    conv_state_dtype: torch.dtype,
) -> bool:
    """判断 fused decode kernel 能否在当前设备上服务该层配置。"""
    # 形状、dtype 与 kernel 可用性 guard 保持原样:
    # 头数必须落在 SUPPORTED_NUM_HEADS 内、token 数与状态 dtype 受限,
    # 且构建时必须已注册 torch.ops._C.fused_kda_decode,否则直接回退。
    if (
        num_heads not in SUPPORTED_NUM_HEADS
        or num_tokens not in SUPPORTED_NUM_TOKENS
        or conv_state_dtype not in SUPPORTED_STATE_DTYPES
        or not hasattr(torch.ops._C, "fused_kda_decode")
    ):
        return False
​
    # 本 PR 核心改动:架构 gate 从 on_gfx950() 放宽为
    # on_gfx950() or on_gfx942(),让 MI325X(gfx942)也能走 fused 路径。
    from vllm.platforms.rocm import on_gfx942, on_gfx950
​
    # gfx942 与 gfx950 同属 CDNA3,共享 wave64、DPP row/quad permute、
    # __builtin_nontemporal_load/store 与 bf16 原语,kernel 主体无需改动。
    return on_gfx950() or on_gfx942()
CMakeLists.txt configuration

构建 gate 所在文件,HIP arch 过滤从 gfx950 改为 gfx942|gfx950,决定 fused kernel 是否被编译、`torch.ops._C.fused_kda_decode` op 是否注册;没有这一层,运行时 gate 永远拿不到 kernel。

# HIP 侧的 FUSED_KDA_DECODE 块:与 CUDA 侧注册同一个 fused_kda_decode op,
# 共用 VLLM_ENABLE_FUSED_KDA_DECODE 开关。kernel 只依赖 CDNA 通用原语,
# 因此同时为 gfx942 与 gfx950 编译;多 arch 构建中只要列表包含两者之一,
# 就会为每个匹配 arch 输出该源,由运行时 gate(kda_decode.py)在其余设备上关闭。
if(VLLM_GPU_LANG STREQUAL "HIP")
    set(FUSED_KDA_DECODE_HIP_ARCHS ${VLLM_GPU_ARCHES})
    # 原实现仅匹配 gfx950,导致 gfx942-only 构建完全不编译 kernel;
    # 改为 gfx942|gfx950 后,MI325X 构建也会注册 torch.ops._C.fused_kda_decode。
    list(FILTER FUSED_KDA_DECODE_HIP_ARCHS INCLUDE REGEX "gfx942|gfx950")
    if(FUSED_KDA_DECODE_HIP_ARCHS)
        set(FUSED_KDA_DECODE_HIP_SRC
            "csrc/libtorch_stable/kimi_k3/fused_kda_decode_kernel_rocm.cu")
        # 后续将该源追加到编译目标并定义 VLLM_ENABLE_FUSED_KDA_DECODE
    endif()
endif()
tests/models/kimi_k3/test_amd_kda_decode.py test-coverage

测试 skip guard 所在文件,`_on_gfx950()` 重命名为 `_on_supported_arch()` 并覆盖 gfx942/gfx950,确保 fused 与 Triton 回退链的对照正确性测试能在 MI325X 上运行。

def _on_supported_arch() -> bool:
    # 原实现名为 _on_gfx950(),仅覆盖 MI355X;
    # 重命名并扩展后覆盖 gfx942 (MI325X) 与 gfx950 (MI355X)。
    if not current_platform.is_rocm():
        return False
    from vllm.platforms.rocm import on_gfx942, on_gfx950
​
    return on_gfx950() or on_gfx942()
​
​
# skip guard 同步放宽,确保 fused kernel 与 Triton 回退链的对照测试
# 可以在 MI325X 上执行,而不是被整模块跳过。
pytestmark = pytest.mark.skipif(
    not _on_supported_arch(),
    reason="The fused KDA decode kernel is only built for gfx942 / gfx950",
)

评论区精华

fork PR 自动 review 策略 other

claude[bot] 提示该 PR 来自 fork,自动 review 被禁用,维护者可用 @claude review 触发一次性 review。

结论:未触发额外 review,维护者 hongxiayang 直接 APPROVED。 · 已解决

pre-commit 失败与修复 style

mergify[bot] 通知 pre-commit 检查失败,要求运行 pre-commit run --all-files 并提交修复。

结论:作者在后续 merge main 的提交中完成同步,CI 通过后合并。 · 已解决

风险与影响

  • 运行时路径切换风险is_fused_kda_decode_supported() 放宽后,所有运行在 gfx942 上的 Kimi-K3 实例都会自动切换 fused 路径,正确性兜底依赖 12 个 kernel 对照测试与 gsm8k parity(5-shot、limit=100)。样本规模有限,若 kernel 在某 gfx942 变体上因驱动或编译差异出现行为漂移,可能导致静默精度变化而无自动回退机制。
  • 收益依赖后续 PR:当前 TPOT 提升仅 0.77%,KDA 融合节省约 4.6 ms/token 的 kernel 时间在 MoE 主导的成本结构下占比很小;#50817 落地后收益会放大,但这也意味着本次改动的实际价值部分押注在后续优化上。
  • multi-arch 构建行为:CMake 注释明确,多 arch 构建中只要 arch 列表包含 gfx942 或 gfx950 就会为每个匹配 arch 输出该源,运行时 gate 负责在其余设备上关闭;未来新增 CDNA 架构时若未同步更新 kda_decode.py 的 gate,可能误启用或误关闭路径。
  • 资源占用:LDS 约 2.8 KiB/block、VGPR 约 130 均落在 CDNA3 限制内,无显式内存或性能风险;CUDA 路径与其它后端完全不受影响。
  • 用户与系统:MI325X(gfx942)上 Kimi-K3 decode 阶段从约 16 个 kernel 的链式执行(含多次 __amd_rocclr_copyBuffer 与 bf16/FP32 转换 kernel)收敛为单个 kda_decode_fusion_kernel,每层 kernel 时间约 74 us -> 7.77 us;端到端 mean TPOT -0.77%、P99 -0.81%、输出吞吐 +0.75%,TTFT 不变(符合 decode-only kernel 预期)。
  • 团队与演进:本 PR 与 #50654、#50817 构成 Kimi-K3 在 MI325X 上的 kernel fusion 优化主线;同时确立了三层 gate(构建/运行时/测试)同步解锁新架构的范式,可作为其它 CDNA 架构启用的参照。
  • 影响范围:严格限定在 ROCm + Kimi-K3 + gfx942/gfx950 组合;CUDA 路径、其它模型与平台零改动,无接口或配置兼容性变化。
运行时门控放宽需正确性兜底 端到端收益受 MoE 主导、待 #50817 兑现 有限验证:gsm8k 仅 5-shot/100 例

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论