[AMD] Support speculative decoding v2 for aiter backend on ROCm/HIP (#17450)

Co-authored-by: kkHuang-amd <wunhuang@amd.com>
Co-authored-by: HaiShaw <hixiao@gmail.com>
This commit is contained in:
Hubert Lu
2026-03-11 17:01:01 -07:00
committed by GitHub
parent acab24a76a
commit 67f02681c9
4 changed files with 341 additions and 24 deletions

View File

@@ -106,6 +106,7 @@ class AiterAttnBackend(AttentionBackend):
model_runner: ModelRunner,
skip_prefill: bool = False,
kv_indptr_buf: Optional[torch.Tensor] = None,
topk: int = 1,
):
super().__init__()
# Lazy import to avoid the initialization of cuda context
@@ -123,6 +124,7 @@ class AiterAttnBackend(AttentionBackend):
self.is_multimodal = model_runner.model_config.is_multimodal
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
self.topk = topk
self.num_head = (
model_runner.model_config.num_attention_heads // get_attention_tp_size()
)
@@ -171,6 +173,7 @@ class AiterAttnBackend(AttentionBackend):
self.mask_indptr = torch.zeros(
(max_bs + 1,), dtype=torch.int64, device=model_runner.device
)
self._kv_indices_scratch: Optional[torch.Tensor] = None
# Create prefill indices updater
if not skip_prefill:
@@ -432,6 +435,74 @@ class AiterAttnBackend(AttentionBackend):
is_causal=is_causal,
)
def _resolve_v2_num_draft_tokens(
self,
extend_seq_lens: Optional[torch.Tensor] = None,
extend_seq_lens_cpu: Optional[list[int]] = None,
) -> int:
"""Resolve fixed per-request extend length for DRAFT_EXTEND_V2."""
num_draft_tokens = self.num_draft_tokens
if num_draft_tokens is None:
if extend_seq_lens is not None and extend_seq_lens.numel() > 0:
# Avoid list scans in hot path when tensor lengths are already available.
num_draft_tokens = int(extend_seq_lens[0].item())
elif extend_seq_lens_cpu:
num_draft_tokens = max(extend_seq_lens_cpu)
else:
raise ValueError(
"DRAFT_EXTEND_V2 requires speculative_num_draft_tokens or "
"non-empty extend_seq_lens/extend_seq_lens_cpu."
)
num_draft_tokens = int(num_draft_tokens)
if extend_seq_lens is not None and extend_seq_lens.numel() > 0:
if not torch.all(extend_seq_lens == num_draft_tokens):
raise ValueError(
"DRAFT_EXTEND_V2 expects fixed extend length per request; got "
f"extend_seq_lens={extend_seq_lens}, expected all == {num_draft_tokens}."
)
if extend_seq_lens_cpu and any(
x != num_draft_tokens for x in extend_seq_lens_cpu
):
raise ValueError(
"DRAFT_EXTEND_V2 expects fixed extend length per request; got "
f"{extend_seq_lens_cpu}, expected all == {num_draft_tokens}."
)
return num_draft_tokens
def _get_kv_indices_scratch(
self, required_tokens: int, device: torch.device
) -> torch.Tensor:
if (
self._kv_indices_scratch is None
or self._kv_indices_scratch.device != device
or self._kv_indices_scratch.numel() < required_tokens
):
self._kv_indices_scratch = torch.empty(
required_tokens, dtype=torch.int32, device=device
)
return self._kv_indices_scratch[:required_tokens]
def _set_uniform_qo_indptr(
self, bs: int, tokens_per_req: int, device: torch.device
) -> torch.Tensor:
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
0,
bs * tokens_per_req + 1,
step=tokens_per_req,
dtype=torch.int32,
device=device,
)
return qo_indptr
def _ensure_spec_v2_topk_supported(self):
if self.topk > 1:
raise NotImplementedError(
"AiterAttnBackend SPEC_V2 path currently supports topk <= 1 only. "
f"Got topk={self.topk}."
)
def mla_fp8_prefill_attn(
self,
q: torch.Tensor,
@@ -508,7 +579,7 @@ class AiterAttnBackend(AttentionBackend):
return output
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init auxiliary variables for triton attention backend."""
"""Init auxiliary variables for aiter attention backend."""
bs = forward_batch.batch_size
kv_indptr = self.kv_indptr
@@ -531,8 +602,8 @@ class AiterAttnBackend(AttentionBackend):
if spec_info is None or forward_batch.forward_mode.is_idle():
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = torch.empty(
forward_batch.seq_lens_sum, dtype=torch.int32, device=self.device
kv_indices = self._get_kv_indices_scratch(
forward_batch.seq_lens_sum, forward_batch.seq_lens.device
)
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
@@ -598,7 +669,97 @@ class AiterAttnBackend(AttentionBackend):
run_graph=False,
)
elif forward_batch.forward_mode.is_draft_extend_v2():
# EAGLE V2: DRAFT_EXTEND_V2 mode - extend draft KV cache with all predicted tokens
self._ensure_spec_v2_topk_supported()
if self.use_mla:
device = forward_batch.seq_lens.device
num_draft_tokens = self._resolve_v2_num_draft_tokens(
extend_seq_lens=forward_batch.extend_seq_lens
)
qo_indptr = self._set_uniform_qo_indptr(bs, num_draft_tokens, device)
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
kv_indices = self._get_kv_indices_scratch(
forward_batch.seq_lens_sum, device
)
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
forward_batch.req_pool_indices,
forward_batch.seq_lens,
kv_indptr,
None,
kv_indices,
self.req_to_token.stride(0),
)
if _use_mla_ps_kernel:
max_seqlen_qo = num_draft_tokens
(
work_metadata,
work_indptr,
work_info_set,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
) = self.make_mla_decode_meta_data_buffer(max_seqlen_qo, bs)
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
qo_indptr,
kv_indptr,
self.kv_last_page_len[:bs],
work_metadata,
work_info_set,
work_indptr,
reduce_indptr,
reduce_final_map,
reduce_partial_map,
max_seqlen_qo,
fast_mode=fast_mode,
max_split_per_batch=num_kv_splits,
intra_batch_mode=intra_batch_mode,
)
self.forward_metadata = ForwardMetadata(
kv_indptr,
kv_indices,
qo_indptr,
self.kv_last_page_len[:bs],
num_draft_tokens,
forward_batch.seq_lens_cpu.max().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,
run_graph=False,
)
else:
self.indices_updater_prefill.update(
forward_batch.req_pool_indices,
forward_batch.seq_lens,
forward_batch.seq_lens_sum,
prefix_lens=None,
encoder_lens=forward_batch.encoder_lens,
spec_info=forward_batch.spec_info,
)
self.forward_metadata = ForwardMetadata(
self.indices_updater_prefill.kv_indptr,
self.indices_updater_prefill.kv_indices,
None,
None,
self.indices_updater_prefill.max_q_len,
self.indices_updater_prefill.max_kv_len,
)
elif forward_batch.forward_mode.is_draft_extend():
# EAGLE V1: DRAFT_EXTEND mode - uses spec_info.accept_length
if self.use_mla:
kv_indices, kv_indptr, qo_indptr, custom_mask = (
spec_info.generate_attn_arg_prefill(
@@ -686,20 +847,19 @@ class AiterAttnBackend(AttentionBackend):
kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs
device = forward_batch.seq_lens.device
qo_indptr = torch.arange(
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=device,
)
kv_indptr = self.kv_indptr
kv_indptr = self.kv_indptr[: bs + 1]
kv_indptr[1 : bs + 1] = torch.cumsum(kv_lens, dim=0)
kv_indptr = kv_indptr[: bs + 1]
kv_indices = torch.empty(
kv_indices = self._get_kv_indices_scratch(
kv_lens_sum,
dtype=torch.int32,
device=device,
device,
)
create_flashinfer_kv_indices_triton[(bs,)](
self.req_to_token,
@@ -1040,7 +1200,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,
)
elif forward_mode.is_target_verify():
@@ -1134,7 +1293,70 @@ class AiterAttnBackend(AttentionBackend):
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 _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,
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,
)
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(
@@ -1314,7 +1536,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,
)
elif forward_mode.is_target_verify():
@@ -1408,8 +1629,78 @@ class AiterAttnBackend(AttentionBackend):
mask_indptr=mask_indptr,
max_extend_len=max_q_len,
)
elif forward_mode.is_draft_extend_v2():
# EAGLE V2: Fixed num_draft_tokens per batch
self._ensure_spec_v2_topk_supported()
seq_lens = seq_lens[:bs]
num_tokens_per_bs = self._resolve_v2_num_draft_tokens()
extend_lens = torch.full(
(bs,), num_tokens_per_bs, dtype=torch.int32, device=seq_lens.device
)
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[1 : bs + 1] = torch.cumsum(extend_lens, dim=0)
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 _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,
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,
)
elif forward_mode.is_draft_extend():
# EAGLE V1: Uses spec_info.accept_length
num_tokens_per_bs = self.speculative_num_steps + 1
seq_lens = seq_lens[:bs]
accept_lens = spec_info.accept_length[:bs]
@@ -1481,6 +1772,14 @@ class AiterAttnBackend(AttentionBackend):
def get_cuda_graph_seq_len_fill_value(self):
return 1
def update_verify_buffers_to_fill_after_draft(
self, spec_info: SpecInput, cuda_graph_bs: Optional[int]
):
# AITER verify path does not require post-draft buffer patching currently.
# This override prevents overlap-plan stream mode from failing with the
# base class NotImplementedError.
pass
def forward_extend(
self,
q: torch.Tensor,
@@ -1528,6 +1827,7 @@ class AiterAttnBackend(AttentionBackend):
forward_batch.forward_mode.is_extend()
and not forward_batch.forward_mode.is_target_verify()
and not forward_batch.forward_mode.is_draft_extend()
and not forward_batch.forward_mode.is_draft_extend_v2()
):
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
if kv_indices.shape[0] == 0 or extend_no_prefix:
@@ -1680,7 +1980,10 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits=num_kv_splits,
)
return o
elif forward_batch.forward_mode.is_draft_extend():
elif (
forward_batch.forward_mode.is_draft_extend()
or forward_batch.forward_mode.is_draft_extend_v2()
):
work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr
@@ -2156,6 +2459,7 @@ class AiterMultiStepDraftBackend:
model_runner,
skip_prefill=True,
kv_indptr_buf=self.kv_indptr[i],
topk=topk,
)
)
self.max_context_len = self.attn_backends[0].max_context_len

View File

@@ -310,7 +310,7 @@ class EagleVerifyInputV2Mixin:
accept_length = torch.empty((bs,), dtype=torch.int32, device=device)
# Sample tokens
if sampling_info.is_all_greedy or _is_npu:
if sampling_info.is_all_greedy or _is_npu or _is_hip:
target_predict = torch.argmax(next_token_logits, dim=-1)
target_predict = target_predict.reshape(bs, self.draft_token_num)
predict, accept_index, accept_length = verify_tree_greedy_func(

View File

@@ -59,6 +59,7 @@ from sglang.srt.utils.common import (
fast_topk,
get_available_gpu_memory,
is_cuda,
is_hip,
is_npu,
next_power_of_2,
)
@@ -66,6 +67,7 @@ from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
_is_npu = is_npu()
_is_cuda = is_cuda()
_is_hip = is_hip()
logger = logging.getLogger(__name__)
@@ -280,18 +282,27 @@ class EagleDraftWorker(BaseDraftWorker):
"npu": EAGLEDraftExtendNpuGraphRunner,
"cuda": EAGLEDraftExtendCudaGraphRunner,
}
supports_hip_aiter_draft_extend_graph = False
if _is_hip:
# Keep import local so non-HIP environments do not require aiter.
from sglang.srt.layers.attention.aiter_backend import (
AiterMultiStepDraftBackend,
)
supports_hip_aiter_draft_extend_graph = isinstance(
self.draft_attn_backend, AiterMultiStepDraftBackend
)
supports_cuda_draft_extend_graph = _is_cuda and (
isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend)
or isinstance(self.draft_attn_backend, TRTLLMMLAMultiStepDraftBackend)
)
# Capture extend
# TODO: support draft extend cuda graph for more attention backends
if self.draft_extend_attn_backend and (
_is_npu
or (
_is_cuda
and isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend)
)
or (
_is_cuda
and isinstance(self.draft_attn_backend, TRTLLMMLAMultiStepDraftBackend)
)
or supports_cuda_draft_extend_graph
or supports_hip_aiter_draft_extend_graph
):
tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id)