Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
63 changes: 39 additions & 24 deletions vllm/config/speculative.py
Original file line number Diff line number Diff line change
Expand Up @@ -1421,45 +1421,60 @@ def verify_equal_vocab_size_if_draft_model(self):
)

@property
def max_num_new_slots_for_drafting(self) -> int:
"""Return the maximum additional drafting slots per request.

The scheduler budget already includes one query slot per decoding request.
Let K be ``num_speculative_tokens``. Standard configurations require:

==================== ============= ======== ================
Algorithm Method Parallel Additional slots
==================== ============= ======== ================
EAGLE3 eagle3 No 0
P-EAGLE eagle3 Yes K - 1
DFlash dflash Yes K
DSpark dspark Yes K - 1
MTP mtp No 0
N-gram ngram No 0
Draft model draft_model No 1
PARD draft_model Yes K
==================== ============= ======== ================
def draft_input_cost_params(self) -> tuple[bool, int]:
"""Return ``(extend_target_batch, fixed_overhead)``.

For a request with ``Q`` scheduled tokens, the number of draft input tokens is
``extend_target_batch * Q + fixed_overhead``.

==================== =========== ======== =================== ==============
Algorithm Method Parallel Extend target batch Fixed overhead
==================== =========== ======== =================== ==============
N-gram ngram No Yes 0
MTP mtp No Yes 0
EAGLE3 eagle3 No Yes 0
P-EAGLE eagle3 Yes Yes K - 1
DFlash dflash Yes No K + 1
DSpark dspark Yes No K
DSpark (speculators) dspark Yes No K + 1
Draft model draft_model No Yes 1
PARD draft_model Yes Yes K
==================== =========== ======== =================== ==============

``K`` is ``num_speculative_tokens``.
"""
num_draft_tokens = self.num_speculative_tokens

if self.use_dflash():
# DFlash uses one bonus query followed by K mask queries.
return num_draft_tokens
# One bonus/anchor query plus K mask queries.
return False, num_draft_tokens + 1

if self.use_dspark():
# Native DSpark uses K queries when the anchor predicts token 0.
# Speculators-format checkpoints retain DFlash's bonus-anchor layout and
# therefore need K + 1 queries.
hf_config = (
self.draft_model_config.hf_config
if self.draft_model_config is not None
else None
)
sample_from_anchor = getattr(hf_config, "sample_from_anchor", True)
return False, num_draft_tokens + int(not sample_from_anchor)

if self.parallel_drafting:
if self.uses_draft_model():
# PARD does not shift the existing input, so all K query
# positions require additional slots.
return num_draft_tokens
return True, num_draft_tokens

# The existing query is reused; only masked queries need new slots.
return num_draft_tokens - 1
return True, num_draft_tokens - 1

if self.uses_draft_model():
# The autoregressive draft-model input retains one unsliced token.
return 1
return True, 1

return 0
return True, 0

def use_gemma4_mtp(self) -> bool:
return (
Expand Down
16 changes: 9 additions & 7 deletions vllm/config/vllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -1825,9 +1825,6 @@ def _set_max_num_scheduled_tokens(self):
worker is configured to handle.
"""
if self.speculative_config is not None:
scheduled_token_delta = (
self.speculative_config.max_num_new_slots_for_drafting
)
max_num_batched_tokens = self.scheduler_config.max_num_batched_tokens
if self.scheduler_config.max_num_scheduled_tokens is None:
self.scheduler_config.max_num_scheduled_tokens = max_num_batched_tokens
Expand All @@ -1851,11 +1848,16 @@ def _set_max_num_scheduled_tokens(self):
" num_speculative_tokens.",
)

if max_num_batched_tokens <= scheduled_token_delta:
extend_target_batch, fixed_overhead = (
self.speculative_config.draft_input_cost_params
)
min_num_draft_input_tokens = int(extend_target_batch) + fixed_overhead
if max_num_batched_tokens < min_num_draft_input_tokens:
raise ValueError(
"VllmConfig does not have enough slots to schedule a token and"
" support the speculative decoding settings."
f" Got {max_num_batched_tokens=} and {scheduled_token_delta=}."
"VllmConfig does not have enough slots for"
" the smallest draft input batch."
f" Got {max_num_batched_tokens=} and"
f" {min_num_draft_input_tokens=}."
)

