Prhub

#42749 [Model][Hardware][AMD]: Part 1/2 -> Enable e2e QK Norm + RoPE + KV Cache runtime fusion for Qwen3-30B-A3B on ROCM_AITER_FA, and ROCM_AITER_UNIFIED_ATTN

原始 PR 作者 jhu960213 合并时间 2026-07-17 06:39 文件变更 15 提交数 62 评论 93 代码增减 +1361 / -4

执行摘要

新增 ROCm AITER QK Norm+RoPE+KV Cache 融合 pass

从原始大PR #39527拆分而来,旨在优化Qwen3-30B-A3B模型在ROCm上的推理性能。QK Norm+RoPE+KV Cache序列存在多次kernel launch和显存读写,融合为一个kernel可显著降低开销。

值得精读,特别是对 ROCm 平台推理优化和 TorchInductor pattern matcher 感兴趣的开发者。该 PR 展示了如何安全地将多次 kernel launch 融合为单个自定义 op,并处理了不同 attention 后端的布局差异。

讨论亮点

关键讨论

  1. Head size 限制必要性:dllehr-amd 质疑在qk_norm_rope_fusion.py中添加 head dim 限制是否必要。jhu960213 解释底层 AITER kernel 仅支持 64/128/256 三种 head dim,限制可防止不支持的 dim 导致崩溃。双方同意保留。
  2. 删除 ROCM_ATTN 空存根:Rohan138 建议不在当前PR中添加 ROCM_ATTN 后端的 noop 实现,留待 Part 2 处理。jhu960213 同意并移除了相关变更,减少了干扰。
  3. 默认关闭融合:Rohan138 指出 O1 优化级别不应默认开启此融合(依赖 Inductor partition)。jhu960213 调整配置使其默认关闭,用户需显式设置 fuse_qk_norm_rope_kvcache=true
  4. SymInt workaround:Rohan138 建议将 SymInt 处理上移至 pattern 预注册阶段。jhu960213 通过自定义 search_pattern 内嵌 SymInt wildcarding 实现,避免了模式匹配时的符号冲突。

实现拆解

  1. 自定义op与模式匹配:在vllm/compilation/passes/fusion/qk_norm_rope_kvcache_fusion.py中注册Torch custom op fused_qk_norm_rope_and_unified_kv_cache_update,并实现QkNormRopeKvCachePattern类,使用 inductor pattern matcher 匹配 unfused 子图(split→RMSNorm→RoPE→unified_kv_cache_update)并替换为 fused op。
  2. AITER kernel封装:在vllm/_aiter_ops.py中新增fused_qk_norm_rope_and_cache静态方法,调用aiter HIP kernel fused_qk_norm_rope_cache_pts_quant_shuffle;再提供do_qk_norm_rope_kvcache_update高层封装,处理fp8视图转换和部分旋转维度逻辑。
  3. 后端抽象接口:在vllm/v1/attention/backend.pyAttentionImplBase中添加fused_qk_norm_rope_kvcache_supporteddo_qk_norm_rope_kvcache_update抽象方法,默认返回False和raise NotImplementedError。
  4. 具体后端实现:在rocm_aiter_fa.pyrocm_aiter_unified_attn.py中分别实现上述方法,根据各自KV cache layout(unbind维度不同)拆分cache并调用_aiter_ops.do_qk_norm_rope_kvcache_update。UNIFIED_ATTN路径设置use_shuffle_layout=False
  5. 配置与测试:在vllm/config/vllm.pycompilation.py中添加fuse_qk_norm_rope_kvcache配置项及rope_kvcache_fusion_max_token_num限制,自动启用逻辑。新增单元测试文件test_rocm_aiter_qk_norm_rope_kvcache_fusion.py,验证fusion前后输出一致性和K cache正确性。
文件 模块 状态 重要度
vllm/compilation/passes/fusion/qk_norm_rope_kvcache_fusion.py 编译融合 added 9.08
tests/compile/passes/test_rocm_aiter_qk_norm_rope_kvcache_fusion.py 测试 added 7.63
vllm/_aiter_ops.py 底层算子 modified 7.52
vllm/v1/attention/backend.py attention 后端 modified 6.73
vllm/v1/attention/backends/rocm_aiter_unified_attn.py attention 后端 modified 6.86
vllm/v1/attention/backends/rocm_aiter_fa.py attention 后端 modified 6.83
vllm/config/vllm.py 配置 modified 6.29
vllm/config/compilation.py 配置 modified 5.75

关键符号

fused_qk_norm_rope_and_unified_kv_cache_update_impl fused_qk_norm_rope_and_unified_kv_cache_update_fake QkNormRopeKvCachePattern fused_qk_norm_rope_kvcache_supported do_qk_norm_rope_kvcache_update fused_qk_norm_rope_and_cache do_qk_norm_rope_kvcache_update

关键源码片段

vllm/compilation/passes/fusion/qk_norm_rope_kvcache_fusion.py core-logic

