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
119 changes: 119 additions & 0 deletions tests/v1/spec_decode/test_dynamic_sd.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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]
Expand Down Expand Up @@ -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)],
Expand Down
25 changes: 25 additions & 0 deletions vllm/config/speculative.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"

Expand Down
14 changes: 12 additions & 2 deletions vllm/config/vllm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
):
Expand All @@ -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
Expand Down Expand Up @@ -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()

Expand Down
17 changes: 14 additions & 3 deletions vllm/v1/core/sched/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading