diff --git a/python/pyproject.toml b/python/pyproject.toml index 24d1a22fe..beacace1c 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -65,6 +65,7 @@ dependencies = [ "tiktoken", "tilelang==0.1.8", "timm==1.0.16", + "tokenspeed_mla==0.1.1", "torch_memory_saver>=0.0.9.post1", "torch==2.11.0", "torchao==0.17.0", diff --git a/python/sglang/srt/layers/attention/attention_registry.py b/python/sglang/srt/layers/attention/attention_registry.py index d5e3286aa..47350e403 100644 --- a/python/sglang/srt/layers/attention/attention_registry.py +++ b/python/sglang/srt/layers/attention/attention_registry.py @@ -62,6 +62,17 @@ def create_trtllm_mla_backend(runner): return TRTLLMMLABackend(runner) +@register_attention_backend("tokenspeed_mla") +def create_tokenspeed_mla_backend(runner): + if not runner.use_mla_backend: + raise ValueError("tokenspeed_mla backend can only be used with MLA models.") + from sglang.srt.layers.attention.tokenspeed_mla_backend import ( + TokenspeedMLABackend, + ) + + return TokenspeedMLABackend(runner) + + @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/tokenspeed_mla_backend.py b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py new file mode 100644 index 000000000..6296900ee --- /dev/null +++ b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py @@ -0,0 +1,247 @@ +# Copyright (c) 2026 LightSeek Foundation +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. + +from __future__ import annotations + +"""Attention backend for the tokenspeed-mla CuTe DSL kernels on Blackwell. + +Subclasses :class:`TRTLLMMLABackend` and overrides only ``_run_decode_kernel`` +and ``_run_prefill_kernel``. All metadata, KV-cache layout, CUDA-graph +plumbing, FP8 quantize/rope, draft-extend padding, and chunked-prefix +dispatch are inherited unchanged from the parent. +""" + +import logging +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.jit_kernel.utils import is_arch_support_pdl +from sglang.srt.layers.attention.trtllm_mla_backend import ( + TRTLLMMLABackend, + TRTLLMMLAMultiStepDraftBackend, + _quantize_fp8_qkv, +) +from sglang.srt.utils import is_tokenspeed_mla_available + +if is_tokenspeed_mla_available(): + import tokenspeed_mla + +if TYPE_CHECKING: + from sglang.srt.layers.radix_attention import RadixAttention + from sglang.srt.model_executor.model_runner import ModelRunner + +logger = logging.getLogger(__name__) + + +# Workspace upper bound for tokenspeed_mla_decode: +# num_sms * num_heads * max_q_len * (kv_lora_rank + 1) * sizeof(float32) +# MAX_Q_LEN=8 covers EAGLE3 num_draft_tokens=4 plus headroom. +_TOKENSPEED_MAX_Q_LEN = 8 + +_g_tokenspeed_workspace: dict[torch.device, torch.Tensor] = {} + + +def _get_tokenspeed_workspace( + device: torch.device, num_heads: int, kv_lora_rank: int +) -> torch.Tensor: + needed = ( + tokenspeed_mla.get_num_sm(device) + * num_heads + * _TOKENSPEED_MAX_Q_LEN + * (kv_lora_rank + 1) + * 4 + ) + existing = _g_tokenspeed_workspace.get(device) + if existing is None or existing.numel() < needed: + _g_tokenspeed_workspace[device] = torch.empty( + needed, dtype=torch.int8, device=device + ) + return _g_tokenspeed_workspace[device] + + +class TokenspeedMLABackend(TRTLLMMLABackend): + """tokenspeed-mla CuTe DSL attention backend (Blackwell SM100, FP8 KV).""" + + def __init__( + self, + model_runner: "ModelRunner", + skip_prefill: bool = False, + kv_indptr_buf: Optional[torch.Tensor] = None, + q_indptr_decode_buf: Optional[torch.Tensor] = None, + ): + super().__init__( + model_runner, + skip_prefill, + kv_indptr_buf, + q_indptr_decode_buf, + ) + + if self.data_type != torch.float8_e4m3fn: + raise ValueError( + "tokenspeed_mla backend requires --kv-cache-dtype fp8_e4m3, " + f"got data_type={self.data_type}." + ) + if self.page_size not in (32, 64): + raise ValueError( + "tokenspeed_mla backend requires page_size in {32, 64}, " + f"got page_size={self.page_size}." + ) + + self._tokenspeed_workspace: Optional[torch.Tensor] = None + + # Pre-JIT the prefill kernel variants. Each cute.compile takes 1-2 min; + # without warm-up the first request trips the 300 s scheduler watchdog. + if is_tokenspeed_mla_available(): + _compile_prefill_kernel = tokenspeed_mla.mla_prefill._compile_prefill_kernel + _compiled_kernels = tokenspeed_mla.mla_prefill._compiled_kernels + head_dim_qk = self.qk_nope_head_dim + self.qk_rope_head_dim + enable_ex2_emulation = tokenspeed_mla.mla_prefill._enable_ex2_emulation() + use_pdl = is_arch_support_pdl() + for is_causal in (True, False): + for return_lse in (True, False): + # Non-causal is only entered from the chunked-prefix + # branch, which always asks for the LSE. + if is_causal is False and return_lse is False: + continue + config = ( + torch.bfloat16, + head_dim_qk, + self.v_head_dim, + is_causal, + return_lse, + use_pdl, + enable_ex2_emulation, + ) + if config in _compiled_kernels: + continue + _compiled_kernels[config] = _compile_prefill_kernel( + torch.bfloat16, + head_dim_qk, + self.v_head_dim, + is_causal, + return_lse, + use_pdl=use_pdl, + enable_ex2_emulation=enable_ex2_emulation, + ) + + def _ensure_workspace(self, device: torch.device) -> torch.Tensor: + if ( + self._tokenspeed_workspace is None + or self._tokenspeed_workspace.device != device + ): + self._tokenspeed_workspace = _get_tokenspeed_workspace( + device, self.num_q_heads, self.kv_lora_rank + ) + return self._tokenspeed_workspace + + def _run_decode_kernel( + self, + query: torch.Tensor, + kv_cache: torch.Tensor, + block_tables: torch.Tensor, + seq_lens: torch.Tensor, + max_seq_len: int, + layer: "RadixAttention", + ) -> torch.Tensor: + k_scale = getattr(layer, "k_scale_float", None) + if k_scale is None: + k_scale = 1.0 + softmax_scale = float(layer.scaling) * float(k_scale) + output_scale = float(k_scale) + + seq_lens_i32 = ( + seq_lens if seq_lens.dtype == torch.int32 else seq_lens.to(torch.int32) + ) + return tokenspeed_mla.tokenspeed_mla_decode( + query=query, + kv_cache=kv_cache, + workspace_buffer=self._ensure_workspace(query.device), + kv_lora_rank=self.kv_lora_rank, + qk_rope_head_dim=self.qk_rope_head_dim, + block_tables=block_tables, + seq_lens=seq_lens_i32, + max_seq_len=int(max_seq_len), + softmax_scale=softmax_scale, + output_scale=output_scale, + enable_pdl=is_arch_support_pdl(), + ) + + def _run_prefill_kernel( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: "RadixAttention", + batch_size: int, + cum_seq_lens_q: torch.Tensor, + max_q_len: int, + seq_lens_kv: torch.Tensor, + cum_seq_lens_kv: torch.Tensor, + max_kv_len: int, + is_causal: bool, + return_lse: bool, + out_buffer: torch.Tensor, + o_sf_scale: float = 1.0, + ): + # Quantize to FP8 for the Blackwell FP8 GEMM speedup (mirrors trtllm-gen). + # The kernel has no per-tensor scale knob for either K or V, so we + # require both ``k_scale_float`` and ``v_scale_float`` to be 1.0. + if self.data_type == torch.float8_e4m3fn: + q, k, v, k_scale, v_scale = _quantize_fp8_qkv(q, k, v, layer) + assert k_scale == 1.0 and v_scale == 1.0, ( + "tokenspeed_mla prefill kernel has no per-tensor K/V scale " + "knob; both k_scale_float and v_scale_float must be 1.0, got " + f"k_scale={k_scale}, v_scale={v_scale}." + ) + + return tokenspeed_mla.tokenspeed_mla_prefill( + query=q, + key=k, + value=v, + seq_lens=seq_lens_kv, + cum_seq_lens=cum_seq_lens_kv, + max_seq_len=int(max_kv_len), + batch_size=int(batch_size), + softmax_scale=float(layer.scaling), + is_causal=is_causal, + return_lse=return_lse, + cum_seq_lens_q=cum_seq_lens_q, + max_seq_len_q=int(max_q_len), + enable_pdl=is_arch_support_pdl(), + ) + + +class TokenspeedMLAMultiStepDraftBackend(TRTLLMMLAMultiStepDraftBackend): + """Multi-step draft backend for tokenspeed_mla used by EAGLE.""" + + def __init__( + self, model_runner: "ModelRunner", topk: int, speculative_num_steps: int + ): + super().__init__(model_runner, topk, speculative_num_steps) + # Parent populates self.attn_backends with TRT-LLM instances; replace + # them with tokenspeed instances sharing the parent's index buffers. + for i in range(self.speculative_num_steps - 1): + self.attn_backends[i] = TokenspeedMLABackend( + model_runner, + skip_prefill=True, + kv_indptr_buf=self.kv_indptr[i], + q_indptr_decode_buf=self.q_indptr_decode, + ) diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index a0fe49418..5ccac1171 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -755,6 +755,109 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): ) return output[:total_tokens, :, :] + def _compute_decode_bmm1_scale(self, layer: RadixAttention) -> float: + """BMM1 scale ``q_scale * k_scale * softmax_scale``. k_scale only + applies when the KV cache stores FP8.""" + q_scale = 1.0 + if self.data_type == torch.float8_e4m3fn: + k_scale = ( + layer.k_scale_float + if getattr(layer, "k_scale_float", None) is not None + else 1.0 + ) + else: + if getattr(layer, "k_scale_float", None) is not None: + logger.warning_once( + "Checkpoint has k_scale but KV cache dtype is not FP8. " + "Ignoring k_scale for BMM1 (k_scale=%.4f, kv_dtype=%s).", + layer.k_scale_float, + self.data_type, + ) + k_scale = 1.0 + return q_scale * k_scale * layer.scaling + + def _run_decode_kernel( + self, + query: torch.Tensor, + kv_cache: torch.Tensor, + block_tables: torch.Tensor, + seq_lens: torch.Tensor, + max_seq_len: int, + layer: RadixAttention, + ) -> torch.Tensor: + """Hook for subclasses to swap the decode/spec-verify kernel.""" + + # Scale computation for TRTLLM MLA kernel BMM1 operation: + # The final BMM1 scale is computed as: q_scale * k_scale * softmax_scale + # Scale components: + # - q_scale: Query scaling factor (set to 1.0 for both FP16/FP8 paths) + # - k_scale: Key scaling factor from model checkpoint. Only applied when KV cache + # stores FP8-quantized values, to compensate for the quantization scaling. + # For BF16/FP16 KV cache, k_scale must be 1.0 since values are unscaled. + # - softmax_scale: Attention softmax scaling = 1/sqrt(head_dim), pre-computed as layer.scaling + bmm1_scale = self._compute_decode_bmm1_scale(layer) + seq_lens_i32 = ( + seq_lens if seq_lens.dtype == torch.int32 else seq_lens.to(torch.int32) + ) + return flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla( + query=query, + kv_cache=kv_cache, + workspace_buffer=self.workspace_buffer, + qk_nope_head_dim=self.qk_nope_head_dim, + kv_lora_rank=self.kv_lora_rank, + qk_rope_head_dim=self.qk_rope_head_dim, + block_tables=block_tables, + seq_lens=seq_lens_i32, + max_seq_len=max_seq_len, + bmm1_scale=bmm1_scale, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), + ) + + def _run_prefill_kernel( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + layer: RadixAttention, + batch_size: int, + cum_seq_lens_q: torch.Tensor, + max_q_len: int, + seq_lens_kv: torch.Tensor, + cum_seq_lens_kv: torch.Tensor, + max_kv_len: int, + is_causal: bool, + return_lse: bool, + out_buffer: torch.Tensor, + o_sf_scale: float = 1.0, + ): + """Hook for subclasses to swap the ragged prefill kernel. Q/K/V arrive + in model-native dtype; subclasses do any kernel-specific quantization. + Returns the output tensor or ``(output, lse)`` if ``return_lse``.""" + q_scale = k_scale = v_scale = 1.0 + if self.data_type == torch.float8_e4m3fn: + q, k, v, k_scale, v_scale = _quantize_fp8_qkv(q, k, v, layer) + return flashinfer.prefill.trtllm_ragged_attention_deepseek( + query=q, + key=k, + value=v, + workspace_buffer=self.workspace_buffer, + batch_size=batch_size, + window_left=-1, + enable_pdl=False, + max_q_len=max_q_len, + bmm1_scale=q_scale * k_scale * layer.scaling, + bmm2_scale=v_scale, + cum_seq_lens_q=cum_seq_lens_q, + cum_seq_lens_kv=cum_seq_lens_kv, + seq_lens=seq_lens_kv, + max_kv_len=max_kv_len, + is_causal=is_causal, + return_lse=return_lse, + o_sf_scale=o_sf_scale, + out=out_buffer, + skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(), + ) + def forward_decode( self, q: torch.Tensor, # q_nope @@ -838,46 +941,13 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): self.init_forward_metadata(forward_batch) metadata = forward_batch.decode_trtllm_mla_metadata - # Scale computation for TRTLLM MLA kernel BMM1 operation: - # The final BMM1 scale is computed as: q_scale * k_scale * softmax_scale - # Scale components: - # - q_scale: Query scaling factor (set to 1.0 for both FP16/FP8 paths) - # - k_scale: Key scaling factor from model checkpoint. Only applied when KV cache - # stores FP8-quantized values, to compensate for the quantization scaling. - # For BF16/FP16 KV cache, k_scale must be 1.0 since values are unscaled. - # - softmax_scale: Attention softmax scaling = 1/sqrt(head_dim), pre-computed as layer.scaling - q_scale = 1.0 - if self.data_type == torch.float8_e4m3fn: - k_scale = ( - layer.k_scale_float - if getattr(layer, "k_scale_float", None) is not None - else 1.0 - ) - else: - if getattr(layer, "k_scale_float", None) is not None: - logger.warning_once( - "Checkpoint has k_scale but KV cache dtype is not FP8. " - "Ignoring k_scale for BMM1 (k_scale=%.4f, kv_dtype=%s).", - layer.k_scale_float, - self.data_type, - ) - k_scale = 1.0 - - bmm1_scale = q_scale * k_scale * layer.scaling - - # Call TRT-LLM kernel - raw_out = flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla( + raw_out = self._run_decode_kernel( query=query, kv_cache=kv_cache, - workspace_buffer=self.workspace_buffer, - qk_nope_head_dim=self.qk_nope_head_dim, - kv_lora_rank=self.kv_lora_rank, - qk_rope_head_dim=self.qk_rope_head_dim, block_tables=metadata.block_kv_indices, - seq_lens=forward_batch.seq_lens.to(torch.int32), + seq_lens=forward_batch.seq_lens, max_seq_len=metadata.max_seq_len_k, - bmm1_scale=bmm1_scale, - skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), + layer=layer, ) # Reshape output directly without slicing @@ -979,25 +1049,8 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id) kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1) - q_scale = 1.0 - if self.data_type == torch.float8_e4m3fn: - k_scale = ( - layer.k_scale_float - if getattr(layer, "k_scale_float", None) is not None - else 1.0 - ) - else: - if getattr(layer, "k_scale_float", None) is not None: - logger.warning_once( - "Checkpoint has k_scale but KV cache dtype is not FP8. " - "Ignoring k_scale for BMM1 (k_scale=%.4f, kv_dtype=%s).", - layer.k_scale_float, - self.data_type, - ) - k_scale = 1.0 q = q.to(self.data_type) - bmm1_scale = q_scale * k_scale * layer.scaling if forward_batch.forward_mode.is_target_verify(): max_seq_len = ( metadata.max_seq_len_k + forward_batch.spec_info.draft_token_num @@ -1054,18 +1107,13 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): assert kv_cache.dtype == self.data_type - raw_out = flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla( + raw_out = self._run_decode_kernel( query=q, kv_cache=kv_cache, - workspace_buffer=self.workspace_buffer, - qk_nope_head_dim=self.qk_nope_head_dim, - kv_lora_rank=self.kv_lora_rank, - qk_rope_head_dim=self.qk_rope_head_dim, block_tables=metadata.block_kv_indices, seq_lens=metadata.seq_lens_k, max_seq_len=max_seq_len, - bmm1_scale=bmm1_scale, - skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(), + layer=layer, ) if needs_unpad: @@ -1087,25 +1135,6 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): k = k.view(-1, layer.tp_k_head_num, layer.head_dim) v = v.view(-1, layer.tp_k_head_num, layer.v_head_dim) - q_scale = k_scale = v_scale = 1.0 - if self.data_type == torch.float8_e4m3fn: - q, k, v, k_scale, v_scale = _quantize_fp8_qkv(q, k, v, layer) - - common_trtllm_args = { - "query": q, - "key": k, - "value": v, - "workspace_buffer": self.workspace_buffer, - "batch_size": forward_batch.batch_size, - "window_left": -1, - "enable_pdl": False, - "max_q_len": self.forward_prefill_metadata.max_seq_len, - "bmm1_scale": q_scale * k_scale * layer.scaling, - "bmm2_scale": v_scale, - "cum_seq_lens_q": self.forward_prefill_metadata.cum_seq_lens, - "skip_softmax_threshold_scale_factor": envs.SGLANG_SKIP_SOFTMAX_PREFILL_THRESHOLD_SCALE_FACTOR.get(), - } - # When chunked prefix cache is enabled, dispatch to different path for ragged attention. if forward_batch.attn_attend_prefix_cache: # MHA for chunked prefix kv cache when running model with MLA @@ -1122,15 +1151,21 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): dtype=self.q_data_type, device=q.device, ) - result = flashinfer.prefill.trtllm_ragged_attention_deepseek( - **common_trtllm_args, - seq_lens=forward_batch.prefix_chunk_seq_lens[chunk_idx], - max_kv_len=forward_batch.prefix_chunk_max_seq_lens[chunk_idx], - o_sf_scale=-1.0, + result = self._run_prefill_kernel( + q=q, + k=k, + v=v, + layer=layer, + batch_size=forward_batch.batch_size, + cum_seq_lens_q=self.forward_prefill_metadata.cum_seq_lens, + max_q_len=self.forward_prefill_metadata.max_seq_len, + seq_lens_kv=forward_batch.prefix_chunk_seq_lens[chunk_idx], cum_seq_lens_kv=forward_batch.prefix_chunk_cu_seq_lens[chunk_idx], + max_kv_len=forward_batch.prefix_chunk_max_seq_lens[chunk_idx], is_causal=False, return_lse=True, - out=out, + out_buffer=out, + o_sf_scale=-1.0, ) # The TRT-LLM ragged attention cubin kernel does not correctly @@ -1159,15 +1194,21 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): device=q.device, dtype=self.q_data_type, ) - return flashinfer.prefill.trtllm_ragged_attention_deepseek( - **common_trtllm_args, - seq_lens=self.forward_prefill_metadata.seq_lens, - max_kv_len=self.forward_prefill_metadata.max_seq_len, - o_sf_scale=1.0, + return self._run_prefill_kernel( + q=q, + k=k, + v=v, + layer=layer, + batch_size=forward_batch.batch_size, + cum_seq_lens_q=self.forward_prefill_metadata.cum_seq_lens, + max_q_len=self.forward_prefill_metadata.max_seq_len, + seq_lens_kv=self.forward_prefill_metadata.seq_lens, cum_seq_lens_kv=self.forward_prefill_metadata.cum_seq_lens, + max_kv_len=self.forward_prefill_metadata.max_seq_len, is_causal=True, return_lse=forward_batch.mha_return_lse, - out=out, + out_buffer=out, + o_sf_scale=1.0, ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index d94f7f174..1d671f012 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -244,6 +244,7 @@ MLA_ATTENTION_BACKENDS = [ "flashmla", "cutlass_mla", "trtllm_mla", + "tokenspeed_mla", "ascend", "nsa", "intel_xpu", @@ -256,6 +257,7 @@ CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS = [ "flashmla", "cutlass_mla", "trtllm_mla", + "tokenspeed_mla", ] TORCH_DTYPE_TO_KV_CACHE_STR = { diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index c1cd0e32c..3c93d41a8 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -134,6 +134,12 @@ def handle_attention_trtllm_mla(attn, forward_batch): return _dispatch_mla_subtype(attn, forward_batch) +def handle_attention_tokenspeed_mla(attn, forward_batch): + # tokenspeed_mla shares the trtllm_mla dispatch pattern: pure prefill goes + # via MHA chunked KV (TRT-LLM ragged), spec decode / decode goes via MLA. + return handle_attention_trtllm_mla(attn, forward_batch) + + def handle_attention_aiter(attn, forward_batch): if forward_batch.forward_mode.is_extend_without_speculative(): return AttnForwardMethod.MHA @@ -183,6 +189,7 @@ AttentionBackendRegistry.register("flashmla", handle_attention_flashmla) AttentionBackendRegistry.register("cutlass_mla", handle_attention_cutlass_mla) AttentionBackendRegistry.register("fa4", handle_attention_fa4) AttentionBackendRegistry.register("trtllm_mla", handle_attention_trtllm_mla) +AttentionBackendRegistry.register("tokenspeed_mla", handle_attention_tokenspeed_mla) AttentionBackendRegistry.register("aiter", handle_attention_aiter) AttentionBackendRegistry.register("nsa", handle_attention_nsa) AttentionBackendRegistry.register("triton", handle_attention_triton) 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 ed33c208a..3093c196b 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 @@ -690,7 +690,7 @@ class DeepseekMLAForwardMixin: ) and forward_batch.attn_backend.kv_cache_dtype == torch.float8_e4m3fn return ( - self.current_attention_backend == "trtllm_mla" + self.current_attention_backend in ("trtllm_mla", "tokenspeed_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 c8ac58c60..40f552498 100644 --- a/python/sglang/srt/models/deepseek_common/utils.py +++ b/python/sglang/srt/models/deepseek_common/utils.py @@ -61,6 +61,7 @@ FORWARD_ABSORB_CORE_ATTENTION_BACKENDS = [ "flashinfer", "cutlass_mla", "trtllm_mla", + "tokenspeed_mla", "ascend", "intel_xpu", ] diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index afec6de73..8aab671ff 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -158,6 +158,7 @@ ATTENTION_BACKEND_CHOICES = [ "flashinfer", "flashmla", "trtllm_mla", + "tokenspeed_mla", "trtllm_mha", "dual_chunk_flash_attn", # AMD specific @@ -2700,6 +2701,25 @@ class ServerArgs: "TensorRT-LLM MLA backend only supports kv-cache-dtype of fp8_e4m3, fp4_e2m1, bf16, or auto." ) + if ( + self.attention_backend == "tokenspeed_mla" + or self.decode_attention_backend == "tokenspeed_mla" + ): + if not is_blackwell_supported(): + raise ValueError( + "tokenspeed_mla backend is only supported on Blackwell GPUs (SM100/SM12x)." + ) + if self.page_size not in [32, 64]: + logger.warning( + f"tokenspeed_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"]: + raise ValueError( + "tokenspeed_mla backend requires kv-cache-dtype=fp8_e4m3, " + f"got {self.kv_cache_dtype}." + ) + 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 4da59b72a..ce6ac5334 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, + "tokenspeed_mla": self._create_tokenspeed_mla_decode_backend, "nsa": self._create_nsa_decode_backend, "ascend": self._create_ascend_decode_backend, "fa4": self._create_fa4_decode_backend, @@ -79,6 +80,7 @@ class DraftBackendFactory: "flashmla": self._create_flashmla_prefill_backend, "trtllm_mha": self._create_trtllm_mha_prefill_backend, "trtllm_mla": self._create_trtllm_mla_prefill_backend, + "tokenspeed_mla": self._create_tokenspeed_mla_prefill_backend, "nsa": self._create_nsa_prefill_backend, "ascend": self._create_ascend_prefill_backend, "fa4": self._create_fa4_prefill_backend, @@ -198,6 +200,20 @@ class DraftBackendFactory: self.draft_model_runner, self.topk, self.speculative_num_steps ) + def _create_tokenspeed_mla_decode_backend(self): + if not get_global_server_args().use_mla_backend: + raise ValueError( + "tokenspeed_mla backend requires MLA model (use_mla_backend=True)." + ) + + from sglang.srt.layers.attention.tokenspeed_mla_backend import ( + TokenspeedMLAMultiStepDraftBackend, + ) + + return TokenspeedMLAMultiStepDraftBackend( + self.draft_model_runner, self.topk, self.speculative_num_steps + ) + def _create_ascend_decode_backend(self): from sglang.srt.hardware_backend.npu.attention.ascend_backend import ( AscendAttnMultiStepDraftBackend, @@ -274,6 +290,18 @@ class DraftBackendFactory: return TRTLLMMLABackend(self.draft_model_runner, skip_prefill=False) + def _create_tokenspeed_mla_prefill_backend(self): + if not get_global_server_args().use_mla_backend: + raise ValueError( + "tokenspeed_mla backend requires MLA model (use_mla_backend=True)." + ) + + from sglang.srt.layers.attention.tokenspeed_mla_backend import ( + TokenspeedMLABackend, + ) + + return TokenspeedMLABackend(self.draft_model_runner, skip_prefill=False) + def _create_ascend_prefill_backend(self): from sglang.srt.hardware_backend.npu.attention.ascend_backend import ( AscendAttnBackend, diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index 3f17160d8..61b6baca6 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -316,6 +316,18 @@ def is_flashinfer_available(): return importlib.util.find_spec("flashinfer") is not None and is_cuda() +@lru_cache(maxsize=1) +def is_tokenspeed_mla_available(): + """ + Check whether the tokenspeed_mla CuTe DSL kernels are available. + Only available on NVIDIA Blackwell (SM100) at the moment. + """ + return ( + importlib.util.find_spec("tokenspeed_mla") is not None + and is_blackwell_supported() + ) + + def is_nvidia_cublas_version_ge_12_9(): """ temporary fix for issue #11272 (cublas 12.9+)