Prhub

#42832 [ROCm][GPT-OSS] Fuse RoPE + static Q FP8 quant on fused RoPE+KV path

原始 PR 作者 akii96 合并时间 2026-06-06 05:22 文件变更 2 提交数 5 评论 17 代码增减 +451 / -1

执行摘要

融合 RoPE 与静态 Q FP8 量化,提升 ROCm GPT-OSS 解码性能

对于 GPT-OSS 风格解码图,RoPE、静态 Q FP8 量化和 KV 缓存更新在热路径上相邻。融合可以消除额外的内核启动和中间内存操作,同时保持注意力所需显式依赖顺序。

该 PR 的设计模式值得精读,特别是 auto_functionalized 在编译 pass 中的正确使用、模式优先级注册机制以及能力检查模式。ROCm 开发者应重点关注,可作为 GPT-OSS 优化路径的参考实现。

讨论亮点
  1. 默认融合范围覆盖:Rohan138 对 vllm/model_executor/models/config.py 中默认 max_token_num 的覆盖提出疑问,认为应在 CLI 中指定。作者移除了该覆盖。

  2. MatcherRotaryEmbedding offsets 参数:Rohan138 指出 matcher_utils.py 中新增的 offsets 参数导致 CI 失败,建议尝试不加。作者确认后完全移除了该修改。

  3. 测试参数化覆盖减少:AI 审查者指出 enable_aiter_triton_rope[True, False] 改为 [True] 减少覆盖。作者修复为 [True, False]

实现拆解

  1. 添加能力检查函数:在 rope_kvcache_fusion.py 中添加 _supports_static_q_fp8_quant_fusion(),通过 current_platform.fp8_dtype()torch.ops._C.static_scaled_fp8_quant 是否存在来判断平台是否支持静态 FP8 量化融合。

  2. 定义新融合模式类 RopeStaticQQuantKVCachePattern:构造函数接收 Attention 层参数,初始化 RoPE 匹配器。get_inputs() 生成带 q_scale 的占位张量。_mk_pattern_with_layer_name_input 定义 pattern(未融合图:RoPE → FP8 quant → KV update)和 replacement(融合图:fused RoPE+KV update → FP8 quant)。

  3. RopeKVCacheFusionPass 中注册:遍历每个 attention 层,通过能力检查后创建模式实例并注册到模式匹配器,注册优先级高于通用 RoPE+KV 模式。

  4. 增加测试覆盖:新增 QKRoPEStaticQKVCacheTestModel 模拟带静态 Q 量化图,test_rope_static_qquant_kvcache_fusion 验证融合后操作数,test_rope_kvcache_fusion_default_keeps_large_ranges_unfused 验证大范围不融合。

  5. 修复与清理:恢复 enable_aiter_triton_rope 参数化为 [True, False],移除 matcher_utils.py 中不必要的 offsets 变更,移除 config.py 中默认 max_token_num 覆盖。

文件 模块 状态 重要度
vllm/compilation/passes/fusion/rope_kvcache_fusion.py 编译融合 modified 8.76
tests/compile/passes/test_rope_kvcache_fusion.py 测试 modified 7.52

关键符号

_supports_static_q_fp8_quant_fusion RopeStaticQQuantKVCachePattern.__init__ RopeStaticQQuantKVCachePattern.get_inputs RopeStaticQQuantKVCachePattern._mk_pattern_with_layer_name_input RopeStaticQQuantKVCachePattern.pattern RopeStaticQQuantKVCachePattern.replacement QKRoPEStaticQKVCacheTestModel.__init__ QKRoPEStaticQKVCacheTestModel.forward QKRoPEStaticQKVCacheTestModel.ops_in_model_before QKRoPEStaticQKVCacheTestModel.ops_in_model_after test_rope_static_qquant_kvcache_fusion test_rope_kvcache_fusion_default_keeps_large_ranges_unfused

关键源码片段

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

核心源码文件,新增 `RopeStaticQQuantKVCachePattern` 类及能力检查函数,实现融合逻辑并注册到编译 pass。

class RopeStaticQQuantKVCachePattern:
    """融合 RoPE + 静态 Q FP8 量化 + KV 缓存更新,保持显式依赖顺序。"""
    FUSED_ROPE_KV_OP = torch.ops.vllm.fused_rope_and_unified_kv_cache_update.default
