[AMD] Derive AITER verify tokens-per-req from input shape (#31221)
This commit is contained in:
@@ -857,6 +857,11 @@ class AiterAttnBackend(AttentionBackend):
|
||||
seq_lens_cpu = (
|
||||
forward_batch.seq_lens.cpu() if in_capture else forward_batch.seq_lens_cpu
|
||||
)
|
||||
verify_tokens_per_req = (
|
||||
forward_batch.input_ids.shape[0] // forward_batch.batch_size
|
||||
if forward_batch.forward_mode.is_target_verify()
|
||||
else None
|
||||
)
|
||||
self._apply_cuda_graph_metadata(
|
||||
bs=forward_batch.batch_size,
|
||||
req_pool_indices=forward_batch.req_pool_indices,
|
||||
@@ -866,6 +871,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
forward_mode=forward_batch.forward_mode,
|
||||
spec_info=forward_batch.spec_info,
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
verify_tokens_per_req=verify_tokens_per_req,
|
||||
)
|
||||
|
||||
# Refill the SWA write-target buffer from the live out_cache_loc and
|
||||
@@ -1196,8 +1202,8 @@ class AiterAttnBackend(AttentionBackend):
|
||||
run_graph=False,
|
||||
)
|
||||
else:
|
||||
draft_num = forward_batch.input_ids.shape[0] // bs
|
||||
bs = len(forward_batch.req_pool_indices)
|
||||
draft_num = spec_info.draft_token_num
|
||||
|
||||
if self._use_unified_verify:
|
||||
page_table, qo_indptr, max_q_len, swa_page_table = (
|
||||
@@ -1500,6 +1506,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
forward_mode: ForwardMode,
|
||||
spec_info: Optional[SpecInput],
|
||||
seq_lens_cpu: Optional[torch.Tensor],
|
||||
verify_tokens_per_req: Optional[int],
|
||||
):
|
||||
|
||||
num_kv_splits = None
|
||||
@@ -1652,11 +1659,17 @@ class AiterAttnBackend(AttentionBackend):
|
||||
|
||||
elif forward_mode.is_target_verify():
|
||||
bs = len(req_pool_indices)
|
||||
assert verify_tokens_per_req is not None
|
||||
# MLA uses a fixed draft length (num_draft_tokens); the non-MLA
|
||||
# unified path derives it per batch from input_ids.
|
||||
tokens_per_req = (
|
||||
self.num_draft_tokens if self.use_mla else verify_tokens_per_req
|
||||
)
|
||||
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,
|
||||
(1 + bs) * tokens_per_req,
|
||||
step=tokens_per_req,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
@@ -1689,9 +1702,9 @@ class AiterAttnBackend(AttentionBackend):
|
||||
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:
|
||||
max_q_len = self.num_draft_tokens
|
||||
if _use_mla_ps_kernel:
|
||||
num_kv_splits = self.max_split_per_batch
|
||||
|
||||
@@ -1735,6 +1748,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
num_kv_splits=num_kv_splits,
|
||||
)
|
||||
else:
|
||||
max_q_len = verify_tokens_per_req
|
||||
if self._use_unified_verify:
|
||||
max_num_blocks_per_seq = (
|
||||
self.max_context_len + self.page_size - 1
|
||||
@@ -1753,7 +1767,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
bs,
|
||||
seq_lens,
|
||||
req_pool_indices,
|
||||
self.num_draft_tokens,
|
||||
verify_tokens_per_req,
|
||||
page_table_dest=page_table,
|
||||
swa_page_table_dest=swa_page_table,
|
||||
)
|
||||
@@ -2278,7 +2292,9 @@ class AiterAttnBackend(AttentionBackend):
|
||||
v=v_unified,
|
||||
out=o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
|
||||
cu_seqlens_q=self.forward_metadata.qo_indptr,
|
||||
seqused_k=forward_batch.seq_lens + self.num_draft_tokens,
|
||||
seqused_k=(
|
||||
forward_batch.seq_lens + self.forward_metadata.max_q_len
|
||||
),
|
||||
max_seqlen_q=self.forward_metadata.max_q_len,
|
||||
max_seqlen_k=max_kv_len,
|
||||
softmax_scale=layer.scaling,
|
||||
|
||||
Reference in New Issue
Block a user