[Spec V1] Split draft-extend phase from EagleDraftInput into new EagleDraftExtendInput (#24859)
This commit is contained in:
@@ -989,17 +989,18 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
self.extend_seq_lens = self._pad_tensor_to_size(self.extend_seq_lens, bs)
|
self.extend_seq_lens = self._pad_tensor_to_size(self.extend_seq_lens, bs)
|
||||||
|
|
||||||
if self.spec_info is not None and self.spec_info.is_draft_input():
|
if self.spec_info is not None and self.spec_info.is_draft_input():
|
||||||
# FIXME(lsyin): remove this isinstance logic
|
|
||||||
spec_info = self.spec_info
|
spec_info = self.spec_info
|
||||||
self.output_cache_loc_backup = self.out_cache_loc
|
self.output_cache_loc_backup = self.out_cache_loc
|
||||||
self.hidden_states_backup = spec_info.hidden_states
|
self.hidden_states_backup = spec_info.hidden_states
|
||||||
if spec_info.topk_p is not None:
|
# spec_info is EagleDraftInput | EagleDraftExtendInput; each carries
|
||||||
|
# a disjoint subset of the fields below, so getattr-guard each one.
|
||||||
|
if getattr(spec_info, "topk_p", None) is not None:
|
||||||
spec_info.topk_p = self._pad_tensor_to_size(spec_info.topk_p, bs)
|
spec_info.topk_p = self._pad_tensor_to_size(spec_info.topk_p, bs)
|
||||||
if spec_info.topk_index is not None:
|
if getattr(spec_info, "topk_index", None) is not None:
|
||||||
spec_info.topk_index = self._pad_tensor_to_size(
|
spec_info.topk_index = self._pad_tensor_to_size(
|
||||||
spec_info.topk_index, bs
|
spec_info.topk_index, bs
|
||||||
)
|
)
|
||||||
if spec_info.num_accepted_drafts is not None:
|
if getattr(spec_info, "num_accepted_drafts", None) is not None:
|
||||||
spec_info.num_accepted_drafts = self._pad_tensor_to_size(
|
spec_info.num_accepted_drafts = self._pad_tensor_to_size(
|
||||||
spec_info.num_accepted_drafts, bs
|
spec_info.num_accepted_drafts, bs
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardMode,
|
ForwardMode,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||||
from sglang.srt.speculative.spec_utils import fast_topk
|
from sglang.srt.speculative.spec_utils import fast_topk
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
require_attn_tp_gather,
|
require_attn_tp_gather,
|
||||||
@@ -360,7 +360,7 @@ class EAGLEDraftExtendCudaGraphRunner:
|
|||||||
else:
|
else:
|
||||||
global_dp_buffer_len = None
|
global_dp_buffer_len = None
|
||||||
|
|
||||||
spec_info = EagleDraftInput(
|
spec_info = EagleDraftExtendInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
num_accepted_drafts=num_accepted_drafts,
|
num_accepted_drafts=num_accepted_drafts,
|
||||||
num_accepted_tokens=num_accepted_tokens,
|
num_accepted_tokens=num_accepted_tokens,
|
||||||
|
|||||||
@@ -240,15 +240,14 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
accepted token logits.
|
accepted token logits.
|
||||||
"""
|
"""
|
||||||
if batch.forward_mode.is_idle():
|
if batch.forward_mode.is_idle():
|
||||||
next_draft_input = EagleDraftInput.create_idle_input(
|
draft_extend_input = EagleDraftExtendInput.create_idle_input(
|
||||||
device=batch.device,
|
device=batch.device,
|
||||||
hidden_size=batch.model_config.spec_hidden_size,
|
hidden_size=batch.model_config.spec_hidden_size,
|
||||||
dtype=batch.model_config.dtype,
|
dtype=batch.model_config.dtype,
|
||||||
topk=self.topk,
|
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
return EagleVerifyOutput.create_idle(
|
return EagleVerifyOutput.create_idle(
|
||||||
next_draft_input=next_draft_input,
|
draft_extend_input=draft_extend_input,
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
device=batch.device,
|
device=batch.device,
|
||||||
spec_steps=self.spec_steps,
|
spec_steps=self.spec_steps,
|
||||||
@@ -545,21 +544,21 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
batch.seq_lens.add_(num_accepted_drafts + 1)
|
batch.seq_lens.add_(num_accepted_drafts + 1)
|
||||||
batch.seq_lens_cpu.add_(num_accepted_tokens_cpu)
|
batch.seq_lens_cpu.add_(num_accepted_tokens_cpu)
|
||||||
|
|
||||||
next_draft_input = EagleDraftInput(
|
draft_extend_input = EagleDraftExtendInput(
|
||||||
hidden_states=batch.spec_info.hidden_states[accept_index],
|
hidden_states=batch.spec_info.hidden_states[accept_index],
|
||||||
num_accepted_drafts=num_accepted_drafts,
|
num_accepted_drafts=num_accepted_drafts,
|
||||||
num_accepted_tokens=num_accepted_drafts + 1,
|
num_accepted_tokens=num_accepted_drafts + 1,
|
||||||
num_accepted_tokens_cpu=num_accepted_tokens_list,
|
num_accepted_tokens_cpu=num_accepted_tokens_list,
|
||||||
|
input_ids=accept_tokens,
|
||||||
|
seq_lens=batch.seq_lens,
|
||||||
|
seq_lens_cpu=batch.seq_lens_cpu,
|
||||||
|
req_pool_indices=batch.req_pool_indices,
|
||||||
)
|
)
|
||||||
|
|
||||||
return EagleVerifyOutput(
|
return EagleVerifyOutput(
|
||||||
next_draft_input=next_draft_input,
|
draft_extend_input=draft_extend_input,
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
accept_tokens=accept_tokens,
|
accept_tokens=accept_tokens,
|
||||||
unfinished_accept_tokens=accept_tokens,
|
|
||||||
seq_lens_for_draft_extend=batch.seq_lens,
|
|
||||||
seq_lens_for_draft_extend_cpu=batch.seq_lens_cpu,
|
|
||||||
req_pool_indices_for_draft_extend=batch.req_pool_indices,
|
|
||||||
num_accepted_drafts_per_req_cpu=num_accepted_drafts_list,
|
num_accepted_drafts_per_req_cpu=num_accepted_drafts_list,
|
||||||
accepted_indices=accept_index,
|
accepted_indices=accept_index,
|
||||||
)
|
)
|
||||||
@@ -614,51 +613,30 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
unfinished_num_accepted_drafts = num_accepted_drafts[
|
unfinished_num_accepted_drafts = num_accepted_drafts[
|
||||||
unfinished_index_device
|
unfinished_index_device
|
||||||
]
|
]
|
||||||
unfinished_accept_tokens = predict[unfinished_accept_index]
|
draft_extend_input = EagleDraftExtendInput(
|
||||||
seq_lens_for_draft_extend = batch.seq_lens[unfinished_index_device]
|
|
||||||
seq_lens_for_draft_extend_cpu = batch.seq_lens_cpu[unfinished_index]
|
|
||||||
req_pool_indices_for_draft_extend = batch.req_pool_indices[
|
|
||||||
unfinished_index_device
|
|
||||||
]
|
|
||||||
next_draft_input = EagleDraftInput(
|
|
||||||
hidden_states=batch.spec_info.hidden_states[
|
hidden_states=batch.spec_info.hidden_states[
|
||||||
unfinished_accept_index
|
unfinished_accept_index
|
||||||
],
|
],
|
||||||
num_accepted_tokens_cpu=draft_input_num_accepted_tokens_cpu,
|
num_accepted_tokens_cpu=draft_input_num_accepted_tokens_cpu,
|
||||||
num_accepted_drafts=unfinished_num_accepted_drafts,
|
num_accepted_drafts=unfinished_num_accepted_drafts,
|
||||||
num_accepted_tokens=unfinished_num_accepted_drafts + 1,
|
num_accepted_tokens=unfinished_num_accepted_drafts + 1,
|
||||||
|
input_ids=predict[unfinished_accept_index],
|
||||||
|
seq_lens=batch.seq_lens[unfinished_index_device],
|
||||||
|
seq_lens_cpu=batch.seq_lens_cpu[unfinished_index],
|
||||||
|
req_pool_indices=batch.req_pool_indices[unfinished_index_device],
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
unfinished_accept_tokens = torch.empty(
|
draft_extend_input = EagleDraftExtendInput.create_idle_input(
|
||||||
(0,), dtype=accept_tokens.dtype, device=accept_tokens.device
|
|
||||||
)
|
|
||||||
seq_lens_for_draft_extend = torch.empty(
|
|
||||||
(0,), dtype=batch.seq_lens.dtype, device=batch.seq_lens.device
|
|
||||||
)
|
|
||||||
seq_lens_for_draft_extend_cpu = torch.empty(
|
|
||||||
(0,), dtype=batch.seq_lens_cpu.dtype
|
|
||||||
)
|
|
||||||
req_pool_indices_for_draft_extend = torch.empty(
|
|
||||||
(0,),
|
|
||||||
dtype=batch.req_pool_indices.dtype,
|
|
||||||
device=batch.req_pool_indices.device,
|
|
||||||
)
|
|
||||||
next_draft_input = EagleDraftInput.create_idle_input(
|
|
||||||
device=batch.device,
|
device=batch.device,
|
||||||
hidden_size=batch.model_config.spec_hidden_size,
|
hidden_size=batch.model_config.spec_hidden_size,
|
||||||
dtype=batch.model_config.dtype,
|
dtype=batch.model_config.dtype,
|
||||||
topk=self.topk,
|
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
|
|
||||||
return EagleVerifyOutput(
|
return EagleVerifyOutput(
|
||||||
next_draft_input=next_draft_input,
|
draft_extend_input=draft_extend_input,
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
accept_tokens=accept_tokens,
|
accept_tokens=accept_tokens,
|
||||||
unfinished_accept_tokens=unfinished_accept_tokens,
|
|
||||||
seq_lens_for_draft_extend=seq_lens_for_draft_extend,
|
|
||||||
seq_lens_for_draft_extend_cpu=seq_lens_for_draft_extend_cpu,
|
|
||||||
req_pool_indices_for_draft_extend=req_pool_indices_for_draft_extend,
|
|
||||||
num_accepted_drafts_per_req_cpu=num_accepted_drafts_list,
|
num_accepted_drafts_per_req_cpu=num_accepted_drafts_list,
|
||||||
accepted_indices=accept_index,
|
accepted_indices=accept_index,
|
||||||
)
|
)
|
||||||
@@ -666,42 +644,33 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
|
|||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
||||||
# The inputs for decode
|
|
||||||
# shape: (b, topk)
|
# shape: (b, topk)
|
||||||
topk_p: torch.Tensor = None
|
topk_p: torch.Tensor = None
|
||||||
topk_index: torch.Tensor = None
|
topk_index: torch.Tensor = None
|
||||||
# shape: (b, hidden_size) when consumed by `draft` forward (one hidden per req);
|
# shape: (b, hidden_size) - one hidden per req, consumed by `draft` forward.
|
||||||
# shape: (total_accepted, hidden_size) when consumed by `draft_extend` forward
|
|
||||||
# (one hidden per accepted token). Workers maintain this invariant locally;
|
|
||||||
# there is no type-level guard. Don't add new readers without checking phase.
|
|
||||||
hidden_states: torch.Tensor = None
|
hidden_states: torch.Tensor = None
|
||||||
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL
|
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL
|
||||||
|
|
||||||
# Inputs for extend
|
# Per-req bonus token (the "+1" target prediction at end of each accept
|
||||||
# shape: (b,)
|
# chain). Written by `EagleDraftExtendInput.prepare_extend_after_decode`;
|
||||||
# `num_accepted_drafts` and `num_accepted_tokens` are kept in sync:
|
# the worker copies it here for next iter's draft.
|
||||||
# `num_accepted_tokens = num_accepted_drafts + 1` (per-req, one bonus per req).
|
|
||||||
# Storing both avoids repeated `+ 1` at every consumer (attn backends, kernels).
|
|
||||||
bonus_tokens: torch.Tensor = None
|
bonus_tokens: torch.Tensor = None
|
||||||
num_accepted_drafts: torch.Tensor = None
|
|
||||||
num_accepted_tokens: torch.Tensor = None
|
|
||||||
# Read by attention backends during draft-extend forward; kept on the
|
|
||||||
# dataclass because the backends access it via `forward_batch.spec_info`.
|
|
||||||
num_accepted_tokens_cpu: List[int] = None
|
|
||||||
|
|
||||||
# Inputs for the attention backends
|
|
||||||
# shape: (b + 1,)
|
# shape: (b + 1,)
|
||||||
kv_indptr: torch.Tensor = None
|
kv_indptr: torch.Tensor = None
|
||||||
kv_indices: torch.Tensor = None
|
kv_indices: torch.Tensor = None
|
||||||
|
|
||||||
# Shape info for padding
|
|
||||||
num_tokens_per_req: int = -1
|
num_tokens_per_req: int = -1
|
||||||
num_tokens_for_logprob_per_req: int = -1
|
num_tokens_for_logprob_per_req: int = -1
|
||||||
|
|
||||||
# Inputs for V2 overlap worker
|
# V2 overlap worker only
|
||||||
future_indices: Optional[FutureIndices] = None
|
future_indices: Optional[FutureIndices] = None
|
||||||
new_seq_lens: Optional[torch.Tensor] = None
|
new_seq_lens: Optional[torch.Tensor] = None
|
||||||
verify_done: Optional[torch.cuda.Event] = None
|
verify_done: Optional[torch.cuda.Event] = None
|
||||||
|
# V2 reuses `EagleDraftInput` across phases (V1 has a separate
|
||||||
|
# `EagleDraftExtendInput` for these). Set during V2's draft-extend.
|
||||||
|
num_accepted_drafts: Optional[torch.Tensor] = None
|
||||||
|
num_accepted_tokens: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
def __post_init__(self):
|
def __post_init__(self):
|
||||||
super().__init__(SpecInputType.EAGLE_DRAFT)
|
super().__init__(SpecInputType.EAGLE_DRAFT)
|
||||||
@@ -741,81 +710,8 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
topk_index=torch.empty((0, topk), device=device, dtype=torch.int64),
|
topk_index=torch.empty((0, topk), device=device, dtype=torch.int64),
|
||||||
capture_hidden_mode=capture_hidden_mode,
|
capture_hidden_mode=capture_hidden_mode,
|
||||||
new_seq_lens=torch.empty((0,), device=device, dtype=torch.int32),
|
new_seq_lens=torch.empty((0,), device=device, dtype=torch.int32),
|
||||||
num_accepted_drafts=torch.empty((0,), device=device, dtype=torch.int32),
|
|
||||||
num_accepted_tokens=torch.empty((0,), device=device, dtype=torch.int32),
|
|
||||||
num_accepted_tokens_cpu=[],
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def prepare_extend_after_decode(
|
|
||||||
self,
|
|
||||||
batch: ScheduleBatch,
|
|
||||||
verify_output: "EagleVerifyOutput",
|
|
||||||
speculative_num_steps: int,
|
|
||||||
):
|
|
||||||
|
|
||||||
if batch.forward_mode.is_idle():
|
|
||||||
return
|
|
||||||
|
|
||||||
# All transient verify->extend handoff state is read off `verify_output`,
|
|
||||||
# not from `self`. The kernel below populates `self.bonus_tokens`
|
|
||||||
# ([bs] per-req) for the next decode round; that is the only state on
|
|
||||||
# `self` that survives past this method.
|
|
||||||
batch.input_ids = verify_output.unfinished_accept_tokens
|
|
||||||
batch.extend_lens = batch.spec_info.num_accepted_tokens_cpu
|
|
||||||
batch.extend_num_tokens = sum(batch.extend_lens)
|
|
||||||
batch.seq_lens = verify_output.seq_lens_for_draft_extend
|
|
||||||
batch.seq_lens_cpu = verify_output.seq_lens_for_draft_extend_cpu
|
|
||||||
batch.req_pool_indices = verify_output.req_pool_indices_for_draft_extend
|
|
||||||
batch.return_logprob = False
|
|
||||||
batch.return_hidden_states = False
|
|
||||||
|
|
||||||
self.capture_hidden_mode = CaptureHiddenMode.LAST
|
|
||||||
self.positions = torch.empty_like(batch.input_ids, dtype=torch.long)
|
|
||||||
self.bonus_tokens = torch.empty_like(
|
|
||||||
self.num_accepted_tokens, dtype=torch.int32
|
|
||||||
)
|
|
||||||
|
|
||||||
create_extend_after_decode_spec_info[(len(batch.seq_lens),)](
|
|
||||||
batch.input_ids,
|
|
||||||
batch.seq_lens,
|
|
||||||
self.num_accepted_tokens,
|
|
||||||
self.positions,
|
|
||||||
self.bonus_tokens,
|
|
||||||
next_power_of_2(max(speculative_num_steps + 1, len(batch.seq_lens))),
|
|
||||||
)
|
|
||||||
|
|
||||||
def generate_attn_arg_prefill(
|
|
||||||
self,
|
|
||||||
req_pool_indices: torch.Tensor,
|
|
||||||
paged_kernel_lens: torch.Tensor,
|
|
||||||
paged_kernel_lens_sum: int,
|
|
||||||
req_to_token: torch.Tensor,
|
|
||||||
):
|
|
||||||
device = req_pool_indices.device
|
|
||||||
bs = self.num_accepted_drafts.numel()
|
|
||||||
qo_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device=device)
|
|
||||||
qo_indptr[1:] = torch.cumsum(self.num_accepted_tokens, dim=0)
|
|
||||||
cum_kv_seq_len = torch.zeros((bs + 1,), dtype=torch.int32, device=device)
|
|
||||||
cum_kv_seq_len[1:] = torch.cumsum(paged_kernel_lens, dim=0)
|
|
||||||
|
|
||||||
if paged_kernel_lens_sum is None:
|
|
||||||
paged_kernel_lens_sum = cum_kv_seq_len[-1]
|
|
||||||
|
|
||||||
kv_indices = torch.empty(
|
|
||||||
paged_kernel_lens_sum, dtype=torch.int32, device=device
|
|
||||||
)
|
|
||||||
|
|
||||||
create_flashinfer_kv_indices_triton[(bs,)](
|
|
||||||
req_to_token,
|
|
||||||
req_pool_indices,
|
|
||||||
paged_kernel_lens,
|
|
||||||
cum_kv_seq_len,
|
|
||||||
None,
|
|
||||||
kv_indices,
|
|
||||||
req_to_token.size(1),
|
|
||||||
)
|
|
||||||
return kv_indices, cum_kv_seq_len, qo_indptr, None
|
|
||||||
|
|
||||||
def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
|
def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
|
||||||
if self.future_indices is not None:
|
if self.future_indices is not None:
|
||||||
self.future_indices.indices = self.future_indices.indices[new_indices]
|
self.future_indices.indices = self.future_indices.indices[new_indices]
|
||||||
@@ -871,29 +767,155 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
|
|||||||
self.topk_index = torch.cat([self.topk_index, spec_info.topk_index])
|
self.topk_index = torch.cat([self.topk_index, spec_info.topk_index])
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class EagleDraftExtendInput(SpecInput):
|
||||||
|
"""Inputs to the draft-extend forward (the per-accepted-token pass after verify).
|
||||||
|
|
||||||
|
Produced by `EagleVerifyInput.verify`, installed on `batch.spec_info` for
|
||||||
|
the draft-extend forward, then replaced with a fresh `EagleDraftInput` for
|
||||||
|
the next iter's draft.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# shape: (total_accepted, hidden_size). Sliced from verify-time hidden_states
|
||||||
|
# by accept_index; consumed by the draft-extend forward.
|
||||||
|
hidden_states: torch.Tensor = None
|
||||||
|
|
||||||
|
# Per-req accept counts. `num_accepted_tokens = num_accepted_drafts + 1`.
|
||||||
|
# Both kept for cuda-graph buffer indexing and the
|
||||||
|
# `create_extend_after_decode_spec_info` kernel.
|
||||||
|
num_accepted_drafts: torch.Tensor = None
|
||||||
|
num_accepted_tokens: torch.Tensor = None
|
||||||
|
# CPU view, read by attention backends during the extend forward.
|
||||||
|
num_accepted_tokens_cpu: List[int] = None
|
||||||
|
|
||||||
|
# Batch-state slices for the draft-extend forward. Set by verify (sliced to
|
||||||
|
# reqs continuing into next iter). `prepare_extend_after_decode` copies
|
||||||
|
# these onto `batch.{input_ids, seq_lens, seq_lens_cpu, req_pool_indices}`.
|
||||||
|
# - input_ids: accept tokens flat over surviving reqs
|
||||||
|
# - seq_lens / _cpu: per-req sequence length (post-accept)
|
||||||
|
# - req_pool_indices: per-req kv-pool slot
|
||||||
|
input_ids: torch.Tensor = None
|
||||||
|
seq_lens: torch.Tensor = None
|
||||||
|
seq_lens_cpu: torch.Tensor = None
|
||||||
|
req_pool_indices: torch.Tensor = None
|
||||||
|
|
||||||
|
# Set by `prepare_extend_after_decode`:
|
||||||
|
# - positions: kernel-written, shape `[total_accepted]`.
|
||||||
|
# - bonus_tokens: kernel-written, shape `[bs]`. The worker reads this
|
||||||
|
# post-extend to populate next iter's `EagleDraftInput.bonus_tokens`.
|
||||||
|
positions: Optional[torch.Tensor] = None
|
||||||
|
bonus_tokens: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.LAST
|
||||||
|
num_tokens_per_req: int = -1
|
||||||
|
num_tokens_for_logprob_per_req: int = 1
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
super().__init__(SpecInputType.EAGLE_DRAFT_EXTEND)
|
||||||
|
|
||||||
|
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
|
||||||
|
return self.num_tokens_per_req, self.num_tokens_for_logprob_per_req
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def create_idle_input(
|
||||||
|
cls,
|
||||||
|
device: torch.device,
|
||||||
|
hidden_size: int,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.LAST,
|
||||||
|
) -> "EagleDraftExtendInput":
|
||||||
|
return cls(
|
||||||
|
hidden_states=torch.empty((0, hidden_size), device=device, dtype=dtype),
|
||||||
|
num_accepted_drafts=torch.empty((0,), device=device, dtype=torch.int32),
|
||||||
|
num_accepted_tokens=torch.empty((0,), device=device, dtype=torch.int32),
|
||||||
|
num_accepted_tokens_cpu=[],
|
||||||
|
input_ids=torch.empty((0,), device=device, dtype=torch.long),
|
||||||
|
seq_lens=torch.empty((0,), device=device, dtype=torch.int32),
|
||||||
|
seq_lens_cpu=torch.empty((0,), dtype=torch.int32),
|
||||||
|
req_pool_indices=torch.empty((0,), device=device, dtype=torch.int64),
|
||||||
|
capture_hidden_mode=capture_hidden_mode,
|
||||||
|
)
|
||||||
|
|
||||||
|
def prepare_extend_after_decode(
|
||||||
|
self,
|
||||||
|
batch: ScheduleBatch,
|
||||||
|
speculative_num_steps: int,
|
||||||
|
):
|
||||||
|
# Caller must have installed `self` as `batch.spec_info` before calling.
|
||||||
|
assert batch.spec_info is self
|
||||||
|
if batch.forward_mode.is_idle():
|
||||||
|
return
|
||||||
|
|
||||||
|
# The kernel below populates `self.positions` and `self.bonus_tokens`;
|
||||||
|
# the worker reads `self.bonus_tokens` to construct next iter's
|
||||||
|
# `EagleDraftInput`.
|
||||||
|
batch.input_ids = self.input_ids
|
||||||
|
batch.extend_lens = self.num_accepted_tokens_cpu
|
||||||
|
batch.extend_num_tokens = sum(batch.extend_lens)
|
||||||
|
batch.seq_lens = self.seq_lens
|
||||||
|
batch.seq_lens_cpu = self.seq_lens_cpu
|
||||||
|
batch.req_pool_indices = self.req_pool_indices
|
||||||
|
batch.return_logprob = False
|
||||||
|
batch.return_hidden_states = False
|
||||||
|
|
||||||
|
self.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||||
|
self.positions = torch.empty_like(batch.input_ids, dtype=torch.long)
|
||||||
|
self.bonus_tokens = torch.empty_like(
|
||||||
|
self.num_accepted_tokens, dtype=torch.int32
|
||||||
|
)
|
||||||
|
|
||||||
|
create_extend_after_decode_spec_info[(len(batch.seq_lens),)](
|
||||||
|
batch.input_ids,
|
||||||
|
batch.seq_lens,
|
||||||
|
self.num_accepted_tokens,
|
||||||
|
self.positions,
|
||||||
|
self.bonus_tokens,
|
||||||
|
next_power_of_2(max(speculative_num_steps + 1, len(batch.seq_lens))),
|
||||||
|
)
|
||||||
|
|
||||||
|
def generate_attn_arg_prefill(
|
||||||
|
self,
|
||||||
|
req_pool_indices: torch.Tensor,
|
||||||
|
paged_kernel_lens: torch.Tensor,
|
||||||
|
paged_kernel_lens_sum: Optional[int],
|
||||||
|
req_to_token: torch.Tensor,
|
||||||
|
):
|
||||||
|
device = req_pool_indices.device
|
||||||
|
bs = self.num_accepted_drafts.numel()
|
||||||
|
qo_indptr = torch.zeros((bs + 1,), dtype=torch.int32, device=device)
|
||||||
|
qo_indptr[1:] = torch.cumsum(self.num_accepted_tokens, dim=0)
|
||||||
|
cum_kv_seq_len = torch.zeros((bs + 1,), dtype=torch.int32, device=device)
|
||||||
|
cum_kv_seq_len[1:] = torch.cumsum(paged_kernel_lens, dim=0)
|
||||||
|
|
||||||
|
if paged_kernel_lens_sum is None:
|
||||||
|
paged_kernel_lens_sum = cum_kv_seq_len[-1]
|
||||||
|
|
||||||
|
kv_indices = torch.empty(
|
||||||
|
paged_kernel_lens_sum, dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
|
||||||
|
create_flashinfer_kv_indices_triton[(bs,)](
|
||||||
|
req_to_token,
|
||||||
|
req_pool_indices,
|
||||||
|
paged_kernel_lens,
|
||||||
|
cum_kv_seq_len,
|
||||||
|
None,
|
||||||
|
kv_indices,
|
||||||
|
req_to_token.size(1),
|
||||||
|
)
|
||||||
|
return kv_indices, cum_kv_seq_len, qo_indptr, None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class EagleVerifyOutput:
|
class EagleVerifyOutput:
|
||||||
# Next iter's persistent draft state, ready to be installed as `batch.spec_info`.
|
# Next iter's draft-extend input, installed as `batch.spec_info` for the
|
||||||
next_draft_input: EagleDraftInput
|
# draft-extend forward.
|
||||||
|
draft_extend_input: EagleDraftExtendInput
|
||||||
# Logit outputs from target worker.
|
# Logit outputs from target worker.
|
||||||
logits_output: LogitsProcessorOutput
|
logits_output: LogitsProcessorOutput
|
||||||
# All accepted tokens flat across all reqs incl. those that finished this
|
# All accepted tokens flat across all reqs incl. those that finished this
|
||||||
# step. Includes the bonus token. Used for output processing.
|
# step. Includes the bonus token. Used for output processing.
|
||||||
accept_tokens: torch.Tensor
|
accept_tokens: torch.Tensor
|
||||||
# Below are transient handoff fields for the next iter's draft-extend pass.
|
|
||||||
# They are scoped to the verify -> prepare_extend_after_decode window only;
|
|
||||||
# `prepare_extend_after_decode` reads them off this object via method arg
|
|
||||||
# rather than smuggling them through `EagleDraftInput`.
|
|
||||||
#
|
|
||||||
# Subset of `accept_tokens` for reqs continuing into next iter's draft-extend
|
|
||||||
# forward (= `accept_tokens` when no req finished; flat over unfinished
|
|
||||||
# reqs only otherwise). Becomes `batch.input_ids` for that forward pass.
|
|
||||||
unfinished_accept_tokens: torch.Tensor
|
|
||||||
# `batch.seq_lens` / `batch.seq_lens_cpu` / `batch.req_pool_indices` to
|
|
||||||
# use for the next iter's draft-extend forward; sliced to surviving reqs.
|
|
||||||
seq_lens_for_draft_extend: torch.Tensor
|
|
||||||
seq_lens_for_draft_extend_cpu: torch.Tensor
|
|
||||||
req_pool_indices_for_draft_extend: torch.Tensor
|
|
||||||
# Accepted token length per sequence in a batch in CPU (full set).
|
# Accepted token length per sequence in a batch in CPU (full set).
|
||||||
num_accepted_drafts_per_req_cpu: List[int]
|
num_accepted_drafts_per_req_cpu: List[int]
|
||||||
# Accepted indices from logits_output.next_token_logits
|
# Accepted indices from logits_output.next_token_logits
|
||||||
@@ -903,21 +925,15 @@ class EagleVerifyOutput:
|
|||||||
def create_idle(
|
def create_idle(
|
||||||
cls,
|
cls,
|
||||||
*,
|
*,
|
||||||
next_draft_input: EagleDraftInput,
|
draft_extend_input: EagleDraftExtendInput,
|
||||||
logits_output: LogitsProcessorOutput,
|
logits_output: LogitsProcessorOutput,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
spec_steps: int,
|
spec_steps: int,
|
||||||
) -> "EagleVerifyOutput":
|
) -> "EagleVerifyOutput":
|
||||||
return cls(
|
return cls(
|
||||||
next_draft_input=next_draft_input,
|
draft_extend_input=draft_extend_input,
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
accept_tokens=torch.empty(0, dtype=torch.long, device=device),
|
accept_tokens=torch.empty(0, dtype=torch.long, device=device),
|
||||||
unfinished_accept_tokens=torch.empty(0, dtype=torch.long, device=device),
|
|
||||||
seq_lens_for_draft_extend=torch.empty(0, dtype=torch.int32, device=device),
|
|
||||||
seq_lens_for_draft_extend_cpu=torch.empty(0, dtype=torch.int32),
|
|
||||||
req_pool_indices_for_draft_extend=torch.empty(
|
|
||||||
0, dtype=torch.int64, device=device
|
|
||||||
),
|
|
||||||
num_accepted_drafts_per_req_cpu=[],
|
num_accepted_drafts_per_req_cpu=[],
|
||||||
accepted_indices=torch.full(
|
accepted_indices=torch.full(
|
||||||
(0, spec_steps + 1), -1, dtype=torch.int32, device=device
|
(0, spec_steps + 1), -1, dtype=torch.int32, device=device
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import (
|
|||||||
EAGLEDraftExtendCudaGraphRunner,
|
EAGLEDraftExtendCudaGraphRunner,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.eagle_info import (
|
from sglang.srt.speculative.eagle_info import (
|
||||||
|
EagleDraftExtendInput,
|
||||||
EagleDraftInput,
|
EagleDraftInput,
|
||||||
EagleVerifyInput,
|
EagleVerifyInput,
|
||||||
EagleVerifyOutput,
|
EagleVerifyOutput,
|
||||||
@@ -480,14 +481,13 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
with self.draft_tp_context(
|
with self.draft_tp_context(
|
||||||
self.draft_model_runner.tp_group
|
self.draft_model_runner.tp_group
|
||||||
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
spec_info = self.draft(batch)
|
verify_input = self.draft(batch)
|
||||||
|
|
||||||
set_time_batch(batch.reqs, "set_spec_draft_end_time", trace_only=True)
|
set_time_batch(batch.reqs, "set_spec_draft_end_time", trace_only=True)
|
||||||
set_time_batch(batch.reqs, "set_spec_verify_start_time", trace_only=True)
|
set_time_batch(batch.reqs, "set_spec_verify_start_time", trace_only=True)
|
||||||
|
|
||||||
logits_output, verify_output, can_run_cuda_graph = self.verify(
|
batch.spec_info = verify_input
|
||||||
batch, spec_info
|
logits_output, verify_output, can_run_cuda_graph = self.verify(batch)
|
||||||
)
|
|
||||||
|
|
||||||
if get_global_tracing_enabled():
|
if get_global_tracing_enabled():
|
||||||
for idx, req in enumerate(batch.reqs):
|
for idx, req in enumerate(batch.reqs):
|
||||||
@@ -503,12 +503,24 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
# NOTE: We should use `check_forward_draft_extend_after_decode`
|
# NOTE: We should use `check_forward_draft_extend_after_decode`
|
||||||
# when DP attention is enabled, but it is slow. Skip it for now.
|
# when DP attention is enabled, but it is slow. Skip it for now.
|
||||||
|
draft_extend_input = verify_output.draft_extend_input
|
||||||
if (
|
if (
|
||||||
self.server_args.enable_dp_attention
|
self.server_args.enable_dp_attention
|
||||||
or verify_output.unfinished_accept_tokens.shape[0] > 0
|
or draft_extend_input.input_ids.shape[0] > 0
|
||||||
):
|
):
|
||||||
# decode is not finished
|
# decode is not finished; stash for extend, then restash
|
||||||
self.forward_draft_extend_after_decode(batch, verify_output)
|
# the next-iter EagleDraftInput it returns.
|
||||||
|
batch.spec_info = draft_extend_input
|
||||||
|
next_draft_input = self.forward_draft_extend_after_decode(batch)
|
||||||
|
batch.spec_info = next_draft_input
|
||||||
|
else:
|
||||||
|
# All reqs finished and dp_attention isn't forcing extend.
|
||||||
|
# Stash an empty EagleDraftInput so next iter's merge_batch
|
||||||
|
# short-circuits on None hidden_states (EagleVerifyInput
|
||||||
|
# has no merge_batch).
|
||||||
|
batch.spec_info = EagleDraftInput(
|
||||||
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
|
)
|
||||||
|
|
||||||
set_time_batch(
|
set_time_batch(
|
||||||
batch.reqs, "set_spec_draft_extend_end_time", trace_only=True
|
batch.reqs, "set_spec_draft_extend_end_time", trace_only=True
|
||||||
@@ -528,7 +540,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def check_forward_draft_extend_after_decode(self, verify_output: EagleVerifyOutput):
|
def check_forward_draft_extend_after_decode(self, verify_output: EagleVerifyOutput):
|
||||||
local_need_forward = verify_output.unfinished_accept_tokens.shape[0] > 0
|
local_need_forward = verify_output.draft_extend_input.input_ids.shape[0] > 0
|
||||||
if not self.server_args.enable_dp_attention:
|
if not self.server_args.enable_dp_attention:
|
||||||
return local_need_forward
|
return local_need_forward
|
||||||
|
|
||||||
@@ -890,7 +902,8 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
# allocator and kv cache pool are shared with target worker
|
# allocator and kv cache pool are shared with target worker
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def verify(self, batch: ScheduleBatch, spec_info: EagleVerifyInput):
|
def verify(self, batch: ScheduleBatch):
|
||||||
|
spec_info: EagleVerifyInput = batch.spec_info
|
||||||
seq_lens_pre_verify = batch.seq_lens.clone()
|
seq_lens_pre_verify = batch.seq_lens.clone()
|
||||||
spec_info.prepare_for_verify(batch, self.page_size)
|
spec_info.prepare_for_verify(batch, self.page_size)
|
||||||
spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
||||||
@@ -900,7 +913,6 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
if not batch.forward_mode.is_idle()
|
if not batch.forward_mode.is_idle()
|
||||||
else ForwardMode.IDLE
|
else ForwardMode.IDLE
|
||||||
)
|
)
|
||||||
batch.spec_info = spec_info
|
|
||||||
|
|
||||||
model_worker_batch = batch.get_model_worker_batch(
|
model_worker_batch = batch.get_model_worker_batch(
|
||||||
seq_lens_cpu_cache=spec_info.seq_lens_cpu
|
seq_lens_cpu_cache=spec_info.seq_lens_cpu
|
||||||
@@ -977,7 +989,6 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE
|
ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE
|
||||||
)
|
)
|
||||||
batch.spec_info = res.next_draft_input
|
|
||||||
|
|
||||||
return logits_output, res, can_run_cuda_graph
|
return logits_output, res, can_run_cuda_graph
|
||||||
|
|
||||||
@@ -1105,20 +1116,20 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
self.capture_for_decode(logits_output, forward_batch.spec_info)
|
self.capture_for_decode(logits_output, forward_batch.spec_info)
|
||||||
|
|
||||||
def forward_draft_extend_after_decode(
|
def forward_draft_extend_after_decode(
|
||||||
self, batch: ScheduleBatch, verify_output: EagleVerifyOutput
|
self, batch: ScheduleBatch
|
||||||
):
|
) -> EagleDraftInput:
|
||||||
assert isinstance(batch.spec_info, EagleDraftInput)
|
draft_extend_input: EagleDraftExtendInput = batch.spec_info
|
||||||
|
|
||||||
# Backup fields that will be modified in-place
|
# Backup fields that will be modified in-place
|
||||||
seq_lens_backup = batch.seq_lens.clone()
|
seq_lens_backup = batch.seq_lens.clone()
|
||||||
seq_lens_cpu_backup = batch.seq_lens_cpu.clone()
|
seq_lens_cpu_backup = batch.seq_lens_cpu.clone()
|
||||||
req_pool_indices_backup = batch.req_pool_indices
|
req_pool_indices_backup = batch.req_pool_indices
|
||||||
num_accepted_drafts_backup = batch.spec_info.num_accepted_drafts.clone()
|
|
||||||
num_accepted_tokens_backup = batch.spec_info.num_accepted_tokens.clone()
|
|
||||||
return_logprob_backup = batch.return_logprob
|
return_logprob_backup = batch.return_logprob
|
||||||
|
|
||||||
input_is_idle = batch.forward_mode.is_idle()
|
input_is_idle = batch.forward_mode.is_idle()
|
||||||
|
|
||||||
if not input_is_idle and verify_output.unfinished_accept_tokens.numel() == 0:
|
if not input_is_idle and draft_extend_input.input_ids.numel() == 0:
|
||||||
|
# All reqs finished this verify; swap to an idle ExtendInput.
|
||||||
batch = batch.copy()
|
batch = batch.copy()
|
||||||
batch.prepare_for_idle()
|
batch.prepare_for_idle()
|
||||||
hidden_size = (
|
hidden_size = (
|
||||||
@@ -1127,19 +1138,19 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
and self.eagle_use_aux_hidden_state
|
and self.eagle_use_aux_hidden_state
|
||||||
else self.model_config.spec_hidden_size
|
else self.model_config.spec_hidden_size
|
||||||
)
|
)
|
||||||
batch.spec_info = EagleDraftInput.create_idle_input(
|
draft_extend_input = EagleDraftExtendInput.create_idle_input(
|
||||||
device=self.device,
|
device=self.device,
|
||||||
hidden_size=hidden_size,
|
hidden_size=hidden_size,
|
||||||
dtype=self.model_config.dtype,
|
dtype=self.model_config.dtype,
|
||||||
topk=self.topk,
|
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
|
batch.spec_info = draft_extend_input
|
||||||
|
|
||||||
batch.spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
# Phase 1: prepare extend (kernel writes draft_extend_input.{positions, bonus_tokens})
|
||||||
batch.spec_info.num_tokens_for_logprob_per_req = 1
|
draft_extend_input.num_tokens_per_req = self.speculative_num_steps + 1
|
||||||
batch.spec_info.prepare_extend_after_decode(
|
draft_extend_input.num_tokens_for_logprob_per_req = 1
|
||||||
|
draft_extend_input.prepare_extend_after_decode(
|
||||||
batch,
|
batch,
|
||||||
verify_output=verify_output,
|
|
||||||
speculative_num_steps=self.speculative_num_steps,
|
speculative_num_steps=self.speculative_num_steps,
|
||||||
)
|
)
|
||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
@@ -1159,7 +1170,7 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
else:
|
else:
|
||||||
forward_batch.seq_lens_sum = batch.seq_lens.sum().item()
|
forward_batch.seq_lens_sum = batch.seq_lens.sum().item()
|
||||||
|
|
||||||
# Run
|
# Phase 2: run draft-extend forward
|
||||||
can_cuda_graph = (
|
can_cuda_graph = (
|
||||||
self.cuda_graph_runner_for_draft_extend
|
self.cuda_graph_runner_for_draft_extend
|
||||||
and self.cuda_graph_runner_for_draft_extend.can_run(forward_batch)
|
and self.cuda_graph_runner_for_draft_extend.can_run(forward_batch)
|
||||||
@@ -1168,11 +1179,10 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
logits_output = self.cuda_graph_runner_for_draft_extend.replay(
|
logits_output = self.cuda_graph_runner_for_draft_extend.replay(
|
||||||
forward_batch
|
forward_batch
|
||||||
)
|
)
|
||||||
forward_batch.spec_info.topk_p, forward_batch.spec_info.topk_index = (
|
# cuda-graph replay populates logits_output.{topk_p, topk_index, hidden_states}.
|
||||||
logits_output.topk_p,
|
topk_p = logits_output.topk_p
|
||||||
logits_output.topk_index,
|
topk_index = logits_output.topk_index
|
||||||
)
|
hidden_states = logits_output.hidden_states
|
||||||
forward_batch.spec_info.hidden_states = logits_output.hidden_states
|
|
||||||
else:
|
else:
|
||||||
forward_batch.can_run_dp_cuda_graph = False
|
forward_batch.can_run_dp_cuda_graph = False
|
||||||
if not forward_batch.forward_mode.is_idle():
|
if not forward_batch.forward_mode.is_idle():
|
||||||
@@ -1185,24 +1195,36 @@ class EAGLEWorker(TpModelWorker):
|
|||||||
logits_output = self.draft_model_runner.forward(
|
logits_output = self.draft_model_runner.forward(
|
||||||
forward_batch, skip_attn_backend_init=True
|
forward_batch, skip_attn_backend_init=True
|
||||||
).logits_output
|
).logits_output
|
||||||
self.capture_for_decode(logits_output, forward_batch.spec_info)
|
# Non-cuda-graph path: compute topk_p / topk_index inline.
|
||||||
|
probs = torch.softmax(logits_output.next_token_logits, dim=-1)
|
||||||
|
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
|
||||||
|
hidden_states = logits_output.hidden_states
|
||||||
|
|
||||||
maybe_detect_nan(
|
maybe_detect_nan(
|
||||||
logits_output.next_token_logits,
|
logits_output.next_token_logits,
|
||||||
f"draft_extend_after_decode (cuda_graph={can_cuda_graph})",
|
f"draft_extend_after_decode (cuda_graph={can_cuda_graph})",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Restore backup.
|
# Phase 3: assemble next-iter EagleDraftInput from extend output
|
||||||
# This is because `seq_lens` can be modified in `prepare_extend_after_decode`
|
next_draft_input = EagleDraftInput(
|
||||||
|
bonus_tokens=draft_extend_input.bonus_tokens,
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
topk_p=topk_p,
|
||||||
|
topk_index=topk_index,
|
||||||
|
capture_hidden_mode=CaptureHiddenMode.FULL,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Restore batch fields. `seq_lens` etc. were modified by
|
||||||
|
# `prepare_extend_after_decode`. Caller installs `next_draft_input` as
|
||||||
|
# `batch.spec_info`.
|
||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
ForwardMode.DECODE if not input_is_idle else ForwardMode.IDLE
|
ForwardMode.DECODE if not input_is_idle else ForwardMode.IDLE
|
||||||
)
|
)
|
||||||
batch.seq_lens = seq_lens_backup
|
batch.seq_lens = seq_lens_backup
|
||||||
batch.seq_lens_cpu = seq_lens_cpu_backup
|
batch.seq_lens_cpu = seq_lens_cpu_backup
|
||||||
batch.req_pool_indices = req_pool_indices_backup
|
batch.req_pool_indices = req_pool_indices_backup
|
||||||
batch.spec_info.num_accepted_drafts = num_accepted_drafts_backup
|
|
||||||
batch.spec_info.num_accepted_tokens = num_accepted_tokens_backup
|
|
||||||
batch.return_logprob = return_logprob_backup
|
batch.return_logprob = return_logprob_backup
|
||||||
|
return next_draft_input
|
||||||
|
|
||||||
def capture_for_decode(
|
def capture_for_decode(
|
||||||
self, logits_output: LogitsProcessorOutput, draft_input: EagleDraftInput
|
self, logits_output: LogitsProcessorOutput, draft_input: EagleDraftInput
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from typing import Dict
|
|||||||
|
|
||||||
from sglang.srt.mem_cache.memory_pool import KVCache
|
from sglang.srt.mem_cache.memory_pool import KVCache
|
||||||
from sglang.srt.speculative.eagle_info import (
|
from sglang.srt.speculative.eagle_info import (
|
||||||
|
EagleDraftExtendInput,
|
||||||
EagleDraftInput,
|
EagleDraftInput,
|
||||||
EagleVerifyInput,
|
EagleVerifyInput,
|
||||||
EagleVerifyOutput,
|
EagleVerifyOutput,
|
||||||
@@ -53,6 +54,14 @@ class FrozenKVMTPDraftInput(EagleDraftInput):
|
|||||||
SpecInput.__init__(self, SpecInputType.FROZEN_KV_MTP_DRAFT)
|
SpecInput.__init__(self, SpecInputType.FROZEN_KV_MTP_DRAFT)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class FrozenKVMTPDraftExtendInput(EagleDraftExtendInput):
|
||||||
|
"""Draft-extend input for Frozen-KV MTP. Tag-only subclass."""
|
||||||
|
|
||||||
|
def __post_init__(self):
|
||||||
|
SpecInput.__init__(self, SpecInputType.FROZEN_KV_MTP_DRAFT_EXTEND)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class FrozenKVMTPVerifyInput(EagleVerifyInput):
|
class FrozenKVMTPVerifyInput(EagleVerifyInput):
|
||||||
"""Verify input for Frozen-KV MTP."""
|
"""Verify input for Frozen-KV MTP."""
|
||||||
@@ -62,21 +71,23 @@ class FrozenKVMTPVerifyInput(EagleVerifyInput):
|
|||||||
|
|
||||||
def verify(self, *args, **kwargs) -> EagleVerifyOutput:
|
def verify(self, *args, **kwargs) -> EagleVerifyOutput:
|
||||||
output = super().verify(*args, **kwargs)
|
output = super().verify(*args, **kwargs)
|
||||||
output.next_draft_input = _to_frozen_kv_mtp_draft_input(output.next_draft_input)
|
output.draft_extend_input = _to_frozen_kv_mtp_draft_extend_input(
|
||||||
|
output.draft_extend_input
|
||||||
|
)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
FrozenKVMTPVerifyOutput = EagleVerifyOutput
|
FrozenKVMTPVerifyOutput = EagleVerifyOutput
|
||||||
|
|
||||||
|
|
||||||
def _to_frozen_kv_mtp_draft_input(
|
def _to_frozen_kv_mtp_draft_extend_input(
|
||||||
draft_input: EagleDraftInput,
|
draft_extend_input: EagleDraftExtendInput,
|
||||||
) -> FrozenKVMTPDraftInput:
|
) -> FrozenKVMTPDraftExtendInput:
|
||||||
if isinstance(draft_input, FrozenKVMTPDraftInput):
|
if isinstance(draft_extend_input, FrozenKVMTPDraftExtendInput):
|
||||||
return draft_input
|
return draft_extend_input
|
||||||
return FrozenKVMTPDraftInput(
|
return FrozenKVMTPDraftExtendInput(
|
||||||
**{
|
**{
|
||||||
field.name: getattr(draft_input, field.name)
|
field.name: getattr(draft_extend_input, field.name)
|
||||||
for field in fields(EagleDraftInput)
|
for field in fields(EagleDraftExtendInput)
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.srt.managers.schedule_batch import ScheduleBatch
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.speculative.frozen_kv_mtp_info import (
|
from sglang.srt.speculative.frozen_kv_mtp_info import (
|
||||||
FrozenKVMTPContext,
|
FrozenKVMTPContext,
|
||||||
|
FrozenKVMTPDraftExtendInput,
|
||||||
FrozenKVMTPDraftInput,
|
FrozenKVMTPDraftInput,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.spec_utils import fast_topk
|
from sglang.srt.speculative.spec_utils import fast_topk
|
||||||
@@ -134,11 +135,8 @@ def select_last_extend_hidden(
|
|||||||
|
|
||||||
|
|
||||||
def select_last_verified_seed(
|
def select_last_verified_seed(
|
||||||
draft_input: FrozenKVMTPDraftInput,
|
draft_input: FrozenKVMTPDraftExtendInput,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
if draft_input.num_accepted_tokens is None:
|
|
||||||
return draft_input.bonus_tokens, draft_input.hidden_states
|
|
||||||
|
|
||||||
counts = draft_input.num_accepted_tokens.to(torch.long)
|
counts = draft_input.num_accepted_tokens.to(torch.long)
|
||||||
last_indices = torch.cumsum(counts, dim=0) - 1
|
last_indices = torch.cumsum(counts, dim=0) - 1
|
||||||
return (
|
return (
|
||||||
|
|||||||
@@ -43,13 +43,13 @@ from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
|||||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
from sglang.srt.observability.req_time_stats import set_time_batch
|
||||||
from sglang.srt.observability.trace import get_global_tracing_enabled
|
from sglang.srt.observability.trace import get_global_tracing_enabled
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.eagle_info import EagleVerifyOutput
|
|
||||||
from sglang.srt.speculative.eagle_utils import (
|
from sglang.srt.speculative.eagle_utils import (
|
||||||
build_tree_kernel_efficient,
|
build_tree_kernel_efficient,
|
||||||
organize_draft_results,
|
organize_draft_results,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.frozen_kv_mtp_info import (
|
from sglang.srt.speculative.frozen_kv_mtp_info import (
|
||||||
FrozenKVMTPContext,
|
FrozenKVMTPContext,
|
||||||
|
FrozenKVMTPDraftExtendInput,
|
||||||
FrozenKVMTPDraftInput,
|
FrozenKVMTPDraftInput,
|
||||||
FrozenKVMTPVerifyInput,
|
FrozenKVMTPVerifyInput,
|
||||||
FrozenKVMTPVerifyOutput,
|
FrozenKVMTPVerifyOutput,
|
||||||
@@ -335,7 +335,7 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
return select_last_extend_hidden(batch, hidden_states)
|
return select_last_extend_hidden(batch, hidden_states)
|
||||||
|
|
||||||
def _select_last_verified_seed(
|
def _select_last_verified_seed(
|
||||||
self, draft_input: FrozenKVMTPDraftInput
|
self, draft_input: FrozenKVMTPDraftExtendInput
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
return select_last_verified_seed(draft_input)
|
return select_last_verified_seed(draft_input)
|
||||||
|
|
||||||
@@ -443,11 +443,12 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
with self.draft_tp_context(
|
with self.draft_tp_context(
|
||||||
self.draft_model_runner.tp_group
|
self.draft_model_runner.tp_group
|
||||||
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
spec_info = self.draft(batch)
|
verify_input = self.draft(batch)
|
||||||
set_time_batch(batch.reqs, "set_spec_draft_end_time", trace_only=True)
|
set_time_batch(batch.reqs, "set_spec_draft_end_time", trace_only=True)
|
||||||
set_time_batch(batch.reqs, "set_spec_verify_start_time", trace_only=True)
|
set_time_batch(batch.reqs, "set_spec_verify_start_time", trace_only=True)
|
||||||
|
|
||||||
logits_output, verify_output, can_run_cuda_graph = self.verify(batch, spec_info)
|
batch.spec_info = verify_input
|
||||||
|
logits_output, verify_output, can_run_cuda_graph = self.verify(batch)
|
||||||
|
|
||||||
if get_global_tracing_enabled():
|
if get_global_tracing_enabled():
|
||||||
for idx, req in enumerate(batch.reqs):
|
for idx, req in enumerate(batch.reqs):
|
||||||
@@ -458,11 +459,15 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
with self.draft_tp_context(
|
with self.draft_tp_context(
|
||||||
self.draft_model_runner.tp_group
|
self.draft_model_runner.tp_group
|
||||||
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
|
||||||
|
draft_extend_input = verify_output.draft_extend_input
|
||||||
if (
|
if (
|
||||||
self.server_args.enable_dp_attention
|
self.server_args.enable_dp_attention
|
||||||
or batch.spec_info.bonus_tokens.numel()
|
or draft_extend_input.input_ids.numel() > 0
|
||||||
):
|
):
|
||||||
self.forward_draft_extend_after_decode(batch, verify_output)
|
# Stash for the seed step; _run_assistant_seed_step swaps in
|
||||||
|
# a fresh FrozenKVMTPDraftInput for next iter.
|
||||||
|
batch.spec_info = draft_extend_input
|
||||||
|
self.forward_draft_extend_after_decode(batch)
|
||||||
set_time_batch(batch.reqs, "set_spec_draft_extend_end_time", trace_only=True)
|
set_time_batch(batch.reqs, "set_spec_draft_extend_end_time", trace_only=True)
|
||||||
|
|
||||||
return GenerationBatchResult(
|
return GenerationBatchResult(
|
||||||
@@ -503,12 +508,13 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
mm_input_embeds=mm_input_embeds,
|
mm_input_embeds=mm_input_embeds,
|
||||||
)
|
)
|
||||||
|
|
||||||
def forward_draft_extend_after_decode(
|
def forward_draft_extend_after_decode(self, batch: ScheduleBatch) -> None:
|
||||||
self, batch: ScheduleBatch, verify_output: EagleVerifyOutput
|
draft_extend_input: FrozenKVMTPDraftExtendInput = batch.spec_info
|
||||||
) -> None:
|
|
||||||
assert isinstance(batch.spec_info, FrozenKVMTPDraftInput)
|
|
||||||
input_is_idle = batch.forward_mode.is_idle()
|
input_is_idle = batch.forward_mode.is_idle()
|
||||||
if not input_is_idle and batch.spec_info.bonus_tokens.numel() == 0:
|
|
||||||
|
if not input_is_idle and draft_extend_input.input_ids.numel() == 0:
|
||||||
|
# All reqs finished; stash an idle FrozenKVMTPDraftInput so the
|
||||||
|
# next-iter draft sees a valid spec_info.
|
||||||
batch = batch.copy()
|
batch = batch.copy()
|
||||||
batch.prepare_for_idle()
|
batch.prepare_for_idle()
|
||||||
batch.spec_info = FrozenKVMTPDraftInput.create_idle_input(
|
batch.spec_info = FrozenKVMTPDraftInput.create_idle_input(
|
||||||
@@ -518,29 +524,32 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
topk=self.topk,
|
topk=self.topk,
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
|
return
|
||||||
|
|
||||||
if batch.forward_mode.is_idle():
|
if batch.forward_mode.is_idle():
|
||||||
return
|
return
|
||||||
|
|
||||||
draft_input = batch.spec_info
|
|
||||||
seq_lens_backup = batch.seq_lens.clone()
|
seq_lens_backup = batch.seq_lens.clone()
|
||||||
seq_lens_cpu_backup = batch.seq_lens_cpu.clone()
|
seq_lens_cpu_backup = batch.seq_lens_cpu.clone()
|
||||||
req_pool_indices_backup = batch.req_pool_indices
|
req_pool_indices_backup = batch.req_pool_indices
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Verify may leave finished requests in ScheduleBatch; seed only
|
# Verify may leave finished requests in ScheduleBatch; seed only
|
||||||
# the unfinished requests carried by `verify_output`.
|
# the unfinished reqs carried by `draft_extend_input`.
|
||||||
batch.seq_lens = verify_output.seq_lens_for_draft_extend
|
batch.seq_lens = draft_extend_input.seq_lens
|
||||||
batch.seq_lens_cpu = verify_output.seq_lens_for_draft_extend_cpu
|
batch.seq_lens_cpu = draft_extend_input.seq_lens_cpu
|
||||||
batch.req_pool_indices = verify_output.req_pool_indices_for_draft_extend
|
batch.req_pool_indices = draft_extend_input.req_pool_indices
|
||||||
|
|
||||||
last_token_ids, last_hidden = self._select_last_verified_seed(draft_input)
|
last_token_ids, last_hidden = self._select_last_verified_seed(
|
||||||
|
draft_extend_input
|
||||||
|
)
|
||||||
|
# `_run_assistant_seed_step` constructs a fresh `FrozenKVMTPDraftInput`
|
||||||
|
# and installs it on `batch.spec_info` for next iter.
|
||||||
self._run_assistant_seed_step(
|
self._run_assistant_seed_step(
|
||||||
batch,
|
batch,
|
||||||
last_token_ids,
|
last_token_ids,
|
||||||
last_hidden,
|
last_hidden,
|
||||||
seq_lens_cpu=verify_output.seq_lens_for_draft_extend_cpu,
|
seq_lens_cpu=draft_extend_input.seq_lens_cpu,
|
||||||
draft_input=draft_input,
|
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
batch.seq_lens = seq_lens_backup
|
batch.seq_lens = seq_lens_backup
|
||||||
@@ -687,7 +696,8 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
score_list, token_list, parents_list, self.speculative_num_draft_tokens
|
score_list, token_list, parents_list, self.speculative_num_draft_tokens
|
||||||
)
|
)
|
||||||
|
|
||||||
def verify(self, batch: ScheduleBatch, spec_info: FrozenKVMTPVerifyInput):
|
def verify(self, batch: ScheduleBatch):
|
||||||
|
spec_info: FrozenKVMTPVerifyInput = batch.spec_info
|
||||||
seq_lens_pre_verify = batch.seq_lens.clone()
|
seq_lens_pre_verify = batch.seq_lens.clone()
|
||||||
spec_info.prepare_for_verify(batch, self.page_size)
|
spec_info.prepare_for_verify(batch, self.page_size)
|
||||||
spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
||||||
@@ -697,7 +707,6 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
if not batch.forward_mode.is_idle()
|
if not batch.forward_mode.is_idle()
|
||||||
else ForwardMode.IDLE
|
else ForwardMode.IDLE
|
||||||
)
|
)
|
||||||
batch.spec_info = spec_info
|
|
||||||
|
|
||||||
model_worker_batch = batch.get_model_worker_batch(
|
model_worker_batch = batch.get_model_worker_batch(
|
||||||
seq_lens_cpu_cache=spec_info.seq_lens_cpu
|
seq_lens_cpu_cache=spec_info.seq_lens_cpu
|
||||||
@@ -766,7 +775,6 @@ class FrozenKVMTPWorker(TpModelWorker):
|
|||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE
|
ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE
|
||||||
)
|
)
|
||||||
batch.spec_info = res.next_draft_input
|
|
||||||
|
|
||||||
del seq_lens_pre_verify
|
del seq_lens_pre_verify
|
||||||
return logits_output, res, can_run_cuda_graph
|
return logits_output, res, can_run_cuda_graph
|
||||||
|
|||||||
@@ -41,7 +41,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardMode,
|
ForwardMode,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||||
from sglang.srt.speculative.multi_layer_eagle_utils import assign_new_state_triton
|
from sglang.srt.speculative.multi_layer_eagle_utils import assign_new_state_triton
|
||||||
from sglang.srt.speculative.spec_utils import fast_topk
|
from sglang.srt.speculative.spec_utils import fast_topk
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -349,7 +349,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner:
|
|||||||
else:
|
else:
|
||||||
global_dp_buffer_len = None
|
global_dp_buffer_len = None
|
||||||
|
|
||||||
spec_info = EagleDraftInput(
|
spec_info = EagleDraftExtendInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
num_accepted_drafts=num_accepted_drafts,
|
num_accepted_drafts=num_accepted_drafts,
|
||||||
num_accepted_tokens=num_accepted_tokens,
|
num_accepted_tokens=num_accepted_tokens,
|
||||||
|
|||||||
@@ -34,6 +34,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.draft_utils import DraftBackendFactory
|
from sglang.srt.speculative.draft_utils import DraftBackendFactory
|
||||||
from sglang.srt.speculative.eagle_info import (
|
from sglang.srt.speculative.eagle_info import (
|
||||||
|
EagleDraftExtendInput,
|
||||||
EagleDraftInput,
|
EagleDraftInput,
|
||||||
EagleVerifyInput,
|
EagleVerifyInput,
|
||||||
EagleVerifyOutput,
|
EagleVerifyOutput,
|
||||||
@@ -271,22 +272,33 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
with self.draft_tp_context(
|
with self.draft_tp_context(
|
||||||
self.mtp_model_runner(0).tp_group
|
self.mtp_model_runner(0).tp_group
|
||||||
), speculative_moe_backend_context():
|
), speculative_moe_backend_context():
|
||||||
spec_info = self.draft(batch)
|
verify_input = self.draft(batch)
|
||||||
logits_output, verify_output, can_run_cuda_graph = self.verify(
|
batch.spec_info = verify_input
|
||||||
batch, spec_info
|
logits_output, verify_output, can_run_cuda_graph = self.verify(batch)
|
||||||
)
|
|
||||||
|
|
||||||
with self.draft_tp_context(
|
with self.draft_tp_context(
|
||||||
self.mtp_model_runner(0).tp_group
|
self.mtp_model_runner(0).tp_group
|
||||||
), speculative_moe_backend_context():
|
), speculative_moe_backend_context():
|
||||||
# NOTE: We should use `check_forward_draft_extend_after_decode`
|
# NOTE: We should use `check_forward_draft_extend_after_decode`
|
||||||
# when DP attention is enabled, but it is slow. Skip it for now.
|
# when DP attention is enabled, but it is slow. Skip it for now.
|
||||||
|
draft_extend_input = verify_output.draft_extend_input
|
||||||
if (
|
if (
|
||||||
self.server_args.enable_dp_attention
|
self.server_args.enable_dp_attention
|
||||||
or verify_output.unfinished_accept_tokens.shape[0] > 0
|
or draft_extend_input.input_ids.shape[0] > 0
|
||||||
):
|
):
|
||||||
# decode is not finished
|
# decode is not finished; stash for extend, then restash
|
||||||
self.forward_draft_extend_after_decode(batch, verify_output)
|
# the next-iter EagleDraftInput it returns.
|
||||||
|
batch.spec_info = draft_extend_input
|
||||||
|
next_draft_input = self.forward_draft_extend_after_decode(batch)
|
||||||
|
batch.spec_info = next_draft_input
|
||||||
|
else:
|
||||||
|
# All reqs finished and dp_attention isn't forcing extend.
|
||||||
|
# Stash an empty EagleDraftInput so next iter's merge_batch
|
||||||
|
# short-circuits on None hidden_states (EagleVerifyInput
|
||||||
|
# has no merge_batch).
|
||||||
|
batch.spec_info = EagleDraftInput(
|
||||||
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
|
)
|
||||||
|
|
||||||
return GenerationBatchResult(
|
return GenerationBatchResult(
|
||||||
logits_output=logits_output,
|
logits_output=logits_output,
|
||||||
@@ -296,7 +308,7 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def check_forward_draft_extend_after_decode(self, verify_output: EagleVerifyOutput):
|
def check_forward_draft_extend_after_decode(self, verify_output: EagleVerifyOutput):
|
||||||
local_need_forward = verify_output.unfinished_accept_tokens.shape[0] > 0
|
local_need_forward = verify_output.draft_extend_input.input_ids.shape[0] > 0
|
||||||
if not self.server_args.enable_dp_attention:
|
if not self.server_args.enable_dp_attention:
|
||||||
return local_need_forward
|
return local_need_forward
|
||||||
|
|
||||||
@@ -471,7 +483,8 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
# allocator and kv cache pool are shared with target worker
|
# allocator and kv cache pool are shared with target worker
|
||||||
pass
|
pass
|
||||||
|
|
||||||
def verify(self, batch: ScheduleBatch, spec_info: EagleVerifyInput):
|
def verify(self, batch: ScheduleBatch):
|
||||||
|
spec_info: EagleVerifyInput = batch.spec_info
|
||||||
spec_info.prepare_for_verify(batch, self.page_size)
|
spec_info.prepare_for_verify(batch, self.page_size)
|
||||||
batch.return_hidden_states = False
|
batch.return_hidden_states = False
|
||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
@@ -479,7 +492,6 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
if not batch.forward_mode.is_idle()
|
if not batch.forward_mode.is_idle()
|
||||||
else ForwardMode.IDLE
|
else ForwardMode.IDLE
|
||||||
)
|
)
|
||||||
batch.spec_info = spec_info
|
|
||||||
|
|
||||||
model_worker_batch = batch.get_model_worker_batch(
|
model_worker_batch = batch.get_model_worker_batch(
|
||||||
seq_lens_cpu_cache=spec_info.seq_lens_cpu
|
seq_lens_cpu_cache=spec_info.seq_lens_cpu
|
||||||
@@ -589,7 +601,6 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE
|
ForwardMode.DECODE if not batch.forward_mode.is_idle() else ForwardMode.IDLE
|
||||||
)
|
)
|
||||||
batch.spec_info = res.next_draft_input
|
|
||||||
|
|
||||||
return logits_output, res, can_run_cuda_graph
|
return logits_output, res, can_run_cuda_graph
|
||||||
|
|
||||||
@@ -653,20 +664,19 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
forward_batch.spec_info.topk_index = torch.cat(topk_index_list, dim=1)
|
forward_batch.spec_info.topk_index = torch.cat(topk_index_list, dim=1)
|
||||||
|
|
||||||
def forward_draft_extend_after_decode(
|
def forward_draft_extend_after_decode(
|
||||||
self, batch: ScheduleBatch, verify_output: EagleVerifyOutput
|
self, batch: ScheduleBatch
|
||||||
):
|
) -> EagleDraftInput:
|
||||||
assert isinstance(batch.spec_info, EagleDraftInput)
|
draft_extend_input: EagleDraftExtendInput = batch.spec_info
|
||||||
|
|
||||||
# Backup fields that will be modified in-place
|
# Backup fields that will be modified in-place
|
||||||
seq_lens_backup = batch.seq_lens.clone()
|
seq_lens_backup = batch.seq_lens.clone()
|
||||||
seq_lens_cpu_backup = batch.seq_lens_cpu.clone()
|
seq_lens_cpu_backup = batch.seq_lens_cpu.clone()
|
||||||
req_pool_indices_backup = batch.req_pool_indices
|
req_pool_indices_backup = batch.req_pool_indices
|
||||||
num_accepted_drafts_backup = batch.spec_info.num_accepted_drafts
|
|
||||||
num_accepted_tokens_backup = batch.spec_info.num_accepted_tokens
|
|
||||||
return_logprob_backup = batch.return_logprob
|
return_logprob_backup = batch.return_logprob
|
||||||
|
|
||||||
input_is_idle = batch.forward_mode.is_idle()
|
input_is_idle = batch.forward_mode.is_idle()
|
||||||
|
|
||||||
if not input_is_idle and verify_output.unfinished_accept_tokens.numel() == 0:
|
if not input_is_idle and draft_extend_input.input_ids.numel() == 0:
|
||||||
batch = batch.copy()
|
batch = batch.copy()
|
||||||
batch.prepare_for_idle()
|
batch.prepare_for_idle()
|
||||||
hidden_size = (
|
hidden_size = (
|
||||||
@@ -674,19 +684,19 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
if self.speculative_algorithm.is_eagle3()
|
if self.speculative_algorithm.is_eagle3()
|
||||||
else self.model_config.hidden_size
|
else self.model_config.hidden_size
|
||||||
)
|
)
|
||||||
batch.spec_info = EagleDraftInput.create_idle_input(
|
draft_extend_input = EagleDraftExtendInput.create_idle_input(
|
||||||
device=self.device,
|
device=self.device,
|
||||||
hidden_size=hidden_size,
|
hidden_size=hidden_size,
|
||||||
dtype=self.model_config.dtype,
|
dtype=self.model_config.dtype,
|
||||||
topk=self.topk,
|
|
||||||
capture_hidden_mode=CaptureHiddenMode.LAST,
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
)
|
)
|
||||||
|
batch.spec_info = draft_extend_input
|
||||||
|
|
||||||
batch.spec_info.num_tokens_per_req = self.speculative_num_steps + 1
|
# Phase 1: prepare extend (kernel writes draft_extend_input.{positions, bonus_tokens})
|
||||||
batch.spec_info.num_tokens_for_logprob_per_req = 1
|
draft_extend_input.num_tokens_per_req = self.speculative_num_steps + 1
|
||||||
batch.spec_info.prepare_extend_after_decode(
|
draft_extend_input.num_tokens_for_logprob_per_req = 1
|
||||||
|
draft_extend_input.prepare_extend_after_decode(
|
||||||
batch,
|
batch,
|
||||||
verify_output=verify_output,
|
|
||||||
speculative_num_steps=self.speculative_num_steps,
|
speculative_num_steps=self.speculative_num_steps,
|
||||||
)
|
)
|
||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
@@ -748,17 +758,23 @@ class MultiLayerEagleWorker(TpModelWorker):
|
|||||||
)
|
)
|
||||||
pt += extend_len
|
pt += extend_len
|
||||||
|
|
||||||
forward_batch.spec_info.topk_p = torch.cat(topk_p_list, dim=1)
|
# Phase 3: assemble next-iter EagleDraftInput from extend output
|
||||||
forward_batch.spec_info.topk_index = torch.cat(topk_index_list, dim=1)
|
next_draft_input = EagleDraftInput(
|
||||||
|
bonus_tokens=draft_extend_input.bonus_tokens,
|
||||||
|
hidden_states=logits_output.hidden_states,
|
||||||
|
topk_p=torch.cat(topk_p_list, dim=1),
|
||||||
|
topk_index=torch.cat(topk_index_list, dim=1),
|
||||||
|
capture_hidden_mode=CaptureHiddenMode.FULL,
|
||||||
|
)
|
||||||
|
|
||||||
# Restore backup.
|
# Restore batch fields. `seq_lens` etc. were modified by
|
||||||
# This is because `seq_lens` can be modified in `prepare_extend_after_decode`
|
# `prepare_extend_after_decode`. Caller installs `next_draft_input` as
|
||||||
|
# `batch.spec_info`.
|
||||||
batch.forward_mode = (
|
batch.forward_mode = (
|
||||||
ForwardMode.DECODE if not input_is_idle else ForwardMode.IDLE
|
ForwardMode.DECODE if not input_is_idle else ForwardMode.IDLE
|
||||||
)
|
)
|
||||||
batch.seq_lens = seq_lens_backup
|
batch.seq_lens = seq_lens_backup
|
||||||
batch.seq_lens_cpu = seq_lens_cpu_backup
|
batch.seq_lens_cpu = seq_lens_cpu_backup
|
||||||
batch.req_pool_indices = req_pool_indices_backup
|
batch.req_pool_indices = req_pool_indices_backup
|
||||||
batch.spec_info.num_accepted_drafts = num_accepted_drafts_backup
|
|
||||||
batch.spec_info.num_accepted_tokens = num_accepted_tokens_backup
|
|
||||||
batch.return_logprob = return_logprob_backup
|
batch.return_logprob = return_logprob_backup
|
||||||
|
return next_draft_input
|
||||||
|
|||||||
@@ -195,8 +195,10 @@ class SpecInputType(IntEnum):
|
|||||||
# NOTE: introduce this to distinguish the SpecInput types of multiple algorithms when asserting in attention backends.
|
# NOTE: introduce this to distinguish the SpecInput types of multiple algorithms when asserting in attention backends.
|
||||||
# If all algorithms can share the same datastrucutre of draft_input and verify_input, consider simplify it
|
# If all algorithms can share the same datastrucutre of draft_input and verify_input, consider simplify it
|
||||||
EAGLE_DRAFT = auto()
|
EAGLE_DRAFT = auto()
|
||||||
|
EAGLE_DRAFT_EXTEND = auto()
|
||||||
EAGLE_VERIFY = auto()
|
EAGLE_VERIFY = auto()
|
||||||
FROZEN_KV_MTP_DRAFT = auto()
|
FROZEN_KV_MTP_DRAFT = auto()
|
||||||
|
FROZEN_KV_MTP_DRAFT_EXTEND = auto()
|
||||||
FROZEN_KV_MTP_VERIFY = auto()
|
FROZEN_KV_MTP_VERIFY = auto()
|
||||||
DFLASH_DRAFT = auto()
|
DFLASH_DRAFT = auto()
|
||||||
DFLASH_VERIFY = auto()
|
DFLASH_VERIFY = auto()
|
||||||
@@ -212,7 +214,9 @@ class SpecInput(ABC):
|
|||||||
# or use another variable name like `draft_input` to substitute `spec_info`
|
# or use another variable name like `draft_input` to substitute `spec_info`
|
||||||
return self.spec_input_type in {
|
return self.spec_input_type in {
|
||||||
SpecInputType.EAGLE_DRAFT,
|
SpecInputType.EAGLE_DRAFT,
|
||||||
|
SpecInputType.EAGLE_DRAFT_EXTEND,
|
||||||
SpecInputType.FROZEN_KV_MTP_DRAFT,
|
SpecInputType.FROZEN_KV_MTP_DRAFT,
|
||||||
|
SpecInputType.FROZEN_KV_MTP_DRAFT_EXTEND,
|
||||||
SpecInputType.DFLASH_DRAFT,
|
SpecInputType.DFLASH_DRAFT,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user