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