[refactor] unify cuda-graph capture/replay across attention backends (#26665)

Co-authored-by: Claude Sonnet 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-05-29 12:46:42 -07:00
committed by GitHub
co-authored by Claude Sonnet 4.6
parent 7fb7b41a3e
commit ff8ed7a302
19 changed files with 1073 additions and 1582 deletions
@@ -469,33 +469,15 @@ class AscendAttnBackend(AttentionBackend):
device=self.device, device=self.device,
) )
def init_forward_metadata_capture_cuda_graph( def _init_cuda_graph_metadata(
self, self,
bs: int, bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], seq_lens: torch.Tensor,
): ) -> "ForwardMetadata":
"""Create and store the per-bs ForwardMetadata for CUDA graph capture."""
metadata = ForwardMetadata() metadata = ForwardMetadata()
metadata.block_tables = self.graph_metadata["block_tables"][:bs, :] metadata.block_tables = self.graph_metadata["block_tables"][:bs, :]
if self.is_dllm_model:
max_len = int(seq_lens[:bs].max().item())
max_seq_pages = (max_len + self.page_size - 1) // self.page_size
metadata.block_tables[:bs, :max_seq_pages].copy_(
(
self.req_to_token[req_pool_indices[:bs], :max_len][
:, :: self.page_size
]
// self.page_size
).to(torch.int32)
)
metadata.block_tables[:bs, max_seq_pages:].fill_(0)
metadata.block_tables[bs:, :].fill_(0)
if self.is_hybrid_swa: if self.is_hybrid_swa:
metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :] metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :]
metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist() metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist()
@@ -515,7 +497,7 @@ class AscendAttnBackend(AttentionBackend):
) )
else: else:
metadata.actual_seq_lengths_q = torch.tensor( metadata.actual_seq_lengths_q = torch.tensor(
[1 + i * 1 for i in range(bs)], [1 + i for i in range(bs)],
dtype=torch.int32, dtype=torch.int32,
device=seq_lens.device, device=seq_lens.device,
) )
@@ -528,13 +510,11 @@ class AscendAttnBackend(AttentionBackend):
metadata.seq_lens_list_cumsum = ( metadata.seq_lens_list_cumsum = (
torch.cumsum(extend_seq_lens_cpu_int, dim=0).int().tolist() torch.cumsum(extend_seq_lens_cpu_int, dim=0).int().tolist()
) )
if ( if (
self.q_head_num_padding is not None self.q_head_num_padding is not None
and self.q_head_num_padding > self.tp_q_head_num and self.q_head_num_padding > self.tp_q_head_num
): ):
# In the MLA architecture, the FIA kernel requires the head count to be a power of 2. dtype = self.model_dtype if self.model_dtype is not None else torch.bfloat16
# Therefore, we pad the head dimension accordingly and initialize an empty tensor for padding.
metadata.nope_padding = torch.empty( metadata.nope_padding = torch.empty(
[ [
bs, bs,
@@ -542,9 +522,7 @@ class AscendAttnBackend(AttentionBackend):
self.q_head_num_padding - self.tp_q_head_num, self.q_head_num_padding - self.tp_q_head_num,
self.kv_lora_rank, self.kv_lora_rank,
], ],
dtype=( dtype=dtype,
self.model_dtype if self.model_dtype is not None else torch.bfloat16
),
device=seq_lens.device, device=seq_lens.device,
) )
metadata.rope_padding = torch.empty( metadata.rope_padding = torch.empty(
@@ -554,16 +532,33 @@ class AscendAttnBackend(AttentionBackend):
self.q_head_num_padding - self.tp_q_head_num, self.q_head_num_padding - self.tp_q_head_num,
self.qk_rope_head_dim, self.qk_rope_head_dim,
], ],
dtype=( dtype=dtype,
self.model_dtype if self.model_dtype is not None else torch.bfloat16
),
device=seq_lens.device, device=seq_lens.device,
) )
self.graph_metadata[bs] = metadata self.graph_metadata[bs] = metadata
self.forward_metadata = metadata return metadata
self.graph_mode = True def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
self._init_cuda_graph_metadata(bs, forward_mode, seq_lens)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
@@ -93,19 +93,16 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
): ):
if forward_mode.is_draft_extend(True): self.init_forward_metadata_replay_cuda_graph(
return bs=bs,
super().init_forward_metadata_capture_cuda_graph( req_pool_indices=req_pool_indices,
bs, seq_lens=seq_lens,
num_tokens, seq_lens_sum=None,
req_pool_indices, encoder_lens=encoder_lens,
seq_lens, forward_mode=forward_mode,
encoder_lens, spec_info=spec_info,
forward_mode, seq_lens_cpu=seq_lens.cpu(),
spec_info,
) )
self.prepare_gdn_inputs(bs, forward_mode, spec_info)
self.graph_mode = True
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
@@ -1492,423 +1492,16 @@ class AiterAttnBackend(AttentionBackend):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
): ):
self.init_forward_metadata_replay_cuda_graph(
num_kv_splits = None bs=bs,
# num_kv_splits_indptr = None req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
work_metadata = None seq_lens_sum=None,
work_info_set = None encoder_lens=encoder_lens,
work_indptr = None forward_mode=forward_mode,
spec_info=spec_info,
reduce_indptr = None seq_lens_cpu=seq_lens.cpu(),
reduce_final_map = None
reduce_partial_map = None
swa_page_table = None
max_kv_len = torch.max(seq_lens).item()
if forward_mode.is_decode_or_idle():
qo_indptr = None
kv_last_page_len = None
max_q_len = None
if spec_info is None or (
self.use_triton_unified_attention and not self.use_mla
):
max_num_blocks_per_seq = (
self.max_context_len + self.page_size - 1
) // self.page_size
if not self.use_triton_unified_attention:
kv_indptr = self.kv_indptr
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
) )
else:
max_q_len = 1
kv_indices = self.cuda_graph_page_table
if self.use_sliding_window_kv_pool:
swa_page_table = self.cuda_graph_swa_page_table
if spec_info is not None:
self._build_unified_page_table_from_spec(
spec_info,
bs,
dest_buf=kv_indices,
swa_dest_buf=swa_page_table,
)
else:
page_indices = self.req_to_token[
req_pool_indices[:bs], :max_kv_len
]
if self.use_sliding_window_kv_pool:
swa_page_indices = (
self.token_to_kv_pool.translate_loc_from_full_to_swa(
page_indices
)
)
page_indices = self._transform_table_1_to_real(page_indices)
swa_page_indices = self._transform_table_1_to_real(
swa_page_indices
)
new_rows = swa_page_indices.shape[0]
new_cols = swa_page_indices.shape[1]
kv_indices[:new_rows, :new_cols].copy_(page_indices)
swa_page_table = self.cuda_graph_swa_page_table
swa_page_table[:new_rows, :new_cols].copy_(swa_page_indices)
elif self.page_size > 1:
page_indices = self._transform_table_1_to_real(page_indices)
new_rows = page_indices.shape[0]
new_cols = page_indices.shape[1]
kv_indices[:new_rows, :new_cols].copy_(page_indices)
qo_indptr = self.qo_indptr_unified_decode[: bs + 1]
kv_indptr = None
else:
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
if self.use_mla:
qo_indptr = self.qo_indptr_[: bs + 1]
qo_indptr[1 : bs + 1] = torch.cumsum(
self.cuda_graph_kv_last_page_len[:bs], dim=0
)
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = 1
if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
kv_last_page_len,
self.work_metadata,
self.work_info_set,
self.work_indptr,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
work_metadata = self.work_metadata
work_info_set = self.work_info_set
work_indptr = self.work_indptr
reduce_indptr = self.reduce_indptr
reduce_final_map = self.reduce_final_map
reduce_partial_map = self.reduce_partial_map
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
max_kv_len,
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
swa_page_table=swa_page_table,
)
elif forward_mode.is_target_verify():
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
(1 + bs) * self.num_draft_tokens,
step=self.num_draft_tokens,
dtype=torch.int32,
device=self.device,
)
if self.use_mla:
kv_lens = seq_lens + self.num_draft_tokens
else:
kv_lens = seq_lens
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
kv_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = self.num_draft_tokens
if self.use_mla:
if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
kv_last_page_len,
self.work_metadata,
self.work_info_set,
self.work_indptr,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
work_metadata = self.work_metadata
work_info_set = self.work_info_set
work_indptr = self.work_indptr
reduce_indptr = self.reduce_indptr
reduce_final_map = self.reduce_final_map
reduce_partial_map = self.reduce_partial_map
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
max_kv_len,
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
)
else:
if self._use_unified_verify:
max_num_blocks_per_seq = (
self.max_context_len + self.page_size - 1
) // self.page_size
page_table = self.cuda_graph_page_table[:bs]
swa_page_table = None
if self.use_sliding_window_kv_pool:
swa_page_table = self.cuda_graph_swa_page_table.view(
-1, max_num_blocks_per_seq
)[:bs]
_page_table, _qo_indptr, _max_q_len, _swa_page_table = (
self._build_verify_unified_metadata(
bs,
seq_lens,
req_pool_indices,
self.num_draft_tokens,
page_table_dest=page_table,
swa_page_table_dest=swa_page_table,
)
)
max_kv_len = max_num_blocks_per_seq * self.page_size
self.forward_metadata = ForwardMetadata(
None,
_page_table,
_qo_indptr,
kv_last_page_len,
_max_q_len,
max_kv_len,
max_extend_len=_max_q_len,
swa_page_table=_swa_page_table,
)
else:
custom_mask = self.cuda_graph_custom_mask
custom_mask[: spec_info.custom_mask.shape[0]] = (
spec_info.custom_mask
)
seq_mask_len = max_q_len * (seq_lens + max_q_len)
mask_indptr = self.mask_indptr
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len[:bs], dim=0)
mask_indptr = mask_indptr[: bs + 1]
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
max_kv_len,
custom_mask=custom_mask,
mask_indptr=mask_indptr,
max_extend_len=max_q_len,
)
elif forward_mode.is_draft_extend_v2():
# EAGLE V2: Uses fixed num_draft_tokens per batch
self._ensure_spec_v2_topk_supported()
num_tokens_per_bs = self._resolve_v2_num_draft_tokens()
qo_indptr = self._set_uniform_qo_indptr(bs, num_tokens_per_bs, self.device)
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = num_tokens_per_bs
if self.use_mla and _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
kv_last_page_len,
self.work_metadata,
self.work_info_set,
self.work_indptr,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
work_metadata = self.work_metadata
work_info_set = self.work_info_set
work_indptr = self.work_indptr
reduce_indptr = self.reduce_indptr
reduce_final_map = self.reduce_final_map
reduce_partial_map = self.reduce_partial_map
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
max_kv_len,
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
)
elif forward_mode.is_draft_extend():
# EAGLE V1: Uses speculative_num_steps + 1
num_tokens_per_bs = self.speculative_num_steps + 1
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
bs * num_tokens_per_bs + 1,
step=num_tokens_per_bs,
dtype=torch.int32,
device=self.device,
)
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
if self.use_mla:
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = num_tokens_per_bs
if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
kv_last_page_len,
self.work_metadata,
self.work_info_set,
self.work_indptr,
self.reduce_indptr,
self.reduce_final_map,
self.reduce_partial_map,
max_q_len,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
work_metadata = self.work_metadata
work_info_set = self.work_info_set
work_indptr = self.work_indptr
reduce_indptr = self.reduce_indptr
reduce_final_map = self.reduce_final_map
reduce_partial_map = self.reduce_partial_map
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
max_kv_len,
work_metadata=work_metadata,
work_info_set=work_info_set,
work_indptr=work_indptr,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
)
else:
# Non-MLA draft_extend cuda graph: use triton extend kernel
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
None,
num_tokens_per_bs,
None,
custom_mask=None,
mask_indptr=None,
max_extend_len=num_tokens_per_bs,
)
else:
raise ValueError(f"Invalid mode: {forward_mode=}")
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
@@ -1934,7 +1527,11 @@ class AiterAttnBackend(AttentionBackend):
reduce_partial_map = None reduce_partial_map = None
swa_page_table = None swa_page_table = None
max_kv_len = seq_lens_cpu.max().item() max_kv_len = (
seq_lens_cpu.max().item()
if seq_lens_cpu is not None
else torch.max(seq_lens).item()
)
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
qo_indptr = None qo_indptr = None
@@ -153,20 +153,18 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
): ):
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle() and spec_info is None:
if spec_info is None: self.init_forward_metadata_replay_cuda_graph(
max_seqlen_pad = self.cuda_graph_kv_indices.shape[1] bs=bs,
req_pool_indices=req_pool_indices,
create_flashmla_kv_indices_triton[(bs,)]( seq_lens=seq_lens,
self.req_to_token, seq_lens_sum=None,
req_pool_indices, encoder_lens=encoder_lens,
seq_lens, forward_mode=forward_mode,
None, spec_info=spec_info,
self.cuda_graph_kv_indices, seq_lens_cpu=None,
self.req_to_token.stride(0),
self.cuda_graph_kv_indices.stride(0),
PAGED_SIZE=PAGE_SIZE,
) )
max_seqlen_pad = self.cuda_graph_kv_indices.shape[1]
self.forward_metadata = CutlassMLADecodeMetadata( self.forward_metadata = CutlassMLADecodeMetadata(
self.cuda_graph_mla_workspace, self.cuda_graph_mla_workspace,
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad], self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
@@ -193,15 +191,11 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor], seq_lens_cpu: Optional[torch.Tensor],
): ):
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
assert seq_lens_cpu is not None
seq_lens = seq_lens[:bs]
create_flashmla_kv_indices_triton[(bs,)]( create_flashmla_kv_indices_triton[(bs,)](
self.req_to_token, self.req_to_token,
req_pool_indices[:bs], req_pool_indices[:bs],
seq_lens, seq_lens[:bs],
None, None,
self.cuda_graph_kv_indices, self.cuda_graph_kv_indices,
self.req_to_token.stride(0), self.req_to_token.stride(0),
@@ -749,47 +749,39 @@ class DeepseekV4AttnBackend(
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
) -> None: ) -> None:
from types import SimpleNamespace
assert req_pool_indices.size(0) == bs assert req_pool_indices.size(0) == bs
assert seq_lens.size(0) == bs assert seq_lens.size(0) == bs
bucket = _GraphBucket.of(forward_mode) bucket = _GraphBucket.of(forward_mode)
raw_type: Optional[type] = None
if bucket == _GraphBucket.DECODE_OR_IDLE: if bucket == _GraphBucket.DECODE_OR_IDLE:
metadata = self.init_forward_metadata_decode( dummy_cache_loc = torch.zeros_like(seq_lens)
max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=torch.zeros_like(seq_lens),
)
raw_type = DSV4RawDecodeMetadata
elif bucket == _GraphBucket.TARGET_VERIFY: elif bucket == _GraphBucket.TARGET_VERIFY:
out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs)
metadata = self.init_forward_metadata_target_verify(
max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc,
use_prefill_cuda_graph=True,
)
raw_type = DSV4RawVerifyMetadata
elif bucket == _GraphBucket.DRAFT_EXTEND:
num_tokens_per_bs = num_tokens // bs
metadata = self.init_forward_metadata_draft_extend(
max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens.tolist(),
num_tokens_per_bs=num_tokens_per_bs,
use_prefill_cuda_graph=True,
)
else: else:
raise NotImplementedError(f"{forward_mode=} not supported yet") dummy_cache_loc = None
self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = metadata self._replay_forward_batch = SimpleNamespace(
self.forward_metadata = metadata out_cache_loc=dummy_cache_loc,
if raw_type is not None: forward_mode=forward_mode,
)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=int(seq_lens.sum().item()),
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
# Preserve _current_capture_raw for on_after_cuda_graph_warmup
metadata = self.forward_metadata
self._current_capture_raw = ( self._current_capture_raw = (
metadata if isinstance(metadata, raw_type) else None metadata
if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata))
else None
) )
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
@@ -892,6 +884,11 @@ class DeepseekV4AttnBackend(
], ],
bucket: _GraphBucket, bucket: _GraphBucket,
) -> None: ) -> None:
if bs not in self.cuda_graph_metadata_of_bucket_and_bs[bucket]:
# First call (from capture): store the new metadata directly.
self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = temp_metadata
self.forward_metadata = temp_metadata
return
chosen_metadata = self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] chosen_metadata = self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs]
chosen_metadata.copy_(temp_metadata) chosen_metadata.copy_(temp_metadata)
self.forward_metadata = chosen_metadata self.forward_metadata = chosen_metadata
@@ -748,47 +748,39 @@ class DeepseekV4HipRadixBackend(
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
) -> None: ) -> None:
from types import SimpleNamespace
assert req_pool_indices.size(0) == bs assert req_pool_indices.size(0) == bs
assert seq_lens.size(0) == bs assert seq_lens.size(0) == bs
bucket = _GraphBucket.of(forward_mode) bucket = _GraphBucket.of(forward_mode)
raw_type: Optional[type] = None
if bucket == _GraphBucket.DECODE_OR_IDLE: if bucket == _GraphBucket.DECODE_OR_IDLE:
metadata = self.init_forward_metadata_decode( dummy_cache_loc = torch.zeros_like(seq_lens)
max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=torch.zeros_like(seq_lens),
)
raw_type = DSV4RawDecodeMetadata
elif bucket == _GraphBucket.TARGET_VERIFY: elif bucket == _GraphBucket.TARGET_VERIFY:
out_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs) dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs)
metadata = self.init_forward_metadata_target_verify(
max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
out_cache_loc=out_cache_loc,
use_prefill_cuda_graph=True,
)
raw_type = DSV4RawVerifyMetadata
elif bucket == _GraphBucket.DRAFT_EXTEND:
num_tokens_per_bs = num_tokens // bs
metadata = self.init_forward_metadata_draft_extend(
max_seq_len=self.MAX_SEQ_LEN_FOR_CAPTURE,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens.tolist(),
num_tokens_per_bs=num_tokens_per_bs,
use_prefill_cuda_graph=True,
)
else: else:
raise NotImplementedError(f"{forward_mode=} not supported yet") dummy_cache_loc = None
self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = metadata self._replay_forward_batch = SimpleNamespace(
self.forward_metadata = metadata out_cache_loc=dummy_cache_loc,
if raw_type is not None: forward_mode=forward_mode,
)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=int(seq_lens.sum().item()),
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
# Preserve _current_capture_raw for on_after_cuda_graph_warmup
metadata = self.forward_metadata
self._current_capture_raw = ( self._current_capture_raw = (
metadata if isinstance(metadata, raw_type) else None metadata
if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata))
else None
) )
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
@@ -891,6 +883,11 @@ class DeepseekV4HipRadixBackend(
], ],
bucket: _GraphBucket, bucket: _GraphBucket,
) -> None: ) -> None:
if bs not in self.cuda_graph_metadata_of_bucket_and_bs[bucket]:
# First call (from capture): store the new metadata directly.
self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] = temp_metadata
self.forward_metadata = temp_metadata
return
chosen_metadata = self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs] chosen_metadata = self.cuda_graph_metadata_of_bucket_and_bs[bucket][bs]
chosen_metadata.copy_(temp_metadata) chosen_metadata.copy_(temp_metadata)
self.forward_metadata = chosen_metadata self.forward_metadata = chosen_metadata
@@ -813,19 +813,21 @@ class DeepseekSparseAttnBackend(
), ),
} }
def init_forward_metadata_capture_cuda_graph( def _build_forward_metadata_cuda_graph(
self, self,
bs: int, bs: int,
num_tokens: int, num_tokens: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor, seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor], seq_lens_cpu: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
out_cache_loc: Optional[torch.Tensor] = None,
actual_forward_mode: Optional["ForwardMode"] = None,
): ):
"""Create and store DSAMetadata for a new batch size during CUDA graph capture."""
self.set_dsa_prefill_impl(forward_batch=None) self.set_dsa_prefill_impl(forward_batch=None)
"""Initialize forward metadata for capturing CUDA graph."""
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
# Normal Decode # Normal Decode
# Get sequence information # Get sequence information
@@ -847,11 +849,11 @@ class DeepseekSparseAttnBackend(
) )
seqlens_expanded = cache_seqlens_int32 seqlens_expanded = cache_seqlens_int32
dsa_extend_seq_lens_list = [1] * num_tokens dsa_extend_seq_lens_list = [1] * bs
if self.dsa_decode_impl == "flashmla_kv": if self.dsa_decode_impl == "flashmla_kv":
flashmla_metadata = self.decode_cuda_graph_metadata[ flashmla_metadata = self.decode_cuda_graph_metadata[
"flashmla_metadata" "flashmla_metadata"
].slice(slice(0, num_tokens + 1)) ].slice(slice(0, bs + 1))
flashmla_metadata.copy_( flashmla_metadata.copy_(
self._compute_flashmla_metadata( self._compute_flashmla_metadata(
cache_seqlens=dsa_cache_seqlens_int32, cache_seqlens=dsa_cache_seqlens_int32,
@@ -969,6 +971,28 @@ class DeepseekSparseAttnBackend(
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
self.forward_metadata = metadata self.forward_metadata = metadata
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Initialize forward metadata for capturing CUDA graph."""
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
bs: int, bs: int,
@@ -985,6 +1009,20 @@ class DeepseekSparseAttnBackend(
"""Initialize forward metadata for replaying CUDA graph.""" """Initialize forward metadata for replaying CUDA graph."""
assert seq_lens_cpu is not None assert seq_lens_cpu is not None
if bs not in self.decode_cuda_graph_metadata:
self._build_forward_metadata_cuda_graph(
bs,
None,
req_pool_indices,
seq_lens,
seq_lens_cpu,
forward_mode,
spec_info,
out_cache_loc,
actual_forward_mode,
)
return
self.set_dsa_prefill_impl(forward_batch=None) self.set_dsa_prefill_impl(forward_batch=None)
seq_lens = seq_lens[:bs] seq_lens = seq_lens[:bs]
@@ -532,16 +532,13 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
), ),
} }
def init_forward_metadata_capture_cuda_graph( def _bind_metadata_buffers(
self, self,
bs: int, bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor, req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[None],
): ):
"""Allocate persistent metadata buffers for CUDA graph capture."""
metadata = DualChunkFlashAttentionMetadata() metadata = DualChunkFlashAttentionMetadata()
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
@@ -580,6 +577,36 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
self.forward_metadata = metadata self.forward_metadata = metadata
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[None],
):
self._bind_metadata_buffers(bs, req_pool_indices, forward_mode)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens.cpu(),
)
# Restore max_seq_len scalars — replay sets actual values but CUDA graph
# needs the safe upper bound baked in at capture time.
if forward_mode.is_decode_or_idle():
md = self.forward_metadata
md.max_seq_len = self.max_context_len
md.max_seq_len_intra = self.max_context_len
md.max_seq_len_succ = self.max_context_len
md.max_seq_len_inter = self.max_context_len
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
bs: int, bs: int,
@@ -1700,43 +1700,37 @@ class FlashAttentionBackend(AttentionBackend):
# For decoder-only models, skip encoder_metadata allocation # For decoder-only models, skip encoder_metadata allocation
self.encoder_metadata = {} self.encoder_metadata = {}
def init_forward_metadata_capture_cuda_graph( def _bind_metadata_buffers(
self, self,
bs: int, bs: int,
num_tokens: int, num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor], encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
): device: torch.device,
"""Initialize forward metadata for capturing CUDA graph.""" ) -> tuple:
metadata = FlashAttentionMetadata() """Create FlashAttentionMetadata with pre-allocated buffer slice refs.
# metadata_expand is needed for Spec Decoding when top k > 1 Assigns all buffer slice references but does NOT fill data values.
Stores the new metadata object(s) in the appropriate lookup dicts.
Returns (metadata, metadata_expand).
"""
metadata = FlashAttentionMetadata()
metadata_expand = FlashAttentionMetadata() metadata_expand = FlashAttentionMetadata()
device = seq_lens.device
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
if spec_info is not None: if spec_info is not None:
# Draft Decode
if self.topk <= 1: if self.topk <= 1:
# When topk = 1, we use the normal decode metadata # Draft Decode topk=1
metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[ metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[
"cache_seqlens" "cache_seqlens"
][:bs] ][:bs]
metadata.max_seq_len_k = seq_lens.max().item() + (
self.speculative_step_id + 1
)
metadata.cu_seqlens_q = self.decode_cuda_graph_metadata[ metadata.cu_seqlens_q = self.decode_cuda_graph_metadata[
"cu_seqlens_q" "cu_seqlens_q"
][: bs + 1] ][: bs + 1]
metadata.cu_seqlens_k = torch.nn.functional.pad( metadata.cu_seqlens_k = self.decode_cuda_graph_metadata[
torch.cumsum( "cu_seqlens_k"
metadata.cache_seqlens_int32, dim=0, dtype=torch.int32 ][: bs + 1]
),
(1, 0),
)
metadata.page_table = self.decode_cuda_graph_metadata[ metadata.page_table = self.decode_cuda_graph_metadata[
"page_table_draft_decode" "page_table_draft_decode"
][:bs, :] ][:bs, :]
@@ -1746,13 +1740,11 @@ class FlashAttentionBackend(AttentionBackend):
][:bs, :] ][:bs, :]
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
else: else:
# When top k > 1, we need two specific draft decode metadata, and then merge states # Draft Decode topk>1: two metadata objects
# 1. The first half of metadata for prefix tokens
metadata.cache_seqlens_int32 = ( metadata.cache_seqlens_int32 = (
self.draft_decode_metadata_topk_normal["cache_seqlens"][:bs] self.draft_decode_metadata_topk_normal["cache_seqlens"][:bs]
) )
metadata.max_seq_len_q = self.topk metadata.max_seq_len_q = self.topk
metadata.max_seq_len_k = seq_lens.max().item()
metadata.cu_seqlens_q = self.draft_decode_metadata_topk_normal[ metadata.cu_seqlens_q = self.draft_decode_metadata_topk_normal[
"cu_seqlens_q" "cu_seqlens_q"
][: bs + 1] ][: bs + 1]
@@ -1763,7 +1755,6 @@ class FlashAttentionBackend(AttentionBackend):
"page_table" "page_table"
][:bs, :] ][:bs, :]
# 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk)
metadata_expand.cache_seqlens_int32 = ( metadata_expand.cache_seqlens_int32 = (
self.draft_decode_metadata_topk_expand["cache_seqlens"][ self.draft_decode_metadata_topk_expand["cache_seqlens"][
: bs * self.topk : bs * self.topk
@@ -1787,16 +1778,15 @@ class FlashAttentionBackend(AttentionBackend):
self.draft_decode_metadata_topk_expand[bs] = metadata_expand self.draft_decode_metadata_topk_expand[bs] = metadata_expand
else: else:
# Normal Decode # Normal Decode
# Get sequence information metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[
metadata.cache_seqlens_int32 = seq_lens.to(torch.int32) "cache_seqlens"
batch_size = len(seq_lens) ][:bs]
device = seq_lens.device metadata.cu_seqlens_q = self.decode_cuda_graph_metadata["cu_seqlens_q"][
metadata.cu_seqlens_k = torch.nn.functional.pad( : bs + 1
torch.cumsum(seq_lens, dim=0, dtype=torch.int32), (1, 0) ]
) metadata.cu_seqlens_k = self.decode_cuda_graph_metadata["cu_seqlens_k"][
# Precompute maximum sequence length : bs + 1
metadata.max_seq_len_k = seq_lens.max().item() ]
# Precompute page table
metadata.page_table = self.decode_cuda_graph_metadata["page_table"][ metadata.page_table = self.decode_cuda_graph_metadata["page_table"][
:bs, : :bs, :
] ]
@@ -1804,70 +1794,32 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_page_table = self.decode_cuda_graph_metadata[ metadata.swa_page_table = self.decode_cuda_graph_metadata[
"swa_page_table" "swa_page_table"
][:bs, :] ][:bs, :]
# Precompute cumulative sequence lengths
metadata.cu_seqlens_q = torch.arange(
0, batch_size + 1, dtype=torch.int32, device=device
)
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
self._maybe_update_local_attn_metadata_for_capture(metadata, batch_size)
# Compute scheduler_metadata into pre-allocated buffer for CUDA graph capture
if self._sched_meta_buf is not None:
sched = self._compute_scheduler_metadata(
batch_size,
max(metadata.max_seq_len_k, 1),
metadata.cache_seqlens_int32,
metadata.cu_seqlens_q,
)
if sched is not None:
n = sched.shape[0]
self._sched_meta_buf[:n] = sched
self._sched_meta_buf[n:] = 0
metadata.scheduler_metadata = self._sched_meta_buf[:n]
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify():
if self.topk <= 1: if self.topk <= 1:
metadata.cache_seqlens_int32 = self.target_verify_metadata[ metadata.cache_seqlens_int32 = self.target_verify_metadata[
"cache_seqlens" "cache_seqlens"
][:bs] ][:bs]
metadata.cache_seqlens_int32.copy_(
(seq_lens + self.speculative_num_draft_tokens)
)
metadata.max_seq_len_q = self.speculative_num_draft_tokens metadata.max_seq_len_q = self.speculative_num_draft_tokens
metadata.max_seq_len_k = ( metadata.cu_seqlens_q = self.target_verify_metadata["cu_seqlens_q"][
seq_lens.max().item() + self.speculative_num_draft_tokens : bs + 1
) ]
metadata.cu_seqlens_q = torch.arange(
0,
bs * self.speculative_num_draft_tokens + 1,
self.speculative_num_draft_tokens,
dtype=torch.int32,
device=device,
)
metadata.cu_seqlens_k = self.target_verify_metadata["cu_seqlens_k"][ metadata.cu_seqlens_k = self.target_verify_metadata["cu_seqlens_k"][
: (bs + 1) : (bs + 1)
] ]
metadata.page_table = self.target_verify_metadata["page_table"][:bs, :] metadata.page_table = self.target_verify_metadata["page_table"][:bs, :]
if self.use_sliding_window_kv_pool: if self.use_sliding_window_kv_pool:
metadata.swa_page_table = self.target_verify_metadata[ metadata.swa_page_table = self.target_verify_metadata[
"swa_page_table" "swa_page_table"
][:bs, :] ][:bs, :]
self.target_verify_metadata[bs] = metadata self.target_verify_metadata[bs] = metadata
else: else:
# When topk > 1, we need two specific target verify metadata, and then merge states # Target Verify topk>1: two (or three with SWA) metadata objects
# 1. The first half of metadata for prefix tokens
metadata.cache_seqlens_int32 = self.target_verify_metadata_topk_normal[ metadata.cache_seqlens_int32 = self.target_verify_metadata_topk_normal[
"cache_seqlens" "cache_seqlens"
][:bs] ][:bs]
metadata.max_seq_len_q = self.speculative_num_draft_tokens metadata.max_seq_len_q = self.speculative_num_draft_tokens
# metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item(), do this in replay
metadata.cu_seqlens_q = self.target_verify_metadata_topk_normal[ metadata.cu_seqlens_q = self.target_verify_metadata_topk_normal[
"cu_seqlens_q" "cu_seqlens_q"
][: bs + 1] ][: bs + 1]
@@ -1878,7 +1830,6 @@ class FlashAttentionBackend(AttentionBackend):
"page_table" "page_table"
][:bs, :] ][:bs, :]
# 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk)
metadata_expand.cache_seqlens_int32 = ( metadata_expand.cache_seqlens_int32 = (
self.target_verify_metadata_topk_expand["cache_seqlens"][ self.target_verify_metadata_topk_expand["cache_seqlens"][
: bs * self.speculative_num_draft_tokens : bs * self.speculative_num_draft_tokens
@@ -1891,7 +1842,6 @@ class FlashAttentionBackend(AttentionBackend):
metadata_expand.cu_seqlens_k = self.target_verify_metadata_topk_expand[ metadata_expand.cu_seqlens_k = self.target_verify_metadata_topk_expand[
"cu_seqlens_k" "cu_seqlens_k"
][: bs * self.speculative_num_draft_tokens + 1] ][: bs * self.speculative_num_draft_tokens + 1]
metadata_expand.page_table = self.target_verify_metadata_topk_expand[ metadata_expand.page_table = self.target_verify_metadata_topk_expand[
"page_table" "page_table"
][: bs * self.speculative_num_draft_tokens] ][: bs * self.speculative_num_draft_tokens]
@@ -1913,7 +1863,6 @@ class FlashAttentionBackend(AttentionBackend):
metadata_swa.cu_seqlens_k = self.target_verify_metadata_topk_swa[ metadata_swa.cu_seqlens_k = self.target_verify_metadata_topk_swa[
"cu_seqlens_k" "cu_seqlens_k"
][: bs * self.speculative_num_draft_tokens + 1] ][: bs * self.speculative_num_draft_tokens + 1]
metadata_swa.page_table = self.target_verify_metadata_topk_swa[ metadata_swa.page_table = self.target_verify_metadata_topk_swa[
"page_table" "page_table"
][: bs * self.speculative_num_draft_tokens] ][: bs * self.speculative_num_draft_tokens]
@@ -1921,33 +1870,20 @@ class FlashAttentionBackend(AttentionBackend):
metadata.swa_spec_metadata = metadata_swa metadata.swa_spec_metadata = metadata_swa
elif forward_mode.is_draft_extend(include_v2=True): elif forward_mode.is_draft_extend(include_v2=True):
num_tokens_per_bs = num_tokens // bs
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][ metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs :bs
] ]
metadata.cache_seqlens_int32.copy_(seq_lens)
num_tokens_per_bs = num_tokens // bs
metadata.max_seq_len_q = num_tokens_per_bs metadata.max_seq_len_q = num_tokens_per_bs
metadata.max_seq_len_k = seq_lens.max().item() metadata.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1]
metadata.cu_seqlens_q = torch.arange(
0,
bs * num_tokens_per_bs + 1,
num_tokens_per_bs,
dtype=torch.int32,
device=device,
)
metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][ metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][
: (bs + 1) : (bs + 1)
] ]
metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :] metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :]
if self.use_sliding_window_kv_pool: if self.use_sliding_window_kv_pool:
metadata.swa_page_table = self.draft_extend_metadata["swa_page_table"][ metadata.swa_page_table = self.draft_extend_metadata["swa_page_table"][
:bs, : :bs, :
] ]
self.draft_extend_metadata[bs] = metadata self.draft_extend_metadata[bs] = metadata
if encoder_lens is not None: if encoder_lens is not None:
@@ -1958,13 +1894,81 @@ class FlashAttentionBackend(AttentionBackend):
metadata.encoder_cu_seqlens_k = self.encoder_metadata[ metadata.encoder_cu_seqlens_k = self.encoder_metadata[
"encoder_cu_seqlens_k" "encoder_cu_seqlens_k"
][: (encoder_bs + 1)] ][: (encoder_bs + 1)]
metadata.encoder_page_table = self.encoder_metadata["encoder_page_table"][ metadata.encoder_page_table = self.encoder_metadata["encoder_page_table"][
:bs, : :bs, :
] ]
self.forward_metadata = metadata return metadata, metadata_expand
self.forward_metadata_spec_decode_expand = metadata_expand
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Initialize forward metadata for capturing CUDA graph."""
seq_lens_cpu = seq_lens.cpu()
self._bind_metadata_buffers(
bs, num_tokens, encoder_lens, forward_mode, spec_info, seq_lens.device
)
if forward_mode.is_decode_or_idle() and spec_info is not None and self.topk > 1:
# topk>1 draft decode: replay needs out_cache_loc which capture doesn't have;
# set forward_metadata directly and let actual CUDA graph replay fill data.
self.forward_metadata = self.draft_decode_metadata_topk_normal[bs]
self.forward_metadata_spec_decode_expand = (
self.draft_decode_metadata_topk_expand[bs]
)
return
if forward_mode.is_target_verify() and self.topk > 1:
# topk>1 target verify: replay needs spec_info.positions and .custom_mask
# which are not populated at capture time.
self.forward_metadata = self.target_verify_metadata_topk_normal[bs]
self.forward_metadata_spec_decode_expand = (
self.target_verify_metadata_topk_expand[bs]
)
return
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
if forward_mode.is_decode_or_idle() and spec_info is None:
# Local attention and scheduler metadata require capture-time slice sizing.
# Both depend on data already filled by replay above.
metadata = self.decode_cuda_graph_metadata[bs]
self._maybe_update_local_attn_metadata_for_capture(metadata, bs)
if self._sched_meta_buf is not None:
sched = self._compute_scheduler_metadata(
bs,
max(metadata.max_seq_len_k, 1),
metadata.cache_seqlens_int32,
metadata.cu_seqlens_q,
)
if sched is not None:
n = sched.shape[0]
self._sched_meta_buf[:n] = sched
self._sched_meta_buf[n:] = 0
metadata.scheduler_metadata = self._sched_meta_buf[:n]
if forward_mode.is_draft_extend(include_v2=True):
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it to
# max(num_accept_tokens_cpu) which is None/empty at capture time,
# falling back to 1. Restore the correct upper bound so the kernel
# sees num_tokens_per_bs (not 1) for all replays of this graph.
self.forward_metadata.max_seq_len_q = num_tokens // bs
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
@@ -557,6 +557,81 @@ class FlashInferAttnBackend(AttentionBackend):
self.cuda_graph_qk_indptr = [x.clone() for x in self.kv_indptr] self.cuda_graph_qk_indptr = [x.clone() for x in self.kv_indptr]
self.cuda_graph_qo_indptr = [x.clone() for x in self.kv_indptr] self.cuda_graph_qo_indptr = [x.clone() for x in self.kv_indptr]
def _create_decode_wrappers(self, bs: int, num_tokens: int) -> list:
return [
BatchDecodeWithPagedKVCacheWrapper(
self.workspace_buffer,
"NHD",
backend=self.decode_backend,
use_cuda_graph=True,
use_tensor_cores=self.decode_use_tensor_cores,
paged_kv_indptr_buffer=self.kv_indptr[i][: num_tokens + 1],
paged_kv_indices_buffer=self.cuda_graph_kv_indices[i],
paged_kv_last_page_len_buffer=self.kv_last_page_len[:num_tokens],
)
for i in range(self.num_wrappers)
]
def _create_prefill_wrappers(self, bs: int, use_custom_mask: bool = False) -> list:
# FlashInfer's prefill wrapper decides mask mode based on whether
# `custom_mask_buf` is initialized (not whether a custom mask is provided).
# For cases like DFLASH draft (ENCODER_ONLY / non-causal) we do NOT use a
# custom mask, so we must avoid initializing `custom_mask_buf`, otherwise
# FlashInfer will treat the (zero) buffer as a real mask and block attention.
wrappers = []
for i in range(self.num_wrappers):
extra = (
{
"custom_mask_buf": self.cuda_graph_custom_mask,
"mask_indptr_buf": self.cuda_graph_qk_indptr[i][: bs + 1],
}
if use_custom_mask
else {}
)
wrappers.append(
BatchPrefillWithPagedKVCacheWrapper(
self.workspace_buffer,
"NHD",
use_cuda_graph=True,
backend=self.prefill_backend,
qo_indptr_buf=self.cuda_graph_qo_indptr[i][: bs + 1],
paged_kv_indptr_buf=self.kv_indptr[i][: bs + 1],
paged_kv_indices_buf=self.cuda_graph_kv_indices[i],
paged_kv_last_page_len_buf=self.kv_last_page_len[:bs],
**extra,
)
)
return wrappers
def _prepare_cuda_graph_metadata(
self,
bs: int,
num_tokens: int,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
) -> None:
if forward_mode.is_decode_or_idle():
decode_wrappers = self._create_decode_wrappers(bs, num_tokens)
self.decode_cuda_graph_metadata[bs] = decode_wrappers
self.forward_metadata = DecodeMetadata(decode_wrappers)
elif (
forward_mode.is_target_verify()
or forward_mode.is_draft_extend()
or forward_mode.is_dllm_extend()
):
use_custom_mask = (
forward_mode.is_target_verify()
and spec_info is not None
and getattr(spec_info, "custom_mask", None) is not None
)
prefill_wrappers = self._create_prefill_wrappers(bs, use_custom_mask)
self.prefill_cuda_graph_metadata[bs] = prefill_wrappers
self.forward_metadata = PrefillMetadata(
prefill_wrappers, forward_mode.is_dllm_extend(), False
)
else:
raise ValueError(f"Invalid mode: {forward_mode=}")
def init_forward_metadata_capture_cuda_graph( def init_forward_metadata_capture_cuda_graph(
self, self,
bs: int, bs: int,
@@ -567,148 +642,24 @@ class FlashInferAttnBackend(AttentionBackend):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
): ):
seq_lens_sum = seq_lens.sum().item()
seq_lens_cpu = seq_lens.cpu()
self._prepare_cuda_graph_metadata(bs, num_tokens, forward_mode, spec_info)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=seq_lens_sum,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
# fast_decode_plan requires _cached_module set by the initial full
# begin_forward call above; install it only after that first plan runs.
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
decode_wrappers = [] for w in self.decode_cuda_graph_metadata[bs]:
for i in range(self.num_wrappers): w.begin_forward = partial(fast_decode_plan, w)
decode_wrappers.append(
BatchDecodeWithPagedKVCacheWrapper(
self.workspace_buffer,
"NHD",
backend=self.decode_backend,
use_cuda_graph=True,
use_tensor_cores=self.decode_use_tensor_cores,
paged_kv_indptr_buffer=self.kv_indptr[i][: num_tokens + 1],
paged_kv_indices_buffer=self.cuda_graph_kv_indices[i],
paged_kv_last_page_len_buffer=self.kv_last_page_len[
:num_tokens
],
)
)
seq_lens_sum = seq_lens.sum().item()
self.indices_updater_decode.update(
req_pool_indices,
seq_lens,
seq_lens.cpu(), # may add a little overhead in capture stage
seq_lens_sum,
decode_wrappers=decode_wrappers,
encoder_lens=encoder_lens,
spec_info=spec_info,
fixed_split_size=None,
disable_split_kv=self.disable_cuda_graph_kv_split,
)
self.decode_cuda_graph_metadata[bs] = decode_wrappers
self.forward_metadata = DecodeMetadata(decode_wrappers)
for i in range(self.num_wrappers):
decode_wrappers[i].begin_forward = partial(
fast_decode_plan, decode_wrappers[i]
)
elif forward_mode.is_target_verify():
# FlashInfer's prefill wrapper decides mask mode based on whether
# `custom_mask_buf` is initialized (not whether a custom mask is provided).
# For cases like DFLASH draft (ENCODER_ONLY / non-causal) we do NOT use a
# custom mask, so we must avoid initializing `custom_mask_buf`, otherwise
# FlashInfer will treat the (zero) buffer as a real mask and block attention.
use_custom_mask = (
spec_info is not None
and getattr(spec_info, "custom_mask", None) is not None
)
prefill_wrappers = []
for i in range(self.num_wrappers):
wrapper_kwargs = {}
if use_custom_mask:
wrapper_kwargs = {
"custom_mask_buf": self.cuda_graph_custom_mask,
"mask_indptr_buf": self.cuda_graph_qk_indptr[i][: bs + 1],
}
prefill_wrappers.append(
BatchPrefillWithPagedKVCacheWrapper(
self.workspace_buffer,
"NHD",
use_cuda_graph=True,
backend=self.prefill_backend,
qo_indptr_buf=self.cuda_graph_qo_indptr[i][: bs + 1],
paged_kv_indptr_buf=self.kv_indptr[i][: bs + 1],
paged_kv_indices_buf=self.cuda_graph_kv_indices[i],
paged_kv_last_page_len_buf=self.kv_last_page_len[:bs],
**wrapper_kwargs,
)
)
seq_lens_sum = seq_lens.sum().item()
self.indices_updater_prefill.update(
req_pool_indices,
seq_lens,
seq_lens.cpu(), # may add a little overhead in capture stage
seq_lens_sum,
prefix_lens=None,
prefill_wrappers=prefill_wrappers,
use_ragged=False,
encoder_lens=encoder_lens,
spec_info=spec_info,
)
self.prefill_cuda_graph_metadata[bs] = prefill_wrappers
self.forward_metadata = PrefillMetadata(prefill_wrappers, False, False)
elif forward_mode.is_draft_extend():
prefill_wrappers = []
for i in range(self.num_wrappers):
prefill_wrappers.append(
BatchPrefillWithPagedKVCacheWrapper(
self.workspace_buffer,
"NHD",
backend=self.prefill_backend,
use_cuda_graph=True,
qo_indptr_buf=self.cuda_graph_qo_indptr[i][: bs + 1],
paged_kv_indptr_buf=self.kv_indptr[i][: bs + 1],
paged_kv_indices_buf=self.cuda_graph_kv_indices[i],
paged_kv_last_page_len_buf=self.kv_last_page_len[:bs],
)
)
seq_lens_sum = seq_lens.sum().item()
self.indices_updater_prefill.update(
req_pool_indices,
seq_lens,
seq_lens.cpu(), # may add a little overhead in capture stage
seq_lens_sum,
prefix_lens=None,
prefill_wrappers=prefill_wrappers,
use_ragged=False,
encoder_lens=encoder_lens,
spec_info=spec_info,
)
self.prefill_cuda_graph_metadata[bs] = prefill_wrappers
self.forward_metadata = PrefillMetadata(prefill_wrappers, False, False)
elif forward_mode.is_dllm_extend():
prefill_wrappers = []
for i in range(self.num_wrappers):
prefill_wrappers.append(
BatchPrefillWithPagedKVCacheWrapper(
self.workspace_buffer,
"NHD",
backend=self.prefill_backend,
use_cuda_graph=True,
qo_indptr_buf=self.cuda_graph_qo_indptr[i][: bs + 1],
paged_kv_indptr_buf=self.kv_indptr[i][: bs + 1],
paged_kv_indices_buf=self.cuda_graph_kv_indices[i],
paged_kv_last_page_len_buf=self.kv_last_page_len[:bs],
)
)
seq_lens_sum = seq_lens.sum().item()
self.indices_updater_prefill.update(
req_pool_indices,
seq_lens,
seq_lens.cpu(), # may add a little overhead in capture stage
seq_lens_sum,
prefix_lens=seq_lens - self.dllm_config.block_size,
prefill_wrappers=prefill_wrappers,
use_ragged=not self.use_paged,
encoder_lens=encoder_lens,
spec_info=None,
)
self.prefill_cuda_graph_metadata[bs] = prefill_wrappers
self.forward_metadata = PrefillMetadata(prefill_wrappers, True, False)
else:
raise ValueError(f"Invalid mode: {forward_mode=}")
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
@@ -733,19 +684,7 @@ class FlashInferAttnBackend(AttentionBackend):
fixed_split_size=None, fixed_split_size=None,
disable_split_kv=self.disable_cuda_graph_kv_split, disable_split_kv=self.disable_cuda_graph_kv_split,
) )
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_cpu[:bs] if seq_lens_cpu is not None else None,
seq_lens_sum,
prefix_lens=None,
prefill_wrappers=self.prefill_cuda_graph_metadata[bs],
use_ragged=False,
encoder_lens=encoder_lens[:bs] if encoder_lens is not None else None,
spec_info=spec_info,
)
elif forward_mode.is_draft_extend():
self.indices_updater_prefill.update( self.indices_updater_prefill.update(
req_pool_indices[:bs], req_pool_indices[:bs],
seq_lens[:bs], seq_lens[:bs],
@@ -384,7 +384,13 @@ class FlashInferMLAAttnBackend(AttentionBackend):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
): ):
seq_lens_sum = seq_lens.sum().item()
seq_lens_cpu = seq_lens.cpu()
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
# Decode: create wrapper, run the initial full begin_forward (False),
# then install the fast plan. After that, call replay so the
# data-update path (update(True)) is also exercised during capture.
decode_wrapper = BatchMLAPagedAttentionWrapper( decode_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer, self.workspace_buffer,
use_cuda_graph=True, use_cuda_graph=True,
@@ -394,8 +400,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
kv_len_arr=self.cuda_graph_kv_lens[:num_tokens], kv_len_arr=self.cuda_graph_kv_lens[:num_tokens],
backend="auto", backend="auto",
) )
seq_lens_sum = seq_lens.sum().item()
self.indices_updater_decode.update( self.indices_updater_decode.update(
req_pool_indices, req_pool_indices,
seq_lens, seq_lens,
@@ -406,9 +410,12 @@ class FlashInferMLAAttnBackend(AttentionBackend):
) )
self.decode_cuda_graph_metadata[bs] = decode_wrapper self.decode_cuda_graph_metadata[bs] = decode_wrapper
self.forward_metadata = DecodeMetadata(decode_wrapper) self.forward_metadata = DecodeMetadata(decode_wrapper)
# fast_mla_decode_plan requires _cached_module set by the initial
# begin_forward above; install it only after that call completes.
decode_wrapper.plan = partial(fast_mla_decode_plan, decode_wrapper) decode_wrapper.plan = partial(fast_mla_decode_plan, decode_wrapper)
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
verify_wrapper = BatchMLAPagedAttentionWrapper( # Prefill: create wrapper and store — replay handles the update call.
prefill_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer, self.workspace_buffer,
use_cuda_graph=True, use_cuda_graph=True,
qo_indptr=self.cuda_graph_qo_indptr[: bs + 1], qo_indptr=self.cuda_graph_qo_indptr[: bs + 1],
@@ -417,43 +424,22 @@ class FlashInferMLAAttnBackend(AttentionBackend):
kv_len_arr=self.cuda_graph_kv_lens[:bs], kv_len_arr=self.cuda_graph_kv_lens[:bs],
backend="auto", backend="auto",
) )
seq_lens_sum = seq_lens.sum().item() self.prefill_cuda_graph_metadata[bs] = prefill_wrapper
self.indices_updater_prefill.update( self.forward_metadata = PrefillMetadata(prefill_wrapper, False)
req_pool_indices,
seq_lens,
seq_lens_sum,
prefix_lens=None,
prefill_wrapper_paged=verify_wrapper,
use_ragged=False,
spec_info=spec_info,
)
self.prefill_cuda_graph_metadata[bs] = verify_wrapper
self.forward_metadata = PrefillMetadata(verify_wrapper, False)
elif forward_mode.is_draft_extend():
draft_extend_wrapper = BatchMLAPagedAttentionWrapper(
self.workspace_buffer,
use_cuda_graph=True,
qo_indptr=self.cuda_graph_qo_indptr[: bs + 1],
kv_indptr=self.cuda_graph_kv_indptr[: bs + 1],
kv_indices=self.cuda_graph_kv_indices,
kv_len_arr=self.cuda_graph_kv_lens[:bs],
backend="auto",
)
seq_lens_sum = seq_lens.sum().item()
self.indices_updater_prefill.update(
req_pool_indices,
seq_lens,
seq_lens_sum,
prefix_lens=None,
prefill_wrapper_paged=draft_extend_wrapper,
use_ragged=False,
spec_info=spec_info,
)
self.prefill_cuda_graph_metadata[bs] = draft_extend_wrapper
self.forward_metadata = PrefillMetadata(draft_extend_wrapper, False)
else: else:
raise ValueError(f"Invalid mode: {forward_mode=}") raise ValueError(f"Invalid mode: {forward_mode=}")
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=seq_lens_sum,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
bs: int, bs: int,
@@ -488,17 +474,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
spec_info=spec_info, spec_info=spec_info,
**self.fast_decode_kwargs, **self.fast_decode_kwargs,
) )
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
self.indices_updater_prefill.update(
req_pool_indices[:bs],
seq_lens[:bs],
seq_lens_sum,
prefix_lens=None,
prefill_wrapper_paged=self.prefill_cuda_graph_metadata[bs],
use_ragged=False,
spec_info=spec_info,
)
elif forward_mode.is_draft_extend():
self.indices_updater_prefill.update( self.indices_updater_prefill.update(
req_pool_indices[:bs], req_pool_indices[:bs],
seq_lens[:bs], seq_lens[:bs],
@@ -4,6 +4,7 @@ Support attention backend for FlashMLA.
from __future__ import annotations from __future__ import annotations
import logging
from dataclasses import dataclass from dataclasses import dataclass
from typing import TYPE_CHECKING, Callable, Optional, Tuple, Union from typing import TYPE_CHECKING, Callable, Optional, Tuple, Union
@@ -22,6 +23,7 @@ if TYPE_CHECKING:
from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpecInput from sglang.srt.speculative.spec_info import SpecInput
logger = logging.getLogger(__name__)
PAGE_SIZE = 64 PAGE_SIZE = 64
@@ -193,83 +195,16 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
): ):
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify():
max_seqlen_pad = triton.cdiv(seq_lens.max().item(), PAGE_SIZE) self.init_forward_metadata_replay_cuda_graph(
bs=bs,
create_flashmla_kv_indices_triton[(bs,)]( req_pool_indices=req_pool_indices,
self.req_to_token, seq_lens=seq_lens,
req_pool_indices, seq_lens_sum=None,
seq_lens, encoder_lens=encoder_lens,
None, forward_mode=forward_mode,
self.cuda_graph_kv_indices, spec_info=spec_info,
self.req_to_token.stride(0), seq_lens_cpu=None,
self.cuda_graph_kv_indices.stride(0),
)
num_q_heads = self.num_q_heads
mla_metadata, num_splits = get_mla_metadata(
seq_lens.to(torch.int32),
num_q_heads,
1,
is_fp8_kvcache=self.is_fp8_kvcache,
)
actual_num_sm_parts = mla_metadata.shape[0]
assert actual_num_sm_parts <= self.cuda_graph_mla_metadata.shape[0], (
f"num_sm_parts {actual_num_sm_parts} exceeds preallocated max "
f"{self.cuda_graph_mla_metadata.shape[0]}"
)
self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata)
self.cuda_graph_num_splits[: bs + 1].copy_(num_splits)
self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[
:actual_num_sm_parts
]
self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1]
self.forward_metadata = FlashMLADecodeMetadata(
self.cuda_graph_mla_metadata_view,
self.cuda_graph_num_splits_view,
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
)
elif forward_mode.is_target_verify():
seq_lens = seq_lens + self.num_draft_tokens
max_seqlen_pad = triton.cdiv(seq_lens.max().item(), PAGE_SIZE)
create_flashmla_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
None,
self.cuda_graph_kv_indices,
self.req_to_token.stride(0),
self.cuda_graph_kv_indices.stride(0),
)
mla_metadata, num_splits = get_mla_metadata(
seq_lens.to(torch.int32),
self.num_draft_tokens * self.num_q_heads,
1,
is_fp8_kvcache=self.is_fp8_kvcache,
)
actual_num_sm_parts = mla_metadata.shape[0]
assert actual_num_sm_parts <= self.cuda_graph_mla_metadata.shape[0]
self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata)
self.cuda_graph_num_splits[: bs + 1].copy_(num_splits)
self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[
:actual_num_sm_parts
]
self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1]
self.forward_metadata = FlashMLADecodeMetadata(
self.cuda_graph_mla_metadata_view,
self.cuda_graph_num_splits_view,
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
) )
else: else:
super().init_forward_metadata_capture_cuda_graph( super().init_forward_metadata_capture_cuda_graph(
@@ -293,11 +228,21 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor], seq_lens_cpu: Optional[torch.Tensor],
): ):
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify():
assert seq_lens_cpu is not None
seq_lens = seq_lens[:bs] seq_lens = seq_lens[:bs]
seq_lens_cpu = seq_lens_cpu[:bs] seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None
max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE)
if forward_mode.is_target_verify():
seq_lens = seq_lens + self.num_draft_tokens
if seq_lens_cpu is not None:
seq_lens_cpu = seq_lens_cpu + self.num_draft_tokens
seq_max = (
seq_lens_cpu.max().item()
if seq_lens_cpu is not None
else seq_lens.max().item()
)
max_seqlen_pad = triton.cdiv(seq_max, PAGE_SIZE)
create_flashmla_kv_indices_triton[(bs,)]( create_flashmla_kv_indices_triton[(bs,)](
self.req_to_token, self.req_to_token,
@@ -308,21 +253,28 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.req_to_token.stride(0), self.req_to_token.stride(0),
self.cuda_graph_kv_indices.stride(0), self.cuda_graph_kv_indices.stride(0),
) )
num_q_heads = self.num_q_heads
q_head_mult = (
self.num_draft_tokens if forward_mode.is_target_verify() else 1
)
mla_metadata, num_splits = get_mla_metadata( mla_metadata, num_splits = get_mla_metadata(
seq_lens.to(torch.int32), seq_lens.to(torch.int32),
num_q_heads, q_head_mult * self.num_q_heads,
1, 1,
is_fp8_kvcache=self.is_fp8_kvcache, is_fp8_kvcache=self.is_fp8_kvcache,
) )
actual_num_sm_parts = mla_metadata.shape[0] actual_num_sm_parts = mla_metadata.shape[0]
assert actual_num_sm_parts <= self.cuda_graph_mla_metadata.shape[0], (
f"num_sm_parts {actual_num_sm_parts} exceeds preallocated max "
f"{self.cuda_graph_mla_metadata.shape[0]}"
)
if actual_num_sm_parts != self.cuda_graph_mla_metadata_view.shape[0]: if (
import logging self.cuda_graph_mla_metadata_view is None
or actual_num_sm_parts != self.cuda_graph_mla_metadata_view.shape[0]
logger = logging.getLogger(__name__) ):
if self.cuda_graph_mla_metadata_view is not None:
logger.warning( logger.warning(
f"num_sm_parts mismatch in CUDA Graph replay: " f"num_sm_parts mismatch in CUDA Graph replay: "
f"capture={self.cuda_graph_mla_metadata_view.shape[0]}, " f"capture={self.cuda_graph_mla_metadata_view.shape[0]}, "
@@ -332,55 +284,17 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[ self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[
:actual_num_sm_parts :actual_num_sm_parts
] ]
# num_splits has shape (bs+1,) — always update for the current bs.
self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1] self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1]
self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata) self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata)
self.cuda_graph_num_splits[: bs + 1].copy_(num_splits) self.cuda_graph_num_splits[: bs + 1].copy_(num_splits)
self.forward_metadata.mla_metadata = self.cuda_graph_mla_metadata_view self.forward_metadata = FlashMLADecodeMetadata(
self.forward_metadata.num_splits = self.cuda_graph_num_splits_view self.cuda_graph_mla_metadata_view,
self.forward_metadata.block_kv_indices = self.cuda_graph_kv_indices[ self.cuda_graph_num_splits_view,
:bs, :max_seqlen_pad self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
]
elif forward_mode.is_target_verify():
seq_lens = seq_lens[:bs] + self.num_draft_tokens
seq_lens_cpu = seq_lens_cpu[:bs] + self.num_draft_tokens
max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE)
create_flashmla_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices[:bs],
seq_lens,
None,
self.cuda_graph_kv_indices,
self.req_to_token.stride(0),
self.cuda_graph_kv_indices.stride(0),
) )
mla_metadata, num_splits = get_mla_metadata(
seq_lens.to(torch.int32),
self.num_draft_tokens * self.num_q_heads,
1,
is_fp8_kvcache=self.is_fp8_kvcache,
)
actual_num_sm_parts = mla_metadata.shape[0]
if actual_num_sm_parts != self.cuda_graph_mla_metadata_view.shape[0]:
self.cuda_graph_mla_metadata_view = self.cuda_graph_mla_metadata[
:actual_num_sm_parts
]
self.cuda_graph_num_splits_view = self.cuda_graph_num_splits[: bs + 1]
self.cuda_graph_mla_metadata[:actual_num_sm_parts].copy_(mla_metadata)
self.cuda_graph_num_splits[: bs + 1].copy_(num_splits)
self.forward_metadata.mla_metadata = self.cuda_graph_mla_metadata_view
self.forward_metadata.num_splits = self.cuda_graph_num_splits_view
self.forward_metadata.block_kv_indices = self.cuda_graph_kv_indices[
:bs, :max_seqlen_pad
]
else: else:
super().init_forward_metadata_replay_cuda_graph( super().init_forward_metadata_replay_cuda_graph(
bs, bs,
@@ -403,8 +403,15 @@ class MambaAttnBackendBase(AttentionBackend):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
): ):
self.forward_metadata = self._capture_metadata( self.init_forward_metadata_replay_cuda_graph(
bs, req_pool_indices, forward_mode, spec_info bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
) )
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
@@ -539,6 +546,9 @@ class MambaAttnBackendBase(AttentionBackend):
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor], seq_lens_cpu: Optional[torch.Tensor],
): ):
if seq_lens_cpu is None:
num_padding = 0
else:
num_padding = torch.count_nonzero( num_padding = torch.count_nonzero(
seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value() seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value()
) )
@@ -576,6 +586,10 @@ class MambaAttnBackendBase(AttentionBackend):
# If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask # If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask
if forward_mode.is_target_verify() and self.topk > 1: if forward_mode.is_target_verify() and self.topk > 1:
if (
spec_info is not None
and getattr(spec_info, "retrieve_next_token", None) is not None
):
bs_without_pad = spec_info.retrieve_next_token.shape[0] bs_without_pad = spec_info.retrieve_next_token.shape[0]
self.retrieve_next_token_list[bs - 1][:bs_without_pad].copy_( self.retrieve_next_token_list[bs - 1][:bs_without_pad].copy_(
spec_info.retrieve_next_token spec_info.retrieve_next_token
@@ -703,13 +717,15 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
): ):
metadata = self._capture_metadata(bs, req_pool_indices, forward_mode, spec_info) self.init_forward_metadata_replay_cuda_graph(
draft_token_num = spec_info.draft_token_num if spec_info is not None else 1 bs=bs,
self.forward_metadata = Mamba2Metadata.prepare_decode( req_pool_indices=req_pool_indices,
metadata, seq_lens=seq_lens,
seq_lens, seq_lens_sum=None,
is_target_verify=forward_mode.is_target_verify(), encoder_lens=encoder_lens,
draft_token_num=draft_token_num, forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
) )
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
@@ -89,9 +89,15 @@ class LightningAttentionBackend(MambaAttnBackendBase):
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]], spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
): ):
metadata = self._capture_metadata(bs, req_pool_indices, forward_mode, spec_info) self.init_forward_metadata_replay_cuda_graph(
self.forward_metadata = BailingLinearMetadata.prepare_decode( bs=bs,
metadata.query_start_loc, metadata.mamba_cache_indices, bs, seq_lens req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
) )
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
@@ -283,11 +283,179 @@ class TritonAttnBackend(AttentionBackend):
MAX_NUM_SEQ=SCHEDULE_SEQ, MAX_NUM_SEQ=SCHEDULE_SEQ,
) )
def _fill_kv_indptr_and_indices(
self,
bs: int,
seq_lens: torch.Tensor,
req_pool_indices: torch.Tensor,
kv_indices: torch.Tensor,
) -> torch.Tensor:
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1:] = torch.cumsum(seq_lens, dim=0)
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
return kv_indptr
def _update_decode_kv_buffers(
self,
bs: int,
seq_lens: torch.Tensor,
req_pool_indices: torch.Tensor,
):
"""Fill KV (and SWA) cuda-graph buffers for decode/idle mode.
Returns ``(kv_indptr, window_kv_indptr, window_kv_lens)`` where
``window_kv_lens`` is ``None`` when sliding-window is disabled.
"""
seq_lens = seq_lens[:bs]
req_pool_indices = req_pool_indices[:bs]
kv_indptr = self._fill_kv_indptr_and_indices(
bs, seq_lens, req_pool_indices, self.cuda_graph_kv_indices
)
window_kv_indptr = self.window_kv_indptr
window_kv_lens = None
if self.sliding_window_size is not None and self.sliding_window_size > 0:
window_kv_indptr, _, window_kv_lens, _ = update_sliding_window_buffer(
self.window_kv_indptr,
self.req_to_token,
self.sliding_window_size,
seq_lens,
req_pool_indices,
bs,
token_to_kv_pool=self.token_to_kv_pool,
window_kv_indices=self.cuda_graph_window_kv_indices,
)
return kv_indptr, window_kv_indptr, window_kv_lens
def _update_target_verify_buffers(
self,
bs: int,
seq_lens: torch.Tensor,
req_pool_indices: torch.Tensor,
spec_info,
):
"""Fill all cuda-graph buffers for target_verify mode.
Returns the ForwardMetadata components:
``(qo_indptr, kv_indptr, custom_mask, mask_indptr,
window_kv_indptr, window_kv_indices, window_num_kv_splits, window_kv_offsets)``
"""
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
(1 + bs) * self.num_draft_tokens,
step=self.num_draft_tokens,
dtype=torch.int32,
device=self.device,
)
kv_indptr = self._fill_kv_indptr_and_indices(
bs, seq_lens, req_pool_indices, self.cuda_graph_kv_indices
)
window_kv_indptr = self.window_kv_indptr
window_kv_indices = None
window_num_kv_splits = None
window_kv_offsets = None
if self.sliding_window_size is not None and self.sliding_window_size > 0:
window_kv_indices = self.cuda_graph_window_kv_indices
window_num_kv_splits = self.cuda_graph_window_num_kv_splits
window_kv_offsets = self.cuda_graph_window_kv_offsets
window_kv_indptr, window_kv_indices, _, window_kv_offsets[:bs] = (
update_sliding_window_buffer(
self.window_kv_indptr,
self.req_to_token,
self.sliding_window_size,
seq_lens[:bs],
req_pool_indices,
bs,
token_to_kv_pool=self.token_to_kv_pool,
window_kv_indices=window_kv_indices,
)
)
custom_mask = self.cuda_graph_custom_mask
if (
spec_info is not None
and getattr(spec_info, "custom_mask", None) is not None
):
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
else:
custom_mask = None
seq_mask_len = self.num_draft_tokens * (seq_lens + self.num_draft_tokens)
mask_indptr = self.mask_indptr[: bs + 1]
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
return (
qo_indptr,
kv_indptr,
custom_mask,
mask_indptr,
window_kv_indptr,
window_kv_indices,
window_num_kv_splits,
window_kv_offsets,
)
def _update_draft_extend_buffers(
self,
bs: int,
seq_lens: torch.Tensor,
req_pool_indices: torch.Tensor,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Fill QO + KV cuda-graph buffers for draft_extend mode.
Returns ``(qo_indptr, kv_indptr, num_tokens_per_bs)``.
"""
seq_lens = seq_lens[:bs]
num_tokens_per_bs = self.speculative_num_steps + 1
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
bs * num_tokens_per_bs + 1,
step=num_tokens_per_bs,
dtype=torch.int32,
device=self.device,
)
if forward_mode.is_draft_extend_v2():
# DRAFT_EXTEND_V2: seq_lens = prefix + extend (bumped by eagle_info_v2).
# Triton extend kernel receives extend K/V as separate tensors, so
# kv_indptr/kv_indices must cover only the prefix portion.
# extend_seq_lens_tensor is only attached to spec_info at real
# replay (eagle_draft_extend_cuda_graph_runner.replay); during the
# capture-time warmup it's absent, so fall back to zeros (matches
# the pre-unification capture path in #26651). Clamp at 0 because
# padded rows (raw_bs..bs) leave seq_lens at the fill value (1)
# while extend_seq_lens stays at num_tokens_per_bs, which would
# otherwise produce negative kv_lens; padded rows reference
# reserved req-pool slot 0 and their output is discarded.
if (
spec_info is not None
and getattr(spec_info, "extend_seq_lens_tensor", None) is not None
):
extend_seq_lens = spec_info.extend_seq_lens_tensor[:bs].to(torch.int32)
else:
extend_seq_lens = torch.zeros(
bs, dtype=torch.int32, device=seq_lens.device
)
kv_lens = torch.clamp(seq_lens - extend_seq_lens, min=0).to(torch.int32)
else:
# DRAFT_EXTEND_V1: seq_lens = prefix only.
kv_lens = seq_lens
kv_indptr = self._fill_kv_indptr_and_indices(
bs, kv_lens, req_pool_indices, self.cuda_graph_kv_indices
)
return qo_indptr, kv_indptr, num_tokens_per_bs
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend.""" """Init auxiliary variables for triton attention backend."""
bs = forward_batch.batch_size bs = forward_batch.batch_size
kv_indptr = self.kv_indptr
window_kv_indptr = self.window_kv_indptr window_kv_indptr = self.window_kv_indptr
window_kv_indices = None window_kv_indices = None
window_num_kv_splits = None window_num_kv_splits = None
@@ -297,19 +465,14 @@ class TritonAttnBackend(AttentionBackend):
if forward_batch.forward_mode.is_decode_or_idle(): if forward_batch.forward_mode.is_decode_or_idle():
if spec_info is None: if spec_info is None:
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = torch.empty( kv_indices = torch.empty(
forward_batch.seq_lens_sum, dtype=torch.int64, device=self.device forward_batch.seq_lens_sum, dtype=torch.int64, device=self.device
) )
create_flashinfer_kv_indices_triton[(bs,)]( kv_indptr = self._fill_kv_indptr_and_indices(
self.req_to_token, bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
kv_indptr, forward_batch.req_pool_indices,
None,
kv_indices, kv_indices,
self.req_to_token.stride(0),
) )
# Sliding window # Sliding window
if ( if (
@@ -371,19 +534,14 @@ class TritonAttnBackend(AttentionBackend):
device=self.device, device=self.device,
) )
# Different with flashinfer kv_indptr and kv_indices construction # Different with flashinfer kv_indptr and kv_indices construction
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = torch.empty( kv_indices = torch.empty(
kv_indptr[-1], dtype=torch.int64, device=self.device forward_batch.seq_lens_sum, dtype=torch.int64, device=self.device
) )
create_flashinfer_kv_indices_triton[(bs,)]( kv_indptr = self._fill_kv_indptr_and_indices(
self.req_to_token, bs,
forward_batch.req_pool_indices,
forward_batch.seq_lens, forward_batch.seq_lens,
kv_indptr, forward_batch.req_pool_indices,
None,
kv_indices, kv_indices,
self.req_to_token.stride(0),
) )
if self.sliding_window_size is not None and self.sliding_window_size > 0: if self.sliding_window_size is not None and self.sliding_window_size > 0:
@@ -435,23 +593,16 @@ class TritonAttnBackend(AttentionBackend):
attn_logits = None attn_logits = None
attn_lse = None attn_lse = None
else: else:
kv_indptr[1 : bs + 1] = torch.cumsum(
forward_batch.extend_prefix_lens, dim=0
)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = torch.empty( kv_indices = torch.empty(
sum(forward_batch.extend_prefix_lens_cpu), sum(forward_batch.extend_prefix_lens_cpu),
dtype=torch.int64, dtype=torch.int64,
device=self.device, device=self.device,
) )
create_flashinfer_kv_indices_triton[(bs,)]( kv_indptr = self._fill_kv_indptr_and_indices(
self.req_to_token, bs,
forward_batch.req_pool_indices,
forward_batch.extend_prefix_lens, forward_batch.extend_prefix_lens,
kv_indptr, forward_batch.req_pool_indices,
None,
kv_indices, kv_indices,
self.req_to_token.stride(0),
) )
# Sliding window # Sliding window
if self.sliding_window_size is not None and self.sliding_window_size > 0: if self.sliding_window_size is not None and self.sliding_window_size > 0:
@@ -578,6 +729,83 @@ class TritonAttnBackend(AttentionBackend):
device=self.device, device=self.device,
) )
def _build_cuda_graph_forward_metadata(
self,
bs: int,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
) -> ForwardMetadata:
"""Construct ForwardMetadata from the current cuda-graph buffer state.
Called by capture after the buffer-update helpers have already run
(either via replay or directly). All fields reference the same
``self.cuda_graph_*`` tensors that the captured graph kernels will
read — the Python object is rebuilt each capture, but the underlying
GPU memory addresses are stable.
"""
swa = self.sliding_window_size is not None and self.sliding_window_size > 0
if forward_mode.is_decode_or_idle():
return ForwardMetadata(
attn_logits=self.cuda_graph_attn_logits,
attn_lse=self.cuda_graph_attn_lse,
max_extend_len=None,
num_kv_splits=self.cuda_graph_num_kv_splits,
kv_indptr=self.kv_indptr[: bs + 1],
kv_indices=self.cuda_graph_kv_indices,
qo_indptr=None,
custom_mask=None,
mask_indptr=None,
window_kv_indptr=self.window_kv_indptr[: bs + 1] if swa else None,
window_kv_indices=self.cuda_graph_window_kv_indices if swa else None,
window_num_kv_splits=(
self.cuda_graph_window_num_kv_splits if swa else None
),
window_kv_offsets=None,
swa_attn_logits=self.cuda_graph_swa_attn_logits,
)
elif forward_mode.is_target_verify():
custom_mask = (
self.cuda_graph_custom_mask
if spec_info is not None
and getattr(spec_info, "custom_mask", None) is not None
else None
)
return ForwardMetadata(
attn_logits=None,
attn_lse=None,
max_extend_len=self.num_draft_tokens,
num_kv_splits=None,
kv_indptr=self.kv_indptr[: bs + 1],
kv_indices=self.cuda_graph_kv_indices,
qo_indptr=self.qo_indptr[: bs + 1],
custom_mask=custom_mask,
mask_indptr=self.mask_indptr[: bs + 1],
window_kv_indptr=self.window_kv_indptr[: bs + 1] if swa else None,
window_kv_indices=self.cuda_graph_window_kv_indices if swa else None,
window_num_kv_splits=(
self.cuda_graph_window_num_kv_splits if swa else None
),
window_kv_offsets=self.cuda_graph_window_kv_offsets if swa else None,
)
elif forward_mode.is_draft_extend(include_v2=True):
return ForwardMetadata(
attn_logits=None,
attn_lse=None,
max_extend_len=self.speculative_num_steps + 1,
num_kv_splits=None,
kv_indptr=self.kv_indptr[: bs + 1],
kv_indices=self.cuda_graph_kv_indices,
qo_indptr=self.qo_indptr[: bs + 1],
custom_mask=None,
mask_indptr=None,
window_kv_indptr=self.window_kv_indptr,
window_kv_indices=None,
window_num_kv_splits=None,
window_kv_offsets=None,
)
else:
raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.")
def init_forward_metadata_capture_cuda_graph( def init_forward_metadata_capture_cuda_graph(
self, self,
bs: int, bs: int,
@@ -589,172 +817,42 @@ class TritonAttnBackend(AttentionBackend):
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
): ):
assert encoder_lens is None, "Not supported" assert encoder_lens is None, "Not supported"
window_kv_indptr = self.window_kv_indptr
window_kv_indices = None
window_num_kv_splits = None
window_kv_offsets = None
swa_attn_logits = None
if forward_mode.is_decode_or_idle():
if spec_info is None:
kv_indptr = self.kv_indptr
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
if (
self.sliding_window_size is not None
and self.sliding_window_size > 0
):
window_kv_indices = self.cuda_graph_window_kv_indices
window_num_kv_splits = self.cuda_graph_window_num_kv_splits
window_kv_indptr, window_kv_indices, _, _ = (
update_sliding_window_buffer_cuda_graph(
self.window_kv_indptr,
window_kv_indices,
self.req_to_token,
self.sliding_window_size,
seq_lens[:bs],
req_pool_indices,
bs,
self.token_to_kv_pool,
)
)
else:
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
attn_logits = self.cuda_graph_attn_logits
swa_attn_logits = self.cuda_graph_swa_attn_logits
attn_lse = self.cuda_graph_attn_lse
max_extend_len = None
num_kv_splits = self.cuda_graph_num_kv_splits
qo_indptr = None
custom_mask = None
mask_indptr = None
elif forward_mode.is_target_verify():
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
(1 + bs) * self.num_draft_tokens,
step=self.num_draft_tokens,
dtype=torch.int32,
device=self.device,
)
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
if self.sliding_window_size is not None and self.sliding_window_size > 0:
window_kv_indices = self.cuda_graph_window_kv_indices
window_num_kv_splits = self.cuda_graph_window_num_kv_splits
window_kv_offsets = self.cuda_graph_window_kv_offsets
window_kv_indptr, window_kv_indices, _, window_kv_offsets[:bs] = (
update_sliding_window_buffer_cuda_graph(
self.window_kv_indptr,
window_kv_indices,
self.req_to_token,
self.sliding_window_size,
seq_lens[:bs],
req_pool_indices,
bs,
self.token_to_kv_pool,
)
)
custom_mask = self.cuda_graph_custom_mask
if (
spec_info is not None
and getattr(spec_info, "custom_mask", None) is not None
):
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
else:
custom_mask = None
seq_mask_len = self.num_draft_tokens * (seq_lens + self.num_draft_tokens)
mask_indptr = self.mask_indptr[: bs + 1]
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
max_extend_len = self.num_draft_tokens
num_kv_splits = None
attn_logits = None
attn_lse = None
elif forward_mode.is_draft_extend(include_v2=True):
num_tokens_per_bs = self.speculative_num_steps + 1
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
bs * num_tokens_per_bs + 1,
step=num_tokens_per_bs,
dtype=torch.int32,
device=self.device,
)
kv_indptr = self.kv_indptr[: bs + 1]
if forward_mode.is_draft_extend_v2():
# DRAFT_EXTEND_V2: seq_lens = prefix + extend (bumped by eagle_info_v2).
# Triton extend kernel receives extend K/V as separate tensors, so
# kv_indptr/kv_indices must cover only the prefix portion.
extend_seq_lens = (
spec_info.extend_seq_lens_tensor[:bs].to(torch.int32)
if spec_info is not None
and getattr(spec_info, "extend_seq_lens_tensor", None) is not None
else torch.zeros(bs, dtype=torch.int32, device=self.device)
)
kv_lens = (seq_lens - extend_seq_lens).to(torch.int32)
else:
# DRAFT_EXTEND_V1: seq_lens = prefix only.
kv_lens = seq_lens
kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
kv_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
custom_mask = None
mask_indptr = None
max_extend_len = num_tokens_per_bs
num_kv_splits = None
attn_logits = None
attn_lse = None
else:
raise ValueError(
f"Invalid forward mode: {forward_mode=} for CUDA Graph capture."
)
# Multi-step speculative decode: kv buffers come from spec_info rather
# than the cuda-graph pool, so replay is not involved for this path.
if forward_mode.is_decode_or_idle() and spec_info is not None:
self.forward_metadata = ForwardMetadata( self.forward_metadata = ForwardMetadata(
attn_logits, attn_logits=self.cuda_graph_attn_logits,
attn_lse, attn_lse=self.cuda_graph_attn_lse,
max_extend_len, max_extend_len=None,
num_kv_splits, num_kv_splits=self.cuda_graph_num_kv_splits,
kv_indptr, kv_indptr=spec_info.kv_indptr,
kv_indices, kv_indices=spec_info.kv_indices,
qo_indptr, qo_indptr=None,
custom_mask, custom_mask=None,
mask_indptr, mask_indptr=None,
window_kv_indptr, window_kv_indptr=self.window_kv_indptr,
window_kv_indices, window_kv_indices=None,
window_num_kv_splits, window_num_kv_splits=None,
window_kv_offsets, window_kv_offsets=None,
swa_attn_logits=swa_attn_logits, swa_attn_logits=self.cuda_graph_swa_attn_logits,
)
return
# Run the same buffer update as replay, then freeze the result into
# a ForwardMetadata whose tensor fields point into the cuda-graph buffers.
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
self.forward_metadata = self._build_cuda_graph_forward_metadata(
bs, forward_mode, spec_info
) )
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
@@ -770,138 +868,23 @@ class TritonAttnBackend(AttentionBackend):
): ):
# NOTE: encoder_lens expected to be zeros or None # NOTE: encoder_lens expected to be zeros or None
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
# Update kv_indptr, kv_indices assert spec_info is None, "Multi-step cuda graph init is not done here."
kv_indptr = self.kv_indptr _, _, window_kv_lens = self._update_decode_kv_buffers(
kv_indices = self.cuda_graph_kv_indices bs, seq_lens, req_pool_indices
num_kv_splits = self.cuda_graph_num_kv_splits
if spec_info is None:
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens[:bs], dim=0)
kv_indptr = kv_indptr[: bs + 1]
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices[:bs],
seq_lens[:bs],
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
num_token = bs
if (
self.sliding_window_size is not None
and self.sliding_window_size > 0
):
window_num_kv_splits = self.cuda_graph_window_num_kv_splits
window_kv_indices = self.cuda_graph_window_kv_indices
_, _, window_kv_lens, _ = update_sliding_window_buffer_cuda_graph(
self.window_kv_indptr,
window_kv_indices,
self.req_to_token,
self.sliding_window_size,
seq_lens[:bs],
req_pool_indices[:bs],
bs,
self.token_to_kv_pool,
) )
self.get_num_kv_splits(self.cuda_graph_num_kv_splits[:bs], seq_lens[:bs])
if window_kv_lens is not None:
self.get_num_kv_splits( self.get_num_kv_splits(
window_num_kv_splits[:num_token], window_kv_lens[:bs] self.cuda_graph_window_num_kv_splits[:bs], window_kv_lens[:bs]
) )
else:
assert False, "Multi-step cuda graph init is not done here."
self.get_num_kv_splits(num_kv_splits[:num_token], seq_lens[:bs])
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify():
# Update qo_indptr, kv_indptr, kv_indices, custom_mask, mask_indptr
bs = len(req_pool_indices) bs = len(req_pool_indices)
qo_indptr = self.qo_indptr[: bs + 1] self._update_target_verify_buffers(
qo_indptr[: bs + 1] = torch.arange( bs, seq_lens, req_pool_indices, spec_info
0,
(1 + bs) * self.num_draft_tokens,
step=self.num_draft_tokens,
dtype=torch.int32,
device=self.device,
) )
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
if self.sliding_window_size is not None and self.sliding_window_size > 0:
window_num_kv_splits = self.cuda_graph_window_num_kv_splits
window_kv_indices = self.cuda_graph_window_kv_indices
window_kv_offsets = self.cuda_graph_window_kv_offsets
_, _, window_kv_lens, window_kv_offsets[:bs] = (
update_sliding_window_buffer_cuda_graph(
self.window_kv_indptr,
window_kv_indices,
self.req_to_token,
self.sliding_window_size,
seq_lens[:bs],
req_pool_indices,
bs,
self.token_to_kv_pool,
)
)
custom_mask = self.cuda_graph_custom_mask
if (
spec_info is not None
and getattr(spec_info, "custom_mask", None) is not None
):
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
else:
custom_mask = None
seq_mask_len = self.num_draft_tokens * (seq_lens + self.num_draft_tokens)
mask_indptr = self.mask_indptr[: bs + 1]
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
elif forward_mode.is_draft_extend(include_v2=True): elif forward_mode.is_draft_extend(include_v2=True):
seq_lens = seq_lens[:bs] self._update_draft_extend_buffers(
num_tokens_per_bs = self.speculative_num_steps + 1 bs, seq_lens, req_pool_indices, forward_mode, spec_info
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
bs * num_tokens_per_bs + 1,
step=num_tokens_per_bs,
dtype=torch.int32,
device=self.device,
)
kv_indptr = self.kv_indptr[: bs + 1]
if forward_mode.is_draft_extend_v2():
# DRAFT_EXTEND_V2: seq_lens = prefix + extend (bumped by eagle_info_v2).
# Triton extend kernel receives extend K/V as separate tensors, so
# kv_indptr/kv_indices must cover only the prefix portion.
# Clamp at 0 because padded rows (raw_bs..bs) leave seq_lens at
# the fill value (1) while extend_seq_lens stays at num_tokens_per_bs,
# which would otherwise produce negative kv_lens; padded rows
# reference reserved req-pool slot 0 and their output is discarded.
assert (
spec_info is not None
and getattr(spec_info, "extend_seq_lens_tensor", None) is not None
), "DRAFT_EXTEND_V2 replay requires spec_info.extend_seq_lens_tensor"
kv_lens = torch.clamp(
seq_lens - spec_info.extend_seq_lens_tensor[:bs].to(torch.int32),
min=0,
).to(torch.int32)
else:
# DRAFT_EXTEND_V1: seq_lens = prefix only.
kv_lens = seq_lens
kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
kv_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
) )
else: else:
raise ValueError( raise ValueError(
@@ -1490,15 +1473,23 @@ def update_sliding_window_buffer(
seq_lens, seq_lens,
req_pool_indices, req_pool_indices,
bs, bs,
device, device=None,
token_to_kv_pool=None, token_to_kv_pool=None,
window_kv_indices=None,
): ):
"""Fill window KV buffers for sliding-window attention.
Pass ``window_kv_indices`` to write into a pre-allocated buffer (CUDA-graph
path); omit it (or pass ``None``) to allocate a fresh tensor (eager path,
requires ``device``).
"""
window_kv_lens = torch.minimum( window_kv_lens = torch.minimum(
seq_lens, seq_lens,
torch.tensor(sliding_window_size), torch.tensor(sliding_window_size),
) )
window_kv_indptr[1 : bs + 1] = torch.cumsum(window_kv_lens, dim=0) window_kv_indptr[1 : bs + 1] = torch.cumsum(window_kv_lens, dim=0)
window_kv_indptr = window_kv_indptr[: bs + 1] window_kv_indptr = window_kv_indptr[: bs + 1]
if window_kv_indices is None:
window_kv_indices = torch.empty( window_kv_indices = torch.empty(
window_kv_indptr[-1], dtype=torch.int64, device=device window_kv_indptr[-1], dtype=torch.int64, device=device
) )
@@ -1512,47 +1503,6 @@ def update_sliding_window_buffer(
window_kv_indices, window_kv_indices,
req_to_token.stride(0), req_to_token.stride(0),
) )
# full to swa index mapping
if hasattr(token_to_kv_pool, "translate_loc_from_full_to_swa"):
kv_last_index = window_kv_indptr[-1]
# Flush before+after: window_kv_indices is a different tensor than out_cache_loc.
token_to_kv_pool.invalidate_loc_cache()
window_kv_indices[:kv_last_index] = (
token_to_kv_pool.translate_loc_from_full_to_swa(
window_kv_indices[:kv_last_index]
)
)
token_to_kv_pool.invalidate_loc_cache()
return window_kv_indptr, window_kv_indices, window_kv_lens, window_kv_start_idx
def update_sliding_window_buffer_cuda_graph(
window_kv_indptr,
window_kv_indices,
req_to_token,
sliding_window_size,
seq_lens,
req_pool_indices,
bs,
token_to_kv_pool=None,
):
window_kv_lens = torch.minimum(
seq_lens,
torch.tensor(sliding_window_size),
)
window_kv_indptr[1 : bs + 1] = torch.cumsum(window_kv_lens, dim=0)
window_kv_indptr = window_kv_indptr[: bs + 1]
window_kv_start_idx = seq_lens - window_kv_lens
create_flashinfer_kv_indices_triton[(bs,)](
req_to_token,
req_pool_indices,
window_kv_lens,
window_kv_indptr,
window_kv_start_idx,
window_kv_indices,
req_to_token.stride(0),
)
# full to swa index mapping
if hasattr(token_to_kv_pool, "translate_loc_from_full_to_swa"): if hasattr(token_to_kv_pool, "translate_loc_from_full_to_swa"):
kv_last_index = window_kv_indptr[-1] kv_last_index = window_kv_indptr[-1]
# Flush before+after: window_kv_indices is a different tensor than out_cache_loc. # Flush before+after: window_kv_indices is a different tensor than out_cache_loc.
@@ -303,42 +303,29 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
), ),
} }
def init_forward_metadata_capture_cuda_graph( def _build_cuda_graph_metadata(
self, self,
bs: int, bs: int,
num_tokens: int, num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode, forward_mode: ForwardMode,
spec_info: Optional[SpecInput], spec_info,
): device: torch.device,
"""Initialize metadata for CUDA graph capture.""" ) -> "TRTLLMMHAMetadata":
"""Create TRTLLMMHAMetadata with pre-allocated buffer slice refs, stored in the dict."""
metadata = TRTLLMMHAMetadata() metadata = TRTLLMMHAMetadata()
device = seq_lens.device
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
if spec_info is not None: if spec_info is not None:
# Draft Decode # Draft Decode (topk = 1)
# Here we only support topk = 1 for now.
metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[ metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[
"cache_seqlens" "cache_seqlens"
][:bs] ][:bs]
metadata.cache_seqlens_int32.copy_(
seq_lens + self.speculative_step_id + 1
)
metadata.max_seq_len_k = seq_lens.max().item() + (
self.speculative_step_id + 1
)
metadata.cu_seqlens_q = self.decode_cuda_graph_metadata["cu_seqlens_q"][ metadata.cu_seqlens_q = self.decode_cuda_graph_metadata["cu_seqlens_q"][
: bs + 1 : bs + 1
] ]
metadata.cu_seqlens_k = torch.nn.functional.pad( metadata.cu_seqlens_k = self.decode_cuda_graph_metadata["cu_seqlens_k"][
torch.cumsum( : bs + 1
metadata.cache_seqlens_int32, dim=0, dtype=torch.int32 ]
),
(1, 0),
)
metadata.page_table = self.decode_cuda_graph_metadata[ metadata.page_table = self.decode_cuda_graph_metadata[
"page_table_draft_decode" "page_table_draft_decode"
][:bs, :] ][:bs, :]
@@ -351,20 +338,15 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
else: else:
# Normal Decode # Normal Decode
# Get sequence information metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[
metadata.cache_seqlens_int32 = seq_lens[:bs].to(torch.int32) "cache_seqlens"
batch_size = len(seq_lens) ][:bs]
metadata.cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(seq_lens, dim=0, dtype=torch.int32), (1, 0)
)
# Precompute maximum sequence length
metadata.max_seq_len_k = seq_lens.max().item()
# Precompute cumulative sequence lengths
metadata.cu_seqlens_q = torch.arange( metadata.cu_seqlens_q = torch.arange(
0, batch_size + 1, dtype=torch.int32, device=device 0, bs + 1, dtype=torch.int32, device=device
)
metadata.cu_seqlens_k = torch.zeros(
bs + 1, dtype=torch.int32, device=device
) )
# Precompute page table
metadata.page_table = self.decode_cuda_graph_metadata["page_table"][ metadata.page_table = self.decode_cuda_graph_metadata["page_table"][
:bs, : :bs, :
] ]
@@ -376,29 +358,18 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
) )
self.decode_cuda_graph_metadata[bs] = metadata self.decode_cuda_graph_metadata[bs] = metadata
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify():
# Target Verify # Target Verify (topk = 1)
# Here we only support topk = 1 for now.
tokens_per_req = num_tokens // bs tokens_per_req = num_tokens // bs
metadata.cache_seqlens_int32 = self.target_verify_metadata["cache_seqlens"][ metadata.cache_seqlens_int32 = self.target_verify_metadata["cache_seqlens"][
:bs :bs
] ]
metadata.cache_seqlens_int32.copy_(seq_lens + tokens_per_req) metadata.cu_seqlens_q = self.target_verify_metadata["cu_seqlens_q"][
: bs + 1
metadata.cu_seqlens_q = torch.arange( ]
0,
bs * tokens_per_req + 1,
tokens_per_req,
dtype=torch.int32,
device=device,
)
metadata.cu_seqlens_k = self.target_verify_metadata["cu_seqlens_k"][ metadata.cu_seqlens_k = self.target_verify_metadata["cu_seqlens_k"][
: (bs + 1) : bs + 1
] ]
metadata.max_seq_len_q = tokens_per_req metadata.max_seq_len_q = tokens_per_req
metadata.max_seq_len_k = seq_lens.max().item() + tokens_per_req
metadata.page_table = self.target_verify_metadata["page_table"][:bs, :] metadata.page_table = self.target_verify_metadata["page_table"][:bs, :]
self._bind_swa_page_table( self._bind_swa_page_table(
metadata, metadata,
@@ -406,29 +377,15 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
"swa_page_table", "swa_page_table",
bs, bs,
) )
self.target_verify_metadata[bs] = metadata self.target_verify_metadata[bs] = metadata
elif forward_mode.is_draft_extend(): elif forward_mode.is_draft_extend():
num_tokens_per_bs = num_tokens // bs
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][ metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
:bs :bs
] ]
metadata.cache_seqlens_int32.copy_(seq_lens) metadata.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1]
num_tokens_per_bs = num_tokens // bs metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][: bs + 1]
metadata.cu_seqlens_q = torch.arange(
0,
bs * num_tokens_per_bs + 1,
num_tokens_per_bs,
dtype=torch.int32,
device=device,
)
metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][
: (bs + 1)
]
num_tokens_per_bs = num_tokens // bs
metadata.max_seq_len_q = num_tokens_per_bs metadata.max_seq_len_q = num_tokens_per_bs
metadata.max_seq_len_k = seq_lens.max().item()
metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :] metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :]
self._bind_swa_page_table( self._bind_swa_page_table(
metadata, metadata,
@@ -436,9 +393,41 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
"swa_page_table", "swa_page_table",
bs, bs,
) )
self.draft_extend_metadata[bs] = metadata self.draft_extend_metadata[bs] = metadata
self.forward_metadata = metadata
return metadata
def init_forward_metadata_capture_cuda_graph(
self,
bs: int,
num_tokens: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
encoder_lens: Optional[torch.Tensor],
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
):
"""Initialize metadata for CUDA graph capture."""
seq_lens_cpu = seq_lens.cpu()
self._build_cuda_graph_metadata(
bs, num_tokens, forward_mode, spec_info, seq_lens.device
)
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=seq_lens_cpu,
)
if forward_mode.is_draft_extend():
# CUDA graph bakes max_seq_len_q as a constant. replay() sets it to
# max(num_accept_tokens_cpu) which is None/empty at capture time,
# falling back to 1. Restore the correct upper bound so the kernel
# sees num_tokens_per_bs (not 1) for all replays of this graph.
self.forward_metadata.max_seq_len_q = num_tokens // bs
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
@@ -444,6 +444,44 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
super().init_cuda_graph_state(max_bs, max_num_tokens, kv_indices_buf) super().init_cuda_graph_state(max_bs, max_num_tokens, kv_indices_buf)
def _init_cuda_graph_metadata(
self,
bs: int,
num_tokens: int,
forward_mode: ForwardMode,
seq_lens: torch.Tensor,
device: torch.device,
):
"""Allocate persistent metadata buffers for CUDA graph capture."""
metadata = TRTLLMMLADecodeMetadata()
if forward_mode.is_target_verify():
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
elif forward_mode.is_draft_extend(include_v2=True):
num_tokens_per_bs = num_tokens // bs
metadata.max_seq_len_q = num_tokens_per_bs
metadata.sum_seq_lens_q = num_tokens_per_bs * bs
metadata.cu_seqlens_q = torch.arange(
0,
bs * num_tokens_per_bs + 1,
num_tokens_per_bs,
dtype=torch.int32,
device=device,
)
metadata.seq_lens_q = torch.full(
(bs,), num_tokens_per_bs, dtype=torch.int32, device=device
)
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
# Capture with full width so future longer sequences are safe during replay.
max_blocks_per_seq = self._calc_padded_blocks(self.max_context_len)
block_kv_indices = self.decode_cuda_graph_kv_indices[:bs, :max_blocks_per_seq]
metadata.block_kv_indices = block_kv_indices
metadata.max_seq_len_k = self.max_context_len
self.decode_cuda_graph_metadata[bs] = metadata
self.forward_decode_metadata = metadata
def init_forward_metadata_capture_cuda_graph( def init_forward_metadata_capture_cuda_graph(
self, self,
bs: int, bs: int,
@@ -472,60 +510,19 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
spec_info, spec_info,
) )
metadata = TRTLLMMLADecodeMetadata() self._init_cuda_graph_metadata(
bs, num_tokens, forward_mode, seq_lens, seq_lens.device
if forward_mode.is_target_verify():
seq_lens = seq_lens + self.num_draft_tokens
metadata.seq_lens_k = torch.zeros(
(bs,), dtype=torch.int32, device=seq_lens.device
) )
metadata.seq_lens_k.copy_(seq_lens.to(dtype=torch.int32)) self.init_forward_metadata_replay_cuda_graph(
elif forward_mode.is_draft_extend(include_v2=True): bs=bs,
num_tokens_per_bs = num_tokens // bs req_pool_indices=req_pool_indices,
metadata.max_seq_len_q = num_tokens_per_bs seq_lens=seq_lens,
metadata.sum_seq_lens_q = num_tokens_per_bs * bs seq_lens_sum=None,
metadata.cu_seqlens_q = torch.arange( encoder_lens=encoder_lens,
0, forward_mode=forward_mode,
bs * num_tokens_per_bs + 1, spec_info=spec_info,
num_tokens_per_bs, seq_lens_cpu=seq_lens.cpu(),
dtype=torch.int32,
device=seq_lens.device,
) )
metadata.seq_lens_q = torch.full(
(bs,), num_tokens_per_bs, dtype=torch.int32, device=seq_lens.device
)
# NOTE(draft_extend seq_len handling):
# forward_batch.seq_lens is the seq_lens of the prev_context + verified tokens.
# To account for pad_draft_extend_query, we need seq_lens = prev_context + max_draft_tokens.
# This will ensure queries align with kvs correctly when calling
# flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla.
seq_lens = seq_lens - metadata.seq_lens_q + metadata.max_seq_len_q
metadata.seq_lens_k = torch.zeros(
(bs,), dtype=torch.int32, device=seq_lens.device
)
metadata.seq_lens_k.copy_(seq_lens.to(dtype=torch.int32))
# Custom fast-path for decode/idle.
# Capture with full width so future longer sequences are safe during replay
max_blocks_per_seq = self._calc_padded_blocks(self.max_context_len)
block_kv_indices = self.decode_cuda_graph_kv_indices[:bs, :max_blocks_per_seq]
create_flashmla_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
None,
block_kv_indices,
self.req_to_token.stride(0),
max_blocks_per_seq,
PAGED_SIZE=self.page_size,
)
metadata.block_kv_indices = block_kv_indices
metadata.max_seq_len_k = self.max_context_len
self.decode_cuda_graph_metadata[bs] = metadata
self.forward_decode_metadata = metadata
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
self, self,
@@ -390,6 +390,39 @@ class WaveAttnBackend(AttentionBackend):
device=self.device, device=self.device,
) )
def _build_cuda_graph_forward_metadata(
self,
bs: int,
forward_mode: ForwardMode,
spec_info: Optional[SpecInput],
) -> ForwardMetadata:
if forward_mode.is_decode_or_idle():
return ForwardMetadata(
attn_logits=self.cuda_graph_attn_logits,
attn_lse=self.cuda_graph_attn_lse,
max_extend_len=None,
num_kv_splits=self.cuda_graph_num_kv_splits,
kv_indptr=self.kv_indptr[: bs + 1],
kv_indices=self.cuda_graph_kv_indices,
qo_indptr=None,
custom_mask=None,
mask_indptr=None,
)
elif forward_mode.is_target_verify():
return ForwardMetadata(
attn_logits=None,
attn_lse=None,
max_extend_len=self.num_draft_tokens,
num_kv_splits=None,
kv_indptr=self.kv_indptr[: bs + 1],
kv_indices=self.cuda_graph_kv_indices,
qo_indptr=self.qo_indptr[: bs + 1],
custom_mask=self.cuda_graph_custom_mask,
mask_indptr=self.mask_indptr[: bs + 1],
)
else:
raise ValueError(f"Invalid forward mode: {forward_mode=} for CUDA Graph.")
def init_forward_metadata_capture_cuda_graph( def init_forward_metadata_capture_cuda_graph(
self, self,
bs: int, bs: int,
@@ -402,76 +435,34 @@ class WaveAttnBackend(AttentionBackend):
): ):
assert encoder_lens is None, "Not supported" assert encoder_lens is None, "Not supported"
if forward_mode.is_decode_or_idle(): # Multi-step speculative decode: kv buffers come from spec_info rather than
if spec_info is None: # the cuda-graph pool, so replay is not involved for this path.
kv_indptr = self.kv_indptr if forward_mode.is_decode_or_idle() and spec_info is not None:
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
else:
kv_indptr, kv_indices = spec_info.kv_indptr, spec_info.kv_indices
attn_logits = self.cuda_graph_attn_logits
attn_lse = self.cuda_graph_attn_lse
max_extend_len = None
num_kv_splits = self.cuda_graph_num_kv_splits
qo_indptr = None
custom_mask = None
mask_indptr = None
elif forward_mode.is_target_verify():
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
(1 + bs) * self.num_draft_tokens,
step=self.num_draft_tokens,
dtype=torch.int32,
device=self.device,
)
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(seq_lens, dim=0)
kv_indices = self.cuda_graph_kv_indices
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
req_pool_indices,
seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
custom_mask = self.cuda_graph_custom_mask
seq_mask_len = self.num_draft_tokens * (seq_lens + self.num_draft_tokens)
mask_indptr = self.mask_indptr[: bs + 1]
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
max_extend_len = self.num_draft_tokens
num_kv_splits = None
attn_logits = None
attn_lse = None
else:
raise ValueError(
f"Invalid forward mode: {forward_mode=} for CUDA Graph capture."
)
self.forward_metadata = ForwardMetadata( self.forward_metadata = ForwardMetadata(
attn_logits, attn_logits=self.cuda_graph_attn_logits,
attn_lse, attn_lse=self.cuda_graph_attn_lse,
max_extend_len, max_extend_len=None,
num_kv_splits, num_kv_splits=self.cuda_graph_num_kv_splits,
kv_indptr, kv_indptr=spec_info.kv_indptr,
kv_indices, kv_indices=spec_info.kv_indices,
qo_indptr, qo_indptr=None,
custom_mask, custom_mask=None,
mask_indptr, mask_indptr=None,
)
return
self.init_forward_metadata_replay_cuda_graph(
bs=bs,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
seq_lens_sum=None,
encoder_lens=encoder_lens,
forward_mode=forward_mode,
spec_info=spec_info,
seq_lens_cpu=None,
)
self.forward_metadata = self._build_cuda_graph_forward_metadata(
bs, forward_mode, spec_info
) )
def init_forward_metadata_replay_cuda_graph( def init_forward_metadata_replay_cuda_graph(
@@ -485,9 +476,7 @@ class WaveAttnBackend(AttentionBackend):
spec_info: Optional[SpecInput], spec_info: Optional[SpecInput],
seq_lens_cpu: Optional[torch.Tensor], seq_lens_cpu: Optional[torch.Tensor],
): ):
# NOTE: encoder_lens expected to be zeros or None
if forward_mode.is_decode_or_idle(): if forward_mode.is_decode_or_idle():
# Update kv_indptr, kv_indices
kv_indptr = self.kv_indptr kv_indptr = self.kv_indptr
kv_indices = self.cuda_graph_kv_indices kv_indices = self.cuda_graph_kv_indices
num_kv_splits = self.cuda_graph_num_kv_splits num_kv_splits = self.cuda_graph_num_kv_splits
@@ -510,7 +499,6 @@ class WaveAttnBackend(AttentionBackend):
num_token = spec_info.kv_indptr.shape[0] - 1 num_token = spec_info.kv_indptr.shape[0] - 1
self.get_num_kv_splits(num_kv_splits[:num_token], seq_lens[:bs]) self.get_num_kv_splits(num_kv_splits[:num_token], seq_lens[:bs])
elif forward_mode.is_target_verify(): elif forward_mode.is_target_verify():
# Update qo_indptr, kv_indptr, kv_indices, custom_mask, mask_indptr
bs = len(req_pool_indices) bs = len(req_pool_indices)
qo_indptr = self.qo_indptr[: bs + 1] qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange( qo_indptr[: bs + 1] = torch.arange(
@@ -7,6 +7,7 @@ import torch
from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS from sglang.srt.layers.attention.attention_registry import ATTENTION_BACKENDS
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.utils import get_device_sm
from sglang.test.test_utils import CustomTestCase from sglang.test.test_utils import CustomTestCase
sys.path.insert(0, str(Path(__file__).resolve().parents[1])) sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
@@ -21,6 +22,9 @@ from sglang.test.kits.attention_unittest.attention_methods.dense_attention impor
replace_backend, replace_backend,
run_dense_fixture_eager, run_dense_fixture_eager,
) )
from sglang.test.kits.attention_unittest.runner_modes.speculative_target_verify_runner import (
_prepare_spec_verify_batch,
)
register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200") register_cuda_ci(est_time=20, stage="base-b", runner_config="4-gpu-b200")
register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large") register_cuda_ci(est_time=20, stage="base-b", runner_config="1-gpu-large")
@@ -53,11 +57,27 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase):
extend_lens=(16,), extend_lens=(16,),
) )
# Mirrors ``runner_fa3_eagle_verify_chain`` in test_fa3.py — the smallest
# case shape that drives a real TARGET_VERIFY CUDA-graph capture through
# FlashAttention's per-bs metadata dicts.
TARGET_VERIFY_CAPTURE_CASE = DenseAttentionCase(
name="tbo_fa3_target_verify_chain_capture",
backend="fa3",
forward_mode=ForwardMode.TARGET_VERIFY,
num_heads=4,
num_kv_heads=4,
page_size=16,
prefix_lens=(4, 7),
extend_lens=(3, 3),
)
def _build_and_wrap(self, case: DenseAttentionCase): def _build_and_wrap(self, case: DenseAttentionCase):
fixture = build_dense_attention_fixture(self, case) fixture = build_dense_attention_fixture(self, case)
try: try:
primary = ATTENTION_BACKENDS["triton"](fixture.runner) primary = ATTENTION_BACKENDS[case.backend](fixture.runner)
children = [ATTENTION_BACKENDS["triton"](fixture.runner) for _ in range(2)] children = [
ATTENTION_BACKENDS[case.backend](fixture.runner) for _ in range(2)
]
except (AssertionError, ImportError, ModuleNotFoundError) as exc: except (AssertionError, ImportError, ModuleNotFoundError) as exc:
self.skipTest(f"tbo child backend unavailable: {exc}") self.skipTest(f"tbo child backend unavailable: {exc}")
wrapper = TboAttnBackend(primary=primary, children=children) wrapper = TboAttnBackend(primary=primary, children=children)
@@ -69,6 +89,56 @@ class TestTboAttnDenseAttentionBackendCorrectness(CustomTestCase):
expected = expected_dense_fixture_output(fixture) expected = expected_dense_fixture_output(fixture)
torch.testing.assert_close(actual, expected, atol=DENSE_ATOL, rtol=DENSE_RTOL) torch.testing.assert_close(actual, expected, atol=DENSE_ATOL, rtol=DENSE_RTOL)
@unittest.skipIf(
get_device_sm() >= 100 or get_device_sm() < 80,
"FA3 backend requires SM 80-90",
)
def test_tbo_target_verify_cuda_graph_capture_delegates_to_primary_capture(self):
"""TBO capture must invoke ``primary.init_forward_metadata_capture_cuda_graph``,
not ``primary.init_forward_metadata_replay_cuda_graph``.
Backends like FlashAttention store per-bs metadata in dicts populated
only by their capture path (via ``_bind_metadata_buffers``). If TBO
short-circuits its capture to its own replay (which delegates to
``primary.replay``), those dicts are empty and replay raises
``KeyError: bs``. Reproduces the deepep-4-gpu-h100 failure where
``flashattention_backend.target_verify_metadata[bs]`` lookup blew up
during ``init_device_graphs``.
Asserts capture completes without raising — numerical correctness of
the captured graph is covered by per-backend spec-verify tests.
"""
case = self.TARGET_VERIFY_CAPTURE_CASE
fixture = self._build_and_wrap(case)
wrapper = fixture.backend
batch = fixture.forward_batch
# Wire TARGET_VERIFY batch state + EAGLE chain (topk=1) spec_info,
# mirroring what the per-backend spec-verify runner sets up.
_prepare_spec_verify_batch(
case,
batch,
topk=1,
spec_kind="eagle",
device=str(batch.seq_lens.device),
)
capture_bs = case.batch_size
num_tokens = sum(case.extend_lens)
wrapper.init_cuda_graph_state(max_bs=capture_bs, max_num_tokens=num_tokens)
# This is the failing call before the fix: TBO.capture delegating to
# primary.replay (instead of primary.capture) reads an unpopulated
# ``target_verify_metadata[bs]`` dict and raises KeyError.
wrapper.init_forward_metadata_capture_cuda_graph(
bs=capture_bs,
num_tokens=num_tokens,
req_pool_indices=batch.req_pool_indices,
seq_lens=batch.seq_lens,
encoder_lens=batch.encoder_lens,
forward_mode=batch.forward_mode,
spec_info=batch.spec_info,
)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()