[Spec] Remove the dead spec V1 scheduler paths (#27977)

This commit is contained in:
Liangsheng Yin
2026-06-11 18:31:13 -07:00
committed by GitHub
parent 2e74ff192c
commit 3ffe72517f
15 changed files with 97 additions and 794 deletions
@@ -564,9 +564,9 @@ class DeepseekSparseAttnBackend(
page_table, repeats=self.speculative_num_draft_tokens, dim=0
)
else:
# DRAFT_EXTEND (v1): V1 worker extends by (num_correct_drafts + 1) per request
# after verification. Lengths vary per request based on how many tokens
# were accepted.
# DRAFT_EXTEND: the draft worker extends by (num_correct_drafts + 1)
# per request after verification. Lengths vary per request based on
# how many tokens were accepted.
page_table = torch.repeat_interleave(
page_table, repeats=forward_batch.extend_seq_lens, dim=0
)
+2 -116
View File
@@ -2,17 +2,14 @@ from __future__ import annotations
import dataclasses
from enum import Enum, auto
from typing import TYPE_CHECKING, List, Optional, Union
from typing import TYPE_CHECKING, List, Optional
import torch
from sglang.srt.environ import envs
if TYPE_CHECKING:
from sglang.srt.layers.logits_processor import LogitsMetadata, LogitsProcessorOutput
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.speculative.eagle_info import EagleVerifyOutput
from sglang.srt.speculative.ngram_info import NgramVerifyInput
from sglang.srt.layers.logits_processor import LogitsMetadata
class LogprobStage(Enum):
@@ -295,117 +292,6 @@ def get_token_ids_logprobs_chunk(
return next_split_pruned_len
def add_output_logprobs_for_spec_v1(
batch: ScheduleBatch,
res: Union[EagleVerifyOutput, NgramVerifyInput],
logits_output: Optional[LogitsProcessorOutput] = None,
):
# Extract args
if logits_output is None:
logits_output = res.logits_output
if hasattr(res, "num_correct_drafts_per_req_cpu"):
num_correct_drafts_per_req_cpu = res.num_correct_drafts_per_req_cpu
else:
# FIXME: Get a NgramVerifyOutput class and use that instead of this hack.
num_correct_drafts_per_req_cpu = res.num_correct_drafts.tolist()
top_logprobs_nums = batch.top_logprobs_nums
token_ids_logprobs = batch.token_ids_logprobs
accept_indices = res.accept_indices
assert len(accept_indices) == len(logits_output.next_token_logits)
temperatures = batch.sampling_info.temperatures
num_draft_tokens = batch.spec_info.draft_token_num
# acceptance indices are the indices in a "flattened" batch.
# dividing it to num_draft_tokens will yield the actual batch index.
temperatures = temperatures[accept_indices // num_draft_tokens]
if envs.SGLANG_RETURN_ORIGINAL_LOGPROB.get():
logprobs = torch.nn.functional.log_softmax(
logits_output.next_token_logits, dim=-1
)
else:
logprobs = torch.nn.functional.log_softmax(
logits_output.next_token_logits / temperatures, dim=-1
)
batch_next_token_ids = res.accept_tokens
num_tokens_per_req = [accept + 1 for accept in num_correct_drafts_per_req_cpu]
# We should repeat top_logprobs_nums to match num_tokens_per_req.
top_logprobs_nums_repeat_interleaved = [
num
for num, num_tokens in zip(top_logprobs_nums, num_tokens_per_req)
for _ in range(num_tokens)
]
token_ids_logprobs_repeat_interleaved = [
token_ids
for token_ids, num_tokens in zip(token_ids_logprobs, num_tokens_per_req)
for _ in range(num_tokens)
]
# Extract logprobs
should_top_logprobs = any(x > 0 for x in top_logprobs_nums)
should_token_ids_logprobs = any(x is not None for x in token_ids_logprobs)
if should_top_logprobs:
(
logits_output.next_token_top_logprobs_val,
logits_output.next_token_top_logprobs_idx,
) = get_top_logprobs(
logprobs,
top_logprobs_nums_repeat_interleaved,
)
if should_token_ids_logprobs:
(
logits_output.next_token_token_ids_logprobs_val,
logits_output.next_token_token_ids_logprobs_idx,
) = get_token_ids_logprobs(
logprobs,
token_ids_logprobs_repeat_interleaved,
)
logits_output.next_token_logprobs = logprobs[
torch.arange(len(batch_next_token_ids), device=batch.sampling_info.device),
batch_next_token_ids,
]
# Add output logprobs to the request
pt = 0
next_token_logprobs = logits_output.next_token_logprobs.tolist()
accept_tokens_list = batch_next_token_ids.tolist()
token_top_logprobs_val = logits_output.next_token_top_logprobs_val
token_top_logprobs_idx = logits_output.next_token_top_logprobs_idx
token_ids_logprobs_val = logits_output.next_token_token_ids_logprobs_val
token_ids_logprobs_idx = logits_output.next_token_token_ids_logprobs_idx
for req, num_tokens in zip(batch.reqs, num_tokens_per_req, strict=True):
for _ in range(num_tokens):
if req.return_logprob:
req.logprob.output_token_logprobs_val.append(next_token_logprobs[pt])
req.logprob.output_token_logprobs_idx.append(accept_tokens_list[pt])
if req.logprob.top_logprobs_num > 0:
assert (
should_top_logprobs
), "Inconsistent state: should_top_logprobs is False"
req.logprob.output_top_logprobs_val.append(
token_top_logprobs_val[pt]
)
req.logprob.output_top_logprobs_idx.append(
token_top_logprobs_idx[pt]
)
if (
req.logprob.token_ids_logprob is not None
and should_token_ids_logprobs
):
req.logprob.output_token_ids_logprobs_val.append(
token_ids_logprobs_val[pt]
)
req.logprob.output_token_ids_logprobs_idx.append(
token_ids_logprobs_idx[pt]
)
pt += 1
def compute_spec_v2_logprobs(
batch,
logits_output,
+3 -3
View File
@@ -95,7 +95,7 @@ def resolve_forward_inputs(batch: ScheduleBatch, future_map: FutureMap) -> None:
# Only the overlap path relays spec extras through the future_map; the
# synchronous (non-overlap) V2 path installs next_draft_input directly.
if batch.enable_overlap and batch.is_spec_v2:
if batch.enable_overlap and not batch.spec_algorithm.is_none():
future_map._resolve_spec_extras(batch)
@@ -326,8 +326,8 @@ class FutureMap:
if indices.shape[0] == 0:
# DP idle: payload is empty stub; lazy-init shape peek would IndexError.
return
# Dispatch by payload type, not spec_algo: spec_v1 (non-overlap spec)
# also passes a token Tensor here.
# Dispatch by payload type, not spec_algo: non-spec decode passes a
# token Tensor here.
# FIXME(lsyin): unify this relay path with a dataclass instead of the
# Tensor / EagleDraftInput type switch.
if isinstance(payload, torch.Tensor):
+5 -39
View File
@@ -2355,25 +2355,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
new_pages = sum(1 for r in requests if r.kv_committed_len % page_size == 0)
return new_pages * page_size
if self.is_spec_v2:
return self._new_tokens_required_next_decode_spec_v2(requests, page_size)
server_args = get_global_server_args()
len_per_topk = server_args.speculative_num_steps or 1
spec_topk = server_args.speculative_eagle_topk or 1
spec_tokens = server_args.speculative_num_draft_tokens
if page_size > 1 and spec_topk > 1:
# last partial page and ceil alignment
len_per_topk = ceil_align(len_per_topk + page_size, page_size)
spec_tokens = ceil_align(spec_tokens, page_size)
elif page_size > 1:
# only page alignment
len_per_topk = ceil_align(len_per_topk, page_size)
spec_tokens = ceil_align(spec_tokens, page_size)
num_tokens = max(len_per_topk * spec_topk, spec_tokens) * len(requests)
return num_tokens
return self._new_tokens_required_next_decode_spec_v2(requests, page_size)
def _new_tokens_required_next_decode_spec_v2(self, requests, page_size):
"""Tight estimate matching eagle_info_v2.prepare_for_decode allocation."""
@@ -2498,12 +2480,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.model_config.vocab_size,
)
@property
def is_spec_v2(self):
# Whether the V2 worker/schema is used. Independent of overlap: the
# non-overlap path also drives the V2 worker, just synchronously.
return self.spec_algorithm.supports_spec_v2()
def mamba_lazy_prealloc_at_boundary(self, mamba_track_interval: int):
"""Allocate a temporary second ping-pong slot for reqs at a track boundary.
@@ -2547,14 +2523,11 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
if hasattr(self, "attn_cp_metadata") and self.attn_cp_metadata is not None:
self.attn_cp_metadata = None
if self.is_spec_v2:
# TODO(spec-v2): all spec v2 should go through this path
if not self.spec_algorithm.is_none():
# Spec decoding: the draft input owns decode preparation
# (allocation, pre-claim, seq-lens bookkeeping).
draft_input: EagleDraftInput = self.spec_info
draft_input.prepare_for_decode(self)
if not self.spec_algorithm.is_none():
# if spec decoding is used, the decode batch is prepared inside
# `forward_batch_speculative_generation` after running draft models.
return
if self.sampling_info.penalizer_orchestrator.is_required:
@@ -2638,8 +2611,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self,
chunked_req_to_exclude: Optional[Union[Req, List[Req]]] = None,
keep_indices: Optional[List[int]] = None,
# FIXME(lsyin): deprecate this API after spec v1 is deprecated
v1_spec_info_filtered: Optional[bool] = False,
):
if keep_indices is None:
if isinstance(chunked_req_to_exclude, Req):
@@ -2706,15 +2677,10 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
self.has_grammar = any(req.grammar for req in self.reqs)
self.sampling_info.filter_batch(keep_indices, keep_indices_device)
# NOTE: spec_info filtered before batch filtering only happens in:
# - Spec v1's verify phase
# - Only for decode batch (running_batch)
has_been_filtered = v1_spec_info_filtered and not self.is_spec_v2
if self.spec_info:
self.spec_info.filter_batch(
new_indices=keep_indices_device,
has_been_filtered=has_been_filtered,
has_been_filtered=False,
)
def merge_batch(self, other: ScheduleBatch):
+17 -24
View File
@@ -1377,8 +1377,7 @@ class Scheduler(
)
def _abort_on_running_timeout(self):
# NOTE: this should be called before a batch is launched,
# as current spec-v1 still filters batch inside verify stage.
# NOTE: this should be called before a batch is launched.
timeout_s = envs.SGLANG_REQ_RUNNING_TIMEOUT.get()
if timeout_s <= 0:
return
@@ -1541,7 +1540,7 @@ class Scheduler(
# TODO(lsyin): support overlap + spec + grammar
need_grammar_sync = (
batch
and batch.is_spec_v2
and not batch.spec_algorithm.is_none()
and batch.has_grammar
and batch.forward_mode.is_decode()
and len(self.result_queue) > 0
@@ -2822,7 +2821,7 @@ class Scheduler(
and new_batch.input_embeds is None
):
# TODO (lianmin): support return_logprob + mixed chunked prefill
self.running_batch.filter_batch(v1_spec_info_filtered=True)
self.running_batch.filter_batch()
if not self.running_batch.is_empty():
self.running_batch.prepare_for_decode()
new_batch.mix_with_running(self.running_batch)
@@ -2867,7 +2866,7 @@ class Scheduler(
"""Update the current running decoding batch."""
initial_bs = batch.batch_size()
batch.filter_batch(v1_spec_info_filtered=True)
batch.filter_batch()
if batch.is_empty():
batch.batch_is_full = False
return batch
@@ -2987,7 +2986,7 @@ class Scheduler(
passes overlap=False.
"""
# 1. snapshot
snapshot_v2_full = batch.is_spec_v2
snapshot_v2_full = not batch.spec_algorithm.is_none()
sched_snapshot = (
{f.name: getattr(batch, f.name) for f in dataclasses.fields(batch)}
if snapshot_v2_full
@@ -3061,7 +3060,7 @@ class Scheduler(
self.future_map.publish, future_indices
)
}
if batch.is_spec_v2
if not batch.spec_algorithm.is_none()
else {}
)
@@ -3069,7 +3068,7 @@ class Scheduler(
batch_result = self.model_worker.forward_batch_generation(
batch, **fwd_kwargs
)
if not batch.is_spec_v2:
if batch.spec_algorithm.is_none():
self.future_map.publish(future_indices, batch.seq_lens + 1)
# Park any refs the worker wants kept alive 2 iters
# (cross-stream tensor lifetime; pinned in the same
@@ -3083,7 +3082,7 @@ class Scheduler(
if batch_result.delay_sample_func is None:
stash_payload = (
batch_result.next_draft_input
if batch.is_spec_v2
if not batch.spec_algorithm.is_none()
else batch_result.next_token_ids
)
self.future_map.stash(future_indices, stash_payload)
@@ -3097,7 +3096,7 @@ class Scheduler(
# Next-iter input_ids relayed via future_map.
batch.input_ids = None
if batch.is_spec_v2:
if not batch.spec_algorithm.is_none():
batch.spec_info = batch_result.next_draft_input
batch.spec_info.future_indices = future_indices
elif self.enable_pdmux and batch.forward_mode.is_split_prefill():
@@ -3108,8 +3107,8 @@ class Scheduler(
batch.req_pool_indices, batch_result.next_token_ids
)
batch.input_ids = None
elif batch.is_spec_v2:
# Non-overlap V2: drive the V2 worker synchronously (no
elif not batch.spec_algorithm.is_none():
# Non-overlap: drive the V2 worker synchronously (no
# future_map relay / on_publish).
resolve_forward_inputs(batch, self.future_map)
with self._forward_isolation(batch, overlap=False):
@@ -3141,17 +3140,11 @@ class Scheduler(
batch, **kwargs
)
if isinstance(batch_result.next_token_ids, torch.Tensor):
if self.spec_algorithm.is_none():
# Non-spec: relay via future_map, gathered next iter.
self.future_map.stash(
batch.req_pool_indices, batch_result.next_token_ids
)
batch.input_ids = None
else:
# Spec_v1 (NGRAM / DFLASH, non-overlap): worker shape
# doesn't match req_pool_indices; relay is unused (worker
# rebuilds input_ids inside verify).
batch.input_ids = batch_result.next_token_ids.to(torch.int64)
# Non-spec: relay via future_map, gathered next iter.
self.future_map.stash(
batch.req_pool_indices, batch_result.next_token_ids
)
batch.input_ids = None
self.update_cache_from_scheduler(batch, batch_result)
# These 2 values are needed for processing the output, but the values can be
@@ -3796,7 +3789,7 @@ class Scheduler(
self.cur_batch = None
if recv_req.mode == "retract" and not self.running_batch.is_empty():
self.running_batch.filter_batch(v1_spec_info_filtered=True)
self.running_batch.filter_batch()
if len(self.running_batch.reqs) != 0:
retracted_reqs = self.running_batch.retract_all(self.server_args)
for req in retracted_reqs:
@@ -631,10 +631,6 @@ class SchedulerBatchResultProcessor:
self.token_to_kv_pool_allocator.free_group_begin()
# Spec V1 handles output_ids, update_finish_state, grammar, and reasoning tokens
# in the verify phase. Non-spec and V2 handle them here in post-processing.
is_spec_v1 = not batch.spec_algorithm.is_none() and not batch.is_spec_v2
for i, req in enumerate(batch.reqs):
req: Req
@@ -645,19 +641,6 @@ class SchedulerBatchResultProcessor:
# And all the over-allocated tokens will be freed in `release_kv_cache`.
continue
if is_spec_v1:
req.time_stats.set_last_decode_finish_time()
self._handle_finish_state_updated_req(
req, batch, result, i, logits_output
)
if req.return_hidden_states and logits_output.hidden_states is not None:
req.hidden_states.append(
logits_output.hidden_states[i].cpu().clone().tolist()
)
if req.grammar is not None:
req.grammar.finished = req.finished()
continue
# Non-spec and V2: full post-processing
next_token_id = next_token_ids[i]
new_accepted_len = 1
@@ -715,31 +698,27 @@ class SchedulerBatchResultProcessor:
next_token_ids: Union[torch.Tensor, List[int]],
) -> Tuple[Union[List[int], List[List[int]]], Optional[List[float]]]:
next_token_logprobs = None
if batch.spec_algorithm.is_none() or batch.is_spec_v2:
if batch.is_spec_v2:
next_token_ids = self._resolve_spec_v2_tokens(result, batch)
elif isinstance(next_token_ids, list):
pass # MLX path: already a list[int], skip torch round-trip
else:
next_token_ids = next_token_ids.tolist()
if not batch.spec_algorithm.is_none():
next_token_ids = self._resolve_spec_v2_tokens(result, batch)
elif isinstance(next_token_ids, list):
pass # MLX path: already a list[int], skip torch round-trip
else:
next_token_ids = next_token_ids.tolist()
if batch.return_logprob:
next_token_logprobs = logits_output.next_token_logprobs.tolist()
if logits_output.next_token_top_logprobs_val:
logits_output.next_token_top_logprobs_val = [
v.tolist() for v in logits_output.next_token_top_logprobs_val
]
logits_output.next_token_top_logprobs_idx = [
x.tolist() for x in logits_output.next_token_top_logprobs_idx
]
if batch.return_logprob:
next_token_logprobs = logits_output.next_token_logprobs.tolist()
if logits_output.next_token_top_logprobs_val:
logits_output.next_token_top_logprobs_val = [
v.tolist() for v in logits_output.next_token_top_logprobs_val
]
logits_output.next_token_top_logprobs_idx = [
x.tolist() for x in logits_output.next_token_top_logprobs_idx
]
if logits_output.next_token_token_ids_logprobs_val:
logits_output.next_token_token_ids_logprobs_val = [
v.tolist()
for v in logits_output.next_token_token_ids_logprobs_val
]
# else: Spec V1 — output_ids, update_finish_state, grammar, and reasoning tokens
# are already handled in the verify phase (eagle_info.py / ngram_info.py).
if logits_output.next_token_token_ids_logprobs_val:
logits_output.next_token_token_ids_logprobs_val = [
v.tolist() for v in logits_output.next_token_token_ids_logprobs_val
]
return next_token_ids, next_token_logprobs
def _apply_decode_logprobs(
@@ -752,9 +731,8 @@ class SchedulerBatchResultProcessor:
next_token_logprobs: list,
logits_output: LogitsProcessorOutput,
) -> None:
# Spec v1 handles logprobs inside its own worker.
# Normalize: non-spec has 1 token, spec v2 has multiple.
if batch.is_spec_v2:
# Normalize: non-spec has 1 token, spec decoding has multiple.
if not batch.spec_algorithm.is_none():
accepted_logprobs = next_token_logprobs[i]
accepted_ids = next_token_id
max_accept = len(accepted_logprobs)
@@ -795,7 +773,7 @@ class SchedulerBatchResultProcessor:
if batch.spec_algorithm.is_none():
# Normal decode: single token
req.grammar.accept_token(next_token_id)
elif batch.is_spec_v2:
else:
# Speculative decode: next_token_id is a list of accepted tokens
for token_id in next_token_id:
req.grammar.accept_token(token_id)
+2 -4
View File
@@ -159,10 +159,8 @@ def get_alloc_len_per_decode(server_args: Optional[ServerArgs] = None) -> int:
if server_args.speculative_algorithm is None:
return 1
# Spec v1:
# 1) alloc topk * num_steps when draft decoding and then restore the allocation
# 2) alloc num_draft_tokens when verifying the drafts
# Sepc v2: allocate max(topk * num_steps, num_draft_tokens)
# Spec decoding allocates max(topk * num_steps, num_draft_tokens) per
# decode step (draft chain and verify block share the reservation).
spec_steps = server_args.speculative_num_steps or 1
spec_topk = server_args.speculative_eagle_topk or 1
@@ -54,19 +54,12 @@ def get_draft_kv_pool(
if draft_worker is None or spec_algorithm.is_ngram():
return None, None
# V2 (EAGLE family) nests the runner under `.draft_worker`; DFLASH /
# FROZEN_KV_MTP expose `.model_runner` directly.
if spec_algorithm.supports_spec_v2():
if server_args.enable_multi_layer_eagle:
draft_runner = draft_worker.draft_worker.draft_runner_list[0]
else:
draft_runner = draft_worker.draft_worker.draft_runner
return draft_runner.token_to_kv_pool, draft_runner.model_config
return (
draft_worker.model_runner.token_to_kv_pool,
draft_worker.model_config,
)
# V2 workers nest the draft runner under `.draft_worker`.
if server_args.enable_multi_layer_eagle:
draft_runner = draft_worker.draft_worker.draft_runner_list[0]
else:
draft_runner = draft_worker.draft_worker.draft_runner
return draft_runner.token_to_kv_pool, draft_runner.model_config
def maybe_register_hicache_draft(
+2 -525
View File
@@ -1,23 +1,13 @@
import copy
import logging
from dataclasses import dataclass
from typing import List, Optional, Tuple
import torch
import torch.nn.functional as F
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
from sglang.srt.distributed import get_tp_group
from sglang.srt.environ import envs
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
from sglang.srt.layers.dp_attention import (
get_attention_tp_group,
is_dp_attention_enabled,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.layers.sampler import apply_custom_logit_processor
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.common import (
alloc_paged_token_slots_extend,
alloc_token_slots,
@@ -29,29 +19,13 @@ from sglang.srt.speculative.eagle_info_v2 import (
EagleDraftInputV2Mixin,
EagleVerifyInputV2Mixin,
)
from sglang.srt.speculative.eagle_utils import verify_tree_greedy_func
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
from sglang.srt.speculative.spec_utils import (
SIMULATE_ACC_LEN,
TREE_SPEC_KERNEL_AVAILABLE,
align_evict_mask_to_page_size,
assign_req_to_token_pool_func,
create_extend_after_decode_spec_info,
create_num_accept_tokens_filter,
filter_finished_cache_loc_kernel,
generate_simulated_accept_index,
get_src_tgt_cache_loc,
get_target_cache_loc,
)
from sglang.srt.utils import is_cuda, is_musa, next_power_of_2
from sglang.srt.utils.async_probe import maybe_detect_nan, maybe_detect_oob
if is_cuda() or is_musa():
from sgl_kernel import (
top_k_renorm_prob,
top_p_renorm_prob,
tree_speculative_sampling_target_only,
)
from sglang.srt.utils import next_power_of_2
from sglang.srt.utils.async_probe import maybe_detect_oob
logger = logging.getLogger(__name__)
@@ -239,464 +213,6 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
return kv_indices, cum_kv_seq_len, qo_indptr, self.custom_mask
def verify(
self,
batch: ScheduleBatch,
logits_output: LogitsProcessorOutput,
token_to_kv_pool_allocator: BaseTokenToKVPoolAllocator,
page_size: int,
vocab_mask: Optional[torch.Tensor] = None, # For grammar
) -> torch.Tensor:
"""
Verify and find accepted tokens based on logits output and batch
(which contains spec decoding information).
WARNING: This API in-place modifies the states of logits_output
This API updates values inside logits_output based on the accepted
tokens. I.e., logits_output.next_token_logits only contains
accepted token logits.
"""
if batch.forward_mode.is_idle():
# hidden_size=None: worker fixup in forward_draft_extend_after_decode
# rebuilds via EagleDraftExtendInput.hidden_size_for(worker)
# (single source incl. EAGLE-3 aux widening).
draft_extend_input = EagleDraftExtendInput.create_idle_input(
device=batch.device,
hidden_size=None,
dtype=None,
capture_hidden_mode=CaptureHiddenMode.LAST,
)
return EagleVerifyOutput.create_idle(
draft_extend_input=draft_extend_input,
logits_output=logits_output,
device=batch.device,
spec_steps=self.spec_steps,
)
bs = self.retrieve_index.shape[0]
candidates = self.draft_token.reshape(bs, self.draft_token_num)
sampling_info = batch.sampling_info
predict_shape = list(logits_output.next_token_logits.shape)[:-1]
predict_shape[-1] += 1
predict = torch.empty(predict_shape, dtype=torch.int32, device=batch.device)
accept_index = torch.full(
(bs, self.spec_steps + 1), -1, dtype=torch.int32, device=batch.device
)
num_correct_drafts = torch.empty((bs,), dtype=torch.int32, device=batch.device)
if bs != len(sampling_info):
sampling_info = copy.deepcopy(sampling_info)
# NOTE: retrieve_index are the indices of the requests that are kept.
sampling_info.filter_batch(
self.retrieve_index.tolist(), self.retrieve_index
)
# Apply the custom logit processors if registered in the sampling info.
if sampling_info.has_custom_logit_processor:
apply_custom_logit_processor(
logits_output.next_token_logits,
sampling_info,
num_tokens_in_batch=self.draft_token_num,
)
# Apply penalty
if (
sampling_info.penalizer_orchestrator.is_required
or sampling_info.logit_bias is not None
):
# This is a relaxed version of penalties for speculative decoding.
sampling_info.penalizer_orchestrator.apply(
logits_output.next_token_logits, repeat=self.draft_token_num
)
if sampling_info.logit_bias is not None:
logits_output.next_token_logits.add_(
torch.repeat_interleave(
sampling_info.logit_bias, self.draft_token_num, dim=0
)
)
# Apply grammar mask
if vocab_mask is not None:
assert self.grammar is not None
self.grammar.apply_vocab_mask(
logits=logits_output.next_token_logits, vocab_mask=vocab_mask
)
# Sample tokens. Force greedy sampling on AMD
is_all_greedy = sampling_info.is_all_greedy
if (not is_all_greedy) and (not TREE_SPEC_KERNEL_AVAILABLE):
logger.warning(
"Tree speculative sampling kernel unavailable (likely AMD/HIP build). "
"Falling back to greedy verification."
)
if is_all_greedy or not TREE_SPEC_KERNEL_AVAILABLE:
target_predict = torch.argmax(logits_output.next_token_logits, dim=-1)
target_predict = target_predict.reshape(bs, self.draft_token_num)
predict, accept_index, num_correct_drafts = verify_tree_greedy_func(
predicts=predict, # mutable
accept_index=accept_index, # mutable
accept_token_num=num_correct_drafts, # mutable
candidates=candidates,
retrieve_index=self.retrieve_index,
retrieve_next_token=self.retrieve_next_token,
retrieve_next_sibling=self.retrieve_next_sibling,
target_predict=target_predict,
topk=self.topk,
)
else:
# apply temperature and get target probs
expanded_temperature = torch.repeat_interleave(
sampling_info.temperatures, self.draft_token_num, dim=0
) # (bs * draft_token_num, 1)
target_probs = F.softmax(
logits_output.next_token_logits / expanded_temperature, dim=-1
) # (bs * draft_token_num, vocab_size)
maybe_detect_nan(target_probs, "verify: target_probs after softmax")
target_probs = top_k_renorm_prob(
target_probs,
torch.repeat_interleave(
sampling_info.top_ks, self.draft_token_num, dim=0
),
) # (bs * draft_token_num, vocab_size)
maybe_detect_nan(target_probs, "verify: target_probs after top_k_renorm")
if sampling_info.need_top_p_sampling:
target_probs = top_p_renorm_prob(
target_probs,
torch.repeat_interleave(
sampling_info.top_ps, self.draft_token_num, dim=0
),
)
maybe_detect_nan(
target_probs, "verify: target_probs after top_p_renorm"
)
target_probs = target_probs.reshape(bs, self.draft_token_num, -1)
draft_probs = torch.zeros(
target_probs.shape, dtype=torch.float32, device=batch.device
)
# coins for rejection sampling
coins = torch.rand_like(
candidates, dtype=torch.float32, device=batch.device
)
# coins for final sampling
coins_for_final_sampling = torch.rand(
(bs,), dtype=torch.float32, device=batch.device
)
tree_speculative_sampling_target_only(
predicts=predict, # mutable
accept_index=accept_index, # mutable
accept_token_num=num_correct_drafts, # mutable
candidates=candidates,
# kwarg LHS retained as `retrive_*` to match sgl_kernel op schema.
retrive_index=self.retrieve_index,
retrive_next_token=self.retrieve_next_token,
retrive_next_sibling=self.retrieve_next_sibling,
uniform_samples=coins,
uniform_samples_for_final_sampling=coins_for_final_sampling,
target_probs=target_probs,
draft_probs=draft_probs,
threshold_single=get_global_server_args().speculative_accept_threshold_single,
threshold_acc=get_global_server_args().speculative_accept_threshold_acc,
deterministic=True,
)
# Sync sampling results across TP ranks: different GPUs may
# produce slightly different target_probs due to floating-point
# non-determinism in softmax/top_k/top_p, causing different
# sampled tokens. Broadcast from rank 0 to ensure consistency.
tp_group = (
get_attention_tp_group()
if is_dp_attention_enabled()
else get_tp_group()
)
if tp_group.world_size > 1:
tp_group.broadcast(predict, src=0)
tp_group.broadcast(accept_index, src=0)
tp_group.broadcast(num_correct_drafts, src=0)
if SIMULATE_ACC_LEN > 0.0:
# Do simulation
accept_index = generate_simulated_accept_index(
accept_index=accept_index,
predict=predict, # mutable
num_correct_drafts=num_correct_drafts, # mutable
bs=bs,
spec_steps=self.spec_steps,
)
# accept_index values index batch.out_cache_loc (size = bs * draft_token_num);
# -1 is the reject sentinel.
maybe_detect_oob(
accept_index,
-1,
bs * self.draft_token_num,
"eagle verify accept_index post-sampling",
)
maybe_detect_oob(
num_correct_drafts,
0,
self.draft_token_num + 1,
"eagle verify num_correct_drafts post-sampling",
)
unfinished_index = []
unfinished_accept_index = []
accept_index_cpu = accept_index.tolist()
predict_cpu = predict.tolist()
has_finished = False
think_end_id = batch.model_config.think_end_id
# Iterate every accepted token and check if req has finished after append the token
# should be checked BEFORE free kv cache slots
for i, (req, accept_index_row) in enumerate(zip(batch.reqs, accept_index_cpu)):
num_accept_tokens = 0
for j, idx in enumerate(accept_index_row):
if idx == -1:
break
num_accept_tokens += 1
id = predict_cpu[idx]
req.output_ids.append(id)
if req.require_reasoning and think_end_id is not None:
req.update_reasoning_tokens(id, think_end_id)
req.update_finish_state()
if not req.finished() and req.grammar is not None:
try:
req.grammar.accept_token(id)
except ValueError as e:
logger.info(
f"{i=}, {req=}\n" f"{accept_index=}\n" f"{predict=}\n"
)
raise e
req.update_finish_state()
if req.finished():
has_finished = True
# set all tokens after finished token to -1 and break
accept_index[i, j + 1 :] = -1
break
# Update KV cache tracking for the accepted tokens
req.kv_committed_len += num_accept_tokens
req.kv_allocated_len = req.kv_committed_len
if not req.finished():
unfinished_index.append(i)
if idx == -1:
unfinished_accept_index.append(accept_index[i, :j])
else:
unfinished_accept_index.append(accept_index[i])
req.spec_verify_ct += 1
num_correct_drafts_this_req = (
sum(1 for idx in accept_index_row if idx != -1) - 1
)
req.spec_num_correct_drafts += num_correct_drafts_this_req
req.update_spec_correct_drafts_histogram(num_correct_drafts_this_req)
if has_finished:
num_correct_drafts = (accept_index != -1).sum(dim=1) - 1
# Free the KV cache for unaccepted tokens
# TODO: fuse them
accept_index = accept_index[accept_index != -1]
accept_tokens = predict[accept_index]
maybe_detect_oob(
accept_tokens,
0,
batch.model_config.vocab_size,
"eagle verify accept_tokens",
)
evict_mask = torch.full_like(self.draft_token, True, dtype=torch.bool)
evict_mask[accept_index] = False
num_correct_drafts_cpu = num_correct_drafts.cpu()
num_accept_tokens_cpu = num_correct_drafts_cpu + 1
# FIXME: this `tolist()` fixes the numerical calculation consistency
# try to unify the tensor representation and list representation
num_correct_drafts_list = num_correct_drafts_cpu.tolist()
num_accept_tokens_list = num_accept_tokens_cpu.tolist()
if page_size == 1:
# TODO: boolean array index leads to a device sync. Remove it.
token_to_kv_pool_allocator.free(batch.out_cache_loc[evict_mask])
else:
if self.topk == 1:
# Only evict full empty page. Do not evict partial empty page
align_evict_mask_to_page_size[len(batch.seq_lens),](
batch.seq_lens,
evict_mask,
page_size,
self.draft_token_num,
next_power_of_2(self.draft_token_num),
)
token_to_kv_pool_allocator.free(batch.out_cache_loc[evict_mask])
else:
# Shift the accepted tokens to the beginning.
# Only evict the last part
src_cache_loc, tgt_cache_loc, to_free_num_slots = get_src_tgt_cache_loc(
batch.seq_lens,
batch.out_cache_loc,
accept_index,
num_correct_drafts,
self.draft_token_num,
page_size,
)
to_free_slots = torch.empty(
(to_free_num_slots.sum().item(),),
dtype=torch.int64,
device=to_free_num_slots.device,
)
# out_cache_loc: [0 1 2, 3 4 5, 6 7 8]
# accept_index: [0 -1 2, 3 4 -1, 6 -1 -1]
# tgt_cache_loc: [0 1 , 3 4 , 6 ]
# to_free_slots: [ 2, 5, 7 8]
# to_free_slots also needs to be page-aligned without the first partial page
#
# split each row of out_cache_loc into two parts.
# 1. the first part goes to tgt_cache_loc. length = num_correct_drafts[i] + 1
# 2. the second part goes to to_free_slots.
get_target_cache_loc[(bs,)](
tgt_cache_loc,
to_free_slots,
num_correct_drafts,
to_free_num_slots,
batch.out_cache_loc,
self.draft_token_num,
next_power_of_2(self.draft_token_num),
next_power_of_2(bs),
)
# Free the kv cache
token_to_kv_pool_allocator.free(to_free_slots)
# Copy the kv cache
batch.token_to_kv_pool_allocator.get_kvcache().move_kv_cache(
tgt_cache_loc, src_cache_loc
)
# Construct EagleVerifyOutput
if not has_finished:
if page_size == 1 or self.topk == 1:
batch.out_cache_loc = batch.out_cache_loc[accept_index]
assign_req_to_token_pool_func(
batch.req_pool_indices,
batch.req_to_token_pool.req_to_token,
batch.seq_lens,
batch.seq_lens + num_correct_drafts + 1,
batch.out_cache_loc,
bs,
)
else:
batch.out_cache_loc = tgt_cache_loc
batch.seq_lens.add_(num_correct_drafts + 1)
batch.seq_lens_cpu.add_(num_accept_tokens_cpu)
draft_extend_input = EagleDraftExtendInput(
hidden_states=(
batch.spec_info.hidden_states[accept_index]
if batch.spec_info.hidden_states is not None
else None
),
num_correct_drafts=num_correct_drafts,
num_accept_tokens=num_correct_drafts + 1,
num_accept_tokens_cpu=num_accept_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(
draft_extend_input=draft_extend_input,
logits_output=logits_output,
accept_tokens=accept_tokens,
num_correct_drafts_per_req_cpu=num_correct_drafts_list,
accept_indices=accept_index,
)
else:
if page_size == 1 or self.topk == 1:
assign_req_to_token_pool_func(
batch.req_pool_indices,
batch.req_to_token_pool.req_to_token,
batch.seq_lens,
batch.seq_lens + num_correct_drafts + 1,
batch.out_cache_loc[accept_index],
bs,
)
batch.seq_lens.add_(num_correct_drafts + 1)
batch.seq_lens_cpu.add_(num_accept_tokens_cpu)
if len(unfinished_accept_index) > 0:
unfinished_accept_index = torch.cat(unfinished_accept_index)
unfinished_index_device = torch.tensor(
unfinished_index, dtype=torch.int64, device=predict.device
)
draft_input_num_correct_drafts_cpu = [
num_correct_drafts_list[i] for i in unfinished_index
]
draft_input_num_accept_tokens_cpu = [
num_accept_tokens_list[i] for i in unfinished_index
]
if page_size == 1 or self.topk == 1:
batch.out_cache_loc = batch.out_cache_loc[unfinished_accept_index]
else:
batch.out_cache_loc = torch.empty(
len(unfinished_index) + sum(draft_input_num_correct_drafts_cpu),
dtype=torch.int64,
device=predict.device,
)
num_accept_tokens_filter = create_num_accept_tokens_filter(
num_correct_drafts,
unfinished_index_device,
batch.seq_lens,
)
batch.seq_lens_cpu.add_(num_accept_tokens_cpu)
filter_finished_cache_loc_kernel[(bs,)](
batch.out_cache_loc,
tgt_cache_loc,
num_correct_drafts,
num_accept_tokens_filter,
next_power_of_2(bs),
next_power_of_2(self.draft_token_num),
)
unfinished_num_correct_drafts = num_correct_drafts[
unfinished_index_device
]
draft_extend_input = EagleDraftExtendInput(
hidden_states=(
batch.spec_info.hidden_states[unfinished_accept_index]
if batch.spec_info.hidden_states is not None
else None
),
num_accept_tokens_cpu=draft_input_num_accept_tokens_cpu,
num_correct_drafts=unfinished_num_correct_drafts,
num_accept_tokens=unfinished_num_correct_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:
# hidden_size=None: worker fixup rebuilds via
# EagleDraftExtendInput.hidden_size_for(worker) (single source).
draft_extend_input = EagleDraftExtendInput.create_idle_input(
device=batch.device,
hidden_size=None,
dtype=None,
capture_hidden_mode=CaptureHiddenMode.LAST,
)
return EagleVerifyOutput(
draft_extend_input=draft_extend_input,
logits_output=logits_output,
accept_tokens=accept_tokens,
num_correct_drafts_per_req_cpu=num_correct_drafts_list,
accept_indices=accept_index,
)
@dataclass
class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
@@ -1007,42 +523,3 @@ class EagleDraftExtendInput(SpecInput):
req_to_token.size(1),
)
return kv_indices, cum_kv_seq_len, qo_indptr, None
@dataclass
class EagleVerifyOutput:
# Next iter's draft-extend input, installed as `batch.spec_info` for the
# draft-extend forward.
draft_extend_input: EagleDraftExtendInput
# Logit outputs from target worker.
logits_output: LogitsProcessorOutput
# All accepted tokens flat across all reqs incl. those that finished this
# step. Includes the bonus token. Used for output processing.
accept_tokens: torch.Tensor
# Accepted token length per sequence in a batch in CPU (full set).
num_correct_drafts_per_req_cpu: List[int]
# Accepted indices from logits_output.next_token_logits
accept_indices: torch.Tensor
# Whether the target verify forward ran a captured cuda graph. Set by
# the worker after `EagleVerifyInput.sample` returns; default kept so
# idle / direct constructions don't have to pass it.
can_run_cuda_graph: bool = False
@classmethod
def create_idle(
cls,
*,
draft_extend_input: EagleDraftExtendInput,
logits_output: LogitsProcessorOutput,
device: torch.device,
spec_steps: int,
) -> "EagleVerifyOutput":
return cls(
draft_extend_input=draft_extend_input,
logits_output=logits_output,
accept_tokens=torch.empty(0, dtype=torch.long, device=device),
num_correct_drafts_per_req_cpu=[],
accept_indices=torch.full(
(0, spec_steps + 1), -1, dtype=torch.int32, device=device
),
)
@@ -151,14 +151,6 @@ class SpeculativeAlgorithm(Enum):
)
return None
def supports_spec_v2(self) -> bool:
return (
self.is_eagle()
or self.is_standalone()
or self.is_ngram()
or self.is_dflash()
)
def need_topk(self) -> bool:
return self.is_eagle() or self.is_standalone()
+20 -5
View File
@@ -4,6 +4,7 @@ should use that classmethod API; do not import from this module directly.
from __future__ import annotations
import logging
from typing import TYPE_CHECKING, Callable, Dict, Optional, Type
import torch
@@ -17,6 +18,8 @@ if TYPE_CHECKING:
WorkerFactory = Callable[["ServerArgs"], Type]
ServerArgsValidator = Callable[["ServerArgs"], None]
logger = logging.getLogger(__name__)
class CustomSpecAlgo:
"""A plugin-registered speculative algorithm. Duck-types
@@ -28,8 +31,12 @@ class CustomSpecAlgo:
branches like ``if spec_algorithm.is_eagle():`` in scheduler /
model_runner). Pass the subclass via ``spec_class=...`` at registration.
Defaults: all ``is_*()`` return ``False`` except ``is_speculative``;
``supports_spec_v2`` follows ``supports_overlap``.
Defaults: all ``is_*()`` return ``False`` except ``is_speculative``.
``supports_overlap=False`` is deprecated: the spec V1 worker path has been
removed, so such algorithms run on the V2 scheduler schema with overlap
disabled (synchronous). Migrate plugin workers to the V2 schema and
overlap scheduling.
"""
def __init__(
@@ -82,14 +89,22 @@ class CustomSpecAlgo:
# Conservative default: the larger KV reserve.
return True
def supports_spec_v2(self) -> bool:
return self.supports_overlap
def create_worker(self, server_args: ServerArgs) -> Type:
if not server_args.disable_overlap_schedule and not self.supports_overlap:
raise ValueError(
f"Speculative algorithm {self.name} does not support overlap scheduling."
)
if not self.supports_overlap:
# Reached only when overlap is disabled, so the algorithm really
# does run synchronously on the V2 schema below.
logger.warning(
"Speculative algorithm %s is registered with "
"supports_overlap=False, which is deprecated: the spec V1 "
"worker path has been removed, and the algorithm now runs on "
"the V2 scheduler schema with overlap disabled (synchronous). "
"Migrate the plugin worker to support overlap scheduling.",
self.name,
)
return self.factory(server_args)
def get_num_tokens_per_bs_for_target_verify(