[AMD] Derive AITER verify tokens-per-req from input shape (#31221)

This commit is contained in:
JemmaFan
2026-08-01 00:29:25 -07:00
committed by GitHub
parent 33ecf4bcd8
commit 2fd78ec2d7
@@ -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,