From 138c2f0421d16f1e4baaa0a9b0007b9ebc8dc3f8 Mon Sep 17 00:00:00 2001 From: Martin Vit Date: Sat, 11 Jul 2026 20:44:09 +0000 Subject: [PATCH] perf: avoid fixed-width draft host synchronization --- .../spec_decode/test_draft_tokens_handler.py | 45 +++++++++++++++++++ vllm/config/speculative.py | 9 ++++ vllm/v1/engine/core.py | 9 +++- vllm/v1/worker/gpu/model_runner.py | 11 ++++- vllm/v1/worker/gpu/spec_decode/utils.py | 28 +++++++++--- 5 files changed, 91 insertions(+), 11 deletions(-) create mode 100644 tests/v1/spec_decode/test_draft_tokens_handler.py diff --git a/tests/v1/spec_decode/test_draft_tokens_handler.py b/tests/v1/spec_decode/test_draft_tokens_handler.py new file mode 100644 index 000000000000..92e660a66321 --- /dev/null +++ b/tests/v1/spec_decode/test_draft_tokens_handler.py @@ -0,0 +1,45 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +from types import SimpleNamespace + +import numpy as np + +from vllm.v1.worker.gpu.spec_decode.utils import DraftTokensHandler + + +def _handler(*, needs_real_draft_tokens: bool) -> DraftTokensHandler: + handler = object.__new__(DraftTokensHandler) + handler.needs_real_draft_tokens = needs_real_draft_tokens + handler.req_ids = [] + handler.draft_tokens_np = None + handler.num_draft_tokens = 0 + return handler + + +def test_fixed_width_drafts_skip_host_copy() -> None: + handler = _handler(needs_real_draft_tokens=False) + input_batch = SimpleNamespace( + req_ids=["req-0", "req-1"], has_structured_output_reqs=False + ) + draft_tokens = np.zeros((2, 2), dtype=np.int32) + + handler.set_draft_tokens(input_batch, draft_tokens) # type: ignore[arg-type] + output = handler.get_draft_tokens() + + assert handler.draft_tokens_np is None + assert output is not None + assert output.req_ids == ["req-0", "req-1"] + assert output.draft_token_ids == [[-1, -1], [-1, -1]] + + +def test_host_draft_ids_trim_negative_suffix() -> None: + handler = _handler(needs_real_draft_tokens=True) + handler.req_ids = ["req-0", "req-1"] + handler.draft_tokens_np = np.array([[11, 12], [21, -1]], dtype=np.int32) + handler.copy_event = SimpleNamespace(synchronize=lambda: None) + + output = handler.get_draft_tokens() + + assert output is not None + assert output.draft_token_ids == [[11, 12], [21]] diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index 1b87da7f61b9..2df68cd5a5cd 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -1360,6 +1360,15 @@ def use_dspark(self) -> bool: def use_causal_cascade(self) -> bool: return self.method == "causal_cascade" + def requires_host_draft_token_ids(self) -> bool: + """Whether async scheduling needs the actual draft ids on the host. + + Fixed-width speculators can use scheduler placeholders and keep draft + ids on the worker. Block speculators may return a shorter prefix using + ``-1`` sentinels, so the scheduler must receive their real draft ids. + """ + return self.method in ("dflash", "dspark", "causal_cascade") + def uses_batch_size_dynamic_speculative_decoding(self) -> bool: return self.num_speculative_tokens_per_batch_size is not None diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py index 7aaddbbb6f38..b19c1d2f059d 100644 --- a/vllm/v1/engine/core.py +++ b/vllm/v1/engine/core.py @@ -160,6 +160,11 @@ def __init__( self.check_for_draft_tokens = ( self.use_spec_decode or vllm_config.model_config.is_diffusion ) + speculative_config = vllm_config.speculative_config + self.requires_host_draft_token_ids = ( + speculative_config is not None + and speculative_config.requires_host_draft_token_ids() + ) if self.scheduler.connector is not None: # type: ignore self.model_executor.init_kv_output_aggregator(self.scheduler.connector) # type: ignore @@ -505,7 +510,7 @@ def step(self) -> tuple[dict[int, EngineCoreOutputs], bool]: scheduler_output, model_output ) if ( - self.check_for_draft_tokens + self.requires_host_draft_token_ids and self.async_scheduling and scheduler_output.total_num_scheduled_tokens > 0 ): @@ -612,7 +617,7 @@ def step_with_batch_queue( scheduler_output, model_output ) if ( - self.check_for_draft_tokens + self.requires_host_draft_token_ids and self.async_scheduling and scheduler_output.total_num_scheduled_tokens > 0 ): diff --git a/vllm/v1/worker/gpu/model_runner.py b/vllm/v1/worker/gpu/model_runner.py index b893f640a8a1..222f836ba545 100644 --- a/vllm/v1/worker/gpu/model_runner.py +++ b/vllm/v1/worker/gpu/model_runner.py @@ -240,8 +240,15 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): "parallel is not supported." ) - # Draft tokens propagation - for spec-dec + struct outputs. - self.draft_tokens_handler = DraftTokensHandler(self.device) + # Draft token propagation for structured outputs and block speculators + # whose returned draft prefix can be shorter than the configured width. + self.draft_tokens_handler = DraftTokensHandler( + self.device, + needs_real_draft_tokens=( + self.speculative_config is not None + and self.speculative_config.requires_host_draft_token_ids() + ), + ) # Pooling models. self.is_pooling_model = self.model_config.runner_type == "pooling" diff --git a/vllm/v1/worker/gpu/spec_decode/utils.py b/vllm/v1/worker/gpu/spec_decode/utils.py index 52c060dda753..5732c563a824 100644 --- a/vllm/v1/worker/gpu/spec_decode/utils.py +++ b/vllm/v1/worker/gpu/spec_decode/utils.py @@ -37,8 +37,13 @@ def limit_draft_tokens( class DraftTokensHandler: - def __init__(self, device: torch.device | None = None): + def __init__( + self, + device: torch.device | None = None, + needs_real_draft_tokens: bool = False, + ): self.device = device + self.needs_real_draft_tokens = needs_real_draft_tokens self.copy_stream = torch.cuda.Stream(device) # Blocking (sleep) event to avoid busy-polling the CUDA driver lock. self.copy_event = torch.cuda.Event(blocking=True) @@ -52,6 +57,14 @@ def set_draft_tokens( ) -> None: self.req_ids = input_batch.req_ids self.num_draft_tokens = draft_tokens.shape[1] + if ( + not self.needs_real_draft_tokens + and not input_batch.has_structured_output_reqs + ): + # Fixed-width speculators use scheduler placeholders. Avoid a D2H + # copy and per-step event synchronization on their decode path. + self.draft_tokens_np = None + return # The scheduler needs the real draft lengths. Some speculators use # -1 as a sentinel for fallback/no-draft slots; sending placeholder @@ -71,14 +84,15 @@ def get_draft_tokens(self) -> DraftTokenIds | None: if self.draft_tokens_np is not None: self.copy_event.synchronize() draft_token_ids = self.draft_tokens_np.tolist() + for token_ids in draft_token_ids: + for i, token_id in enumerate(token_ids): + if token_id < 0: + del token_ids[i:] + break else: - # This case only happens when async scheduling is disabled. + # Fixed-width async scheduling only needs the draft count. The + # worker retains the actual token ids for verification. draft_token_ids = [[-1] * self.num_draft_tokens for _ in self.req_ids] - for token_ids in draft_token_ids: - for i, token_id in enumerate(token_ids): - if token_id < 0: - del token_ids[i:] - break return DraftTokenIds(self.req_ids, draft_token_ids)