Support draft extend cuda graph for tokenspeed_mla attention backend (#25489)

This commit is contained in:
Qiaolin Yu
2026-05-18 11:26:16 -07:00
committed by GitHub
parent f5049709b3
commit 1f185c6ba8
2 changed files with 10 additions and 14 deletions
@@ -106,10 +106,14 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
)
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():
self._tokenspeed_workspace = _get_tokenspeed_workspace(
self.device, self.num_q_heads, self.kv_lora_rank
)
# 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.
_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
@@ -142,16 +146,6 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
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,
@@ -173,7 +167,7 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
return tokenspeed_mla.tokenspeed_mla_decode(
query=query,
kv_cache=kv_cache,
workspace_buffer=self._ensure_workspace(query.device),
workspace_buffer=self._tokenspeed_workspace,
kv_lora_rank=self.kv_lora_rank,
qk_rope_head_dim=self.qk_rope_head_dim,
block_tables=block_tables,
@@ -12,6 +12,7 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_r
from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner import (
EAGLEDraftNpuGraphRunner,
)
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
from sglang.srt.layers.attention.triton_backend import TritonAttnBackend
from sglang.srt.layers.attention.trtllm_mla_backend import (
TRTLLMMLABackend,
@@ -314,6 +315,7 @@ class EagleDraftWorker(BaseDraftWorker):
supports_cuda_draft_extend_graph = (_is_cuda or _is_musa) and (
isinstance(self.draft_extend_attn_backend, TritonAttnBackend)
or isinstance(self.draft_extend_attn_backend, TRTLLMMLABackend)
or isinstance(self.draft_extend_attn_backend, TokenspeedMLABackend)
)
# Capture extend
# TODO: support draft extend cuda graph for more attention backends