From e583740e619ec2aa58f7fa1a342d172ed938555c Mon Sep 17 00:00:00 2001 From: Andy Friedrich Date: Thu, 27 Aug 2026 21:13:10 +0000 Subject: [PATCH] [ROCm][Compile] Pattern-match MLA dual RMSNorm + FP8 group quant Rebased onto main after #53540 landed the fused_mla_dual_rms_norm_group_quant custom op. The op registration this PR previously carried is dropped; only the accessor the pattern matcher needs is added here. Adds MLADualRMSGroupQuantPattern, the group-quant sibling of the existing MLADualRMSPerTokenQuantPattern, so the MLA FP8 path picks up the fused AITER kernel through the compile pass rather than a hand-wired call site -- covering DeepSeek-R1 / MLA, which #53540 does not touch. Co-Authored-By: Claude Opus 5 (1M context) Signed-off-by: Andy Friedrich --- .../passes/test_fuse_mla_dual_rms_norm.py | 142 ++++++++++++++++++ vllm/_aiter_ops.py | 4 + .../passes/fusion/rocm_aiter_fusion.py | 138 +++++++++++++++++ 3 files changed, 284 insertions(+) diff --git a/tests/compile/passes/test_fuse_mla_dual_rms_norm.py b/tests/compile/passes/test_fuse_mla_dual_rms_norm.py index 6f20d4e15874..56b664d11a58 100644 --- a/tests/compile/passes/test_fuse_mla_dual_rms_norm.py +++ b/tests/compile/passes/test_fuse_mla_dual_rms_norm.py @@ -154,6 +154,148 @@ def test_fuse_mla_dual_rms_norm( backend.check_after_ops(model.ops_in_model_after()) +GROUP_SIZE = 128 + + +class MLADualRMSNormFp8GroupTestModel(torch.nn.Module): + """ + Minimal model reproducing the FP8 MLA attention path with *group* quant: + linear -> split([q_dim, kv_dim]) + +-- q_c (getitem 0) -> rocm_aiter_rmsnorm_fp8_group_quant -> dequant + +-- kv_lora (getitem 1) -> split([kv_c_dim, k_pe_dim]) + +-- kv_c (getitem 0) -> rms_norm (bf16) + +-- k_pe + """ + + def __init__( + self, + hidden_size: int, + q_dim: int = Q_DIM, + kv_c_dim: int = KV_C_DIM, + k_pe_dim: int = K_PE_DIM, + eps: float = EPS, + group_size: int = GROUP_SIZE, + ): + super().__init__() + self.q_dim = q_dim + self.kv_dim = kv_c_dim + k_pe_dim + self.kv_c_dim = kv_c_dim + self.k_pe_dim = k_pe_dim + self.eps = eps + self.group_size = group_size + + self.proj = torch.nn.Linear(hidden_size, q_dim + self.kv_dim, bias=False) + self.q_weight = torch.nn.Parameter(torch.ones(q_dim)) + self.kv_norm = RMSNorm(kv_c_dim, eps=eps) + + def _dequant(self, x_fp8: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: + # Group: an (M, N // group_size) scale broadcast across each group. + scale = scale.repeat_interleave(self.group_size, dim=-1) + return (x_fp8.to(torch.float32) * scale).to(torch.bfloat16) + + def forward(self, x: torch.Tensor): + # Avoid graph input being a direct arg to a matched pattern node + x = torch.relu(x) + + projected = self.proj(x) + + q_c, kv_lora = projected.split([self.q_dim, self.kv_dim], dim=-1) + kv_c, k_pe = kv_lora.split([self.kv_c_dim, self.k_pe_dim], dim=-1) + + q_fp8, q_scale = torch.ops.vllm.rocm_aiter_rmsnorm_fp8_group_quant( + q_c, self.q_weight, self.eps, self.group_size + ) + kv_normed = self.kv_norm(kv_c) + + return self._dequant(q_fp8, q_scale), kv_normed, k_pe + + def ops_in_model_before(self): + return [ + torch.ops.vllm.rocm_aiter_rmsnorm_fp8_group_quant.default, + torch.ops.vllm_ir.rms_norm.default, + ] + + def ops_in_model_after(self): + return [torch.ops.vllm.fused_mla_dual_rms_norm_group_quant.default] + + +@pytest.mark.parametrize("dtype", [torch.bfloat16]) +@pytest.mark.parametrize("hidden_size", [7168]) +@pytest.mark.skipif( + not is_aiter_found_and_supported(), + reason="Only test on ROCm with AITER installed and supported", +) +def test_fuse_mla_dual_rms_norm_fp8_group( + dtype: torch.dtype, + hidden_size: int, + monkeypatch: pytest.MonkeyPatch, +): + torch._dynamo.reset() + + vllm_config = VllmConfig( + model_config=ModelConfig(dtype=dtype), + compilation_config=CompilationConfig( + mode=CompilationMode.VLLM_COMPILE, + custom_ops=["+rms_norm"], + pass_config=PassConfig( + fuse_mla_dual_rms_norm=True, + eliminate_noops=True, + ), + ), + ) + + with vllm.config.set_current_vllm_config(vllm_config), monkeypatch.context() as m: + from vllm.compilation.passes.fusion.rocm_aiter_fusion import ( + MLADualRMSNormFusionPass, + ) + + torch.set_default_device("cuda") + torch.set_default_dtype(dtype) + torch.manual_seed(42) + + m.setenv("VLLM_ROCM_USE_AITER", "1") + rocm_aiter_ops.refresh_env_variables() + + fusion_pass = MLADualRMSNormFusionPass(vllm_config) + passes = [ + NoOpEliminationPass(vllm_config), + fusion_pass, + PostCleanupPass(vllm_config), + ] + backend = TestBackend(*passes) + model = MLADualRMSNormFp8GroupTestModel(hidden_size) + + x = torch.randn(4, hidden_size) + torch._dynamo.mark_dynamic(x, 0) + + with torch.inference_mode(): + outputs_unfused = model(x) + + model_fused = torch.compile(model, backend=backend) + outputs_fused = model_fused(x) + + q_deq_u, kv_normed_u, k_pe_u = outputs_unfused + q_deq_f, kv_normed_f, k_pe_f = outputs_fused + + torch.testing.assert_close(k_pe_u, k_pe_f, atol=0, rtol=0) + + torch.testing.assert_close(kv_normed_u, kv_normed_f, atol=1e-2, rtol=1e-2) + + E4M3_STEP = 0.125 + exact_frac = (q_deq_u == q_deq_f).float().mean().item() + assert exact_frac > 0.99, ( + f"q: only {exact_frac:.4f} of elements bit-exact; scales likely differ" + ) + torch.testing.assert_close(q_deq_u, q_deq_f, atol=1e-2, rtol=E4M3_STEP) + + assert fusion_pass.matched_count == 1, ( + f"Expected 1 fused pair, got {fusion_pass.matched_count}" + ) + + backend.check_before_ops(model.ops_in_model_before()) + backend.check_after_ops(model.ops_in_model_after()) + + class MLADualRMSNormFp8PerTokenTestModel(torch.nn.Module): """ Minimal model reproducing the FP8 MLA attention path with *per-token* quant: diff --git a/vllm/_aiter_ops.py b/vllm/_aiter_ops.py index 4c215e4fe341..3c0558ea45bc 100644 --- a/vllm/_aiter_ops.py +++ b/vllm/_aiter_ops.py @@ -2411,6 +2411,10 @@ def get_fused_mla_dual_rms_norm_op() -> OpOverload: def get_fused_mla_dual_rms_norm_per_token_quant_op() -> OpOverload: return torch.ops.vllm.fused_mla_dual_rms_norm_per_token_quant.default + @staticmethod + def get_fused_mla_dual_rms_norm_group_quant_op() -> OpOverload: + return torch.ops.vllm.fused_mla_dual_rms_norm_group_quant.default + @staticmethod def fused_qk_rmsnorm_group_quant( q: torch.Tensor, diff --git a/vllm/compilation/passes/fusion/rocm_aiter_fusion.py b/vllm/compilation/passes/fusion/rocm_aiter_fusion.py index 67ed51e89322..2f25a466e9b6 100644 --- a/vllm/compilation/passes/fusion/rocm_aiter_fusion.py +++ b/vllm/compilation/passes/fusion/rocm_aiter_fusion.py @@ -1063,6 +1063,142 @@ def _replacement( return _replacement +class MLADualRMSGroupQuantPattern( + VllmPatternReplacement[ + ..., + tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + ], + ] +): + """ + Fuse the MLA FP8 attention path -- q-latent RMSNorm + FP8 *group* quant + plus kv-latent RMSNorm -- into AITER's ``fused_qk_rmsnorm_group_quant``. + + Group-quant sibling of :class:`MLADualRMSPerTokenQuantPattern`. With a + group-quantized FP8 ``q_b_proj``, the earlier ``RocmAiterRMSNormQuantFusionPass`` + folds the q side into ``rocm_aiter_rmsnorm_fp8_group_quant`` and leaves the + kv side a plain ``vllm_ir.rms_norm``. This pattern matches that asymmetric + pair:: + + gemm -> split_with_sizes([q_dim, kv_dim]) + +-- q_c -> rocm_aiter_rmsnorm_fp8_group_quant -> (q_fp8, q_scale) + +-- kv_lora -> split_with_sizes([kv_c_dim, k_pe_dim]) + +-- kv_c -> vllm_ir.rms_norm -> kv_normed (bf16) + +-- k_pe + """ + + GROUP_QUANT_OP = rocm_aiter_ops.get_rmsnorm_group_fused_quant_op() + FUSED_OP = rocm_aiter_ops.get_fused_mla_dual_rms_norm_group_quant_op() + + def __init__(self, epsilon: float, group_size: int = 128) -> None: + self._epsilon = epsilon + self._group_size = group_size + + def get_inputs(self) -> list[torch.Tensor]: + q_dim, kv_c_dim, k_pe_dim = 256, 128, 64 + return [ + self.empty_bf16(5, q_dim + kv_c_dim + k_pe_dim), + self.empty_bf16(q_dim), + self.empty_bf16(kv_c_dim), + ] + + @property + def pattern( + self, + ) -> Callable[ + ..., + tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + ], + ]: + eps = self._epsilon + group_size = self._group_size + group_quant_op = self.GROUP_QUANT_OP + + def _pattern( + projected: torch.Tensor, + q_weight: torch.Tensor, + kv_weight: torch.Tensor, + ) -> tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + ]: + q_dim = q_weight.shape[0] + kv_dim = projected.shape[-1] - q_dim + kv_c_dim = kv_weight.shape[0] + k_pe_dim = kv_dim - kv_c_dim + q_c, kv_lora = projected.split([q_dim, kv_dim], dim=-1) + kv_c, k_pe = kv_lora.split([kv_c_dim, k_pe_dim], dim=-1) + q_quant = group_quant_op( + x=q_c, + weight=q_weight, + variance_epsilon=eps, + group_size=group_size, + ) + kv_normed = vllm.ir.ops.rms_norm(kv_c, kv_weight, eps) + return q_quant[0], q_quant[1], kv_normed, k_pe + + return _pattern + + @property + def replacement( + self, + ) -> Callable[ + ..., + tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + ], + ]: + eps = self._epsilon + group_size = self._group_size + fused_op = self.FUSED_OP + + def _replacement( + projected: torch.Tensor, + q_weight: torch.Tensor, + kv_weight: torch.Tensor, + ) -> tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + ]: + q_dim = q_weight.shape[0] + kv_dim = projected.shape[-1] - q_dim + kv_c_dim = kv_weight.shape[0] + k_pe_dim = kv_dim - kv_c_dim + q_c, kv_lora = projected.split([q_dim, kv_dim], dim=-1) + kv_c, k_pe = kv_lora.split([kv_c_dim, k_pe_dim], dim=-1) + at = fused_op( + q_c, + q_weight, + kv_c, + kv_weight, + eps, + eps, + group_size, + # (M, N // group_size) scales, matching what the + # rocm_aiter_rmsnorm_fp8_group_quant producer emits. + False, + ) + # q_fp8, q_scale, kv_normed, k_pe + return at[0], at[1], at[2], k_pe + + return _replacement + + class MLADualRMSNormFusionPass(VllmFusionPatternMatcherPass): """ Post-grad PatternMatcher pass that fuses paired q / kv RMS norms in @@ -1079,6 +1215,8 @@ class MLADualRMSNormFusionPass(VllmFusionPatternMatcherPass): def __init__(self, config: VllmConfig) -> None: super().__init__(config, "mla_dual_rms_norm_fusion_pass") + # Scalars are matched exactly; every in-tree producer emits group_size=128. for epsilon in [1e-5, 1e-6]: self.register(MLADualRMSNormPattern(epsilon)) self.register(MLADualRMSPerTokenQuantPattern(epsilon)) + self.register(MLADualRMSGroupQuantPattern(epsilon))