Skip to content
Closed
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
34 changes: 34 additions & 0 deletions tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
OffloadingConnectorScheduler,
RequestOffloadState,
)
from vllm.v1.core.sched.output import SchedulerOutput
from vllm.v1.kv_cache_interface import (
FullAttentionSpec,
KVCacheGroupSpec,
Expand Down Expand Up @@ -64,6 +65,39 @@ def test_scheduler_reports_allocation_failure(request_runner):
assert reduced[_ConnectorMetricName.ALLOCATION_FAILURE] == 1


def test_store_defers_chunk_until_block_id_is_tracked(request_runner):
"""Store the ready prefix, then retry a chunk once its GPU block ID arrives."""
runner = request_runner(
block_size=4,
num_gpu_blocks=10,
async_scheduling=True,
)
runner.new_request(token_ids=[0] * 12)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)

req_status = runner.connector_scheduler._req_status["0"]
req_status.update_offload_keys()
group_state = req_status.group_states[0]
group_state.block_ids.extend([1, 2])

scheduler_output = SchedulerOutput.make_empty()
scheduler_output.num_scheduled_tokens = {"0": 12}
first_jobs = runner.connector_scheduler._build_store_jobs(scheduler_output)

assert len(first_jobs) == 1
assert next(iter(first_jobs.values())).src_spec.block_ids.tolist() == [1, 2]
assert group_state.next_stored_chunk_idx == 2

group_state.block_ids.append(3)
second_jobs = runner.connector_scheduler._build_store_jobs(scheduler_output)

assert len(second_jobs) == 1
assert next(iter(second_jobs.values())).src_spec.block_ids.tolist() == [3]
assert group_state.next_stored_chunk_idx == 3


@pytest.mark.parametrize("async_scheduling", [True, False])
@pytest.mark.parametrize("prompt_offset", [-1, -2])
def test_last_block_offloaded_at_request_finish(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -327,16 +327,15 @@ def storable_chunks(
num_chunks = max(0, num_chunks - 1)
return num_chunks

def advance_stored_idx(self, num_offloadable_tokens: int) -> None:
# max(): at the prefill->decode transition of a chunk-aligned prompt,
# storable_chunks drops by one (the eagle exclusion kicks in), and the
# index must not move backwards past already-stored chunks.
for group_config, group_state in zip(
self.config.kv_group_configs, self.group_states
def advance_stored_idx(self, num_chunks_by_group: Sequence[int]) -> None:
# Keep the cursor monotonic: the EAGLE exclusion can lower a group's
# storable boundary at the prefill-to-decode transition.
for num_chunks, group_state in zip(
num_chunks_by_group, self.group_states, strict=True
):
group_state.next_stored_chunk_idx = max(
group_state.next_stored_chunk_idx,
self.storable_chunks(group_config, num_offloadable_tokens),
num_chunks,
)

def update_num_hit_chunks(self, num_cached_tokens: int) -> None:
Expand Down Expand Up @@ -950,14 +949,38 @@ def _build_store_jobs(

# Filter out chunks skipped due to sliding window attention / SSM
# or unreachable by the load path's alignment constraints.
new_offload_keys: list[OffloadKey] = []
group_chunk_ends: list[int] = []
for group_config, group_state in zip(
self.config.kv_group_configs, req_status.group_states
):
num_chunks = req_status.storable_chunks(
num_storable_chunks = req_status.storable_chunks(
group_config, num_offloadable_tokens
)
num_tracked_chunks = len(group_state.block_ids) // blocks_per_chunk
num_chunks = min(
num_storable_chunks,
len(group_state.offload_keys),
num_tracked_chunks,
)
group_chunk_ends.append(num_chunks)
if num_chunks < num_storable_chunks:
logger.debug(
"Request %s deferring group %d offload chunks: "
"storable=%d keys=%d tracked_blocks=%d",
req_id,
group_config.group_idx,
num_storable_chunks,
len(group_state.offload_keys),
len(group_state.block_ids),
)

new_offload_keys: list[OffloadKey] = []
for group_config, group_state, num_chunks in zip(
self.config.kv_group_configs,
req_status.group_states,
group_chunk_ends,
strict=True,
):
start_chunk_idx = group_state.next_stored_chunk_idx
if num_chunks <= start_chunk_idx:
continue
Expand All @@ -972,7 +995,6 @@ def _build_store_jobs(
+ blocks_per_chunk
- 1 : num_chunks * blocks_per_chunk : blocks_per_chunk
]
assert len(offload_keys) == len(offload_block_ids)

alignment_chunk_count = group_config.alignment_chunk_count
tail = group_config.sliding_window_size_in_chunks
Expand All @@ -996,7 +1018,7 @@ def _build_store_jobs(
new_offload_keys.append(offload_key)

if not new_offload_keys:
req_status.advance_stored_idx(num_offloadable_tokens)
req_status.advance_stored_idx(group_chunk_ends)
self._maybe_cleanup_finished_req(req_id, req_status)
continue

Expand All @@ -1012,7 +1034,7 @@ def _build_store_jobs(
continue

if not store_output.keys_to_store:
req_status.advance_stored_idx(num_offloadable_tokens)
req_status.advance_stored_idx(group_chunk_ends)
self._maybe_cleanup_finished_req(req_id, req_status)
continue

Expand All @@ -1025,15 +1047,15 @@ def _build_store_jobs(
src_block_ids: list[int] = []
sliding_window_block_ids: list[int] = []
non_sliding_window_block_ids: list[int] = []
for group_config, group_state in zip(
self.config.kv_group_configs, req_status.group_states
for group_config, group_state, num_chunks in zip(
self.config.kv_group_configs,
req_status.group_states,
group_chunk_ends,
strict=True,
):
is_sliding_window = (
group_config.sliding_window_size_in_chunks is not None
)
num_chunks = req_status.storable_chunks(
group_config, num_offloadable_tokens
)
start_chunk_idx = group_state.next_stored_chunk_idx
block_ids = group_state.block_ids
num_group_blocks = 0
Expand Down