​
    def _mk_pattern_with_layer_name_input(self, _ln):
        def pattern(qkv, positions, cos_sin_cache, q_scale, layer_name):
            # 未融合的计算图:依次执行 RoPE → FP8 quant → KV cache update
            q, k, v = qkv.split([self.q_size, self.k_size, self.v_size], dim=-1)
            q, k = self.rope_matcher(positions, q, k, cos_sin_cache) # 执行 RoPE
            # 对 Q 做静态 FP8 量化
            q_fp8 = torch.empty(q.shape, device=q.device, dtype=FP8_DTYPE)
            _, q_fp8 = auto_functionalized(
                torch.ops._C.static_scaled_fp8_quant.default,
                result=q_fp8, input=q, scale=q_scale, group_shape=(-1, -1),
            )
            q_view = q_fp8.view(-1, self.num_heads, self.head_size)
            k_view = k.view(-1, self.num_kv_heads, self.head_size)
            v_view = v.view(-1, self.num_kv_heads, self.head_size_v)
            kv_cache_dummy = torch.ops.vllm.unified_kv_cache_update(k_view, v_view, layer_name)
            return kv_cache_dummy, q_view, k_view, v_view
​
        def replacement(qkv, positions, cos_sin_cache, q_scale, layer_name):
            # 融合后的计算图:先融合 RoPE+KV,再对 Q 做 FP8 quant
            q, k, v = qkv.split([self.q_size, self.k_size, self.v_size], dim=-1)
            q_view = q.view(-1, self.num_heads, self.head_size)
            k_view = k.view(-1, self.num_kv_heads, self.head_size)
            v_view = v.view(-1, self.num_kv_heads, self.head_size_v)
            # 执行融合的 RoPE + KV cache update(自动函数化)
            rope_kv_results = auto_functionalized(
                self.FUSED_ROPE_KV_OP,
                result=v_view,
                query=q_view,
                key=k_view,
            )
            q_rope = rope_kv_results[0] # 融合后的 Q
            # 对融合后的 Q 施加静态 FP8 量化
            q_fp8 = torch.empty(q_rope.shape, device=q_rope.device, dtype=FP8_DTYPE)
            _, q_fp8 = auto_functionalized(
                torch.ops._C.static_scaled_fp8_quant.default,
                result=q_fp8, input=q_rope, scale=q_scale, group_shape=(-1, -1),
            )
            return [
                v_view, # kv_cache_dummy(占位)
                q_fp8.view(-1, self.num_heads, self.head_size),
                k_view,
                v_view,
            ]
        return pattern, replacement

评论区精华

默认融合范围覆盖 设计

Rohan138 对 `vllm/model_executor/models/config.py` 中默认 `max_token_num` 的覆盖提出疑问:"Why are we changing the default here? If 16384 gives better perf on gpt-oss, we can specify it in the CLI"

结论:作者移除覆盖,恢复默认行为。 · 已解决

MatcherRotaryEmbedding offsets 参数 正确性

Rohan138 指出 `matcher_utils.py` 中新增的 `offsets` 参数导致 CI 失败:"This change is causing the CI failures",建议尝试不加该参数。

结论:作者确认后完全移除了对 `matcher_utils.py` 的修改,融合仍正常触发。 · 已解决

测试参数化覆盖减少 测试

AI 审查者指出 `enable_aiter_triton_rope` 从 `[True, False]` 改为 `[True]` 显著减少测试覆盖。

结论:作者修复,恢复为 `[True, False]`。 · 已解决

风险与影响

  • 平台特定性:仅对 ROCm 平台生效,能力检查确保其他平台安全跳过。
  • 依赖静态 FP8 量化操作:若 torch.ops._C.static_scaled_fp8_quantfp8_dtype 不可用,融合不会激活,无错误风险。
  • 核心编译 pass 变更:新增的可选融合路径受能力保护,不影响现有行为,但仍需确保回归测试覆盖。
  • 测试覆盖有限:测试仅覆盖特定参数组合(如 head_size=64),未覆盖所有 shape。
  • 用户:ROCm GPT-OSS 用户 decode 阶段自动获得性能提升(吞吐量 +1-4%,TPOT -1-6%),无需配置变更。
  • 系统:不改变公共 API、配置项或依赖。
  • 团队:新增融合模式作为可复用模板,未来可推广到其他模型或操作。
  • 测试:新增单元测试仅在 ROCm CI 中执行,不影响其他 CI 流程。
核心路径变更 ROCm 特定 依赖 AITER 测试覆盖有限

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论