[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:
co-authored by
Claude Sonnet 4.6
parent
7fb7b41a3e
commit
ff8ed7a302
@@ -469,33 +469,15 @@ class AscendAttnBackend(AttentionBackend):
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def init_forward_metadata_capture_cuda_graph(
|
||||
def _init_cuda_graph_metadata(
|
||||
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],
|
||||
):
|
||||
seq_lens: torch.Tensor,
|
||||
) -> "ForwardMetadata":
|
||||
"""Create and store the per-bs ForwardMetadata for CUDA graph capture."""
|
||||
metadata = ForwardMetadata()
|
||||
|
||||
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:
|
||||
metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :]
|
||||
metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist()
|
||||
@@ -515,7 +497,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
)
|
||||
else:
|
||||
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,
|
||||
device=seq_lens.device,
|
||||
)
|
||||
@@ -528,13 +510,11 @@ class AscendAttnBackend(AttentionBackend):
|
||||
metadata.seq_lens_list_cumsum = (
|
||||
torch.cumsum(extend_seq_lens_cpu_int, dim=0).int().tolist()
|
||||
)
|
||||
|
||||
if (
|
||||
self.q_head_num_padding is not None
|
||||
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.
|
||||
# Therefore, we pad the head dimension accordingly and initialize an empty tensor for padding.
|
||||
dtype = self.model_dtype if self.model_dtype is not None else torch.bfloat16
|
||||
metadata.nope_padding = torch.empty(
|
||||
[
|
||||
bs,
|
||||
@@ -542,9 +522,7 @@ class AscendAttnBackend(AttentionBackend):
|
||||
self.q_head_num_padding - self.tp_q_head_num,
|
||||
self.kv_lora_rank,
|
||||
],
|
||||
dtype=(
|
||||
self.model_dtype if self.model_dtype is not None else torch.bfloat16
|
||||
),
|
||||
dtype=dtype,
|
||||
device=seq_lens.device,
|
||||
)
|
||||
metadata.rope_padding = torch.empty(
|
||||
@@ -554,16 +532,33 @@ class AscendAttnBackend(AttentionBackend):
|
||||
self.q_head_num_padding - self.tp_q_head_num,
|
||||
self.qk_rope_head_dim,
|
||||
],
|
||||
dtype=(
|
||||
self.model_dtype if self.model_dtype is not None else torch.bfloat16
|
||||
),
|
||||
dtype=dtype,
|
||||
device=seq_lens.device,
|
||||
)
|
||||
|
||||
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(
|
||||
self,
|
||||
|
||||
@@ -93,19 +93,16 @@ class AscendGDNAttnBackend(AscendMambaAttnBackendBase):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
|
||||
):
|
||||
if forward_mode.is_draft_extend(True):
|
||||
return
|
||||
super().init_forward_metadata_capture_cuda_graph(
|
||||
bs,
|
||||
num_tokens,
|
||||
req_pool_indices,
|
||||
seq_lens,
|
||||
encoder_lens,
|
||||
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=None,
|
||||
encoder_lens=encoder_lens,
|
||||
forward_mode=forward_mode,
|
||||
spec_info=spec_info,
|
||||
seq_lens_cpu=seq_lens.cpu(),
|
||||
)
|
||||
self.prepare_gdn_inputs(bs, forward_mode, spec_info)
|
||||
self.graph_mode = True
|
||||
|
||||
def init_forward_metadata_replay_cuda_graph(
|
||||
self,
|
||||
|
||||
@@ -1492,423 +1492,16 @@ class AiterAttnBackend(AttentionBackend):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[SpecInput],
|
||||
):
|
||||
|
||||
num_kv_splits = None
|
||||
# num_kv_splits_indptr = None
|
||||
|
||||
work_metadata = None
|
||||
work_info_set = None
|
||||
work_indptr = None
|
||||
|
||||
reduce_indptr = None
|
||||
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=}")
|
||||
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(
|
||||
self,
|
||||
@@ -1934,7 +1527,11 @@ class AiterAttnBackend(AttentionBackend):
|
||||
reduce_partial_map = 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():
|
||||
qo_indptr = None
|
||||
|
||||
@@ -153,24 +153,22 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[SpecInput],
|
||||
):
|
||||
if forward_mode.is_decode_or_idle():
|
||||
if spec_info is None:
|
||||
max_seqlen_pad = self.cuda_graph_kv_indices.shape[1]
|
||||
|
||||
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),
|
||||
PAGED_SIZE=PAGE_SIZE,
|
||||
)
|
||||
self.forward_metadata = CutlassMLADecodeMetadata(
|
||||
self.cuda_graph_mla_workspace,
|
||||
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
|
||||
)
|
||||
if forward_mode.is_decode_or_idle() and spec_info is None:
|
||||
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,
|
||||
)
|
||||
max_seqlen_pad = self.cuda_graph_kv_indices.shape[1]
|
||||
self.forward_metadata = CutlassMLADecodeMetadata(
|
||||
self.cuda_graph_mla_workspace,
|
||||
self.cuda_graph_kv_indices[:bs, :max_seqlen_pad],
|
||||
)
|
||||
else:
|
||||
super().init_forward_metadata_capture_cuda_graph(
|
||||
bs,
|
||||
@@ -193,15 +191,11 @@ class CutlassMLABackend(FlashInferMLAAttnBackend):
|
||||
spec_info: Optional[SpecInput],
|
||||
seq_lens_cpu: Optional[torch.Tensor],
|
||||
):
|
||||
|
||||
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,)](
|
||||
self.req_to_token,
|
||||
req_pool_indices[:bs],
|
||||
seq_lens,
|
||||
seq_lens[:bs],
|
||||
None,
|
||||
self.cuda_graph_kv_indices,
|
||||
self.req_to_token.stride(0),
|
||||
|
||||
@@ -749,48 +749,40 @@ class DeepseekV4AttnBackend(
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[SpecInput],
|
||||
) -> None:
|
||||
from types import SimpleNamespace
|
||||
|
||||
assert req_pool_indices.size(0) == bs
|
||||
assert seq_lens.size(0) == bs
|
||||
|
||||
bucket = _GraphBucket.of(forward_mode)
|
||||
raw_type: Optional[type] = None
|
||||
if bucket == _GraphBucket.DECODE_OR_IDLE:
|
||||
metadata = self.init_forward_metadata_decode(
|
||||
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
|
||||
dummy_cache_loc = torch.zeros_like(seq_lens)
|
||||
elif bucket == _GraphBucket.TARGET_VERIFY:
|
||||
out_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,
|
||||
)
|
||||
dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs)
|
||||
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.forward_metadata = metadata
|
||||
if raw_type is not None:
|
||||
self._current_capture_raw = (
|
||||
metadata if isinstance(metadata, raw_type) else None
|
||||
)
|
||||
self._replay_forward_batch = SimpleNamespace(
|
||||
out_cache_loc=dummy_cache_loc,
|
||||
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 = (
|
||||
metadata
|
||||
if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata))
|
||||
else None
|
||||
)
|
||||
|
||||
def init_forward_metadata_replay_cuda_graph(
|
||||
self,
|
||||
@@ -892,6 +884,11 @@ class DeepseekV4AttnBackend(
|
||||
],
|
||||
bucket: _GraphBucket,
|
||||
) -> 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.copy_(temp_metadata)
|
||||
self.forward_metadata = chosen_metadata
|
||||
|
||||
@@ -748,48 +748,40 @@ class DeepseekV4HipRadixBackend(
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[SpecInput],
|
||||
) -> None:
|
||||
from types import SimpleNamespace
|
||||
|
||||
assert req_pool_indices.size(0) == bs
|
||||
assert seq_lens.size(0) == bs
|
||||
|
||||
bucket = _GraphBucket.of(forward_mode)
|
||||
raw_type: Optional[type] = None
|
||||
if bucket == _GraphBucket.DECODE_OR_IDLE:
|
||||
metadata = self.init_forward_metadata_decode(
|
||||
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
|
||||
dummy_cache_loc = torch.zeros_like(seq_lens)
|
||||
elif bucket == _GraphBucket.TARGET_VERIFY:
|
||||
out_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,
|
||||
)
|
||||
dummy_cache_loc = torch.zeros(num_tokens, **self.cuda_int32_kwargs)
|
||||
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.forward_metadata = metadata
|
||||
if raw_type is not None:
|
||||
self._current_capture_raw = (
|
||||
metadata if isinstance(metadata, raw_type) else None
|
||||
)
|
||||
self._replay_forward_batch = SimpleNamespace(
|
||||
out_cache_loc=dummy_cache_loc,
|
||||
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 = (
|
||||
metadata
|
||||
if isinstance(metadata, (DSV4RawDecodeMetadata, DSV4RawVerifyMetadata))
|
||||
else None
|
||||
)
|
||||
|
||||
def init_forward_metadata_replay_cuda_graph(
|
||||
self,
|
||||
@@ -891,6 +883,11 @@ class DeepseekV4HipRadixBackend(
|
||||
],
|
||||
bucket: _GraphBucket,
|
||||
) -> 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.copy_(temp_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,
|
||||
bs: int,
|
||||
num_tokens: int,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
encoder_lens: Optional[torch.Tensor],
|
||||
seq_lens_cpu: Optional[torch.Tensor],
|
||||
forward_mode: ForwardMode,
|
||||
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)
|
||||
|
||||
"""Initialize forward metadata for capturing CUDA graph."""
|
||||
if forward_mode.is_decode_or_idle():
|
||||
# Normal Decode
|
||||
# Get sequence information
|
||||
@@ -847,11 +849,11 @@ class DeepseekSparseAttnBackend(
|
||||
)
|
||||
|
||||
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":
|
||||
flashmla_metadata = self.decode_cuda_graph_metadata[
|
||||
"flashmla_metadata"
|
||||
].slice(slice(0, num_tokens + 1))
|
||||
].slice(slice(0, bs + 1))
|
||||
flashmla_metadata.copy_(
|
||||
self._compute_flashmla_metadata(
|
||||
cache_seqlens=dsa_cache_seqlens_int32,
|
||||
@@ -969,6 +971,28 @@ class DeepseekSparseAttnBackend(
|
||||
self.decode_cuda_graph_metadata[bs] = 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(
|
||||
self,
|
||||
bs: int,
|
||||
@@ -985,6 +1009,20 @@ class DeepseekSparseAttnBackend(
|
||||
"""Initialize forward metadata for replaying CUDA graph."""
|
||||
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)
|
||||
|
||||
seq_lens = seq_lens[:bs]
|
||||
|
||||
@@ -532,16 +532,13 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
|
||||
),
|
||||
}
|
||||
|
||||
def init_forward_metadata_capture_cuda_graph(
|
||||
def _bind_metadata_buffers(
|
||||
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],
|
||||
):
|
||||
"""Allocate persistent metadata buffers for CUDA graph capture."""
|
||||
metadata = DualChunkFlashAttentionMetadata()
|
||||
|
||||
if forward_mode.is_decode_or_idle():
|
||||
@@ -580,6 +577,36 @@ class DualChunkFlashAttentionBackend(AttentionBackend):
|
||||
|
||||
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(
|
||||
self,
|
||||
bs: int,
|
||||
|
||||
@@ -1700,43 +1700,37 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
# For decoder-only models, skip encoder_metadata allocation
|
||||
self.encoder_metadata = {}
|
||||
|
||||
def init_forward_metadata_capture_cuda_graph(
|
||||
def _bind_metadata_buffers(
|
||||
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."""
|
||||
metadata = FlashAttentionMetadata()
|
||||
device: torch.device,
|
||||
) -> tuple:
|
||||
"""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()
|
||||
|
||||
device = seq_lens.device
|
||||
if forward_mode.is_decode_or_idle():
|
||||
if spec_info is not None:
|
||||
# Draft Decode
|
||||
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[
|
||||
"cache_seqlens"
|
||||
][:bs]
|
||||
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"
|
||||
][: bs + 1]
|
||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||
torch.cumsum(
|
||||
metadata.cache_seqlens_int32, dim=0, dtype=torch.int32
|
||||
),
|
||||
(1, 0),
|
||||
)
|
||||
metadata.cu_seqlens_k = self.decode_cuda_graph_metadata[
|
||||
"cu_seqlens_k"
|
||||
][: bs + 1]
|
||||
metadata.page_table = self.decode_cuda_graph_metadata[
|
||||
"page_table_draft_decode"
|
||||
][:bs, :]
|
||||
@@ -1746,13 +1740,11 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
][:bs, :]
|
||||
self.decode_cuda_graph_metadata[bs] = metadata
|
||||
else:
|
||||
# When top k > 1, we need two specific draft decode metadata, and then merge states
|
||||
# 1. The first half of metadata for prefix tokens
|
||||
# Draft Decode topk>1: two metadata objects
|
||||
metadata.cache_seqlens_int32 = (
|
||||
self.draft_decode_metadata_topk_normal["cache_seqlens"][:bs]
|
||||
)
|
||||
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[
|
||||
"cu_seqlens_q"
|
||||
][: bs + 1]
|
||||
@@ -1763,7 +1755,6 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
"page_table"
|
||||
][:bs, :]
|
||||
|
||||
# 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk)
|
||||
metadata_expand.cache_seqlens_int32 = (
|
||||
self.draft_decode_metadata_topk_expand["cache_seqlens"][
|
||||
: bs * self.topk
|
||||
@@ -1787,16 +1778,15 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
self.draft_decode_metadata_topk_expand[bs] = metadata_expand
|
||||
else:
|
||||
# Normal Decode
|
||||
# Get sequence information
|
||||
metadata.cache_seqlens_int32 = seq_lens.to(torch.int32)
|
||||
batch_size = len(seq_lens)
|
||||
device = seq_lens.device
|
||||
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 page table
|
||||
metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[
|
||||
"cache_seqlens"
|
||||
][:bs]
|
||||
metadata.cu_seqlens_q = self.decode_cuda_graph_metadata["cu_seqlens_q"][
|
||||
: bs + 1
|
||||
]
|
||||
metadata.cu_seqlens_k = self.decode_cuda_graph_metadata["cu_seqlens_k"][
|
||||
: bs + 1
|
||||
]
|
||||
metadata.page_table = self.decode_cuda_graph_metadata["page_table"][
|
||||
:bs, :
|
||||
]
|
||||
@@ -1804,70 +1794,32 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata.swa_page_table = self.decode_cuda_graph_metadata[
|
||||
"swa_page_table"
|
||||
][: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._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():
|
||||
if self.topk <= 1:
|
||||
metadata.cache_seqlens_int32 = self.target_verify_metadata[
|
||||
"cache_seqlens"
|
||||
][: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_k = (
|
||||
seq_lens.max().item() + self.speculative_num_draft_tokens
|
||||
)
|
||||
|
||||
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_q = self.target_verify_metadata["cu_seqlens_q"][
|
||||
: bs + 1
|
||||
]
|
||||
metadata.cu_seqlens_k = self.target_verify_metadata["cu_seqlens_k"][
|
||||
: (bs + 1)
|
||||
]
|
||||
|
||||
metadata.page_table = self.target_verify_metadata["page_table"][:bs, :]
|
||||
|
||||
if self.use_sliding_window_kv_pool:
|
||||
metadata.swa_page_table = self.target_verify_metadata[
|
||||
"swa_page_table"
|
||||
][:bs, :]
|
||||
|
||||
self.target_verify_metadata[bs] = metadata
|
||||
else:
|
||||
# When topk > 1, we need two specific target verify metadata, and then merge states
|
||||
# 1. The first half of metadata for prefix tokens
|
||||
# Target Verify topk>1: two (or three with SWA) metadata objects
|
||||
metadata.cache_seqlens_int32 = self.target_verify_metadata_topk_normal[
|
||||
"cache_seqlens"
|
||||
][:bs]
|
||||
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[
|
||||
"cu_seqlens_q"
|
||||
][: bs + 1]
|
||||
@@ -1878,7 +1830,6 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
"page_table"
|
||||
][:bs, :]
|
||||
|
||||
# 2. The second half of metadata for draft tokens (per_batch_num_tokens = topk)
|
||||
metadata_expand.cache_seqlens_int32 = (
|
||||
self.target_verify_metadata_topk_expand["cache_seqlens"][
|
||||
: bs * self.speculative_num_draft_tokens
|
||||
@@ -1891,7 +1842,6 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata_expand.cu_seqlens_k = self.target_verify_metadata_topk_expand[
|
||||
"cu_seqlens_k"
|
||||
][: bs * self.speculative_num_draft_tokens + 1]
|
||||
|
||||
metadata_expand.page_table = self.target_verify_metadata_topk_expand[
|
||||
"page_table"
|
||||
][: bs * self.speculative_num_draft_tokens]
|
||||
@@ -1913,7 +1863,6 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata_swa.cu_seqlens_k = self.target_verify_metadata_topk_swa[
|
||||
"cu_seqlens_k"
|
||||
][: bs * self.speculative_num_draft_tokens + 1]
|
||||
|
||||
metadata_swa.page_table = self.target_verify_metadata_topk_swa[
|
||||
"page_table"
|
||||
][: bs * self.speculative_num_draft_tokens]
|
||||
@@ -1921,33 +1870,20 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata.swa_spec_metadata = metadata_swa
|
||||
|
||||
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"][
|
||||
: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_k = seq_lens.max().item()
|
||||
|
||||
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_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1]
|
||||
metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][
|
||||
: (bs + 1)
|
||||
]
|
||||
metadata.page_table = self.draft_extend_metadata["page_table"][:bs, :]
|
||||
|
||||
if self.use_sliding_window_kv_pool:
|
||||
metadata.swa_page_table = self.draft_extend_metadata["swa_page_table"][
|
||||
:bs, :
|
||||
]
|
||||
|
||||
self.draft_extend_metadata[bs] = metadata
|
||||
|
||||
if encoder_lens is not None:
|
||||
@@ -1958,13 +1894,81 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
metadata.encoder_cu_seqlens_k = self.encoder_metadata[
|
||||
"encoder_cu_seqlens_k"
|
||||
][: (encoder_bs + 1)]
|
||||
|
||||
metadata.encoder_page_table = self.encoder_metadata["encoder_page_table"][
|
||||
:bs, :
|
||||
]
|
||||
|
||||
self.forward_metadata = metadata
|
||||
self.forward_metadata_spec_decode_expand = metadata_expand
|
||||
return metadata, 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(
|
||||
self,
|
||||
|
||||
@@ -557,6 +557,81 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
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]
|
||||
|
||||
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(
|
||||
self,
|
||||
bs: int,
|
||||
@@ -567,148 +642,24 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
forward_mode: ForwardMode,
|
||||
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():
|
||||
decode_wrappers = []
|
||||
for i in range(self.num_wrappers):
|
||||
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=}")
|
||||
for w in self.decode_cuda_graph_metadata[bs]:
|
||||
w.begin_forward = partial(fast_decode_plan, w)
|
||||
|
||||
def init_forward_metadata_replay_cuda_graph(
|
||||
self,
|
||||
@@ -733,19 +684,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
fixed_split_size=None,
|
||||
disable_split_kv=self.disable_cuda_graph_kv_split,
|
||||
)
|
||||
elif forward_mode.is_target_verify():
|
||||
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():
|
||||
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
|
||||
self.indices_updater_prefill.update(
|
||||
req_pool_indices[:bs],
|
||||
seq_lens[:bs],
|
||||
|
||||
@@ -384,7 +384,13 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[SpecInput],
|
||||
):
|
||||
seq_lens_sum = seq_lens.sum().item()
|
||||
seq_lens_cpu = seq_lens.cpu()
|
||||
|
||||
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(
|
||||
self.workspace_buffer,
|
||||
use_cuda_graph=True,
|
||||
@@ -394,8 +400,6 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
kv_len_arr=self.cuda_graph_kv_lens[:num_tokens],
|
||||
backend="auto",
|
||||
)
|
||||
|
||||
seq_lens_sum = seq_lens.sum().item()
|
||||
self.indices_updater_decode.update(
|
||||
req_pool_indices,
|
||||
seq_lens,
|
||||
@@ -406,9 +410,12 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
)
|
||||
self.decode_cuda_graph_metadata[bs] = 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)
|
||||
elif forward_mode.is_target_verify():
|
||||
verify_wrapper = BatchMLAPagedAttentionWrapper(
|
||||
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
|
||||
# Prefill: create wrapper and store — replay handles the update call.
|
||||
prefill_wrapper = BatchMLAPagedAttentionWrapper(
|
||||
self.workspace_buffer,
|
||||
use_cuda_graph=True,
|
||||
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],
|
||||
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=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)
|
||||
self.prefill_cuda_graph_metadata[bs] = prefill_wrapper
|
||||
self.forward_metadata = PrefillMetadata(prefill_wrapper, False)
|
||||
else:
|
||||
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(
|
||||
self,
|
||||
bs: int,
|
||||
@@ -488,17 +474,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
spec_info=spec_info,
|
||||
**self.fast_decode_kwargs,
|
||||
)
|
||||
elif forward_mode.is_target_verify():
|
||||
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():
|
||||
elif forward_mode.is_target_verify() or forward_mode.is_draft_extend():
|
||||
self.indices_updater_prefill.update(
|
||||
req_pool_indices[:bs],
|
||||
seq_lens[:bs],
|
||||
|
||||
@@ -4,6 +4,7 @@ Support attention backend for FlashMLA.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
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.speculative.spec_info import SpecInput
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
PAGE_SIZE = 64
|
||||
|
||||
@@ -193,83 +195,16 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[SpecInput],
|
||||
):
|
||||
if forward_mode.is_decode_or_idle():
|
||||
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),
|
||||
)
|
||||
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],
|
||||
if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify():
|
||||
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,
|
||||
)
|
||||
else:
|
||||
super().init_forward_metadata_capture_cuda_graph(
|
||||
@@ -293,60 +228,21 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
||||
spec_info: Optional[SpecInput],
|
||||
seq_lens_cpu: Optional[torch.Tensor],
|
||||
):
|
||||
if forward_mode.is_decode_or_idle():
|
||||
assert seq_lens_cpu is not None
|
||||
if forward_mode.is_decode_or_idle() or forward_mode.is_target_verify():
|
||||
seq_lens = seq_lens[:bs]
|
||||
seq_lens_cpu = seq_lens_cpu[:bs]
|
||||
max_seqlen_pad = triton.cdiv(seq_lens_cpu.max().item(), PAGE_SIZE)
|
||||
seq_lens_cpu = seq_lens_cpu[:bs] if seq_lens_cpu is not None else None
|
||||
|
||||
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),
|
||||
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()
|
||||
)
|
||||
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]
|
||||
|
||||
if actual_num_sm_parts != self.cuda_graph_mla_metadata_view.shape[0]:
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
logger.warning(
|
||||
f"num_sm_parts mismatch in CUDA Graph replay: "
|
||||
f"capture={self.cuda_graph_mla_metadata_view.shape[0]}, "
|
||||
f"replay={actual_num_sm_parts}. "
|
||||
f"This may indicate batch size changed between capture and replay."
|
||||
)
|
||||
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
|
||||
]
|
||||
|
||||
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)
|
||||
max_seqlen_pad = triton.cdiv(seq_max, PAGE_SIZE)
|
||||
|
||||
create_flashmla_kv_indices_triton[(bs,)](
|
||||
self.req_to_token,
|
||||
@@ -358,29 +254,47 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
||||
self.cuda_graph_kv_indices.stride(0),
|
||||
)
|
||||
|
||||
q_head_mult = (
|
||||
self.num_draft_tokens if forward_mode.is_target_verify() else 1
|
||||
)
|
||||
mla_metadata, num_splits = get_mla_metadata(
|
||||
seq_lens.to(torch.int32),
|
||||
self.num_draft_tokens * self.num_q_heads,
|
||||
q_head_mult * 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], (
|
||||
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 (
|
||||
self.cuda_graph_mla_metadata_view is None
|
||||
or actual_num_sm_parts != self.cuda_graph_mla_metadata_view.shape[0]
|
||||
):
|
||||
if self.cuda_graph_mla_metadata_view is not None:
|
||||
logger.warning(
|
||||
f"num_sm_parts mismatch in CUDA Graph replay: "
|
||||
f"capture={self.cuda_graph_mla_metadata_view.shape[0]}, "
|
||||
f"replay={actual_num_sm_parts}. "
|
||||
f"This may indicate batch size changed between capture and replay."
|
||||
)
|
||||
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]
|
||||
# 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_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
|
||||
]
|
||||
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:
|
||||
super().init_forward_metadata_replay_cuda_graph(
|
||||
bs,
|
||||
|
||||
@@ -403,8 +403,15 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
|
||||
):
|
||||
self.forward_metadata = self._capture_metadata(
|
||||
bs, req_pool_indices, 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=None,
|
||||
encoder_lens=encoder_lens,
|
||||
forward_mode=forward_mode,
|
||||
spec_info=spec_info,
|
||||
seq_lens_cpu=None,
|
||||
)
|
||||
|
||||
def init_forward_metadata_replay_cuda_graph(
|
||||
@@ -539,9 +546,12 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
spec_info: Optional[SpecInput],
|
||||
seq_lens_cpu: Optional[torch.Tensor],
|
||||
):
|
||||
num_padding = torch.count_nonzero(
|
||||
seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value()
|
||||
)
|
||||
if seq_lens_cpu is None:
|
||||
num_padding = 0
|
||||
else:
|
||||
num_padding = torch.count_nonzero(
|
||||
seq_lens_cpu == self.get_cuda_graph_seq_len_fill_value()
|
||||
)
|
||||
# Make sure forward metadata is correctly handled for padding reqs
|
||||
req_pool_indices[bs - num_padding :] = 0
|
||||
mamba_indices = self.req_to_token_pool.get_mamba_indices(req_pool_indices)
|
||||
@@ -576,13 +586,17 @@ 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 forward_mode.is_target_verify() and self.topk > 1:
|
||||
bs_without_pad = spec_info.retrieve_next_token.shape[0]
|
||||
self.retrieve_next_token_list[bs - 1][:bs_without_pad].copy_(
|
||||
spec_info.retrieve_next_token
|
||||
)
|
||||
self.retrieve_next_sibling_list[bs - 1][:bs_without_pad].copy_(
|
||||
spec_info.retrieve_next_sibling
|
||||
)
|
||||
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]
|
||||
self.retrieve_next_token_list[bs - 1][:bs_without_pad].copy_(
|
||||
spec_info.retrieve_next_token
|
||||
)
|
||||
self.retrieve_next_sibling_list[bs - 1][:bs_without_pad].copy_(
|
||||
spec_info.retrieve_next_sibling
|
||||
)
|
||||
return ForwardMetadata(
|
||||
query_start_loc=self.query_start_loc_list[bs - 1],
|
||||
mamba_cache_indices=self.state_indices_list[bs - 1],
|
||||
@@ -703,13 +717,15 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
|
||||
):
|
||||
metadata = self._capture_metadata(bs, req_pool_indices, forward_mode, spec_info)
|
||||
draft_token_num = spec_info.draft_token_num if spec_info is not None else 1
|
||||
self.forward_metadata = Mamba2Metadata.prepare_decode(
|
||||
metadata,
|
||||
seq_lens,
|
||||
is_target_verify=forward_mode.is_target_verify(),
|
||||
draft_token_num=draft_token_num,
|
||||
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,
|
||||
)
|
||||
|
||||
def init_forward_metadata_replay_cuda_graph(
|
||||
|
||||
@@ -89,9 +89,15 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[Union[EagleDraftInput, EagleVerifyInput]],
|
||||
):
|
||||
metadata = self._capture_metadata(bs, req_pool_indices, forward_mode, spec_info)
|
||||
self.forward_metadata = BailingLinearMetadata.prepare_decode(
|
||||
metadata.query_start_loc, metadata.mamba_cache_indices, bs, 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=None,
|
||||
)
|
||||
|
||||
def init_forward_metadata_replay_cuda_graph(
|
||||
|
||||
@@ -283,11 +283,179 @@ class TritonAttnBackend(AttentionBackend):
|
||||
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):
|
||||
"""Init auxiliary variables for triton attention backend."""
|
||||
|
||||
bs = forward_batch.batch_size
|
||||
kv_indptr = self.kv_indptr
|
||||
window_kv_indptr = self.window_kv_indptr
|
||||
window_kv_indices = None
|
||||
window_num_kv_splits = None
|
||||
@@ -297,19 +465,14 @@ class TritonAttnBackend(AttentionBackend):
|
||||
|
||||
if forward_batch.forward_mode.is_decode_or_idle():
|
||||
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(
|
||||
forward_batch.seq_lens_sum, dtype=torch.int64, device=self.device
|
||||
)
|
||||
create_flashinfer_kv_indices_triton[(bs,)](
|
||||
self.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
kv_indptr = self._fill_kv_indptr_and_indices(
|
||||
bs,
|
||||
forward_batch.seq_lens,
|
||||
kv_indptr,
|
||||
None,
|
||||
forward_batch.req_pool_indices,
|
||||
kv_indices,
|
||||
self.req_to_token.stride(0),
|
||||
)
|
||||
# Sliding window
|
||||
if (
|
||||
@@ -371,19 +534,14 @@ class TritonAttnBackend(AttentionBackend):
|
||||
device=self.device,
|
||||
)
|
||||
# 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_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,)](
|
||||
self.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
kv_indptr = self._fill_kv_indptr_and_indices(
|
||||
bs,
|
||||
forward_batch.seq_lens,
|
||||
kv_indptr,
|
||||
None,
|
||||
forward_batch.req_pool_indices,
|
||||
kv_indices,
|
||||
self.req_to_token.stride(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_lse = None
|
||||
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(
|
||||
sum(forward_batch.extend_prefix_lens_cpu),
|
||||
dtype=torch.int64,
|
||||
device=self.device,
|
||||
)
|
||||
create_flashinfer_kv_indices_triton[(bs,)](
|
||||
self.req_to_token,
|
||||
forward_batch.req_pool_indices,
|
||||
kv_indptr = self._fill_kv_indptr_and_indices(
|
||||
bs,
|
||||
forward_batch.extend_prefix_lens,
|
||||
kv_indptr,
|
||||
None,
|
||||
forward_batch.req_pool_indices,
|
||||
kv_indices,
|
||||
self.req_to_token.stride(0),
|
||||
)
|
||||
# Sliding window
|
||||
if self.sliding_window_size is not None and self.sliding_window_size > 0:
|
||||
@@ -578,6 +729,83 @@ class TritonAttnBackend(AttentionBackend):
|
||||
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(
|
||||
self,
|
||||
bs: int,
|
||||
@@ -589,172 +817,42 @@ class TritonAttnBackend(AttentionBackend):
|
||||
spec_info: Optional[SpecInput],
|
||||
):
|
||||
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
|
||||
# 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(
|
||||
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=spec_info.kv_indptr,
|
||||
kv_indices=spec_info.kv_indices,
|
||||
qo_indptr=None,
|
||||
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,
|
||||
swa_attn_logits=self.cuda_graph_swa_attn_logits,
|
||||
)
|
||||
return
|
||||
|
||||
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."
|
||||
)
|
||||
|
||||
self.forward_metadata = ForwardMetadata(
|
||||
attn_logits,
|
||||
attn_lse,
|
||||
max_extend_len,
|
||||
num_kv_splits,
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
qo_indptr,
|
||||
custom_mask,
|
||||
mask_indptr,
|
||||
window_kv_indptr,
|
||||
window_kv_indices,
|
||||
window_num_kv_splits,
|
||||
window_kv_offsets,
|
||||
swa_attn_logits=swa_attn_logits,
|
||||
# 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(
|
||||
@@ -770,138 +868,23 @@ class TritonAttnBackend(AttentionBackend):
|
||||
):
|
||||
# NOTE: encoder_lens expected to be zeros or None
|
||||
if forward_mode.is_decode_or_idle():
|
||||
# Update kv_indptr, kv_indices
|
||||
kv_indptr = self.kv_indptr
|
||||
kv_indices = self.cuda_graph_kv_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),
|
||||
assert spec_info is None, "Multi-step cuda graph init is not done here."
|
||||
_, _, window_kv_lens = self._update_decode_kv_buffers(
|
||||
bs, seq_lens, req_pool_indices
|
||||
)
|
||||
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.cuda_graph_window_num_kv_splits[:bs], window_kv_lens[:bs]
|
||||
)
|
||||
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(
|
||||
window_num_kv_splits[:num_token], 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():
|
||||
# Update qo_indptr, kv_indptr, kv_indices, custom_mask, mask_indptr
|
||||
bs = len(req_pool_indices)
|
||||
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,
|
||||
self._update_target_verify_buffers(
|
||||
bs, seq_lens, req_pool_indices, spec_info
|
||||
)
|
||||
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):
|
||||
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,
|
||||
)
|
||||
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),
|
||||
self._update_draft_extend_buffers(
|
||||
bs, seq_lens, req_pool_indices, forward_mode, spec_info
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
@@ -1490,18 +1473,26 @@ def update_sliding_window_buffer(
|
||||
seq_lens,
|
||||
req_pool_indices,
|
||||
bs,
|
||||
device,
|
||||
device=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(
|
||||
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_indices = torch.empty(
|
||||
window_kv_indptr[-1], dtype=torch.int64, device=device
|
||||
)
|
||||
if window_kv_indices is None:
|
||||
window_kv_indices = torch.empty(
|
||||
window_kv_indptr[-1], dtype=torch.int64, device=device
|
||||
)
|
||||
window_kv_start_idx = seq_lens - window_kv_lens
|
||||
create_flashinfer_kv_indices_triton[(bs,)](
|
||||
req_to_token,
|
||||
@@ -1512,47 +1503,6 @@ def update_sliding_window_buffer(
|
||||
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"):
|
||||
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"):
|
||||
kv_last_index = window_kv_indptr[-1]
|
||||
# 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,
|
||||
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."""
|
||||
spec_info,
|
||||
device: torch.device,
|
||||
) -> "TRTLLMMHAMetadata":
|
||||
"""Create TRTLLMMHAMetadata with pre-allocated buffer slice refs, stored in the dict."""
|
||||
metadata = TRTLLMMHAMetadata()
|
||||
device = seq_lens.device
|
||||
|
||||
if forward_mode.is_decode_or_idle():
|
||||
if spec_info is not None:
|
||||
# Draft Decode
|
||||
# Here we only support topk = 1 for now.
|
||||
# Draft Decode (topk = 1)
|
||||
metadata.cache_seqlens_int32 = self.decode_cuda_graph_metadata[
|
||||
"cache_seqlens"
|
||||
][: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"][
|
||||
: bs + 1
|
||||
]
|
||||
metadata.cu_seqlens_k = torch.nn.functional.pad(
|
||||
torch.cumsum(
|
||||
metadata.cache_seqlens_int32, dim=0, dtype=torch.int32
|
||||
),
|
||||
(1, 0),
|
||||
)
|
||||
metadata.cu_seqlens_k = self.decode_cuda_graph_metadata["cu_seqlens_k"][
|
||||
: bs + 1
|
||||
]
|
||||
metadata.page_table = self.decode_cuda_graph_metadata[
|
||||
"page_table_draft_decode"
|
||||
][:bs, :]
|
||||
@@ -351,20 +338,15 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
self.decode_cuda_graph_metadata[bs] = metadata
|
||||
else:
|
||||
# Normal Decode
|
||||
# Get sequence information
|
||||
metadata.cache_seqlens_int32 = seq_lens[:bs].to(torch.int32)
|
||||
batch_size = len(seq_lens)
|
||||
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.cache_seqlens_int32 = self.decode_cuda_graph_metadata[
|
||||
"cache_seqlens"
|
||||
][:bs]
|
||||
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"][
|
||||
:bs, :
|
||||
]
|
||||
@@ -376,29 +358,18 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
)
|
||||
self.decode_cuda_graph_metadata[bs] = metadata
|
||||
elif forward_mode.is_target_verify():
|
||||
# Target Verify
|
||||
# Here we only support topk = 1 for now.
|
||||
# Target Verify (topk = 1)
|
||||
tokens_per_req = num_tokens // bs
|
||||
metadata.cache_seqlens_int32 = self.target_verify_metadata["cache_seqlens"][
|
||||
:bs
|
||||
]
|
||||
metadata.cache_seqlens_int32.copy_(seq_lens + tokens_per_req)
|
||||
|
||||
metadata.cu_seqlens_q = torch.arange(
|
||||
0,
|
||||
bs * tokens_per_req + 1,
|
||||
tokens_per_req,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
|
||||
metadata.cu_seqlens_q = self.target_verify_metadata["cu_seqlens_q"][
|
||||
: bs + 1
|
||||
]
|
||||
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_k = seq_lens.max().item() + tokens_per_req
|
||||
|
||||
metadata.page_table = self.target_verify_metadata["page_table"][:bs, :]
|
||||
self._bind_swa_page_table(
|
||||
metadata,
|
||||
@@ -406,29 +377,15 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
"swa_page_table",
|
||||
bs,
|
||||
)
|
||||
|
||||
self.target_verify_metadata[bs] = metadata
|
||||
elif forward_mode.is_draft_extend():
|
||||
num_tokens_per_bs = num_tokens // bs
|
||||
metadata.cache_seqlens_int32 = self.draft_extend_metadata["cache_seqlens"][
|
||||
:bs
|
||||
]
|
||||
metadata.cache_seqlens_int32.copy_(seq_lens)
|
||||
num_tokens_per_bs = num_tokens // bs
|
||||
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.cu_seqlens_q = self.draft_extend_metadata["cu_seqlens_q"][: bs + 1]
|
||||
metadata.cu_seqlens_k = self.draft_extend_metadata["cu_seqlens_k"][: bs + 1]
|
||||
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, :]
|
||||
self._bind_swa_page_table(
|
||||
metadata,
|
||||
@@ -436,9 +393,41 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
||||
"swa_page_table",
|
||||
bs,
|
||||
)
|
||||
|
||||
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(
|
||||
self,
|
||||
|
||||
@@ -444,6 +444,44 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
|
||||
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(
|
||||
self,
|
||||
bs: int,
|
||||
@@ -472,60 +510,19 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
spec_info,
|
||||
)
|
||||
|
||||
metadata = TRTLLMMLADecodeMetadata()
|
||||
|
||||
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))
|
||||
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=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,
|
||||
self._init_cuda_graph_metadata(
|
||||
bs, num_tokens, forward_mode, seq_lens, 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(),
|
||||
)
|
||||
|
||||
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(
|
||||
self,
|
||||
|
||||
@@ -390,6 +390,39 @@ class WaveAttnBackend(AttentionBackend):
|
||||
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(
|
||||
self,
|
||||
bs: int,
|
||||
@@ -402,76 +435,34 @@ class WaveAttnBackend(AttentionBackend):
|
||||
):
|
||||
assert encoder_lens is None, "Not supported"
|
||||
|
||||
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),
|
||||
)
|
||||
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),
|
||||
# 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(
|
||||
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=spec_info.kv_indptr,
|
||||
kv_indices=spec_info.kv_indices,
|
||||
qo_indptr=None,
|
||||
custom_mask=None,
|
||||
mask_indptr=None,
|
||||
)
|
||||
return
|
||||
|
||||
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(
|
||||
attn_logits,
|
||||
attn_lse,
|
||||
max_extend_len,
|
||||
num_kv_splits,
|
||||
kv_indptr,
|
||||
kv_indices,
|
||||
qo_indptr,
|
||||
custom_mask,
|
||||
mask_indptr,
|
||||
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(
|
||||
@@ -485,9 +476,7 @@ class WaveAttnBackend(AttentionBackend):
|
||||
spec_info: Optional[SpecInput],
|
||||
seq_lens_cpu: Optional[torch.Tensor],
|
||||
):
|
||||
# NOTE: encoder_lens expected to be zeros or None
|
||||
if forward_mode.is_decode_or_idle():
|
||||
# Update kv_indptr, kv_indices
|
||||
kv_indptr = self.kv_indptr
|
||||
kv_indices = self.cuda_graph_kv_indices
|
||||
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
|
||||
self.get_num_kv_splits(num_kv_splits[:num_token], seq_lens[:bs])
|
||||
elif forward_mode.is_target_verify():
|
||||
# Update qo_indptr, kv_indptr, kv_indices, custom_mask, mask_indptr
|
||||
bs = len(req_pool_indices)
|
||||
qo_indptr = self.qo_indptr[: bs + 1]
|
||||
qo_indptr[: bs + 1] = torch.arange(
|
||||
|
||||
Reference in New Issue
Block a user