From 55efcf46b96cebba02bd04c87f7b6e03b8b3e1f6 Mon Sep 17 00:00:00 2001 From: feed <144643411+feednetinfra@users.noreply.github.com> Date: Thu, 13 Aug 2026 11:16:35 +0800 Subject: [PATCH] Preserve constant effective K schedule semantics Co-authored-by: OpenAI Codex Signed-off-by: feed <144643411+feednetinfra@users.noreply.github.com> --- tests/v1/spec_decode/test_dynamic_sd.py | 119 ++++++++++++++++++++++++ vllm/config/speculative.py | 25 +++++ vllm/config/vllm.py | 14 ++- vllm/v1/core/sched/scheduler.py | 17 +++- 4 files changed, 170 insertions(+), 5 deletions(-) diff --git a/tests/v1/spec_decode/test_dynamic_sd.py b/tests/v1/spec_decode/test_dynamic_sd.py index 8d46b241c628..7035930380f7 100644 --- a/tests/v1/spec_decode/test_dynamic_sd.py +++ b/tests/v1/spec_decode/test_dynamic_sd.py @@ -7,7 +7,10 @@ import pytest from tests.v1.core.utils import create_requests, create_scheduler +from vllm.config import CUDAGraphMode from vllm.v1.core.sched.scheduler import Scheduler +from vllm.v1.cudagraph_dispatcher import CudagraphDispatcher +from vllm.v1.outputs import DraftTokenIds, ModelRunnerOutput from vllm.v1.spec_decode.dynamic.utils import build_dynamic_sd_schedule_lookup from vllm.v1.structured_output import StructuredOutputManager @@ -95,6 +98,35 @@ def test_dynamic_sd_clamps_k_to_runtime_max(): assert dynamic_sd_lookup[256] == 3 +@pytest.mark.parametrize( + ("schedule", "runtime_max_k", "expected_constant_k", "expected_variable"), + [ + (None, 3, None, False), + ([(1, 16, 3)], 3, 3, False), + ([(1, 16, 3), (32, 64, 3)], 3, 3, False), + ([(1, 16, 5), (32, 64, 4)], 3, 3, False), + ([(1, 16, 3), (32, 64, 1)], 3, None, True), + ([(1, 16, 0)], 3, None, True), + ([1], 3, None, True), + ], +) +def test_speculative_config_classifies_effective_schedule_shape( + schedule: object, + runtime_max_k, + expected_constant_k, + expected_variable, +): + config = _make_scheduler_with_dynamic_sd( + [(1, 16, runtime_max_k)], + runtime_num_speculative_tokens=runtime_max_k, + ).vllm_config.speculative_config + assert config is not None + config.num_speculative_tokens_per_batch_size = schedule # type: ignore[assignment] + + assert config.constant_num_speculative_tokens() == expected_constant_k + assert config.uses_variable_speculative_decoding() is expected_variable + + def test_dynamic_sd_rejects_invalid_schedule_entry(): with pytest.raises(ValueError, match="3-item sequence"): _make_lookup([(1, 16, 3), (32, 64)]) # type: ignore[list-item] @@ -216,6 +248,93 @@ def test_dynamic_sd_is_disabled_with_data_parallel(caplog_vllm): assert output.num_spec_tokens_to_schedule == 3 +def test_constant_schedule_is_kept_with_data_parallel(caplog_vllm): + with caplog_vllm.at_level(logging.WARNING, logger="vllm"): + scheduler = create_scheduler( + max_num_seqs=16, + max_num_batched_tokens=160, + num_speculative_tokens=3, + num_speculative_tokens_per_batch_size=[ + (1, 8, 2), + (9, 16, 2), + ], + data_parallel_size=2, + ) + + speculative_config = scheduler.vllm_config.speculative_config + assert speculative_config is not None + assert speculative_config.num_speculative_tokens_per_batch_size is not None + assert scheduler.dynamic_sd_lookup is not None + assert scheduler.num_spec_tokens == 2 + assert scheduler.num_uniform_spec_tokens == 2 + assert "Dynamic speculative decoding is not supported" not in caplog_vllm.text + + output = _add_requests_and_schedule(scheduler, 16) + assert output.num_spec_tokens_to_schedule == 2 + + +def test_constant_schedule_pads_first_decode_step_with_effective_k(monkeypatch): + monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "0") + scheduler = create_scheduler( + num_speculative_tokens=3, + num_speculative_tokens_per_batch_size=[(1, 16, 2)], + enable_prefix_caching=True, + block_size=16, + ) + assert scheduler.vllm_config.compilation_config.cudagraph_mode.has_full_cudagraphs() + assert scheduler.vllm_config.num_speculative_tokens == 2 + assert scheduler.num_spec_tokens == 2 + assert scheduler.num_uniform_spec_tokens == 2 + compilation_config = scheduler.vllm_config.compilation_config + compilation_config.cudagraph_capture_sizes = [3, 6] + compilation_config.max_cudagraph_capture_size = 6 + dispatcher = CudagraphDispatcher(scheduler.vllm_config) + dispatcher.initialize_cudagraph_keys( + compilation_config.cudagraph_mode, + uniform_decode_query_len=3, + ) + runtime_mode, batch_descriptor = dispatcher.dispatch( + num_tokens=3, + uniform_decode=True, + ) + assert runtime_mode == CUDAGraphMode.FULL + assert batch_descriptor.uniform + running_request, cache_hit_request = create_requests( + num_requests=2, + num_tokens=33, + same_prompt=True, + max_tokens=16, + ) + + scheduler.add_request(running_request) + output = scheduler.schedule() + scheduler.update_from_output( + output, + ModelRunnerOutput( + req_ids=[running_request.request_id], + req_id_to_index={running_request.request_id: 0}, + sampled_token_ids=[[100]], + logprobs=None, + prompt_logprobs_dict={}, + pooler_output=[], + ), + ) + scheduler.update_draft_token_ids( + DraftTokenIds([running_request.request_id], [[1, 2]]) + ) + + scheduler.add_request(cache_hit_request) + output = scheduler.schedule() + + assert output.num_spec_tokens_to_schedule == 2 + assert output.scheduled_spec_decode_tokens[running_request.request_id] == [1, 2] + assert output.num_scheduled_tokens[cache_hit_request.request_id] == 3 + assert output.scheduled_spec_decode_tokens[cache_hit_request.request_id] == [ + -1, + -1, + ] + + def test_scheduler_uses_static_k_when_no_requests_are_scheduled(): scheduler = _make_scheduler_with_dynamic_sd( [(1, 16, 3), (64, 128, 2), (256, 4096, 0)], diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index be70da78b350..63761ea7d4e6 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -1493,6 +1493,31 @@ def use_dspark(self) -> bool: def uses_dynamic_speculative_decoding(self) -> bool: return self.num_speculative_tokens_per_batch_size is not None + def constant_num_speculative_tokens(self) -> int | None: + """Return a positive K when every schedule entry resolves to it.""" + schedule = self.num_speculative_tokens_per_batch_size + if not schedule or any( + not isinstance(entry, list | tuple) or len(entry) != 3 for entry in schedule + ): + return None + + max_k = self.num_speculative_tokens + try: + scheduled_k = {min(int(entry[2]), max_k) for entry in schedule} + except (TypeError, ValueError): + return None + if len(scheduled_k) != 1: + return None + + constant_k = next(iter(scheduled_k)) + return constant_k if constant_k > 0 else None + + def uses_variable_speculative_decoding(self) -> bool: + """Whether the schedule can change the target verification width.""" + return self.uses_dynamic_speculative_decoding() and ( + self.constant_num_speculative_tokens() is None + ) + def uses_draft_model(self) -> bool: return self.method == "draft_model" diff --git a/vllm/config/vllm.py b/vllm/config/vllm.py index 67ac2c2b748e..bda3d1f18dda 100644 --- a/vllm/config/vllm.py +++ b/vllm/config/vllm.py @@ -926,7 +926,7 @@ def _maybe_override_dynamic_sd_cudagraph_mode(self) -> None: speculative_config = self.speculative_config if ( speculative_config is None - or not speculative_config.uses_dynamic_speculative_decoding() + or not speculative_config.uses_variable_speculative_decoding() or not self.compilation_config.cudagraph_mode.has_full_cudagraphs() or self.use_v2_model_runner ): @@ -941,11 +941,20 @@ def _maybe_override_dynamic_sd_cudagraph_mode(self) -> None: ) self.compilation_config.cudagraph_mode = CUDAGraphMode.PIECEWISE + def _normalize_constant_speculative_schedule(self) -> None: + speculative_config = self.speculative_config + if speculative_config is None: + return + + constant_k = speculative_config.constant_num_speculative_tokens() + if constant_k is not None: + speculative_config.num_speculative_tokens = constant_k + def _maybe_disable_dynamic_sd_for_data_parallel(self) -> None: speculative_config = self.speculative_config if ( speculative_config is None - or not speculative_config.uses_dynamic_speculative_decoding() + or not speculative_config.uses_variable_speculative_decoding() or self.parallel_config.data_parallel_size <= 1 ): return @@ -1391,6 +1400,7 @@ def has_blocked_weights(): "optimization level defaults." ) + self._normalize_constant_speculative_schedule() self._maybe_disable_dynamic_sd_for_data_parallel() self._maybe_override_dynamic_sd_cudagraph_mode() diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index 4ebc248b8449..b935eede3727 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -245,9 +245,17 @@ def __init__( speculative_config = vllm_config.speculative_config self.use_eagle = False self.num_spec_tokens = vllm_config.num_speculative_tokens + self.num_uniform_spec_tokens = self.num_spec_tokens self.num_lookahead_tokens = vllm_config.num_lookahead_tokens self.dynamic_sd_lookup: list[int] | None = None + self.variable_speculative_tokens = False if speculative_config is not None: + self.variable_speculative_tokens = ( + speculative_config.uses_variable_speculative_decoding() + ) + constant_k = speculative_config.constant_num_speculative_tokens() + if constant_k is not None: + self.num_uniform_spec_tokens = constant_k if speculative_config.num_speculative_tokens_per_batch_size: self.dynamic_sd_lookup = build_dynamic_sd_schedule_lookup( speculative_config.num_speculative_tokens_per_batch_size, @@ -890,12 +898,15 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: # preserve full cudagraph for this step. # Not for diffusion where draft tokens can't be padded. if ( - (self.num_spec_tokens > 0 and self.dynamic_sd_lookup is None) + ( + self.num_uniform_spec_tokens > 0 + and not self.variable_speculative_tokens + ) and self.num_sampled_tokens_per_step > 0 and num_new_tokens == 1 and (scheduled_running_reqs and not prefill_scheduled) ): - num_new_tokens = 1 + self.num_spec_tokens + num_new_tokens = 1 + self.num_uniform_spec_tokens if ( num_new_tokens > request_token_budget or num_computed_tokens + num_new_tokens > self.max_model_len @@ -1085,7 +1096,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: if pad_spec_decode: scheduled_spec_decode_tokens[request_id] = [ -1 - ] * self.num_spec_tokens + ] * self.num_uniform_spec_tokens # Only track requests that will still be prefilling after this chunk. if num_computed_tokens + num_new_tokens < request.num_tokens: self._inflight_prefills.add(request)