From 9d4dc4ca2fefae1a19648d7989f1dffb351e1b29 Mon Sep 17 00:00:00 2001 From: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Date: Wed, 17 Jun 2026 04:49:47 +0800 Subject: [PATCH] [Kernel] Support GLM-5 dimensions for TRT-LLM ragged MLA prefill (#43525) Signed-off-by: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com> Signed-off-by: Matthew Bonanni Co-authored-by: Matthew Bonanni --- docs/design/attention_backends.md | 6 +- tests/v1/attention/test_mla_backends.py | 12 +++- .../v1/attention/test_mla_prefill_registry.py | 2 - .../v1/attention/test_mla_prefill_selector.py | 69 ++++++++++++------- .../generate_attention_backend_docs.py | 66 +++++++++++++++--- .../v1/attention/backends/mla/prefill/base.py | 27 ++++++-- .../backends/mla/prefill/flashinfer.py | 15 +++- .../backends/mla/prefill/selector.py | 45 ++++++------ .../backends/mla/prefill/tokenspeed_mla.py | 15 +++- .../backends/mla/prefill/trtllm_ragged.py | 20 +++++- 10 files changed, 200 insertions(+), 77 deletions(-) diff --git a/docs/design/attention_backends.md b/docs/design/attention_backends.md index 6f8feeb887a..8a261c502ed 100644 --- a/docs/design/attention_backends.md +++ b/docs/design/attention_backends.md @@ -214,9 +214,9 @@ hardware and configuration. | Backend | Description | Dtypes | Compute Cap. | Notes | | ------- | ----------- | ------ | ------------ | ----- | | `FLASH_ATTN`‡ | FlashAttention varlen (FA2/FA3/FA4) | fp16, bf16 | Any | FA4 on SM100+, FA3 on SM90, FA2 otherwise | -| `TRTLLM_RAGGED` | TensorRT-LLM ragged attention | fp16, bf16 | 10.x | DeepSeek R1 dims only | -| `FLASHINFER` | FlashInfer CUTLASS backend | fp16, bf16 | 10.x | DeepSeek R1 dims only | -| `TOKENSPEED_MLA` | | fp16, bf16 | 10.x | DeepSeek R1 dims only | +| `TRTLLM_RAGGED` | TensorRT-LLM ragged attention | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) or (qk_nope_head_dim=192, qk_rope_head_dim=64, v_head_dim=256) only | +| `FLASHINFER` | FlashInfer CUTLASS backend | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) only | +| `TOKENSPEED_MLA` | | fp16, bf16 | 10.x | (qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128) only | > **‡** Automatic selection tries FlashAttention first. On Blackwell > (SM100), the fallback order is TRT-LLM Ragged, FlashInfer, then diff --git a/tests/v1/attention/test_mla_backends.py b/tests/v1/attention/test_mla_backends.py index 109e56cb383..1ef4f96617e 100644 --- a/tests/v1/attention/test_mla_backends.py +++ b/tests/v1/attention/test_mla_backends.py @@ -765,7 +765,8 @@ def test_backend_correctness( if not backends_to_test: pytest.skip(f"No backends support kv_cache_dtype={kv_cache_dtype}") - # Skip prefill backends that can't satisfy capability/deps/R1 constraints. + # Skip prefill backends that can't satisfy capability/deps/dimension constraints. + from vllm.v1.attention.backends.mla.prefill.base import MLADimensions from vllm.v1.attention.backends.mla.prefill.selector import ( MLAPrefillSelectorConfig, ) @@ -773,7 +774,14 @@ def test_backend_correctness( try: prefill_invalid_reasons = prefill_backend.get_class().validate_configuration( current_platform.get_device_capability(), - MLAPrefillSelectorConfig(dtype=torch.bfloat16, is_r1_compatible=True), + MLAPrefillSelectorConfig( + dtype=torch.bfloat16, + mla_dimensions=MLADimensions( + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + ), + ), ) except ImportError: prefill_invalid_reasons = ["ImportError"] diff --git a/tests/v1/attention/test_mla_prefill_registry.py b/tests/v1/attention/test_mla_prefill_registry.py index 4b701b8c13b..668c17c3f55 100644 --- a/tests/v1/attention/test_mla_prefill_registry.py +++ b/tests/v1/attention/test_mla_prefill_registry.py @@ -16,7 +16,6 @@ class CustomMLAPrefillBackend(MLAPrefillBackend): """Mock custom MLA prefill backend for testing.""" supported_dtypes = [torch.bfloat16, torch.float16] - requires_r1_mla_dimensions = False @staticmethod def get_name() -> str: @@ -83,7 +82,6 @@ def test_register_custom_backend_as_decorator(): @register_mla_prefill_backend(MLAPrefillBackendEnum.CUSTOM) class DecoratedPrefillBackend(MLAPrefillBackend): supported_dtypes = [torch.bfloat16] - requires_r1_mla_dimensions = False @staticmethod def get_name() -> str: diff --git a/tests/v1/attention/test_mla_prefill_selector.py b/tests/v1/attention/test_mla_prefill_selector.py index d5c80c80c03..54e68e03f26 100644 --- a/tests/v1/attention/test_mla_prefill_selector.py +++ b/tests/v1/attention/test_mla_prefill_selector.py @@ -9,12 +9,12 @@ import torch from vllm.config import AttentionConfig, ModelConfig, VllmConfig from vllm.platforms.interface import DeviceCapability +from vllm.v1.attention.backends.mla.prefill.base import MLADimensions from vllm.v1.attention.backends.mla.prefill.registry import MLAPrefillBackendEnum from vllm.v1.attention.backends.mla.prefill.selector import ( MLAPrefillSelectorConfig, _auto_select_mla_prefill_backend, get_mla_prefill_backend, - is_deepseek_r1_mla_compatible, ) @@ -149,11 +149,14 @@ class TestAutoSelectMLAPrefillBackend: """Tests for fallback and error paths in auto-selection.""" def test_blackwell_falls_back_to_trtllm(self): - vllm_config = _make_vllm_config() capability = DeviceCapability(major=10, minor=0) selector_config = MLAPrefillSelectorConfig( dtype=torch.bfloat16, - is_r1_compatible=is_deepseek_r1_mla_compatible(vllm_config), + mla_dimensions=MLADimensions( + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + ), ) try: @@ -177,11 +180,14 @@ class TestAutoSelectMLAPrefillBackend: assert backend.get_name() == "TRTLLM_RAGGED" def test_all_fail_raises_error(self): - vllm_config = _make_vllm_config() capability = DeviceCapability(major=10, minor=0) selector_config = MLAPrefillSelectorConfig( dtype=torch.bfloat16, - is_r1_compatible=is_deepseek_r1_mla_compatible(vllm_config), + mla_dimensions=MLADimensions( + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + ), ) def mock_get_class(backend_enum): # noqa: ARG001 @@ -201,28 +207,26 @@ class TestAutoSelectMLAPrefillBackend: class TestBackendValidation: """Tests for backend validation logic.""" - def test_r1_dimension_requirement(self): + def test_backend_supported_dimension_validation(self): try: from vllm.v1.attention.backends.mla.prefill.flashinfer import ( FlashInferPrefillBackend, ) + from vllm.v1.attention.backends.mla.prefill.trtllm_ragged import ( + TrtllmRaggedPrefillBackend, + ) except ImportError: - pytest.skip("FlashInfer prefill backend not available") + pytest.skip("MLA prefill backend not available") return - assert FlashInferPrefillBackend.requires_r1_mla_dimensions is True - - vllm_config = _make_vllm_config( - model_config=_make_mock_model_config( - qk_nope_head_dim=128, - qk_rope_head_dim=64, - v_head_dim=128, - ) - ) capability = DeviceCapability(major=10, minor=0) selector_config = MLAPrefillSelectorConfig( dtype=torch.bfloat16, - is_r1_compatible=is_deepseek_r1_mla_compatible(vllm_config), + mla_dimensions=MLADimensions( + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + ), ) with patch.object(FlashInferPrefillBackend, "is_available", return_value=True): @@ -232,16 +236,13 @@ class TestBackendValidation: ) assert len(invalid_reasons) == 0 - vllm_config_invalid = _make_vllm_config( - model_config=_make_mock_model_config( + selector_config_invalid = MLAPrefillSelectorConfig( + dtype=torch.bfloat16, + mla_dimensions=MLADimensions( qk_nope_head_dim=64, qk_rope_head_dim=64, v_head_dim=128, - ) - ) - selector_config_invalid = MLAPrefillSelectorConfig( - dtype=torch.bfloat16, - is_r1_compatible=is_deepseek_r1_mla_compatible(vllm_config_invalid), + ), ) with patch.object(FlashInferPrefillBackend, "is_available", return_value=True): @@ -250,7 +251,25 @@ class TestBackendValidation: selector_config_invalid, ) assert len(invalid_reasons) == 1 - assert "DeepSeek R1 MLA dimensions" in invalid_reasons[0] + assert "supported MLA dimensions" in invalid_reasons[0] + + selector_config_glm5 = MLAPrefillSelectorConfig( + dtype=torch.bfloat16, + mla_dimensions=MLADimensions( + qk_nope_head_dim=192, + qk_rope_head_dim=64, + v_head_dim=256, + ), + ) + + with patch.object( + TrtllmRaggedPrefillBackend, "is_available", return_value=True + ): + invalid_reasons = TrtllmRaggedPrefillBackend.validate_configuration( + capability, + selector_config_glm5, + ) + assert invalid_reasons == [] class TestMLAPrefillBackendParsing: diff --git a/tools/pre_commit/generate_attention_backend_docs.py b/tools/pre_commit/generate_attention_backend_docs.py index 7bc7f1de4b3..91720911a63 100644 --- a/tools/pre_commit/generate_attention_backend_docs.py +++ b/tools/pre_commit/generate_attention_backend_docs.py @@ -384,6 +384,49 @@ def parse_mla_prefill_priorities() -> dict[str, list[str]]: return priorities +def parse_mla_dimensions_call(node: ast.AST) -> str | None: + """Parse an MLADimensions(...) call into a compact display string.""" + if not isinstance(node, ast.Call): + return None + + func = node.func + if not isinstance(func, ast.Name) or func.id != "MLADimensions": + return None + + dimensions: dict[str, int] = {} + for keyword in node.keywords: + if ( + keyword.arg is not None + and isinstance(keyword.value, ast.Constant) + and isinstance(keyword.value.value, int) + ): + dimensions[keyword.arg] = keyword.value.value + + qk_nope_head_dim = dimensions.get("qk_nope_head_dim") + qk_rope_head_dim = dimensions.get("qk_rope_head_dim") + v_head_dim = dimensions.get("v_head_dim") + if qk_nope_head_dim is None or qk_rope_head_dim is None or v_head_dim is None: + return None + + return ( + f"(qk_nope_head_dim={qk_nope_head_dim}, " + f"qk_rope_head_dim={qk_rope_head_dim}, v_head_dim={v_head_dim})" + ) + + +def parse_supported_mla_dimensions(node: ast.AST | None) -> list[str]: + """Parse a supported_mla_dimensions class variable.""" + if not isinstance(node, ast.List): + return [] + + supported_dimensions = [] + for element in node.elts: + dimensions = parse_mla_dimensions_call(element) + if dimensions is not None: + supported_dimensions.append(dimensions) + return supported_dimensions + + def parse_mla_prefill_backend_file(class_path: str) -> dict[str, Any] | None: """Parse a single MLA prefill backend file to extract its properties. @@ -409,20 +452,20 @@ def parse_mla_prefill_backend_file(class_path: str) -> dict[str, Any] | None: info: dict[str, Any] = { "compute_capability": "Any", - "requires_r1_dims": False, + "supported_mla_dimensions": [], "dtypes": "fp16, bf16", # Default from base class } # Parse class variables for item in class_node.body: - if isinstance(item, ast.Assign): - for target in item.targets: - if ( - isinstance(target, ast.Name) - and target.id == "requires_r1_mla_dimensions" - and isinstance(item.value, ast.Constant) - ): - info["requires_r1_dims"] = item.value.value + if ( + isinstance(item, ast.AnnAssign) + and isinstance(item.target, ast.Name) + and item.target.id == "supported_mla_dimensions" + ): + info["supported_mla_dimensions"] = parse_supported_mla_dimensions( + item.value + ) # Parse supported_dtypes class variable if ( @@ -515,8 +558,9 @@ def parse_mla_prefill_backends() -> list[dict[str, Any]]: marker = "‡" notes = "" - if backend_info.get("requires_r1_dims"): - notes = "DeepSeek R1 dims only" + supported_mla_dimensions = backend_info.get("supported_mla_dimensions", []) + if supported_mla_dimensions: + notes = " or ".join(supported_mla_dimensions) + " only" elif backend_name == "FLASH_ATTN": notes = "FA4 on SM100+, FA3 on SM90, FA2 otherwise" diff --git a/vllm/v1/attention/backends/mla/prefill/base.py b/vllm/v1/attention/backends/mla/prefill/base.py index 3a73fd907c2..ff478aec4ad 100644 --- a/vllm/v1/attention/backends/mla/prefill/base.py +++ b/vllm/v1/attention/backends/mla/prefill/base.py @@ -3,6 +3,7 @@ """Abstract base class for MLA prefill backends.""" from abc import ABC, abstractmethod +from dataclasses import dataclass from typing import TYPE_CHECKING, ClassVar import torch @@ -19,6 +20,20 @@ if TYPE_CHECKING: ) +@dataclass(frozen=True, kw_only=True) +class MLADimensions: + qk_nope_head_dim: int + qk_rope_head_dim: int + v_head_dim: int + + def __str__(self) -> str: + return ( + f"(qk_nope_head_dim={self.qk_nope_head_dim}, " + f"qk_rope_head_dim={self.qk_rope_head_dim}, " + f"v_head_dim={self.v_head_dim})" + ) + + class MLAPrefillBackend(ABC): """Abstract base class for MLA prefill backends.""" @@ -26,7 +41,7 @@ class MLAPrefillBackend(ABC): torch.float16, torch.bfloat16, ] - requires_r1_mla_dimensions: ClassVar[bool] = False + supported_mla_dimensions: ClassVar[list[MLADimensions]] = [] @staticmethod @abstractmethod @@ -71,10 +86,14 @@ class MLAPrefillBackend(ABC): if not cls.is_available(): invalid_reasons.append("required dependencies not available") - if cls.requires_r1_mla_dimensions and not selector_config.is_r1_compatible: + if ( + cls.supported_mla_dimensions + and selector_config.mla_dimensions not in cls.supported_mla_dimensions + ): + supported = ", ".join(str(dims) for dims in cls.supported_mla_dimensions) invalid_reasons.append( - "model does not have DeepSeek R1 MLA dimensions " - "(qk_nope_head_dim=128, qk_rope_head_dim=64, v_head_dim=128)" + "Model does not have supported MLA dimensions " + f"(got {selector_config.mla_dimensions}; supported: {supported})" ) return invalid_reasons diff --git a/vllm/v1/attention/backends/mla/prefill/flashinfer.py b/vllm/v1/attention/backends/mla/prefill/flashinfer.py index 92e26a0768e..557c16f97f0 100644 --- a/vllm/v1/attention/backends/mla/prefill/flashinfer.py +++ b/vllm/v1/attention/backends/mla/prefill/flashinfer.py @@ -2,12 +2,15 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """FlashInfer backend for MLA prefill.""" -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, ClassVar import torch import vllm.envs as envs -from vllm.v1.attention.backends.mla.prefill.base import MLAPrefillBackend +from vllm.v1.attention.backends.mla.prefill.base import ( + MLADimensions, + MLAPrefillBackend, +) from vllm.v1.attention.backends.utils import ( PerLayerParameters, get_per_layer_parameters, @@ -33,7 +36,13 @@ _DEFAULT_NUM_CHUNKS = 32 class FlashInferPrefillBackend(MLAPrefillBackend): """FlashInfer backend for MLA prefill.""" - requires_r1_mla_dimensions = True + supported_mla_dimensions: ClassVar[list[MLADimensions]] = [ + MLADimensions( + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + ), + ] @staticmethod def get_name() -> str: diff --git a/vllm/v1/attention/backends/mla/prefill/selector.py b/vllm/v1/attention/backends/mla/prefill/selector.py index 816f4fd4b73..e100c098acb 100644 --- a/vllm/v1/attention/backends/mla/prefill/selector.py +++ b/vllm/v1/attention/backends/mla/prefill/selector.py @@ -13,6 +13,7 @@ import torch from vllm.logger import init_logger from vllm.platforms.interface import DeviceCapability +from vllm.v1.attention.backends.mla.prefill.base import MLADimensions from vllm.v1.attention.backends.mla.prefill.registry import MLAPrefillBackendEnum if TYPE_CHECKING: @@ -31,24 +32,17 @@ class MLAPrefillSelectorConfig(NamedTuple): """ dtype: torch.dtype - is_r1_compatible: bool + mla_dimensions: MLADimensions = MLADimensions( + qk_nope_head_dim=0, + qk_rope_head_dim=0, + v_head_dim=0, + ) - -def is_deepseek_r1_mla_compatible(vllm_config: "VllmConfig") -> bool: - """Check if model has DeepSeek R1 compatible MLA dimensions. - - DeepSeek R1 MLA dimensions are: - - qk_nope_head_dim = 128 - - qk_rope_head_dim = 64 - - v_head_dim = 128 - """ - if vllm_config.model_config is None: - return False - hf_text_config = vllm_config.model_config.hf_text_config - qk_nope_head_dim = getattr(hf_text_config, "qk_nope_head_dim", 1) - qk_rope_head_dim = getattr(hf_text_config, "qk_rope_head_dim", 1) - v_head_dim = getattr(hf_text_config, "v_head_dim", 1) - return qk_nope_head_dim == 128 and qk_rope_head_dim == 64 and v_head_dim == 128 + def __repr__(self): + return ( + f"MLAPrefillSelectorConfig(dtype={self.dtype}, " + f"mla_dimensions={self.mla_dimensions})" + ) def _get_mla_prefill_backend_priorities( @@ -101,10 +95,19 @@ def get_mla_prefill_backend( attention_config = vllm_config.attention_config - selector_config = MLAPrefillSelectorConfig( - dtype=vllm_config.model_config.dtype, - is_r1_compatible=is_deepseek_r1_mla_compatible(vllm_config), - ) + model_config = vllm_config.model_config + if model_config is None: + selector_config = MLAPrefillSelectorConfig(dtype=torch.get_default_dtype()) + else: + hf_text_config = model_config.hf_text_config + selector_config = MLAPrefillSelectorConfig( + dtype=model_config.dtype, + mla_dimensions=MLADimensions( + qk_nope_head_dim=getattr(hf_text_config, "qk_nope_head_dim", 0), + qk_rope_head_dim=getattr(hf_text_config, "qk_rope_head_dim", 0), + v_head_dim=getattr(hf_text_config, "v_head_dim", 0), + ), + ) if attention_config.mla_prefill_backend is not None: selected_backend = attention_config.mla_prefill_backend diff --git a/vllm/v1/attention/backends/mla/prefill/tokenspeed_mla.py b/vllm/v1/attention/backends/mla/prefill/tokenspeed_mla.py index 21f1c25be77..1f041f37317 100644 --- a/vllm/v1/attention/backends/mla/prefill/tokenspeed_mla.py +++ b/vllm/v1/attention/backends/mla/prefill/tokenspeed_mla.py @@ -2,11 +2,14 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """TokenSpeed CuTe DSL backend for MLA prefill.""" -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, ClassVar import torch -from vllm.v1.attention.backends.mla.prefill.base import MLAPrefillBackend +from vllm.v1.attention.backends.mla.prefill.base import ( + MLADimensions, + MLAPrefillBackend, +) if TYPE_CHECKING: from vllm.config import VllmConfig @@ -19,7 +22,13 @@ if TYPE_CHECKING: class TokenspeedMLAPrefillBackend(MLAPrefillBackend): """TokenSpeed CuTe DSL backend for MLA prefill.""" - requires_r1_mla_dimensions = True + supported_mla_dimensions: ClassVar[list[MLADimensions]] = [ + MLADimensions( + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + ), + ] @staticmethod def get_name() -> str: diff --git a/vllm/v1/attention/backends/mla/prefill/trtllm_ragged.py b/vllm/v1/attention/backends/mla/prefill/trtllm_ragged.py index 1d4102ef3c2..90f721272dc 100644 --- a/vllm/v1/attention/backends/mla/prefill/trtllm_ragged.py +++ b/vllm/v1/attention/backends/mla/prefill/trtllm_ragged.py @@ -2,12 +2,15 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project """TRT-LLM Ragged backend for MLA prefill.""" -from typing import TYPE_CHECKING +from typing import TYPE_CHECKING, ClassVar import torch import vllm.envs as envs -from vllm.v1.attention.backends.mla.prefill.base import MLAPrefillBackend +from vllm.v1.attention.backends.mla.prefill.base import ( + MLADimensions, + MLAPrefillBackend, +) from vllm.v1.worker.workspace import current_workspace_manager if TYPE_CHECKING: @@ -21,7 +24,18 @@ if TYPE_CHECKING: class TrtllmRaggedPrefillBackend(MLAPrefillBackend): """TRT-LLM Ragged backend for MLA prefill.""" - requires_r1_mla_dimensions = True + supported_mla_dimensions: ClassVar[list[MLADimensions]] = [ + MLADimensions( + qk_nope_head_dim=128, + qk_rope_head_dim=64, + v_head_dim=128, + ), + MLADimensions( + qk_nope_head_dim=192, + qk_rope_head_dim=64, + v_head_dim=256, + ), + ] @staticmethod def get_name() -> str: