diff --git a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py index 6296900ee..8e4fcad80 100644 --- a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py +++ b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py @@ -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, diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 8a815eb5e..d5d64b844 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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