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
32 changes: 32 additions & 0 deletions tests/v1/core/test_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -1304,6 +1304,38 @@ def test_draft_slots_budgeted_per_scheduled_request(tmp_path, monkeypatch):
assert scheduler.schedule().num_scheduled_tokens == {"0": 10, "1": 4}


@pytest.mark.parametrize(
("num_requests", "num_tokens", "expected"),
[
(2, 8, {"0": 8, "1": 8}),
(5, 1, {"0": 1, "1": 1, "2": 1, "3": 1}),
],
)
def test_dspark_uses_separate_draft_input_budget(
tmp_path, monkeypatch, num_requests, num_tokens, expected
):
monkeypatch.setenv("VLLM_USE_V2_MODEL_RUNNER", "0")
(tmp_path / "config.json").write_text(
'{"architectures": ["OPTForCausalLM"], "model_type": "opt"}'
)
scheduler = create_scheduler(
model=str(tmp_path),
max_num_seqs=16,
max_num_batched_tokens=16,
num_speculative_tokens=4,
parallel_drafting=True,
skip_tokenizer_init=True,
)
speculative_config = scheduler.vllm_config.speculative_config
assert speculative_config is not None
speculative_config.method = "dspark"

for request in create_requests(num_requests=num_requests, num_tokens=num_tokens):
scheduler.add_request(request)

assert scheduler.schedule().num_scheduled_tokens == expected


# Note - these test cases mirror some of those in test_rejection_sampler.py
@pytest.mark.parametrize(
"spec_tokens,output_tokens,expected",
Expand Down
21 changes: 18 additions & 3 deletions vllm/v1/core/sched/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -496,8 +496,16 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
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
dspark_draft_tokens = (
spec.num_speculative_tokens if spec is not None and spec.use_dspark() else 0
)
draft_slots = (
spec.max_num_new_slots_for_drafting
if spec is not None and not spec.use_dspark()
else 0
)
input_budget = self.scheduler_config.max_num_batched_tokens
draft_input_budget = input_budget
if self._pause_state == PauseState.PAUSED_ALL:
# Do not schedule any requests when paused.
token_budget = 0
Expand Down Expand Up @@ -525,7 +533,7 @@ 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:
if input_budget <= draft_slots or draft_input_budget < dspark_draft_tokens:
break

if (
Expand Down Expand Up @@ -662,6 +670,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
restored = num_scheduled_tokens.pop(preempted_req_id)
token_budget += restored
input_budget += restored + draft_slots
draft_input_budget += dspark_draft_tokens
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 @@ -700,6 +709,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_input_budget -= dspark_draft_tokens
req_index += 1

# Speculative decode related.
Expand Down Expand Up @@ -750,7 +760,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:
if (
input_budget <= draft_slots
or draft_input_budget < dspark_draft_tokens
):
break
# Paused streaming sessions (WAITING_FOR_STREAMING_REQ) are not
# in `running` but still hold a model-runner request slot.
Expand Down Expand Up @@ -1133,6 +1146,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_input_budget -= dspark_draft_tokens
request.status = RequestStatus.RUNNING
request.num_computed_tokens = num_computed_tokens
if pad_spec_decode:
Expand Down Expand Up @@ -1173,6 +1187,7 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:

assert token_budget >= 0
assert input_budget >= 0
assert draft_input_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