新增的核心融合 pass 和 custom op,是整个 PR 的核心逻辑所在。

# vllm/compilation/passes/fusion/qk_norm_rope_kvcache_fusion.py
# 自定义 op 的实现:融合 QK Norm + RoPE + KV cache 更新def fused_qk_norm_rope_and_unified_kv_cache_update_impl(
    q_out: torch.Tensor,
    k_out: torch.Tensor,
    qkv: torch.Tensor,
    positions: torch.Tensor,
    q_weight: torch.Tensor,
    k_weight: torch.Tensor,
    rms_norm_eps: float,
    cos_sin_cache: torch.Tensor,
    is_neox: bool,
    layer_name: str = "",
) -> torch.Tensor:
    """
    执行融合后的 QK Norm + RoPE + KV cache 更新。
    从全局上下文中获取当前 attention layer 和 kv cache,
    然后委托给 backend 的 do_qk_norm_rope_kvcache_update 方法。
    """
    _, attn_layer, kv_cache, layer_slot_mapping = get_attention_context(layer_name)
    if layer_slot_mapping is not None:
        # 正式运行时,调用具体后端实现的融合方法
        attn_layer.impl.do_qk_norm_rope_kvcache_update(
            attn_layer,
            qkv,
            q_out,
            k_out,
            positions,
            q_weight,
            k_weight,
            rms_norm_eps,
            cos_sin_cache,
            is_neox,
            kv_cache,
            layer_slot_mapping,
        )
    else:
        # profiling 阶段:仅填充预分配输出占位
        q_out.zero_()
        k_out.zero_()
    return torch.empty(0, device=qkv.device, dtype=qkv.dtype)# fake 实现用于元数据传播
def fused_qk_norm_rope_and_unified_kv_cache_update_fake(
    q_out: torch.Tensor,
    k_out: torch.Tensor,
    qkv: torch.Tensor,
    positions: torch.Tensor,
    q_weight: torch.Tensor,
    k_weight: torch.Tensor,
    rms_norm_eps: float,
    cos_sin_cache: torch.Tensor,
    is_neox: bool,
    layer_name: str = "",
) -> torch.Tensor:
    return torch.empty(0, device=qkv.device, dtype=qkv.dtype)# 注册自定义 op
direct_register_custom_op(
    op_name="fused_qk_norm_rope_and_unified_kv_cache_update",
    op_func=fused_qk_norm_rope_and_unified_kv_cache_update_impl,
    mutates_args=["q_out", "k_out"],
    fake_impl=fused_qk_norm_rope_and_unified_kv_cache_update_fake,
)

评论区精华

Fusion head size guard 与后端兼容性 设计

dllehr-amd 质疑在 `qk_norm_rope_fusion.py` 中添加 head dim 限制的必要性。jhu960213 解释底层 AITER kernel 仅支持 64/128/256,限制可防止不支持的 dim 导致 kernel 崩溃。

结论:双方同意保留限制,作为防御性措施。 · 已解决

删除 ROCM_ATTN 中的 noop 代码 设计

Rohan138 建议移除 ROCM_ATTN 后端的空实现存根,留待 Part 2 再添加。jhu960213 同意并移除,避免引入无用代码。

结论:ROCM_ATTN 后端的融合实现移至后续 PR。 · 已解决

默认关闭 fuse_qk_norm_rope_kvcache 性能

Rohan138 指出 O1 优化级别不应默认开启此融合,因其依赖 Inductor partition。jhu960213 调整配置使其默认关闭。

结论:融合 pass 默认关闭,用户需显式设置 `fuse_qk_norm_rope_kvcache=true` 启用。 · 已解决

风险与影响

  1. 回归风险:新 fusion pass 可能与其他编译 pass(如 RopeKVCacheFusionPass)冲突,导致意外模式被匹配或漏匹配。已在测试中覆盖非融合场景。
  2. 性能影响:融合 kernel 仅支持有限 head dim,不匹配时静默跳过,不会造成错误但无法获得加速。用户需注意 rope_kvcache_fusion_max_token_num 限制。
  3. 兼容性:该融合专用 ROCm AITER 后端,NVIDIA 或其他平台上不会开启。但不影响其他后端的正常路径。
  4. 配置复杂性:用户需通过 --compilation-config 传递 JSON 配置开启,使用门槛较高。

用户影响:Qwen3-30B-A3B 模型在 ROCm 平台上的推理延迟和吞吐量有显著改善(具体数据参考 PR 截图)。用户需手动设置编译配置以启用融合。
系统影响:新增编译 pass 和 custom op,但不影响未使用融合的模型。
团队影响:PR 被拆分为 Part 1/2,Part 1 聚焦 AITER 后端,Part 2 将添加 ROCM_ATTN 支持。降低了单次 review 负担。

ROCm 专用 新 fusion pass 依赖 Inductor 分区 需要手动启用

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论