Skip to content
54 changes: 39 additions & 15 deletions tests/v1/core/test_async_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -701,16 +701,29 @@ def test_requires_kv_delivery_defaults_to_producer_role():
assert scheduler.requires_kv_delivery is expected, role


@pytest.mark.parametrize("kv_role", ["kv_producer", "kv_consumer"])
def test_kv_pressure_preempt_mid_handoff(kv_role: str):
"""P/D race: a request is KV-pressure preempted while the output of its
final prefill chunk -- the hand-off token that would finish it -- is in
flight.

On a producer, that output must be dropped so the request recomputes;
delivering it would finish the request and hand off blocks the preemption
already freed, so the consumer pulls garbage. A consumer hands nothing off,
so it keeps the lossless deliver-stale path.
@pytest.mark.parametrize(
("kv_role", "defer_free"),
[
("kv_producer", False),
("kv_consumer", False),
("kv_consumer", True),
],
)
def test_kv_pressure_preempt_mid_handoff(kv_role: str, defer_free: bool):
"""P/D race: KV pressure hits while the output of a request's final
prefill chunk -- the hand-off token that would finish it -- is in flight.

When the victim's blocks free immediately, the request is preempted. On a
producer, that output must be dropped so the request recomputes; delivering
it would finish the request and hand off blocks the preemption already
freed, so the consumer pulls garbage. A consumer hands nothing off, so it
keeps the lossless deliver-stale path.

A consumer with overlapping batches (async scheduling or PP) instead fences
the victim's free behind its in-flight output, so the allocation retry
stops instead of preempting; the request then finishes from that output
once it lands. The gate depends on the platform (async scheduling is
force-disabled on CPU), so force the flag to cover both paths everywhere.
"""
is_producer = kv_role == "kv_producer"
scheduler = create_scheduler(
Expand All @@ -722,10 +735,13 @@ def test_kv_pressure_preempt_mid_handoff(kv_role: str):
max_num_batched_tokens=512,
)
assert scheduler.requires_kv_delivery is is_producer
# The production gate requires overlapping batches, which async
# scheduling only provides off-CPU.
scheduler.defer_block_free = defer_free

# 32-token prompts fill 2 blocks each, exhausting the usable pool, so the
# next decode allocation preempts the tail of the running queue (the handoff
# request) while its prefill output is still in flight.
# next decode allocation targets the tail of the running queue (the
# handoff request) while its prefill output is still in flight.
decoder = create_requests(
num_requests=1, num_tokens=32, max_tokens=8, req_ids=["decoder"]
)[0]
Expand All @@ -739,9 +755,17 @@ def test_kv_pressure_preempt_mid_handoff(kv_role: str):
assert handoff.num_output_placeholders == 1

scheduler.schedule()
assert handoff.status == RequestStatus.PREEMPTED
assert handoff.num_stale_output_tokens == handoff.num_prompt_tokens
assert handoff.drop_stale_output is is_producer
if defer_free:
# The victim's blocks are fenced behind its in-flight output, so
# preempting it could not satisfy the allocation; the retry stops
# instead of preempting.
assert handoff.status == RequestStatus.RUNNING
assert handoff.num_stale_output_tokens == 0
else:
# Blocks free immediately, so the handoff request is preempted.
assert handoff.status == RequestStatus.PREEMPTED
assert handoff.num_stale_output_tokens == handoff.num_prompt_tokens
assert handoff.drop_stale_output is is_producer

scheduler.update_from_output(sched_output, _make_model_runner_output(sched_output))

Expand Down
74 changes: 72 additions & 2 deletions tests/v1/core/test_deferred_block_free.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,13 +48,17 @@ def _make_model_runner_output(
)


def _create_deferring_scheduler():
def _create_deferring_scheduler(scheduling_policy="fcfs"):
"""Async scheduler with deferred block freeing forced on.

