[attn backend] Integrate tokenspeed_mla prefill/decode kernels (fp8 kv cache, blackwell) (#24925)

This commit is contained in:
Qiaolin Yu
2026-05-13 17:36:17 -07:00
committed by GitHub
parent 22d3f3996c
commit 7618ad7075
11 changed files with 462 additions and 92 deletions
+1
View File
@@ -65,6 +65,7 @@ dependencies = [
"tiktoken", "tiktoken",
"tilelang==0.1.8", "tilelang==0.1.8",
"timm==1.0.16", "timm==1.0.16",
"tokenspeed_mla==0.1.1",
"torch_memory_saver>=0.0.9.post1", "torch_memory_saver>=0.0.9.post1",
"torch==2.11.0", "torch==2.11.0",
"torchao==0.17.0", "torchao==0.17.0",
@@ -62,6 +62,17 @@ def create_trtllm_mla_backend(runner):
return TRTLLMMLABackend(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") @register_attention_backend("aiter")
def create_aiter_backend(runner): def create_aiter_backend(runner):
from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend
@@ -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,
)
@@ -755,6 +755,109 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
) )
return output[:total_tokens, :, :] 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( def forward_decode(
self, self,
q: torch.Tensor, # q_nope q: torch.Tensor, # q_nope
@@ -838,46 +941,13 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
self.init_forward_metadata(forward_batch) self.init_forward_metadata(forward_batch)
metadata = forward_batch.decode_trtllm_mla_metadata metadata = forward_batch.decode_trtllm_mla_metadata
# Scale computation for TRTLLM MLA kernel BMM1 operation: raw_out = self._run_decode_kernel(
# 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(
query=query, query=query,
kv_cache=kv_cache, 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, 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, max_seq_len=metadata.max_seq_len_k,
bmm1_scale=bmm1_scale, layer=layer,
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
) )
# Reshape output directly without slicing # 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) 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) 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) q = q.to(self.data_type)
bmm1_scale = q_scale * k_scale * layer.scaling
if forward_batch.forward_mode.is_target_verify(): if forward_batch.forward_mode.is_target_verify():
max_seq_len = ( max_seq_len = (
metadata.max_seq_len_k + forward_batch.spec_info.draft_token_num 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 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, query=q,
kv_cache=kv_cache, 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, block_tables=metadata.block_kv_indices,
seq_lens=metadata.seq_lens_k, seq_lens=metadata.seq_lens_k,
max_seq_len=max_seq_len, max_seq_len=max_seq_len,
bmm1_scale=bmm1_scale, layer=layer,
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
) )
if needs_unpad: if needs_unpad:
@@ -1087,25 +1135,6 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
k = k.view(-1, layer.tp_k_head_num, layer.head_dim) 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) 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. # When chunked prefix cache is enabled, dispatch to different path for ragged attention.
if forward_batch.attn_attend_prefix_cache: if forward_batch.attn_attend_prefix_cache:
# MHA for chunked prefix kv cache when running model with MLA # MHA for chunked prefix kv cache when running model with MLA
@@ -1122,15 +1151,21 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
dtype=self.q_data_type, dtype=self.q_data_type,
device=q.device, device=q.device,
) )
result = flashinfer.prefill.trtllm_ragged_attention_deepseek( result = self._run_prefill_kernel(
**common_trtllm_args, q=q,
seq_lens=forward_batch.prefix_chunk_seq_lens[chunk_idx], k=k,
max_kv_len=forward_batch.prefix_chunk_max_seq_lens[chunk_idx], v=v,
o_sf_scale=-1.0, 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], 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, is_causal=False,
return_lse=True, return_lse=True,
out=out, out_buffer=out,
o_sf_scale=-1.0,
) )
# The TRT-LLM ragged attention cubin kernel does not correctly # The TRT-LLM ragged attention cubin kernel does not correctly
@@ -1159,15 +1194,21 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
device=q.device, device=q.device,
dtype=self.q_data_type, dtype=self.q_data_type,
) )
return flashinfer.prefill.trtllm_ragged_attention_deepseek( return self._run_prefill_kernel(
**common_trtllm_args, q=q,
seq_lens=self.forward_prefill_metadata.seq_lens, k=k,
max_kv_len=self.forward_prefill_metadata.max_seq_len, v=v,
o_sf_scale=1.0, 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, cum_seq_lens_kv=self.forward_prefill_metadata.cum_seq_lens,
max_kv_len=self.forward_prefill_metadata.max_seq_len,
is_causal=True, is_causal=True,
return_lse=forward_batch.mha_return_lse, return_lse=forward_batch.mha_return_lse,
out=out, out_buffer=out,
o_sf_scale=1.0,
) )
@@ -244,6 +244,7 @@ MLA_ATTENTION_BACKENDS = [
"flashmla", "flashmla",
"cutlass_mla", "cutlass_mla",
"trtllm_mla", "trtllm_mla",
"tokenspeed_mla",
"ascend", "ascend",
"nsa", "nsa",
"intel_xpu", "intel_xpu",
@@ -256,6 +257,7 @@ CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACKENDS = [
"flashmla", "flashmla",
"cutlass_mla", "cutlass_mla",
"trtllm_mla", "trtllm_mla",
"tokenspeed_mla",
] ]
TORCH_DTYPE_TO_KV_CACHE_STR = { TORCH_DTYPE_TO_KV_CACHE_STR = {
@@ -134,6 +134,12 @@ def handle_attention_trtllm_mla(attn, forward_batch):
return _dispatch_mla_subtype(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): def handle_attention_aiter(attn, forward_batch):
if forward_batch.forward_mode.is_extend_without_speculative(): if forward_batch.forward_mode.is_extend_without_speculative():
return AttnForwardMethod.MHA return AttnForwardMethod.MHA
@@ -183,6 +189,7 @@ AttentionBackendRegistry.register("flashmla", handle_attention_flashmla)
AttentionBackendRegistry.register("cutlass_mla", handle_attention_cutlass_mla) AttentionBackendRegistry.register("cutlass_mla", handle_attention_cutlass_mla)
AttentionBackendRegistry.register("fa4", handle_attention_fa4) AttentionBackendRegistry.register("fa4", handle_attention_fa4)
AttentionBackendRegistry.register("trtllm_mla", handle_attention_trtllm_mla) AttentionBackendRegistry.register("trtllm_mla", handle_attention_trtllm_mla)
AttentionBackendRegistry.register("tokenspeed_mla", handle_attention_tokenspeed_mla)
AttentionBackendRegistry.register("aiter", handle_attention_aiter) AttentionBackendRegistry.register("aiter", handle_attention_aiter)
AttentionBackendRegistry.register("nsa", handle_attention_nsa) AttentionBackendRegistry.register("nsa", handle_attention_nsa)
AttentionBackendRegistry.register("triton", handle_attention_triton) AttentionBackendRegistry.register("triton", handle_attention_triton)
@@ -690,7 +690,7 @@ class DeepseekMLAForwardMixin:
) and forward_batch.attn_backend.kv_cache_dtype == torch.float8_e4m3fn ) and forward_batch.attn_backend.kv_cache_dtype == torch.float8_e4m3fn
return ( return (
self.current_attention_backend == "trtllm_mla" self.current_attention_backend in ("trtllm_mla", "tokenspeed_mla")
and ( and (
forward_batch.forward_mode.is_decode_or_idle() forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify() or forward_batch.forward_mode.is_target_verify()
@@ -61,6 +61,7 @@ FORWARD_ABSORB_CORE_ATTENTION_BACKENDS = [
"flashinfer", "flashinfer",
"cutlass_mla", "cutlass_mla",
"trtllm_mla", "trtllm_mla",
"tokenspeed_mla",
"ascend", "ascend",
"intel_xpu", "intel_xpu",
] ]
+20
View File
@@ -158,6 +158,7 @@ ATTENTION_BACKEND_CHOICES = [
"flashinfer", "flashinfer",
"flashmla", "flashmla",
"trtllm_mla", "trtllm_mla",
"tokenspeed_mla",
"trtllm_mha", "trtllm_mha",
"dual_chunk_flash_attn", "dual_chunk_flash_attn",
# AMD specific # 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." "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 ( if (
self.attention_backend == "trtllm_mha" self.attention_backend == "trtllm_mha"
or self.decode_attention_backend == "trtllm_mha" or self.decode_attention_backend == "trtllm_mha"
@@ -53,6 +53,7 @@ class DraftBackendFactory:
"flashmla": self._create_flashmla_decode_backend, "flashmla": self._create_flashmla_decode_backend,
"trtllm_mha": self._create_trtllm_mha_decode_backend, "trtllm_mha": self._create_trtllm_mha_decode_backend,
"trtllm_mla": self._create_trtllm_mla_decode_backend, "trtllm_mla": self._create_trtllm_mla_decode_backend,
"tokenspeed_mla": self._create_tokenspeed_mla_decode_backend,
"nsa": self._create_nsa_decode_backend, "nsa": self._create_nsa_decode_backend,
"ascend": self._create_ascend_decode_backend, "ascend": self._create_ascend_decode_backend,
"fa4": self._create_fa4_decode_backend, "fa4": self._create_fa4_decode_backend,
@@ -79,6 +80,7 @@ class DraftBackendFactory:
"flashmla": self._create_flashmla_prefill_backend, "flashmla": self._create_flashmla_prefill_backend,
"trtllm_mha": self._create_trtllm_mha_prefill_backend, "trtllm_mha": self._create_trtllm_mha_prefill_backend,
"trtllm_mla": self._create_trtllm_mla_prefill_backend, "trtllm_mla": self._create_trtllm_mla_prefill_backend,
"tokenspeed_mla": self._create_tokenspeed_mla_prefill_backend,
"nsa": self._create_nsa_prefill_backend, "nsa": self._create_nsa_prefill_backend,
"ascend": self._create_ascend_prefill_backend, "ascend": self._create_ascend_prefill_backend,
"fa4": self._create_fa4_prefill_backend, "fa4": self._create_fa4_prefill_backend,
@@ -198,6 +200,20 @@ class DraftBackendFactory:
self.draft_model_runner, self.topk, self.speculative_num_steps 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): def _create_ascend_decode_backend(self):
from sglang.srt.hardware_backend.npu.attention.ascend_backend import ( from sglang.srt.hardware_backend.npu.attention.ascend_backend import (
AscendAttnMultiStepDraftBackend, AscendAttnMultiStepDraftBackend,
@@ -274,6 +290,18 @@ class DraftBackendFactory:
return TRTLLMMLABackend(self.draft_model_runner, skip_prefill=False) 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): def _create_ascend_prefill_backend(self):
from sglang.srt.hardware_backend.npu.attention.ascend_backend import ( from sglang.srt.hardware_backend.npu.attention.ascend_backend import (
AscendAttnBackend, AscendAttnBackend,
+12
View File
@@ -316,6 +316,18 @@ def is_flashinfer_available():
return importlib.util.find_spec("flashinfer") is not None and is_cuda() 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(): def is_nvidia_cublas_version_ge_12_9():
""" """
temporary fix for issue #11272 (cublas 12.9+) temporary fix for issue #11272 (cublas 12.9+)