def _set_cudagraph_sizes(self):
Expand Down
53 changes: 40 additions & 13 deletions vllm/v1/core/sched/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -255,6 +255,10 @@ def __init__(
# reserve between a chunk boundary and the prefill end.
self.num_prefill_lookahead = 0
self.dynamic_sd_lookup: list[int] | None = None
# Draft input cost per request: optionally extend the target batch,
# then add a fixed overhead.
self.draft_extend_target_batch = False
self.draft_fixed_overhead = 0
if speculative_config is not None:
if speculative_config.num_speculative_tokens_per_batch_size:
self.dynamic_sd_lookup = build_dynamic_sd_schedule_lookup(
Expand All @@ -269,6 +273,10 @@ def __init__(
if speculative_config.use_multi_module_mtp()
else 1
)
(
self.draft_extend_target_batch,
self.draft_fixed_overhead,
) = speculative_config.draft_input_cost_params

# Create the KV cache manager.
if hash_block_size is None:
Expand Down Expand Up @@ -473,6 +481,24 @@ def _reserve_prefill_lookahead(
num_new_tokens -= self.num_prefill_lookahead - remaining
return max(num_new_tokens, 0)

def _get_num_draft_input_tokens(self, num_target_tokens: int) -> int:
"""Return the draft input size for one scheduled request."""
return (
self.draft_extend_target_batch * num_target_tokens
+ self.draft_fixed_overhead
)

def _get_request_token_budget(
self, token_budget: int, draft_budget: int
) -> int:
"""Limit a request's token budget by the available draft capacity."""
remaining_draft_budget = draft_budget - self.draft_fixed_overhead
if remaining_draft_budget < 0:
return 0
if self.draft_extend_target_batch:
return min(token_budget, remaining_draft_budget)
return token_budget

def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
self.current_step += 1
# NOTE(woosuk) on the scheduling algorithm:
Expand All @@ -494,9 +520,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
req_to_new_blocks: dict[str, KVCacheBlocks] = {}
num_scheduled_tokens: dict[str, int] = {}
token_budget = self.max_num_scheduled_tokens
spec = self.vllm_config.speculative_config
draft_slots = spec.max_num_new_slots_for_drafting if spec is not None else 0
input_budget = self.scheduler_config.max_num_batched_tokens
draft_budget = self.scheduler_config.max_num_batched_tokens
if self._pause_state == PauseState.PAUSED_ALL:
# Do not schedule any requests when paused.
token_budget = 0
Expand Down Expand Up @@ -524,7 +548,10 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
req_index = 0
while req_index < len(self.running) and token_budget > 0:
request = self.running[req_index]
if input_budget <= draft_slots:
request_token_budget = self._get_request_token_budget(
token_budget, draft_budget
)
if request_token_budget == 0:
break

if (
Expand Down Expand Up @@ -562,9 +589,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
)
if 0 < self.scheduler_config.long_prefill_token_threshold < num_new_tokens:
num_new_tokens = self.scheduler_config.long_prefill_token_threshold
num_new_tokens = min(
num_new_tokens, token_budget, input_budget - draft_slots
)
num_new_tokens = min(num_new_tokens, request_token_budget)

# Make sure the input position does not exceed the max model len.
# This is necessary when using spec decoding.
Expand Down Expand Up @@ -660,7 +685,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
scheduled_running_reqs.remove(preempted_req)
restored = num_scheduled_tokens.pop(preempted_req_id)
token_budget += restored
input_budget += restored + draft_slots
draft_budget += self._get_num_draft_input_tokens(restored)
req_to_new_blocks.pop(preempted_req_id)
scheduled_spec_decode_tokens.pop(preempted_req_id, None)
preempted_encoder_inputs = scheduled_encoder_inputs.pop(
Expand Down Expand Up @@ -698,7 +723,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
req_to_new_blocks[request_id] = new_blocks
num_scheduled_tokens[request_id] = num_new_tokens
token_budget -= num_new_tokens
input_budget -= num_new_tokens + draft_slots
draft_budget -= self._get_num_draft_input_tokens(num_new_tokens)
req_index += 1

# Speculative decode related.
Expand Down Expand Up @@ -749,7 +774,10 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
step_skipped_waiting = create_request_queue(self.policy)

while (self.waiting or self.skipped_waiting) and token_budget > 0:
if input_budget <= draft_slots:
request_token_budget = self._get_request_token_budget(
token_budget, draft_budget
)
if request_token_budget == 0:
break
# Paused streaming sessions (WAITING_FOR_STREAMING_REQ) are not
# in `running` but still hold a model-runner request slot.
Expand Down Expand Up @@ -924,7 +952,6 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
# compute to a cadence-aligned step.
break
else:
request_token_budget = min(token_budget, input_budget - draft_slots)
# Number of tokens to be scheduled.
# We use `request.num_tokens` instead of
# `request.num_prompt_tokens` to consider the resumed
Expand Down Expand Up @@ -1131,7 +1158,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
)
num_scheduled_tokens[request_id] = num_new_tokens
token_budget -= num_new_tokens
input_budget -= num_new_tokens + draft_slots
draft_budget -= self._get_num_draft_input_tokens(num_new_tokens)
request.status = RequestStatus.RUNNING
request.num_computed_tokens = num_computed_tokens
if pad_spec_decode:
Expand Down Expand Up @@ -1171,7 +1198,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
assert total_num_scheduled_tokens <= self.max_num_scheduled_tokens

assert token_budget >= 0
assert input_budget >= 0
assert draft_budget >= 0
assert len(self.running) <= self.max_num_running_reqs
# Since some requests in the RUNNING queue may not be scheduled in
# this step, the total number of scheduled requests can be smaller than
Expand Down
Loading