The production gate additionally requires a PD KV-consumer connector;
the mechanism itself is independent of it.
"""
scheduler = create_scheduler(model=MODEL, async_scheduling=True)
scheduler = create_scheduler(
model=MODEL,
async_scheduling=True,
scheduling_policy=scheduling_policy,
)
scheduler.defer_block_free = True
return scheduler

Expand All @@ -79,6 +83,24 @@ def _setup_request_with_inflight_step(scheduler, max_tokens: int = 5):
return request, out0, out1


def _fail_one_allocation(scheduler, call_number: int):
real_allocate = scheduler.kv_cache_manager.allocate_slots
calls = 0

def allocate(*args, **kwargs):
nonlocal calls
calls += 1
if calls == call_number:
return None
return real_allocate(*args, **kwargs)

return patch.object(
scheduler.kv_cache_manager,
"allocate_slots",
side_effect=allocate,
)


def test_gate_enabled_for_async_consumer():
# Overlapping batches + consumer-side connector enables the gate. Async
# scheduling (which would give >1 concurrent batches) is force-disabled on
Expand Down Expand Up @@ -224,6 +246,54 @@ def test_preempt_defers_free_and_clears_bookkeeping():
assert pool.get_num_free_blocks() == num_free_initially


@pytest.mark.parametrize(
("policy", "failed_call"),
[
("fcfs", 1),
("priority", 2),
],
)
def test_allocation_retry_waits_for_fence_then_succeeds(policy, failed_call):
scheduler = _create_deferring_scheduler(policy)
requests = create_requests(
num_requests=3,
num_tokens=NUM_PROMPT_TOKENS,
max_tokens=10,
stop_token_ids=[STOP_TOKEN_ID],
)
for request in requests:
scheduler.add_request(request)
out0 = scheduler.schedule()
out1 = scheduler.schedule()

worst, trigger, tail = requests
worst.priority, trigger.priority, tail.priority = 9, 0, 1

with _fail_one_allocation(scheduler, failed_call) as allocate:
blocked = scheduler.schedule()
assert allocate.call_count == failed_call
assert not blocked.preempted_req_ids
retry_trigger = requests[failed_call - 1]
assert retry_trigger.request_id not in blocked.num_scheduled_tokens

for output in (out0, out1, blocked):
if output.total_num_scheduled_tokens:
scheduler.update_from_output(
output,
_make_model_runner_output(output),
)

expected_victim = worst if policy == "priority" else tail
with _fail_one_allocation(scheduler, failed_call) as allocate:
resumed = scheduler.schedule()

assert allocate.call_count > failed_call
assert expected_victim.request_id in resumed.preempted_req_ids
assert retry_trigger.request_id in resumed.num_scheduled_tokens
if policy == "priority":
assert tail.request_id in resumed.num_scheduled_tokens


def test_multiple_deferred_frees_drain_in_order():
scheduler = _create_deferring_scheduler()
pool = scheduler.kv_cache_manager.block_pool
Expand Down
3 changes: 3 additions & 0 deletions tests/v1/core/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
SpeculativeConfig,
VllmConfig,
)
from vllm.config.scheduler import SchedulerPolicy
from vllm.multimodal.inputs import (
MultiModalFeatureSpec,
MultiModalKwargsItem,
Expand Down Expand Up @@ -74,6 +75,7 @@ def create_scheduler(
use_v2_model_runner: bool | None = None,
kv_cache_spec: KVCacheSpec | None = None,
per_request_spec_decode_metrics: str = "none",
scheduling_policy: SchedulerPolicy = "fcfs",
) -> Scheduler | AsyncScheduler:
"""Create scheduler under test.

Expand Down Expand Up @@ -113,6 +115,7 @@ def create_scheduler(
is_encoder_decoder=model_config.is_encoder_decoder,
# Ensure admission/preemption mechanics are deterministic
watermark=0.0,
policy=scheduling_policy,
)
# Cache config, optionally force APC
cache_config = CacheConfig(
Expand Down
26 changes: 16 additions & 10 deletions vllm/v1/core/sched/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -745,13 +745,16 @@ def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput:
self.running,
key=lambda r: (r.priority, r.arrival_time),
)
# Record the index of the preemption victim to
# maintain accurate loop state.
else:
preempted_req = self.running[-1]

# A deferred free will not help with immediate allocation.
if not self._request_blocks_can_be_freed(preempted_req):
break

if self.policy == SchedulingPolicy.PRIORITY:
victim_index = self.running.index(preempted_req)
del self.running[victim_index]
# Decrement the loop cursor if the removed request
# preceded the current iteration, preventing the
# silent omission of the subsequent request.
if victim_index < req_index:
req_index -= 1

Expand Down Expand Up @@ -2605,15 +2608,18 @@ def set_pause_state(self, pause_state: PauseState) -> None:
logger.info("setting pause state to %s", pause_state.name)
self._pause_state = pause_state

def _request_blocks_can_be_freed(self, request: Request) -> bool:
# We must defer freeing blocks if an async kv connector may
# write to them immediately (not ordered with GPU stream).
return not self.defer_block_free or (
request.last_sched_seq <= self.processed_step_seq
)

def _free_request_blocks(self, request: Request):
"""Free the request's KV blocks, deferring the return to the block
pool when an in-flight GPU step may still write them.
"""
if not self.defer_block_free or (
# Last scheduled step already processed: no in-flight write remains
# (always the case for a normal finish), so free now.
request.last_sched_seq <= self.processed_step_seq
):
if self._request_blocks_can_be_freed(request):
self.kv_cache_manager.free(request)
return
blocks = self.kv_cache_manager.pop_blocks_for_free(request)
Expand Down
Loading