[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:
co-authored by
kkHuang-amd
HaiShaw
parent
acab24a76a
commit
67f02681c9
@@ -106,6 +106,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
model_runner: ModelRunner,
|
model_runner: ModelRunner,
|
||||||
skip_prefill: bool = False,
|
skip_prefill: bool = False,
|
||||||
kv_indptr_buf: Optional[torch.Tensor] = None,
|
kv_indptr_buf: Optional[torch.Tensor] = None,
|
||||||
|
topk: int = 1,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
# Lazy import to avoid the initialization of cuda context
|
# 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.is_multimodal = model_runner.model_config.is_multimodal
|
||||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
||||||
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
|
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
|
||||||
|
self.topk = topk
|
||||||
self.num_head = (
|
self.num_head = (
|
||||||
model_runner.model_config.num_attention_heads // get_attention_tp_size()
|
model_runner.model_config.num_attention_heads // get_attention_tp_size()
|
||||||
)
|
)
|
||||||
@@ -171,6 +173,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
self.mask_indptr = torch.zeros(
|
self.mask_indptr = torch.zeros(
|
||||||
(max_bs + 1,), dtype=torch.int64, device=model_runner.device
|
(max_bs + 1,), dtype=torch.int64, device=model_runner.device
|
||||||
)
|
)
|
||||||
|
self._kv_indices_scratch: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# Create prefill indices updater
|
# Create prefill indices updater
|
||||||
if not skip_prefill:
|
if not skip_prefill:
|
||||||
@@ -432,6 +435,74 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
is_causal=is_causal,
|
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(
|
def mla_fp8_prefill_attn(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -508,7 +579,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
return output
|
return output
|
||||||
|
|
||||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
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
|
bs = forward_batch.batch_size
|
||||||
kv_indptr = self.kv_indptr
|
kv_indptr = self.kv_indptr
|
||||||
@@ -531,8 +602,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
if spec_info is None or forward_batch.forward_mode.is_idle():
|
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[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
|
||||||
kv_indptr = kv_indptr[: bs + 1]
|
kv_indptr = kv_indptr[: bs + 1]
|
||||||
kv_indices = torch.empty(
|
kv_indices = self._get_kv_indices_scratch(
|
||||||
forward_batch.seq_lens_sum, dtype=torch.int32, device=self.device
|
forward_batch.seq_lens_sum, forward_batch.seq_lens.device
|
||||||
)
|
)
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
create_flashinfer_kv_indices_triton[(bs,)](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
@@ -598,7 +669,97 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
run_graph=False,
|
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():
|
elif forward_batch.forward_mode.is_draft_extend():
|
||||||
|
# EAGLE V1: DRAFT_EXTEND mode - uses spec_info.accept_length
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
kv_indices, kv_indptr, qo_indptr, custom_mask = (
|
kv_indices, kv_indptr, qo_indptr, custom_mask = (
|
||||||
spec_info.generate_attn_arg_prefill(
|
spec_info.generate_attn_arg_prefill(
|
||||||
@@ -686,20 +847,19 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs
|
kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs
|
||||||
device = forward_batch.seq_lens.device
|
device = forward_batch.seq_lens.device
|
||||||
|
|
||||||
qo_indptr = torch.arange(
|
qo_indptr = self.qo_indptr[: bs + 1]
|
||||||
|
qo_indptr[: bs + 1] = torch.arange(
|
||||||
0,
|
0,
|
||||||
(1 + bs) * draft_num,
|
(1 + bs) * draft_num,
|
||||||
step=draft_num,
|
step=draft_num,
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=device,
|
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[1 : bs + 1] = torch.cumsum(kv_lens, dim=0)
|
||||||
kv_indptr = kv_indptr[: bs + 1]
|
kv_indices = self._get_kv_indices_scratch(
|
||||||
kv_indices = torch.empty(
|
|
||||||
kv_lens_sum,
|
kv_lens_sum,
|
||||||
dtype=torch.int32,
|
device,
|
||||||
device=device,
|
|
||||||
)
|
)
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
create_flashinfer_kv_indices_triton[(bs,)](
|
||||||
self.req_to_token,
|
self.req_to_token,
|
||||||
@@ -1040,7 +1200,6 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map=reduce_final_map,
|
reduce_final_map=reduce_final_map,
|
||||||
reduce_partial_map=reduce_partial_map,
|
reduce_partial_map=reduce_partial_map,
|
||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
# num_kv_splits_indptr=num_kv_splits_indptr,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
@@ -1134,7 +1293,70 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
mask_indptr=mask_indptr,
|
mask_indptr=mask_indptr,
|
||||||
max_extend_len=max_q_len,
|
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():
|
elif forward_mode.is_draft_extend():
|
||||||
|
# EAGLE V1: Uses speculative_num_steps + 1
|
||||||
num_tokens_per_bs = self.speculative_num_steps + 1
|
num_tokens_per_bs = self.speculative_num_steps + 1
|
||||||
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(
|
||||||
@@ -1314,7 +1536,6 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
reduce_final_map=reduce_final_map,
|
reduce_final_map=reduce_final_map,
|
||||||
reduce_partial_map=reduce_partial_map,
|
reduce_partial_map=reduce_partial_map,
|
||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
# num_kv_splits_indptr=num_kv_splits_indptr,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
elif forward_mode.is_target_verify():
|
elif forward_mode.is_target_verify():
|
||||||
@@ -1408,8 +1629,78 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
mask_indptr=mask_indptr,
|
mask_indptr=mask_indptr,
|
||||||
max_extend_len=max_q_len,
|
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():
|
elif forward_mode.is_draft_extend():
|
||||||
|
# EAGLE V1: Uses spec_info.accept_length
|
||||||
num_tokens_per_bs = self.speculative_num_steps + 1
|
num_tokens_per_bs = self.speculative_num_steps + 1
|
||||||
seq_lens = seq_lens[:bs]
|
seq_lens = seq_lens[:bs]
|
||||||
accept_lens = spec_info.accept_length[:bs]
|
accept_lens = spec_info.accept_length[:bs]
|
||||||
@@ -1481,6 +1772,14 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
def get_cuda_graph_seq_len_fill_value(self):
|
def get_cuda_graph_seq_len_fill_value(self):
|
||||||
return 1
|
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(
|
def forward_extend(
|
||||||
self,
|
self,
|
||||||
q: torch.Tensor,
|
q: torch.Tensor,
|
||||||
@@ -1528,6 +1827,7 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
forward_batch.forward_mode.is_extend()
|
forward_batch.forward_mode.is_extend()
|
||||||
and not forward_batch.forward_mode.is_target_verify()
|
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()
|
||||||
|
and not forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
):
|
):
|
||||||
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
|
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
|
||||||
if kv_indices.shape[0] == 0 or extend_no_prefix:
|
if kv_indices.shape[0] == 0 or extend_no_prefix:
|
||||||
@@ -1680,7 +1980,10 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
num_kv_splits=num_kv_splits,
|
num_kv_splits=num_kv_splits,
|
||||||
)
|
)
|
||||||
return o
|
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_metadata = self.forward_metadata.work_metadata
|
||||||
work_indptr = self.forward_metadata.work_indptr
|
work_indptr = self.forward_metadata.work_indptr
|
||||||
@@ -2156,6 +2459,7 @@ class AiterMultiStepDraftBackend:
|
|||||||
model_runner,
|
model_runner,
|
||||||
skip_prefill=True,
|
skip_prefill=True,
|
||||||
kv_indptr_buf=self.kv_indptr[i],
|
kv_indptr_buf=self.kv_indptr[i],
|
||||||
|
topk=topk,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
self.max_context_len = self.attn_backends[0].max_context_len
|
self.max_context_len = self.attn_backends[0].max_context_len
|
||||||
|
|||||||
@@ -310,7 +310,7 @@ class EagleVerifyInputV2Mixin:
|
|||||||
accept_length = torch.empty((bs,), dtype=torch.int32, device=device)
|
accept_length = torch.empty((bs,), dtype=torch.int32, device=device)
|
||||||
|
|
||||||
# Sample tokens
|
# 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 = torch.argmax(next_token_logits, dim=-1)
|
||||||
target_predict = target_predict.reshape(bs, self.draft_token_num)
|
target_predict = target_predict.reshape(bs, self.draft_token_num)
|
||||||
predict, accept_index, accept_length = verify_tree_greedy_func(
|
predict, accept_index, accept_length = verify_tree_greedy_func(
|
||||||
|
|||||||
@@ -59,6 +59,7 @@ from sglang.srt.utils.common import (
|
|||||||
fast_topk,
|
fast_topk,
|
||||||
get_available_gpu_memory,
|
get_available_gpu_memory,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
|
is_hip,
|
||||||
is_npu,
|
is_npu,
|
||||||
next_power_of_2,
|
next_power_of_2,
|
||||||
)
|
)
|
||||||
@@ -66,6 +67,7 @@ from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
|||||||
|
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
|
_is_hip = is_hip()
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -280,18 +282,27 @@ class EagleDraftWorker(BaseDraftWorker):
|
|||||||
"npu": EAGLEDraftExtendNpuGraphRunner,
|
"npu": EAGLEDraftExtendNpuGraphRunner,
|
||||||
"cuda": EAGLEDraftExtendCudaGraphRunner,
|
"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
|
# Capture extend
|
||||||
# TODO: support draft extend cuda graph for more attention backends
|
# TODO: support draft extend cuda graph for more attention backends
|
||||||
if self.draft_extend_attn_backend and (
|
if self.draft_extend_attn_backend and (
|
||||||
_is_npu
|
_is_npu
|
||||||
or (
|
or supports_cuda_draft_extend_graph
|
||||||
_is_cuda
|
or supports_hip_aiter_draft_extend_graph
|
||||||
and isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend)
|
|
||||||
)
|
|
||||||
or (
|
|
||||||
_is_cuda
|
|
||||||
and isinstance(self.draft_attn_backend, TRTLLMMLAMultiStepDraftBackend)
|
|
||||||
)
|
|
||||||
):
|
):
|
||||||
tic = time.perf_counter()
|
tic = time.perf_counter()
|
||||||
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
|
||||||
|
|||||||
@@ -1,9 +1,9 @@
|
|||||||
import os
|
|
||||||
import unittest
|
import unittest
|
||||||
from types import SimpleNamespace
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import requests
|
import requests
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.test.ci.ci_register import register_amd_ci
|
from sglang.test.ci.ci_register import register_amd_ci
|
||||||
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
|
from sglang.test.few_shot_gsm8k import run_eval as run_eval_few_shot_gsm8k
|
||||||
@@ -87,6 +87,10 @@ class TestDeepseekR1MXFP4MTP(CustomTestCase):
|
|||||||
def setUpClass(cls):
|
def setUpClass(cls):
|
||||||
cls.model = DEEPSEEK_R1_MODEL_PATH
|
cls.model = DEEPSEEK_R1_MODEL_PATH
|
||||||
cls.base_url = DEFAULT_URL_FOR_TEST
|
cls.base_url = DEFAULT_URL_FOR_TEST
|
||||||
|
|
||||||
|
envs.SGLANG_ENABLE_SPEC_V2.set(True)
|
||||||
|
envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.set(True)
|
||||||
|
|
||||||
other_args = [
|
other_args = [
|
||||||
"--tp",
|
"--tp",
|
||||||
"8",
|
"8",
|
||||||
@@ -113,8 +117,6 @@ class TestDeepseekR1MXFP4MTP(CustomTestCase):
|
|||||||
@classmethod
|
@classmethod
|
||||||
def tearDownClass(cls):
|
def tearDownClass(cls):
|
||||||
kill_process_tree(cls.process.pid)
|
kill_process_tree(cls.process.pid)
|
||||||
if "SGLANG_ENABLE_SPEC_V2" in os.environ:
|
|
||||||
del os.environ["SGLANG_ENABLE_SPEC_V2"]
|
|
||||||
|
|
||||||
def test_a_gsm8k(
|
def test_a_gsm8k(
|
||||||
self,
|
self,
|
||||||
|
|||||||
Reference in New Issue
Block a user