From 885a319fe953eadd9d5452a1d4f759ffc3a7cbf0 Mon Sep 17 00:00:00 2001 From: Jason Cook Date: Tue, 8 Sep 2026 15:31:03 -0400 Subject: [PATCH 1/2] Coalesce the final prefill chunk of prompts that end off the block grid prefill_checkpoint_plan required the planned span to end on a block boundary, so a prompt whose length is not a multiple of 512 never coalesced its last span: the scheduler fell back to the reachable-boundary ladder (for 16,228 tokens: 8192, 4096, 2048, 1024, 512, 356) with one recurrent checkpoint per chunk. Allow the span to end at the prompt end instead. The interior retained states are exported in one pass as before and the final state lands in the block that holds the last token, which is the column an ordinary unaligned tail already uses; continuation_layout derives that column from the last token rather than assuming alignment. With this, 16,228 tokens schedule as 8192 + 8036 and 32,319 as three 8192 chunks plus 7743. The single-token tail of a prompt ending one past a block boundary still has no interior state and splits ordinarily. Signed-off-by: Jason Cook --- .../core/test_recurrent_prefill_checkpoint.py | 76 ++++++++++++++++++- vllm/v1/core/recurrent_prefill_checkpoint.py | 21 +++-- 2 files changed, 89 insertions(+), 8 deletions(-) diff --git a/tests/v1/core/test_recurrent_prefill_checkpoint.py b/tests/v1/core/test_recurrent_prefill_checkpoint.py index 359692a6f3ae..3c6a26609dd1 100644 --- a/tests/v1/core/test_recurrent_prefill_checkpoint.py +++ b/tests/v1/core/test_recurrent_prefill_checkpoint.py @@ -46,7 +46,7 @@ def test_fresh_and_continuation_plans_export_exact_retained_boundaries(): ) -@pytest.mark.parametrize("end", [8193, 16385, 24576]) +@pytest.mark.parametrize("end", [8193, 24576]) def test_unrepresentable_or_over_budget_final_spans_use_ordinary_splitting(end): assert ( prefill_checkpoint_plan( @@ -61,6 +61,46 @@ def test_unrepresentable_or_over_budget_final_spans_use_ordinary_splitting(end): ) +def test_unaligned_prompt_end_coalesces_only_as_the_final_span(): + publications = (12288, 14336, 15360, 15872) + assert prefill_checkpoint_plan( + start=8192, + end=16228, + prompt=16228, + num_tokens=16228, + block_size=512, + publications=publications, + ) == (8192, 16228, publications) + # An unaligned end that is not the prompt end still splits ordinarily. + assert ( + prefill_checkpoint_plan( + start=8192, + end=16000, + prompt=16228, + num_tokens=16228, + block_size=512, + publications=publications, + ) + is None + ) + # A single-token tail has no interior state to retain. + assert ( + prefill_checkpoint_plan( + start=16384, + end=16385, + prompt=16385, + num_tokens=16385, + block_size=512, + publications=(14336, 15872, 16384), + ) + is None + ) + assert checkpoint_metadata((8192, 16228, publications), 8192, 16228, 512, 4) == ( + [4096, 6144, 7168, 7680], + [23, 27, 29, 30], + ) + + def test_worker_checkpoint_rows_follow_request_order_and_reject_span_drift(): batch = NS( req_ids=["b", "a"], @@ -358,6 +398,40 @@ def test_dcp_geometry_retains_required_states_without_extra_passes(prompt, dcp): assert len({blocks[column].block_id for column in columns}) == len(columns) +@pytest.mark.parametrize("prompt", [16228, 16383, 32319, 6244]) +def test_dcp4_unaligned_prompt_tail_coalesces_into_one_chunk(prompt): + cache, manager, scheduler, request = cache_fixture(prompt) + start = (prompt - 1) // 8192 * 8192 + expected = tuple( + sorted( + p + for p in manager._expand_reachable_boundaries([prompt - 1]) + if start < p < prompt + ) + ) + assert 1 <= len(expected) <= 4 + if start: + while request.num_computed_tokens < start: + assert cache.allocate_slots(request, 8192) is not None + if request.num_computed_tokens == 0: + scheduler._record_coalescing_origin(request, 0, 0, 0, False) + request.num_computed_tokens += 8192 + tail = prompt - start + plan = scheduler._recurrent_checkpoint_plan(request, start, prompt) + assert plan == (start, prompt, expected) + assert scheduler._mamba_block_aligned_split(request, tail) == tail + assert ( + cache.allocate_slots( + request, tail, num_lookahead_tokens=3, recurrent_checkpoint_plan=plan + ) + is not None + ) + blocks = manager.req_to_blocks[request.request_id] + columns = [p // 512 - 1 for p in expected] + [(prompt - 1) // 512] + assert all(not blocks[column].is_null for column in columns) + assert len({blocks[column].block_id for column in columns}) == len(columns) + + def test_two_checkpoint_capacity_keeps_safe_dcp4_fallback(): _, manager, scheduler, request = cache_fixture(8192, checkpoints=2) assert manager.hit_alignment_tokens == 512 diff --git a/vllm/v1/core/recurrent_prefill_checkpoint.py b/vllm/v1/core/recurrent_prefill_checkpoint.py index f64180cab63e..e421b476e4d2 100644 --- a/vllm/v1/core/recurrent_prefill_checkpoint.py +++ b/vllm/v1/core/recurrent_prefill_checkpoint.py @@ -54,14 +54,21 @@ def prefill_checkpoint_plan( publications: tuple[int, ...], shared_prefix_boundary: int = 0, ) -> CheckpointPlan | None: - """Export exact retained states inside a bounded pure-prompt chunk.""" + """Export exact retained states inside a bounded pure-prompt chunk. + + The span must start on a block boundary and end either on a block + boundary or at the prompt end. A prompt whose length is not a block + multiple therefore coalesces its final chunk too: the interior retained + states are exported in the same pass and the final state lands in the + tail block, exactly as an ordinary unaligned final chunk stores it. + """ if ( start < 0 or start % block_size or not start < end <= prompt or num_tokens != prompt or end - start > 8192 - or end % block_size + or (end % block_size and end != prompt) ): return None required = {position for position in publications if start < position < end} @@ -113,12 +120,12 @@ def continuation_layout( """ start, end, targets = plan validate_plan(plan, start, end, block_size) - if start <= 0 or start % block_size or end % block_size: - raise ValueError( - "continuation requires aligned nonzero source and final states" - ) + if start <= 0 or start % block_size: + raise ValueError("continuation requires an aligned nonzero source state") source = start // block_size - 1 - final = end // block_size - 1 + # The final state occupies the block holding the last token, so an + # unaligned prompt end maps to the same column an ordinary tail uses. + final = (end - 1) // block_size if speculative_blocks < 0 or len(blocks) != source + 1 + speculative_blocks: raise ValueError("continuation table does not end at its expected reserve") if blocks[source].is_null or blocks[source].ref_cnt < 1: From 7e73d9253b4663f8820fcfbb6dd14c505a0f3286 Mon Sep 17 00:00:00 2001 From: Jason Cook Date: Tue, 8 Sep 2026 15:37:50 -0400 Subject: [PATCH 2/2] Let coalesced prefill share a step with decoding requests Checkpoint coalescing was admitted only when the engine held exactly one request, so any decoding request in the system sent every new prompt down the aligned-split ladder. Checkpoint plans are built per request and the worker maps them by packed row, so decodes do not interfere. Require instead that at most one request still has prompt tokens to compute; two prefills in flight keep the ordinary splitting as before. Signed-off-by: Jason Cook --- .../core/test_recurrent_prefill_checkpoint.py | 21 ++++++++++++++++ vllm/v1/core/sched/scheduler.py | 24 +++++++++++++++---- 2 files changed, 41 insertions(+), 4 deletions(-) diff --git a/tests/v1/core/test_recurrent_prefill_checkpoint.py b/tests/v1/core/test_recurrent_prefill_checkpoint.py index 3c6a26609dd1..cf0f0ac25867 100644 --- a/tests/v1/core/test_recurrent_prefill_checkpoint.py +++ b/tests/v1/core/test_recurrent_prefill_checkpoint.py @@ -432,6 +432,27 @@ def test_dcp4_unaligned_prompt_tail_coalesces_into_one_chunk(prompt): assert len({blocks[column].block_id for column in columns}) == len(columns) +def test_coalescing_stays_exclusive_to_prompt_work_not_to_decodes(): + from vllm.v1.core.sched.scheduler import Scheduler + + scheduler = Scheduler.__new__(Scheduler) + scheduler._kda_coalescing_enabled = True + prefill = NS(num_computed_tokens=0, num_prompt_tokens=16228) + decode = NS(num_computed_tokens=8300, num_prompt_tokens=8192) + scheduler.running = [decode, decode, prefill] + scheduler.waiting = [] + scheduler.skipped_waiting = [] + assert scheduler._kda_coalescing_prefill_exclusive() + scheduler.waiting = [NS(num_computed_tokens=0, num_prompt_tokens=4096)] + assert not scheduler._kda_coalescing_prefill_exclusive() + scheduler.waiting = [] + scheduler.running = [prefill, NS(num_computed_tokens=8192, num_prompt_tokens=12288)] + assert not scheduler._kda_coalescing_prefill_exclusive() + scheduler._kda_coalescing_enabled = False + scheduler.running = [prefill] + assert not scheduler._kda_coalescing_prefill_exclusive() + + def test_two_checkpoint_capacity_keeps_safe_dcp4_fallback(): _, manager, scheduler, request = cache_fixture(8192, checkpoints=2) assert manager.hit_alignment_tokens == 512 diff --git a/vllm/v1/core/sched/scheduler.py b/vllm/v1/core/sched/scheduler.py index 0e3f694f274f..4025c255db55 100644 --- a/vllm/v1/core/sched/scheduler.py +++ b/vllm/v1/core/sched/scheduler.py @@ -557,6 +557,25 @@ def __init__( ), ) + def _kda_coalescing_prefill_exclusive(self) -> bool: + """True when at most one request still has prompt tokens to compute. + + Checkpoint plans are built per request and the worker maps them by + packed row, so decoding requests may share the step with the one + prefilling request. Two requests with prompt work in flight keep the + ordinary aligned splitting, as before. + """ + if not self._kda_coalescing_enabled: + return False + prefilling = 0 + for queue in (self.running, self.waiting, self.skipped_waiting): + for request in queue: + if request.num_computed_tokens < request.num_prompt_tokens: + prefilling += 1 + if prefilling > 1: + return False + return True + def _recurrent_checkpoint_plan( self, request: Request, start: int, end: int ) -> CheckpointPlan | None: @@ -890,10 +909,7 @@ def _has_waiting_boundary_logits(self) -> bool: def schedule(self, throttle_prefills: bool = False) -> SchedulerOutput: self.current_step += 1 - self._kda_coalescing_exclusive = ( - self._kda_coalescing_enabled - and len(self.running) + len(self.waiting) + len(self.skipped_waiting) == 1 - ) + self._kda_coalescing_exclusive = self._kda_coalescing_prefill_exclusive() # NOTE(woosuk) on the scheduling algorithm: # There's no "decoding phase" nor "prefill phase" in the scheduler. # Each request just has the num_computed_tokens and