Skip to content
Merged
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
45 changes: 45 additions & 0 deletions tests/v1/spec_decode/test_draft_tokens_handler.py
Original file line number Diff line number Diff line change
@@ -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]]
9 changes: 9 additions & 0 deletions vllm/config/speculative.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
9 changes: 7 additions & 2 deletions vllm/v1/engine/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
):
Expand Down Expand Up @@ -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
):
Expand Down
11 changes: 9 additions & 2 deletions vllm/v1/worker/gpu/model_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
28 changes: 21 additions & 7 deletions vllm/v1/worker/gpu/spec_decode/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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
Expand All @@ -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)


Expand Down
Loading