diff --git a/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py b/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py index 9a8783b084e..397451413c1 100644 --- a/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py +++ b/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py @@ -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 diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py index ef016bb02eb..ef87e36147d 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py @@ -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() diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py index 48dd61b50b8..6c75bda0c4c 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py @@ -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 diff --git a/vllm/v1/kv_offload/base.py b/vllm/v1/kv_offload/base.py index 6f80ba6252b..f4930cb7e6d 100644 --- a/vllm/v1/kv_offload/base.py +++ b/vllm/v1/kv_offload/base.py @@ -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 diff --git a/vllm/v1/kv_offload/cpu/manager.py b/vllm/v1/kv_offload/cpu/manager.py index 9751f616dcf..39e64933b5c 100644 --- a/vllm/v1/kv_offload/cpu/manager.py +++ b/vllm/v1/kv_offload/cpu/manager.py @@ -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 diff --git a/vllm/v1/kv_offload/cpu/policies/arc.py b/vllm/v1/kv_offload/cpu/policies/arc.py index e2af991c0cc..5b01815c2d7 100644 --- a/vllm/v1/kv_offload/cpu/policies/arc.py +++ b/vllm/v1/kv_offload/cpu/policies/arc.py @@ -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: diff --git a/vllm/v1/kv_offload/cpu/policies/base.py b/vllm/v1/kv_offload/cpu/policies/base.py index ee4916956e6..0febfe90d61 100644 --- a/vllm/v1/kv_offload/cpu/policies/base.py +++ b/vllm/v1/kv_offload/cpu/policies/base.py @@ -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. + """ diff --git a/vllm/v1/kv_offload/cpu/policies/lru.py b/vllm/v1/kv_offload/cpu/policies/lru.py index bf9a8b66e65..51680d8bcc5 100644 --- a/vllm/v1/kv_offload/cpu/policies/lru.py +++ b/vllm/v1/kv_offload/cpu/policies/lru.py @@ -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: