diff --git a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py index 2518d0323..601a80cea 100644 --- a/python/sglang/srt/layers/attention/flashinfer_mla_backend.py +++ b/python/sglang/srt/layers/attention/flashinfer_mla_backend.py @@ -197,7 +197,6 @@ class FlashInferMLAAttnBackend(AttentionBackend): skip_prefill: bool = False, kv_indptr_buf: Optional[torch.Tensor] = None, q_indptr_decode_buf: Optional[torch.Tensor] = None, - skip_init_workspace_buffer: bool = False, ): super().__init__() @@ -205,7 +204,6 @@ class FlashInferMLAAttnBackend(AttentionBackend): self.max_context_len = model_runner.model_config.context_len self.device = model_runner.device self.skip_prefill = skip_prefill - self.skip_init_workspace_buffer = skip_init_workspace_buffer self.enable_chunk_kv = ( not skip_prefill and get_global_server_args().disaggregation_mode != "decode" @@ -215,18 +213,15 @@ class FlashInferMLAAttnBackend(AttentionBackend): self.page_size = model_runner.page_size # Allocate buffers - if skip_init_workspace_buffer: - self.workspace_buffer = None - else: - global global_workspace_buffer - if global_workspace_buffer is None: - # different from flashinfer zero_init_global_workspace_buffer - global_workspace_buffer = torch.empty( - envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(), - dtype=torch.uint8, - device=model_runner.device, - ) - self.workspace_buffer = global_workspace_buffer + global global_workspace_buffer + if global_workspace_buffer is None: + # different from flashinfer zero_init_global_workspace_buffer + global_workspace_buffer = torch.empty( + envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(), + dtype=torch.uint8, + device=model_runner.device, + ) + self.workspace_buffer = global_workspace_buffer max_bs = model_runner.req_to_token_pool.size if kv_indptr_buf is None: @@ -248,53 +243,42 @@ class FlashInferMLAAttnBackend(AttentionBackend): else: self.q_indptr_decode = q_indptr_decode_buf - if skip_init_workspace_buffer: - self.fmha_backend = None - self.prefill_wrapper_ragged = None - self.prefill_wrapper_paged = None - self.prefill_wrapper_verify = None - self.decode_wrapper = None - self.indices_updater_prefill = None - self.indices_updater_decode = None + if is_sm100_supported(): + self.fmha_backend = "cutlass" else: - if is_sm100_supported(): - self.fmha_backend = "cutlass" - else: - self.fmha_backend = "auto" + self.fmha_backend = "auto" - self.prefill_wrapper_ragged = BatchPrefillWithRaggedKVCacheWrapper( - self.workspace_buffer, "NHD", backend=self.fmha_backend + self.prefill_wrapper_ragged = BatchPrefillWithRaggedKVCacheWrapper( + self.workspace_buffer, "NHD", backend=self.fmha_backend + ) + + if not self.skip_prefill: + self.prefill_wrapper_paged = BatchMLAPagedAttentionWrapper( + self.workspace_buffer, + backend="auto", ) - if not self.skip_prefill: - self.prefill_wrapper_paged = BatchMLAPagedAttentionWrapper( - self.workspace_buffer, - backend="auto", - ) - - # FlashinferMLA backend uses mla wrapper for target verify - self.prefill_wrapper_verify = BatchMLAPagedAttentionWrapper( - self.workspace_buffer, - backend="auto", - ) - - self.decode_wrapper = BatchMLAPagedAttentionWrapper( - self.workspace_buffer, backend="auto" + # FlashinferMLA backend uses mla wrapper for target verify + self.prefill_wrapper_verify = BatchMLAPagedAttentionWrapper( + self.workspace_buffer, + backend="auto", ) - # Create indices updater - if not skip_prefill: - self.indices_updater_prefill = FlashInferMLAIndicesUpdaterPrefill( - model_runner, self - ) - if self.enable_chunk_kv: - self.mha_chunk_kv_cache = FlashInferMhaChunkKVRunner( - model_runner, self - ) + self.decode_wrapper = BatchMLAPagedAttentionWrapper( + self.workspace_buffer, backend="auto" + ) - self.indices_updater_decode = FlashInferMLAIndicesUpdaterDecode( + # Create indices updater + if not skip_prefill: + self.indices_updater_prefill = FlashInferMLAIndicesUpdaterPrefill( model_runner, self ) + if self.enable_chunk_kv: + self.mha_chunk_kv_cache = FlashInferMhaChunkKVRunner(model_runner, self) + + self.indices_updater_decode = FlashInferMLAIndicesUpdaterDecode( + model_runner, self + ) # Other metadata self.forward_metadata: Union[PrefillMetadata, DecodeMetadata] = None diff --git a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py index a787d2773..6296900ee 100644 --- a/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py +++ b/python/sglang/srt/layers/attention/tokenspeed_mla_backend.py @@ -92,7 +92,6 @@ class TokenspeedMLABackend(TRTLLMMLABackend): skip_prefill, kv_indptr_buf, q_indptr_decode_buf, - skip_init_workspace_buffer=True, ) if self.data_type != torch.float8_e4m3fn: diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index d6b2f3da7..5ccac1171 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -263,14 +263,12 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): skip_prefill: bool = False, kv_indptr_buf: Optional[torch.Tensor] = None, q_indptr_decode_buf: Optional[torch.Tensor] = None, - skip_init_workspace_buffer: bool = False, ): super().__init__( model_runner, skip_prefill, kv_indptr_buf, q_indptr_decode_buf, - skip_init_workspace_buffer=True, ) config = model_runner.model_config @@ -296,17 +294,14 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): # Workspace allocation self.workspace_size = DEFAULT_WORKSPACE_SIZE_MB * 1024 * 1024 - if skip_init_workspace_buffer: - self.workspace_buffer = None - else: - global global_zero_init_workspace_buffer - if global_zero_init_workspace_buffer is None: - global_zero_init_workspace_buffer = torch.zeros( - self.workspace_size, - dtype=torch.uint8, - device=model_runner.device, - ) - self.workspace_buffer = global_zero_init_workspace_buffer + global global_zero_init_workspace_buffer + if global_zero_init_workspace_buffer is None: + global_zero_init_workspace_buffer = torch.zeros( + self.workspace_size, + dtype=torch.uint8, + device=model_runner.device, + ) + self.workspace_buffer = global_zero_init_workspace_buffer # CUDA graph state self.decode_cuda_graph_metadata = {} @@ -383,10 +378,6 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): return block_kv_indices - def init_mha_chunk_metadata(self, forward_batch: "ForwardBatch") -> None: - """Skip parent's flashinfer wrapper plan().""" - return None - def init_cuda_graph_state( self, max_bs: int, @@ -681,6 +672,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): else: return super().init_forward_metadata(forward_batch) + def init_mha_chunk_metadata(self, forward_batch: ForwardBatch): + super().init_mha_chunk_metadata(forward_batch, disable_flashinfer_ragged=True) + def pad_draft_extend_query( self, q: torch.Tensor,