diff --git a/docs_new/docs/advanced_features/attention_backend.mdx b/docs_new/docs/advanced_features/attention_backend.mdx index d9fea017c..293859818 100644 --- a/docs_new/docs/advanced_features/attention_backend.mdx +++ b/docs_new/docs/advanced_features/attention_backend.mdx @@ -235,6 +235,15 @@ The support matrix is split into two parts: MHA (standard attention) and MLA (mu ✅ ❌ + + **CuteDSL MLA (Blackwell)** + 32 or 64 + ✅ + ❌ + ✅ + ✅ + ❌ + **FA3 (FlashAttention 3)** n/a @@ -307,13 +316,14 @@ Speculative decoding topk: `topk` is the number of draft tokens sampled per step Page size controls how many tokens are grouped into a KV cache block. For the prefix cache to take effect, the number of tokens must fill at least one complete page. For example, if your prompt is only 32 tokens and `page_size = 64`, it won't fill a complete page and cannot be matched in the prefix cache (pages cannot be padded). With 65 tokens and `page_size = 64`, only the first page of 64 tokens will be cached and matched; the remaining 1 token is discarded. Use `page_size = 1` for maximum prefix reuse (token-level matching). Note that higher page sizes generally improve attention kernel performance, so prefer `page_size > 1` when prefix cache reuse is not critical. -Many backends that do not natively operate on pages can emulate `page_size > 1` at the wrapper layer by expanding page tables to per-token indices. The "Page Size > 1 (native)" column indicates true in-kernel paging. Some backends require fixed native page sizes and cannot be reduced/emulated differently: TRTLLM MHA (16/32/64), TRTLLM MLA (32/64), FlashMLA (64), Cutlass MLA (128), Ascend (128). +Many backends that do not natively operate on pages can emulate `page_size > 1` at the wrapper layer by expanding page tables to per-token indices. The "Page Size > 1 (native)" column indicates true in-kernel paging. Some backends require fixed native page sizes and cannot be reduced/emulated differently: TRTLLM MHA (16/32/64), TRTLLM MLA (32/64), CuteDSL MLA (32/64), FlashMLA (64), Cutlass MLA (128), Ascend (128). MLA page-size constraints: - FlashInfer MLA: page_size = 1. - FlashMLA: page_size = 64. - Cutlass MLA: page_size = 128. - TRTLLM MLA: page_size ∈ {32, 64}. +- CuteDSL MLA: page_size ∈ {32, 64} (decode-only; prefill falls back to `trtllm_mla` when unset). ### GDN Attention Backends diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index 6d86c0189..0efb24bfe 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -74,6 +74,15 @@ def create_tokenspeed_mla_backend(runner): return TokenspeedMLABackend(runner) +@register_attention_backend("cutedsl_mla") +def create_cutedsl_mla_backend(runner): + if not runner.use_mla_backend: + raise ValueError("cutedsl_mla backend can only be used with MLA models.") + from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLABackend + + return TRTLLMMLABackend(runner, backend="cute-dsl") + + @register_attention_backend("aiter") def create_aiter_backend(runner): from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 4a5fed11c..493185647 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -229,6 +229,11 @@ def _quantize_fp8_qkv(q, k, v, layer): global_zero_init_workspace_buffer = None +# cute-dsl needs its own workspace: it overwrites the buffer with split-KV +# partials, which corrupts the trtllm-gen multiCtasKv counters that rely on the +# zero-init buffer (they share it under attention-backend=cutedsl_mla, where +# draft-extend falls back to trtllm-gen) and deadlocks the reduction. +global_cute_dsl_workspace_buffer = None @dataclass @@ -263,6 +268,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): skip_prefill: bool = False, kv_indptr_buf: Optional[torch.Tensor] = None, q_indptr_decode_buf: Optional[torch.Tensor] = None, + backend: str = "trtllm-gen", ): super().__init__( model_runner, @@ -286,6 +292,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): self.kv_cache_dim = self.kv_lora_rank + self.qk_rope_head_dim # Runtime parameters + self.backend = backend self.scaling = config.scaling self.data_type = model_runner.kv_cache_dtype self.q_data_type = model_runner.dtype @@ -294,14 +301,26 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): # Workspace allocation self.workspace_size = DEFAULT_WORKSPACE_SIZE_MB * 1024 * 1024 - global global_zero_init_workspace_buffer - if global_zero_init_workspace_buffer is None: - global_zero_init_workspace_buffer = torch.zeros( - self.workspace_size, - dtype=torch.uint8, - device=model_runner.device, - ) - self.workspace_buffer = global_zero_init_workspace_buffer + if self.backend == "cute-dsl": + # Separate buffer from trtllm-gen (see note above); safe to share + # among cute-dsl instances. + global global_cute_dsl_workspace_buffer + if global_cute_dsl_workspace_buffer is None: + global_cute_dsl_workspace_buffer = torch.zeros( + self.workspace_size, + dtype=torch.int8, + device=model_runner.device, + ) + self.workspace_buffer = global_cute_dsl_workspace_buffer + else: + global global_zero_init_workspace_buffer + if global_zero_init_workspace_buffer is None: + global_zero_init_workspace_buffer = torch.zeros( + self.workspace_size, + dtype=torch.int8, + device=model_runner.device, + ) + self.workspace_buffer = global_zero_init_workspace_buffer # CUDA graph state self.decode_cuda_graph_metadata = {} @@ -807,6 +826,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): seq_lens_i32 = ( seq_lens if seq_lens.dtype == torch.int32 else seq_lens.to(torch.int32) ) + extra_kwargs = {"backend": self.backend} if self.backend != "trtllm-gen" else {} return flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla( query=query, kv_cache=kv_cache, @@ -819,6 +839,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): max_seq_len=max_seq_len, bmm1_scale=bmm1_scale, skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), + **extra_kwargs, ) def _run_prefill_kernel( @@ -1224,7 +1245,11 @@ class TRTLLMMLAMultiStepDraftBackend(FlashInferMLAMultiStepDraftBackend): """Multi-step draft backend for TRT-LLM MLA used by EAGLE.""" def __init__( - self, model_runner: "ModelRunner", topk: int, speculative_num_steps: int + self, + model_runner: "ModelRunner", + topk: int, + speculative_num_steps: int, + backend: str = "trtllm-gen", ): super().__init__(model_runner, topk, speculative_num_steps) @@ -1234,6 +1259,7 @@ class TRTLLMMLAMultiStepDraftBackend(FlashInferMLAMultiStepDraftBackend): skip_prefill=True, kv_indptr_buf=self.kv_indptr[i], q_indptr_decode_buf=self.q_indptr_decode, + backend=backend, ) def init_forward_metadata(self, forward_batch: ForwardBatch): diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 3e124ce8f..6d9db720a 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -250,6 +250,7 @@ MLA_ATTENTION_BACKENDS = [ "fa4", "triton", "flashmla", + "cutedsl_mla", "cutlass_mla", "trtllm_mla", "tokenspeed_mla", @@ -264,6 +265,7 @@ CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS = [ "fa3", "fa4", "flashmla", + "cutedsl_mla", "cutlass_mla", "trtllm_mla", "tokenspeed_mla", diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 364f065b4..f366159ee 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -698,7 +698,8 @@ class DeepseekMLAForwardMixin: ) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn return ( - self.current_attention_backend in ("trtllm_mla", "tokenspeed_mla") + self.current_attention_backend + in ("trtllm_mla", "tokenspeed_mla", "cutedsl_mla") and ( forward_batch.forward_mode.is_decode_or_idle() or forward_batch.forward_mode.is_target_verify() diff --git a/python/sglang/srt/models/deepseek_common/utils.py b/python/sglang/srt/models/deepseek_common/utils.py index 60e3b20a0..418e65959 100644 --- a/python/sglang/srt/models/deepseek_common/utils.py +++ b/python/sglang/srt/models/deepseek_common/utils.py @@ -62,6 +62,7 @@ FORWARD_ABSORB_CORE_ATTENTION_BACKENDS = [ "flashinfer", "cutlass_mla", "trtllm_mla", + "cutedsl_mla", "tokenspeed_mla", "ascend", "intel_xpu", diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index ec851bb03..9d27b3366 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -166,6 +166,7 @@ ATTENTION_BACKEND_CHOICES = [ "flashinfer", "flashmla", "trtllm_mla", + "cutedsl_mla", "tokenspeed_mla", "trtllm_mha", "dual_chunk_flash_attn", @@ -2822,6 +2823,35 @@ class ServerArgs: f"got {self.kv_cache_dtype}." ) + if ( + self.attention_backend == "cutedsl_mla" + or self.decode_attention_backend == "cutedsl_mla" + or self.prefill_attention_backend == "cutedsl_mla" + ): + assert ( + self.prefill_attention_backend != "cutedsl_mla" + ), "CuteDSL MLA only supports decoding for now" + if not is_sm100_supported(): + raise ValueError( + "CuteDSL MLA backend is only supported on Blackwell GPUs (SM100). Please use a different backend." + ) + if self.page_size not in [32, 64]: + logger.warning( + f"CuteDSL MLA only supports page_size of 32 or 64, changing page_size from {self.page_size} to 64." + ) + self.page_size = 64 + if self.kv_cache_dtype not in [ + "fp8_e4m3", + "bf16", + "bfloat16", + "auto", + ]: + raise ValueError( + "CuteDSL MLA backend only supports kv-cache-dtype of fp8_e4m3, bf16, or auto." + ) + if self.prefill_attention_backend is None: + self.prefill_attention_backend = "trtllm_mla" + if ( self.attention_backend == "trtllm_mha" or self.decode_attention_backend == "trtllm_mha" diff --git a/python/sglang/srt/speculative/draft_utils.py b/python/sglang/srt/speculative/draft_utils.py index b465612b6..10bc78954 100644 --- a/python/sglang/srt/speculative/draft_utils.py +++ b/python/sglang/srt/speculative/draft_utils.py @@ -53,6 +53,7 @@ class DraftBackendFactory: "flashmla": self._create_flashmla_decode_backend, "trtllm_mha": self._create_trtllm_mha_decode_backend, "trtllm_mla": self._create_trtllm_mla_decode_backend, + "cutedsl_mla": self._create_cutedsl_mla_decode_backend, "tokenspeed_mla": self._create_tokenspeed_mla_decode_backend, "dsa": self._create_dsa_decode_backend, "nsa": self._create_dsa_decode_backend, # Deprecated alias for "dsa" @@ -81,6 +82,8 @@ class DraftBackendFactory: "flashmla": self._create_flashmla_prefill_backend, "trtllm_mha": self._create_trtllm_mha_prefill_backend, "trtllm_mla": self._create_trtllm_mla_prefill_backend, + # cute-dsl MLA only supports decode; draft-extend falls back to trtllm-gen. + "cutedsl_mla": self._create_trtllm_mla_prefill_backend, "tokenspeed_mla": self._create_tokenspeed_mla_prefill_backend, "dsa": self._create_dsa_prefill_backend, "nsa": self._create_dsa_prefill_backend, # Deprecated alias for "dsa" @@ -188,7 +191,7 @@ class DraftBackendFactory: self.draft_model_runner, self.topk, self.speculative_num_steps ) - def _create_trtllm_mla_decode_backend(self): + def _create_trtllm_mla_decode_backend(self, backend: str = "trtllm-gen"): if not get_global_server_args().use_mla_backend: raise ValueError( "trtllm_mla backend requires MLA model (use_mla_backend=True)." @@ -199,9 +202,15 @@ class DraftBackendFactory: ) return TRTLLMMLAMultiStepDraftBackend( - self.draft_model_runner, self.topk, self.speculative_num_steps + self.draft_model_runner, + self.topk, + self.speculative_num_steps, + backend=backend, ) + def _create_cutedsl_mla_decode_backend(self): + return self._create_trtllm_mla_decode_backend(backend="cute-dsl") + def _create_tokenspeed_mla_decode_backend(self): if not get_global_server_args().use_mla_backend: raise ValueError(