执行摘要
- 一句话:新增ROCm AITER QK Norm+RoPE+KV Cache融合pass
- 推荐动作:值得精读,特别是对 ROCm 平台推理优化和 TorchInductor pattern matcher 感兴趣的开发者。该 PR 展示了如何安全地将多次 kernel launch 融合为单个自定义 op,并处理了不同 attention 后端的布局差异。
功能与动机
从原始大PR #39527拆分而来,旨在优化Qwen3-30B-A3B模型在ROCm上的推理性能。QK Norm+RoPE+KV Cache序列存在多次kernel launch和显存读写,融合为一个kernel可显著降低开销。
实现拆解
- 自定义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。
- 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视图转换和部分旋转维度逻辑。
- 后端抽象接口:在
vllm/v1/attention/backend.py的AttentionImplBase中添加fused_qk_norm_rope_kvcache_supported和do_qk_norm_rope_kvcache_update抽象方法,默认返回False和raise NotImplementedError。
- 具体后端实现:在
rocm_aiter_fa.py和rocm_aiter_unified_attn.py中分别实现上述方法,根据各自KV cache layout(unbind维度不同)拆分cache并调用_aiter_ops.do_qk_norm_rope_kvcache_update。UNIFIED_ATTN路径设置use_shuffle_layout=False。
- 配置与测试:在
vllm/config/vllm.py和compilation.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(模块 编译融合;类别 source;类型 core-logic;符号 fused_qk_norm_rope_and_unified_kv_cache_update_impl, fused_qk_norm_rope_and_unified_kv_cache_update_fake, QkNormRopeKvCachePattern, init): 新增的核心融合pass和custom op,是整个PR的核心逻辑所在。
tests/compile/passes/test_rocm_aiter_qk_norm_rope_kvcache_fusion.py(模块 测试;类别 test;类型 test-coverage;符号 QKNormRoPEKVCacheTestModel, init, build_attn_metadata, forward): 新增的单元测试,验证融合pass的正确性和数值精度。
vllm/_aiter_ops.py(模块 底层算子;类别 source;类型 core-logic;符号 fused_qk_norm_rope_and_cache, do_qk_norm_rope_kvcache_update): 封装了AITER融合内核的调用,是后端实现与底层kernel的桥梁。
vllm/v1/attention/backend.py(模块 attention后端;类别 source;类型 core-logic;符号 fused_qk_norm_rope_kvcache_supported, do_qk_norm_rope_kvcache_update): 定义了融合后端的抽象接口,是整个融合扩展的基础。
vllm/v1/attention/backends/rocm_aiter_unified_attn.py(模块 attention后端;类别 source;类型 core-logic;符号 fused_qk_norm_rope_kvcache_supported, do_qk_norm_rope_kvcache_update): 实现了UNIFIED_ATTN后端的融合方法,处理K/V-first layout。
vllm/v1/attention/backends/rocm_aiter_fa.py(模块 attention后端;类别 source;类型 core-logic;符号 fused_qk_norm_rope_kvcache_supported, do_qk_norm_rope_kvcache_update): 实现了FA后端的融合方法,使用 unbind(1) 分割 KV cache。
vllm/config/vllm.py(模块 配置;类别 source;类型 config-change;符号 enable_qk_norm_rope_kvcache): 配置自动启用逻辑和开关默认值。
vllm/config/compilation.py(模块 配置;类别 source;类型 config-change): 添加 fusion 开关的编译配置。
关键符号: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
新增的核心融合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,
)
评论区精华
关键讨论
- Head size 限制必要性:dllehr-amd 质疑在
qk_norm_rope_fusion.py中添加 head dim 限制是否必要。jhu960213 解释底层 AITER kernel 仅支持 64/128/256 三种 head dim,限制可防止不支持的 dim 导致崩溃。双方同意保留。
- 删除 ROCM_ATTN 空存根:Rohan138 建议不在当前PR中添加 ROCM_ATTN 后端的 noop 实现,留待 Part 2 处理。jhu960213 同意并移除了相关变更,减少了干扰。
- 默认关闭融合:Rohan138 指出 O1 优化级别不应默认开启此融合(依赖 Inductor partition)。jhu960213 调整配置使其默认关闭,用户需显式设置
fuse_qk_norm_rope_kvcache=true。
- SymInt workaround:Rohan138 建议将 SymInt 处理上移至 pattern 预注册阶段。jhu960213 通过自定义 search_pattern 内嵌 SymInt wildcarding 实现,避免了模式匹配时的符号冲突。
- Fusion head size guard 与后端兼容性 (design): 双方同意保留限制,作为防御性措施。
- 删除 ROCM_ATTN 中的 noop 代码 (design): ROCM_ATTN 后端的融合实现移至后续 PR。
- 默认关闭 fuse_qk_norm_rope_kvcache (performance): 融合 pass 默认关闭,用户需显式设置
fuse_qk_norm_rope_kvcache=true 启用。
风险与影响
关联脉络
- PR #39527 [Optimize] Optimize Qwen3-30B with ROCm AITER fusions: 此PR是#39527的拆分Part 1,原始PR包含更广泛的ROCm融合优化。
参与讨论