Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
142 changes: 142 additions & 0 deletions tests/compile/passes/test_fuse_mla_dual_rms_norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
4 changes: 4 additions & 0 deletions vllm/_aiter_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -2449,6 +2449,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,
Expand Down
138 changes: 138 additions & 0 deletions vllm/compilation/passes/fusion/rocm_aiter_fusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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))
Loading