diff --git a/tests/v1/core/test_async_scheduler.py b/tests/v1/core/test_async_scheduler.py index 7f256f025e6d..0a21e896b030 100644 --- a/tests/v1/core/test_async_scheduler.py +++ b/tests/v1/core/test_async_scheduler.py @@ -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( @@ -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] @@ -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)) diff --git a/tests/v1/core/test_deferred_block_free.py b/tests/v1/core/test_deferred_block_free.py index 13789a396015..ebbf492e208d 100644 --- a/tests/v1/core/test_deferred_block_free.py +++ b/tests/v1/core/test_deferred_block_free.py @@ -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 @@ -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 @@ -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 diff --git a/tests/v1/core/utils.py b/tests/v1/core/utils.py index d3090e2ddace..3c4a8960a3fa 100644 --- a/tests/v1/core/utils.py +++ b/tests/v1/core/utils.py @@ -17,6 +17,7 @@ SpeculativeConfig, VllmConfig, ) +from vllm.config.scheduler import SchedulerPolicy from vllm.multimodal.inputs import ( MultiModalFeatureSpec, MultiModalKwargsItem, @@ -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. @@ -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( diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index 01423212999f..8b2a551b5579 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -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 @@ -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)