[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:
co-authored by
gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
Or Ozeri
parent
6548560496
commit
5bd8c71e79
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
|
||||
"""
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user