[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 = (
|
seq_lens_cpu = (
|
||||||
forward_batch.seq_lens.cpu() if in_capture else forward_batch.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(
|
self._apply_cuda_graph_metadata(
|
||||||
bs=forward_batch.batch_size,
|
bs=forward_batch.batch_size,
|
||||||
req_pool_indices=forward_batch.req_pool_indices,
|
req_pool_indices=forward_batch.req_pool_indices,
|
||||||
@@ -866,6 +871,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
forward_mode=forward_batch.forward_mode,
|
forward_mode=forward_batch.forward_mode,
|
||||||
spec_info=forward_batch.spec_info,
|
spec_info=forward_batch.spec_info,
|
||||||
seq_lens_cpu=seq_lens_cpu,
|
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
|
# Refill the SWA write-target buffer from the live out_cache_loc and
|
||||||
@@ -1196,8 +1202,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
run_graph=False,
|
run_graph=False,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
draft_num = forward_batch.input_ids.shape[0] // bs
|
||||||
bs = len(forward_batch.req_pool_indices)
|
bs = len(forward_batch.req_pool_indices)
|
||||||
draft_num = spec_info.draft_token_num
|
|
||||||
|
|
||||||
if self._use_unified_verify:
|
if self._use_unified_verify:
|
||||||
page_table, qo_indptr, max_q_len, swa_page_table = (
|
page_table, qo_indptr, max_q_len, swa_page_table = (
|
||||||
@@ -1500,6 +1506,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
forward_mode: ForwardMode,
|
forward_mode: ForwardMode,
|
||||||
spec_info: Optional[SpecInput],
|
spec_info: Optional[SpecInput],
|
||||||
seq_lens_cpu: Optional[torch.Tensor],
|
seq_lens_cpu: Optional[torch.Tensor],
|
||||||
|
verify_tokens_per_req: Optional[int],
|
||||||
):
|
):
|
||||||
|
|
||||||
num_kv_splits = None
|
num_kv_splits = None
|
||||||
@@ -1652,11 +1659,17 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
bs = len(req_pool_indices)
|
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 = self.qo_indptr[: bs + 1]
|
||||||
qo_indptr[: bs + 1] = torch.arange(
|
qo_indptr[: bs + 1] = torch.arange(
|
||||||
0,
|
0,
|
||||||
(1 + bs) * self.num_draft_tokens,
|
(1 + bs) * tokens_per_req,
|
||||||
step=self.num_draft_tokens,
|
step=tokens_per_req,
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
)
|
)
|
||||||
@@ -1689,9 +1702,9 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
self.req_to_token.stride(0),
|
self.req_to_token.stride(0),
|
||||||
)
|
)
|
||||||
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
|
||||||
max_q_len = self.num_draft_tokens
|
|
||||||
|
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
|
max_q_len = self.num_draft_tokens
|
||||||
if _use_mla_ps_kernel:
|
if _use_mla_ps_kernel:
|
||||||
num_kv_splits = self.max_split_per_batch
|
num_kv_splits = self.max_split_per_batch
|
||||||
|
|
||||||
@@ -1735,6 +1748,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
max_q_len = verify_tokens_per_req
|
||||||
if self._use_unified_verify:
|
if self._use_unified_verify:
|
||||||
max_num_blocks_per_seq = (
|
max_num_blocks_per_seq = (
|
||||||
self.max_context_len + self.page_size - 1
|
self.max_context_len + self.page_size - 1
|
||||||
@@ -1753,7 +1767,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
bs,
|
bs,
|
||||||
seq_lens,
|
seq_lens,
|
||||||
req_pool_indices,
|
req_pool_indices,
|
||||||
self.num_draft_tokens,
|
verify_tokens_per_req,
|
||||||
page_table_dest=page_table,
|
page_table_dest=page_table,
|
||||||
swa_page_table_dest=swa_page_table,
|
swa_page_table_dest=swa_page_table,
|
||||||
)
|
)
|
||||||
@@ -2278,7 +2292,9 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
v=v_unified,
|
v=v_unified,
|
||||||
out=o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
|
out=o.view(-1, layer.tp_q_head_num, layer.v_head_dim),
|
||||||
cu_seqlens_q=self.forward_metadata.qo_indptr,
|
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_q=self.forward_metadata.max_q_len,
|
||||||
max_seqlen_k=max_kv_len,
|
max_seqlen_k=max_kv_len,
|
||||||
softmax_scale=layer.scaling,
|
softmax_scale=layer.scaling,
|
||||||
|
|||||||
Reference in New Issue
Block a user