[kv_offload] Implement reset_cache() for the offloading connector (#41956)

Signed-off-by: Martin Hickey <martin.hickey@ie.ibm.com>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Co-authored-by: Or Ozeri <or@ozery.com>
This commit is contained in:
Martin Hickey
2026-05-14 16:00:10 +03:00
committed by GitHub
co-authored by gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Or Ozeri
parent 6548560496
commit 5bd8c71e79
8 changed files with 153 additions and 0 deletions
@@ -919,3 +919,80 @@ def test_complete_store_called_per_job(request_runner, async_scheduling: bool):
# Finish: no store pending -> no further call.
runner.run(decoded_tokens=[EOS_TOKEN_ID])
assert runner.manager.complete_store.call_count == 0
@pytest.mark.parametrize("async_scheduling", [True, False])
def test_reset_cache(request_runner, async_scheduling: bool):
"""reset_cache flushes in-flight loads, calls manager.reset_cache(), resets
next_stored_block_idx for active requests and clears job tracking."""
block_size = 4
block_size_factor = 3
offloaded_block_size = block_size * block_size_factor
num_gpu_blocks = 100
runner = request_runner(
block_size=block_size,
num_gpu_blocks=num_gpu_blocks,
async_scheduling=async_scheduling,
block_size_factor=block_size_factor,
)
# Store 1 offloaded block (3 GPU blocks) to CPU.
runner.new_request(token_ids=[0] * offloaded_block_size)
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output(keys)
)
runner.run(decoded_tokens=[EOS_TOKEN_ID], expected_stored=(0, 1, 2))
# Reset GPU prefix cache then start a request that loads from CPU.
# Leave the load in-flight so that reset_cache must flush it.
runner.scheduler.reset_prefix_cache()
runner.new_request(token_ids=[0] * offloaded_block_size)
runner.connector_scheduler._maximal_prefix_lookup = lambda key, req_context: 1
runner.manager.prepare_store.side_effect = lambda keys, req_context: (
generate_store_output([])
)
runner.run(decoded_tokens=[], complete_transfers=False)
# Capture in-flight load job IDs before reset.
load_job_ids = {
jid
for jid, status in runner.connector_scheduler._jobs.items()
if not status.is_store
}
assert load_job_ids, "expected in-flight load jobs before reset"
# Record job counter to verify the reset counter is set correctly.
job_counter_before_reset = runner.connector_scheduler._job_counter
# After update_state_after_alloc, next_stored_block_idx is advanced to
# skip the loaded prefix; reset_cache must bring it back to 0.
for req_status in runner.connector_scheduler._req_status.values():
for group_state in req_status.group_states:
assert group_state.next_stored_block_idx > 0
# Reset the cache
runner.connector_scheduler.reset_cache()
# manager.reset_cache() must be called exactly once.
runner.manager.reset_cache.assert_called_once()
# In-flight load jobs must be queued for flushing to prevent CUDA stream
# races between old loads and new post-reset stores.
assert load_job_ids <= runner.connector_scheduler._current_batch_jobs_to_flush
# All internal job tracking must be cleared.
assert not runner.connector_scheduler._jobs
assert not runner.connector_scheduler._block_id_to_pending_jobs
if runner.connector_scheduler._blocks_being_loaded is not None:
assert not runner.connector_scheduler._blocks_being_loaded
# Job reset counter must equal the job counter so that completions for
# pre-reset jobs arriving from workers are silently discarded.
assert runner.connector_scheduler._stale_job_threshold == job_counter_before_reset
# next_stored_block_idx must be reset to 0 for every active request so
# that post-reset stores restart from block 0.
for req_status in runner.connector_scheduler._req_status.values():
for group_state in req_status.group_states:
assert group_state.next_stored_block_idx == 0
@@ -223,6 +223,9 @@ class OffloadingConnectorScheduler:
# Job ID counter shared by loads and stores.
self._job_counter: int = 0
# Threshold value for stale jobs. All job ids >= _stale_job_threshold are
# active jobs.
self._stale_job_threshold: int = 0
self._jobs: dict[int, TransferJobStatus] = {}
# block_id -> pending store job_ids. Used to track jobs that needs
@@ -802,6 +805,13 @@ class OffloadingConnectorScheduler:
meta = OffloadingWorkerMetadata()
for job_id, count in meta.completed_jobs.items():
assert count > 0
if job_id < self._stale_job_threshold:
logger.debug(
"Skipping stale completed job %d (pre-reset counter: %d)",
job_id,
self._stale_job_threshold,
)
continue
job_status = self._jobs[job_id]
job_status.pending_count -= count
if job_status.pending_count > 0:
@@ -882,5 +892,33 @@ class OffloadingConnectorScheduler:
lora_name=None,
)
def reset_cache(self) -> None:
"""Reset the offloading manager cache, evicting all stored blocks."""
# reset_cache cannot be called in the middle of a schedule step
assert not self._current_batch_load_jobs
assert not self._current_batch_jobs_to_flush
# Flush all in-flight jobs
self._current_batch_jobs_to_flush.update(self._jobs.keys())
# Reset offloading manager cache
self.manager.reset_cache()
# Reset store progress so active requests re-offload from block 0
for status in self._req_status.values():
for group_state in status.group_states:
group_state.next_stored_block_idx = 0
# Discard jobs and save job_counter to be able to discard worker responses
self._stale_job_threshold = self._job_counter
self._jobs.clear()
self._block_id_to_pending_jobs.clear()
# Note: _current_batch_jobs_to_flush is intentionally NOT cleared.
# The load flush IDs collected above must be delivered to workers.
if self._blocks_being_loaded is not None:
self._blocks_being_loaded.clear()
def shutdown(self) -> None:
self.manager.shutdown()
@@ -174,6 +174,11 @@ class OffloadingConnector(KVConnectorBase_V1, SupportsHMA):
def get_required_kvcache_layout(cls, vllm_config: VllmConfig) -> str | None:
return "HND"
def reset_cache(self) -> bool | None:
assert self.connector_scheduler is not None
self.connector_scheduler.reset_cache()
return True
def get_kv_connector_stats(self) -> KVConnectorStats | None:
if self.connector_worker is None:
return None # We only emit stats from the worker-side
+4
View File
@@ -218,6 +218,10 @@ class OffloadingManager(ABC):
"""
return ()
def reset_cache(self) -> None:
"""Evict all tracked blocks and reset internal state."""
return
def shutdown(self) -> None:
"""Shutdown the manager and release any resources."""
return
+11
View File
@@ -223,6 +223,17 @@ class CPUOffloadingManager(OffloadingManager):
)
)
def reset_cache(self) -> None:
# Clear ALL blocks unconditionally. The scheduler's _stale_job_threshold
# guarantees that complete_load / complete_store are never called for
# pre-reset jobs, so no lazy cleanup is needed. The scheduler also
# flushes in-flight load job IDs to the workers before any new stores
# can begin, preventing a cross-direction data race on reused offload block IDs.
self._policy.clear()
self._free_list.clear()
self._num_allocated_blocks = 0
def take_events(self) -> Iterable[OffloadingEvent]:
if self.events is not None:
yield from self.events
+7
View File
@@ -94,6 +94,13 @@ class ARCCachePolicy(CachePolicy):
# move to MRU position (end) to keep it fresh in the ghost list
self.b2.move_to_end(key)
def clear(self) -> None:
self.t1.clear()
self.t2.clear()
self.b1.clear()
self.b2.clear()
self.target_t1_size = 0.0
def evict(
self, n: int, protected: set[OffloadKey]
) -> list[tuple[OffloadKey, BlockStatus]] | None:
+8
View File
@@ -74,3 +74,11 @@ class CachePolicy(ABC):
For ARC: ghost list cleanup (trimming to cache_capacity) is performed
at the end of a successful eviction.
"""
@abstractmethod
def clear(self) -> None:
"""
Remove ALL blocks regardless of ref_cnt.
Ghost lists and adaptive state are also reset.
"""
+3
View File
@@ -28,6 +28,9 @@ class LRUCachePolicy(CachePolicy):
if key in self.blocks:
self.blocks.move_to_end(key)
def clear(self) -> None:
self.blocks.clear()
def evict(
self, n: int, protected: set[OffloadKey]
) -> list[tuple[OffloadKey, BlockStatus]] | None: