[attn backend] avoid initing parent class's workspace buffer (#25321)

This commit is contained in:
Qiaolin Yu
2026-05-16 03:30:33 -07:00
committed by GitHub
parent aec4022e58
commit 2f81718773
3 changed files with 70 additions and 47 deletions
@@ -197,6 +197,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
skip_prefill: bool = False, skip_prefill: bool = False,
kv_indptr_buf: Optional[torch.Tensor] = None, kv_indptr_buf: Optional[torch.Tensor] = None,
q_indptr_decode_buf: Optional[torch.Tensor] = None, q_indptr_decode_buf: Optional[torch.Tensor] = None,
skip_init_workspace_buffer: bool = False,
): ):
super().__init__() super().__init__()
@@ -204,6 +205,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.max_context_len = model_runner.model_config.context_len self.max_context_len = model_runner.model_config.context_len
self.device = model_runner.device self.device = model_runner.device
self.skip_prefill = skip_prefill self.skip_prefill = skip_prefill
self.skip_init_workspace_buffer = skip_init_workspace_buffer
self.enable_chunk_kv = ( self.enable_chunk_kv = (
not skip_prefill not skip_prefill
and get_global_server_args().disaggregation_mode != "decode" and get_global_server_args().disaggregation_mode != "decode"
@@ -213,15 +215,18 @@ class FlashInferMLAAttnBackend(AttentionBackend):
self.page_size = model_runner.page_size self.page_size = model_runner.page_size
# Allocate buffers # Allocate buffers
global global_workspace_buffer if skip_init_workspace_buffer:
if global_workspace_buffer is None: self.workspace_buffer = None
# different from flashinfer zero_init_global_workspace_buffer else:
global_workspace_buffer = torch.empty( global global_workspace_buffer
envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(), if global_workspace_buffer is None:
dtype=torch.uint8, # different from flashinfer zero_init_global_workspace_buffer
device=model_runner.device, global_workspace_buffer = torch.empty(
) envs.SGLANG_FLASHINFER_WORKSPACE_SIZE.get(),
self.workspace_buffer = global_workspace_buffer dtype=torch.uint8,
device=model_runner.device,
)
self.workspace_buffer = global_workspace_buffer
max_bs = model_runner.req_to_token_pool.size max_bs = model_runner.req_to_token_pool.size
if kv_indptr_buf is None: if kv_indptr_buf is None:
@@ -243,42 +248,53 @@ class FlashInferMLAAttnBackend(AttentionBackend):
else: else:
self.q_indptr_decode = q_indptr_decode_buf self.q_indptr_decode = q_indptr_decode_buf
if is_sm100_supported(): if skip_init_workspace_buffer:
self.fmha_backend = "cutlass" 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
else: else:
self.fmha_backend = "auto" if is_sm100_supported():
self.fmha_backend = "cutlass"
else:
self.fmha_backend = "auto"
self.prefill_wrapper_ragged = BatchPrefillWithRaggedKVCacheWrapper( self.prefill_wrapper_ragged = BatchPrefillWithRaggedKVCacheWrapper(
self.workspace_buffer, "NHD", backend=self.fmha_backend self.workspace_buffer, "NHD", backend=self.fmha_backend
)
if not self.skip_prefill:
self.prefill_wrapper_paged = BatchMLAPagedAttentionWrapper(
self.workspace_buffer,
backend="auto",
) )
# FlashinferMLA backend uses mla wrapper for target verify if not self.skip_prefill:
self.prefill_wrapper_verify = BatchMLAPagedAttentionWrapper( self.prefill_wrapper_paged = BatchMLAPagedAttentionWrapper(
self.workspace_buffer, self.workspace_buffer,
backend="auto", 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"
) )
self.decode_wrapper = BatchMLAPagedAttentionWrapper( # Create indices updater
self.workspace_buffer, backend="auto" 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
)
# Create indices updater self.indices_updater_decode = FlashInferMLAIndicesUpdaterDecode(
if not skip_prefill:
self.indices_updater_prefill = FlashInferMLAIndicesUpdaterPrefill(
model_runner, self 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 # Other metadata
self.forward_metadata: Union[PrefillMetadata, DecodeMetadata] = None self.forward_metadata: Union[PrefillMetadata, DecodeMetadata] = None
@@ -92,6 +92,7 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
skip_prefill, skip_prefill,
kv_indptr_buf, kv_indptr_buf,
q_indptr_decode_buf, q_indptr_decode_buf,
skip_init_workspace_buffer=True,
) )
if self.data_type != torch.float8_e4m3fn: if self.data_type != torch.float8_e4m3fn:
@@ -263,12 +263,14 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
skip_prefill: bool = False, skip_prefill: bool = False,
kv_indptr_buf: Optional[torch.Tensor] = None, kv_indptr_buf: Optional[torch.Tensor] = None,
q_indptr_decode_buf: Optional[torch.Tensor] = None, q_indptr_decode_buf: Optional[torch.Tensor] = None,
skip_init_workspace_buffer: bool = False,
): ):
super().__init__( super().__init__(
model_runner, model_runner,
skip_prefill, skip_prefill,
kv_indptr_buf, kv_indptr_buf,
q_indptr_decode_buf, q_indptr_decode_buf,
skip_init_workspace_buffer=True,
) )
config = model_runner.model_config config = model_runner.model_config
@@ -294,14 +296,17 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
# Workspace allocation # Workspace allocation
self.workspace_size = DEFAULT_WORKSPACE_SIZE_MB * 1024 * 1024 self.workspace_size = DEFAULT_WORKSPACE_SIZE_MB * 1024 * 1024
global global_zero_init_workspace_buffer if skip_init_workspace_buffer:
if global_zero_init_workspace_buffer is None: self.workspace_buffer = None
global_zero_init_workspace_buffer = torch.zeros( else:
self.workspace_size, global global_zero_init_workspace_buffer
dtype=torch.uint8, if global_zero_init_workspace_buffer is None:
device=model_runner.device, global_zero_init_workspace_buffer = torch.zeros(
) self.workspace_size,
self.workspace_buffer = global_zero_init_workspace_buffer dtype=torch.uint8,
device=model_runner.device,
)
self.workspace_buffer = global_zero_init_workspace_buffer
# CUDA graph state # CUDA graph state
self.decode_cuda_graph_metadata = {} self.decode_cuda_graph_metadata = {}
@@ -378,6 +383,10 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
return block_kv_indices 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( def init_cuda_graph_state(
self, self,
max_bs: int, max_bs: int,
@@ -672,9 +681,6 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
else: else:
return super().init_forward_metadata(forward_batch) 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( def pad_draft_extend_query(
self, self,
q: torch.Tensor, q: torch.Tensor,