[AMD] Fix EAGLE3 speculative decoding with aiter attention backend (#19362)

This commit is contained in:
Hubert Lu
2026-03-03 16:12:13 -08:00
committed by GitHub
parent e6411ba315
commit 441045a7bf
2 changed files with 107 additions and 107 deletions

View File

@@ -968,31 +968,34 @@ class AiterAttnBackend(AttentionBackend):
)
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:
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),
)
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = self.num_draft_tokens
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.kv_cache_dtype == fp8_dtype:
if self.use_mla:
if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
@@ -1035,37 +1038,11 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
)
else:
# Non-MLA target_verify cuda graph: use triton extend kernel metadata
draft_num = self.num_draft_tokens
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
(1 + bs) * draft_num,
step=draft_num,
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
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.custom_mask
seq_mask_len = draft_num * (seq_lens + draft_num)
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]
@@ -1074,12 +1051,12 @@ class AiterAttnBackend(AttentionBackend):
kv_indptr,
kv_indices,
qo_indptr,
None,
draft_num,
None,
kv_last_page_len,
max_q_len,
kv_indptr[-1].item(),
custom_mask=custom_mask,
mask_indptr=mask_indptr,
max_extend_len=draft_num,
max_extend_len=max_q_len,
)
elif forward_mode.is_draft_extend():
num_tokens_per_bs = self.speculative_num_steps + 1
@@ -1290,64 +1267,71 @@ class AiterAttnBackend(AttentionBackend):
kv_indices,
self.req_to_token.stride(0),
)
if not self.use_mla:
# Non-MLA: update custom_mask and mask_indptr for triton extend kernel
custom_mask = self.cuda_graph_custom_mask
custom_mask[: spec_info.custom_mask.shape[0]] = spec_info.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)
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = self.num_draft_tokens
# if self.kv_cache_dtype == fp8_dtype:
if _use_mla_ps_kernel:
if self.use_mla:
if _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
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,
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,
kv_indptr[-1].item(),
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:
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[: bs + 1]
mask_indptr[1 : bs + 1] = torch.cumsum(seq_mask_len, dim=0)
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,
kv_indptr[-1].item(),
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,
# num_kv_splits_indptr=num_kv_splits_indptr,
)
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
kv_last_page_len,
max_q_len,
kv_indptr[-1].item(),
custom_mask=custom_mask,
mask_indptr=mask_indptr,
max_extend_len=max_q_len,
)
elif forward_mode.is_draft_extend():
num_tokens_per_bs = self.speculative_num_steps + 1
@@ -1371,7 +1355,7 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = num_tokens_per_bs
if _use_mla_ps_kernel:
if self.use_mla and _use_mla_ps_kernel:
num_kv_splits = self.max_split_per_batch
@@ -1413,7 +1397,6 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
)
else: