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