[attn backend] Integrate tokenspeed_mla prefill/decode kernels (fp8 kv cache, blackwell) (#24925)
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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, :, :]
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -61,6 +61,7 @@ FORWARD_ABSORB_CORE_ATTENTION_BACKENDS = [
|
||||
"flashinfer",
|
||||
"cutlass_mla",
|
||||
"trtllm_mla",
|
||||
"tokenspeed_mla",
|
||||
"ascend",
|
||||
"intel_xpu",
|
||||
]
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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+)
|
||||
|
||||
Reference in New Issue
Block a user