diff --git a/docs/features/sleep_mode.md b/docs/features/sleep_mode.md index 2616f374ea5a..b4647e435209 100644 --- a/docs/features/sleep_mode.md +++ b/docs/features/sleep_mode.md @@ -100,6 +100,10 @@ VLLM_SERVER_DEV_MODE=1 vllm serve Qwen/Qwen3-0.6B \ --port 8000 ``` +`sleep(level=1, clear_connector_cache=False)` keeps a KV connector's offloaded blocks, so +requests after the wake-up are served from them instead of re-prefilled. Block hashes do not +cover the weights, so leave the default in place whenever the weights are replaced. + Below is an example of how to sleep and wake up a model in level 1. ```bash diff --git a/docs/training/async_rl.md b/docs/training/async_rl.md index 9e75a24eaa1a..3f18686dac6a 100644 --- a/docs/training/async_rl.md +++ b/docs/training/async_rl.md @@ -28,6 +28,10 @@ The `mode` parameter controls how in-flight requests are handled: The `clear_cache` parameter controls whether to clear the KV cache and prefix cache after pausing. +`clear_cache=True` also resets a configured KV connector's cache. `clear_connector_cache=False` +keeps it while still clearing the local caches, so later requests reuse the retained blocks. Block +hashes do not cover the weights, so only do this when the pause does not replace them. + ### resume_generation ```python diff --git a/rust/proto/control.proto b/rust/proto/control.proto index ef3218d35c91..3379bd93a1a5 100644 --- a/rust/proto/control.proto +++ b/rust/proto/control.proto @@ -126,6 +126,7 @@ enum PauseMode { message PauseGenerationRequest { PauseMode mode = 1; optional bool clear_cache = 2; + optional bool clear_connector_cache = 3; } message PauseGenerationResponse {} @@ -138,6 +139,7 @@ message IsPausedResponse { bool paused = 1; } message SleepRequest { optional uint32 level = 1; PauseMode mode = 2; + optional bool clear_connector_cache = 3; } message SleepResponse {} diff --git a/rust/src/engine-core-client/examples/external_engine_utility_call.rs b/rust/src/engine-core-client/examples/external_engine_utility_call.rs index b1826bb76e69..fcee5eeea8b9 100644 --- a/rust/src/engine-core-client/examples/external_engine_utility_call.rs +++ b/rust/src/engine-core-client/examples/external_engine_utility_call.rs @@ -110,7 +110,7 @@ async fn main() -> Result<()> { if args.skip_sleep_wake { println!("sleep_wake=skipped"); } else { - client.sleep(args.sleep_level, args.sleep_mode).await.with_context(|| { + client.sleep(args.sleep_level, args.sleep_mode, true).await.with_context(|| { format!( "failed to call sleep utility with level={} mode={}", args.sleep_level, args.sleep_mode diff --git a/rust/src/engine-core-client/src/client.rs b/rust/src/engine-core-client/src/client.rs index aa6aacf0e017..299cf827df54 100644 --- a/rust/src/engine-core-client/src/client.rs +++ b/rust/src/engine-core-client/src/client.rs @@ -906,8 +906,19 @@ impl EngineCoreClient { } /// Put the engine to sleep. - pub async fn sleep(&self, level: u32, mode: PauseMode) -> Result<()> { - self.call_utility::<(), _>("sleep", (level, mode)).await?; + pub async fn sleep( + &self, + level: u32, + mode: PauseMode, + clear_connector_cache: bool, + ) -> Result<()> { + // Only the opt-out sends the third argument, so a default call still + // drives an engine that predates it. + if clear_connector_cache { + self.call_utility::<(), _>("sleep", (level, mode)).await?; + } else { + self.call_utility::<(), _>("sleep", (level, mode, false)).await?; + } Ok(()) } @@ -928,8 +939,18 @@ impl EngineCoreClient { } /// Pause the scheduler so generation can be halted - pub async fn pause_scheduler(&self, mode: PauseMode, clear_cache: bool) -> Result<()> { - self.call_utility::<(), _>("pause_scheduler", (mode, clear_cache)).await?; + pub async fn pause_scheduler( + &self, + mode: PauseMode, + clear_cache: bool, + clear_connector_cache: bool, + ) -> Result<()> { + if clear_connector_cache { + self.call_utility::<(), _>("pause_scheduler", (mode, clear_cache)).await?; + } else { + self.call_utility::<(), _>("pause_scheduler", (mode, clear_cache, false)) + .await?; + } Ok(()) } diff --git a/rust/src/server/src/grpc/control.rs b/rust/src/server/src/grpc/control.rs index 680ebdfbdaf2..a52ddf8f9ace 100644 --- a/rust/src/server/src/grpc/control.rs +++ b/rust/src/server/src/grpc/control.rs @@ -316,9 +316,10 @@ impl pb::control_server::Control for ControlServiceImpl { let request = request.into_inner(); let mode = pause_mode(request.mode)?; let clear_cache = request.clear_cache.unwrap_or(true); + let clear_connector_cache = request.clear_connector_cache.unwrap_or(true); let _guard = self.rl_lock.lock().await; self.client() - .pause_scheduler(mode, clear_cache) + .pause_scheduler(mode, clear_cache, clear_connector_cache) .await .map_err(|error| utility_status("pause_generation", error))?; Ok(Response::new(pb::PauseGenerationResponse {})) @@ -356,9 +357,10 @@ impl pb::control_server::Control for ControlServiceImpl { let request = request.into_inner(); let mode = pause_mode(request.mode)?; let level = request.level.unwrap_or(1); + let clear_connector_cache = request.clear_connector_cache.unwrap_or(true); let _guard = self.rl_lock.lock().await; self.client() - .sleep(level, mode) + .sleep(level, mode, clear_connector_cache) .await .map_err(|error| utility_status("sleep", error))?; Ok(Response::new(pb::SleepResponse {})) diff --git a/rust/src/server/src/routes/pause.rs b/rust/src/server/src/routes/pause.rs index 1654c85b0200..e2f752a1ed58 100644 --- a/rust/src/server/src/routes/pause.rs +++ b/rust/src/server/src/routes/pause.rs @@ -19,6 +19,8 @@ pub(crate) struct PauseParams { mode: PauseMode, #[serde(default = "default_clear_cache")] clear_cache: bool, + #[serde(default = "default_clear_connector_cache")] + clear_connector_cache: bool, } #[derive(Serialize)] @@ -35,6 +37,10 @@ const fn default_clear_cache() -> bool { true } +const fn default_clear_connector_cache() -> bool { + true +} + fn invalid_query(error: QueryRejection) -> ApiError { ApiError::invalid_request(error.body_text(), Some("mode")) } @@ -52,7 +58,11 @@ pub async fn pause( state .engine_core_client() - .pause_scheduler(params.mode, params.clear_cache) + .pause_scheduler( + params.mode, + params.clear_cache, + params.clear_connector_cache, + ) .await .map_err(|error| utility_call_error("pause", error))?; diff --git a/rust/src/server/src/routes/sleep.rs b/rust/src/server/src/routes/sleep.rs index d64a0850b6e4..80af3d7d3435 100644 --- a/rust/src/server/src/routes/sleep.rs +++ b/rust/src/server/src/routes/sleep.rs @@ -25,12 +25,18 @@ pub(crate) struct SleepParams { level: u32, #[serde(default)] mode: PauseMode, + #[serde(default = "default_clear_connector_cache")] + clear_connector_cache: bool, } const fn default_sleep_level() -> u32 { 1 } +const fn default_clear_connector_cache() -> bool { + true +} + fn invalid_query(error: QueryRejection) -> ApiError { ApiError::invalid_request(error.body_text(), Some("mode")) } @@ -44,7 +50,7 @@ pub async fn sleep( state .engine_core_client() - .sleep(params.level, params.mode) + .sleep(params.level, params.mode, params.clear_connector_cache) .await .map_err(|error| utility_call_error("sleep", error))?; diff --git a/rust/src/server/src/routes/tests.rs b/rust/src/server/src/routes/tests.rs index 89ce35ce6845..378ab875de63 100644 --- a/rust/src/server/src/routes/tests.rs +++ b/rust/src/server/src/routes/tests.rs @@ -5854,6 +5854,48 @@ async fn sleep_route_uses_python_compatible_default_query_values() { engine_task.await.expect("mock engine task"); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn sleep_route_sends_connector_opt_out_only_when_requested() { + // Only the opt-out widens the call, so the default still drives an older engine. + let (app, engine_task) = test_admin_app_with_engine_script(|dealer, push| { + boxed_test_future(async move { + let utility = recv_engine_message(dealer).await; + let payload = decode_value(&utility[1]).expect("decode utility payload"); + let array = payload.as_array().expect("utility payload array"); + let call_id = array[1].as_u64().expect("call id"); + + assert_eq!(array[2], Value::from("sleep")); + assert_eq!( + array[3], + Value::Array(vec![ + Value::from(1_u64), + Value::from("abort"), + Value::from(false), + ]) + ); + + send_outputs(push, utility_outputs(call_id, utility_none_result())).await; + }) + }) + .await; + + let response = app + .clone() + .call( + Request::builder() + .method("POST") + .uri("/sleep?clear_connector_cache=false") + .body(Body::empty()) + .expect("build request"), + ) + .await + .expect("call app"); + + assert_eq!(response.status(), StatusCode::OK); + engine_task.await.expect("mock engine task"); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] async fn release_kv_cache_memory_route_sends_expected_utility_call() { @@ -6064,6 +6106,48 @@ async fn pause_route_uses_python_compatible_default_query_values() { ); } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[serial] +async fn pause_route_sends_connector_opt_out_only_when_requested() { + // Only the opt-out widens the call, so the default still drives an older engine. + let (app, engine_task) = test_admin_app_with_engine_script(|dealer, push| { + boxed_test_future(async move { + let utility = recv_engine_message(dealer).await; + let payload = decode_value(&utility[1]).expect("decode utility payload"); + let array = payload.as_array().expect("utility payload array"); + let call_id = array[1].as_u64().expect("call id"); + + assert_eq!(array[2], Value::from("pause_scheduler")); + assert_eq!( + array[3], + Value::Array(vec![ + Value::from("abort"), + Value::from(true), + Value::from(false), + ]) + ); + + send_outputs(push, utility_outputs(call_id, utility_none_result())).await; + }) + }) + .await; + + let response = app + .clone() + .call( + Request::builder() + .method("POST") + .uri("/pause?clear_connector_cache=false") + .body(Body::empty()) + .expect("build request"), + ) + .await + .expect("call app"); + + assert_eq!(response.status(), StatusCode::OK); + engine_task.await.expect("mock engine task"); +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[serial] async fn pause_route_rejects_invalid_mode() { diff --git a/tests/v1/core/test_scheduler.py b/tests/v1/core/test_scheduler.py index 4951014d1ba1..c361468c0d54 100644 --- a/tests/v1/core/test_scheduler.py +++ b/tests/v1/core/test_scheduler.py @@ -50,6 +50,7 @@ KVCacheGroupSpec, MambaSpec, ) +from vllm.v1.metrics.stats import PrefixCacheStats from vllm.v1.outputs import ( DraftTokenIds, ECConnectorOutput, @@ -1563,6 +1564,47 @@ def test_kv_cache_release_rejects_nonresident_memory(sleeping_tags): assert core.model_executor.sleeping_tags == sleeping_tags +@pytest.mark.parametrize("clear_connector_cache", [True, False]) +def test_pause_clear_connector_cache_opt_out(clear_connector_cache): + """The tier survives the pause, so a caller keeping the weights may keep it.""" + scheduler = create_scheduler(enable_prefix_caching=True) + scheduler.connector = Mock( + supports_retained_cache_on_pause=True, + **{"has_pending_push_work.return_value": False}, + ) + scheduler.connector_prefix_cache_stats = PrefixCacheStats() + + core = object.__new__(EngineCore) + core.scheduler = scheduler + core.model_executor = Mock(is_sleeping=False) + core.mm_receiver_cache = None + core.batch_queue = None + + core.pause_scheduler( + mode="abort", + clear_cache=True, + clear_connector_cache=clear_connector_cache, + ) + + assert scheduler.connector.reset_cache.called is clear_connector_cache + + +def test_pause_refuses_retention_a_connector_cannot_take(): + """Keeping the cache is only safe where the connector expects the resume.""" + scheduler = create_scheduler(enable_prefix_caching=True) + scheduler.connector = Mock(supports_retained_cache_on_pause=False) + + core = object.__new__(EngineCore) + core.scheduler = scheduler + core.model_executor = Mock(is_sleeping=False) + + with pytest.raises(ValueError, match="cannot keep its cache"): + core.pause_scheduler(clear_cache=True, clear_connector_cache=False) + + scheduler.connector.reset_cache.assert_not_called() + assert scheduler.pause_state == PauseState.UNPAUSED + + def test_reset_connector_cache_no_connector_is_no_op_success(): """``reset_connector_cache`` must return True when no connector is configured. diff --git a/tests/v1/engine/test_engine_core.py b/tests/v1/engine/test_engine_core.py index e16e0550c21d..87def6d918ce 100644 --- a/tests/v1/engine/test_engine_core.py +++ b/tests/v1/engine/test_engine_core.py @@ -691,6 +691,46 @@ def test_kv_cache_release_rejects_unsafe_state(pause_state, has_requests, has_ba core.model_executor.discard.assert_not_called() +@pytest.mark.parametrize("deferred", [False, True]) +@pytest.mark.parametrize( + "kwargs,expected", [({}, True), (dict(clear_connector_cache=False), False)] +) +def test_pause_forwards_connector_flag(deferred: bool, kwargs, expected): + """Immediate and deferred pauses both forward it, and omitting it clears.""" + core = _pausable_engine_core_proc() + core.engines_running = deferred + seen: list[bool] = [] + core._reset_caches = lambda reset_connector=True: seen.append(reset_connector) + + result = EngineCoreProc.pause_scheduler( + core, mode="keep", clear_cache=True, **kwargs + ) + if deferred: + assert seen == [] + core.engines_running = False + core._notify_idle_state_callbacks() + assert result.result(timeout=0) is None + + assert seen == [expected] + + +@pytest.mark.parametrize("level", [1, 2]) +def test_sleep_forwards_connector_flag(level: int): + """sleep() is the entry point RL callers use, not pause_scheduler.""" + core = object.__new__(EngineCore) + core.model_executor = MagicMock() + core.scheduler = MagicMock() + core.scheduler.has_requests.return_value = False + core.batch_queue = None + seen: list[bool] = [] + core._reset_caches = lambda reset_connector=True: seen.append(reset_connector) + + EngineCore.sleep(core, level=level, clear_connector_cache=False) + + assert seen == [False] + core.model_executor.sleep.assert_called_once_with(level) + + @pytest.mark.parametrize("deferred", [False, True]) def test_pause_synchronizes_device_before_cache_reset(deferred: bool): """A resolved pause promises an idle device: the barrier must run before @@ -699,7 +739,7 @@ def test_pause_synchronizes_device_before_cache_reset(deferred: bool): core.engines_running = deferred order: list[str] = [] core.model_executor.collective_rpc.side_effect = lambda method: order.append(method) - core._reset_caches = lambda: order.append("reset_caches") + core._reset_caches = lambda **kwargs: order.append("reset_caches") result = EngineCoreProc.pause_scheduler(core, mode="keep", clear_cache=True) if deferred: 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 035fb6b0054f..8c8357f412ef 100644 --- a/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py +++ b/tests/v1/kv_connector/unit/offloading_connector/test_scheduler.py @@ -989,6 +989,48 @@ def test_offloading_connector(request_runner, async_scheduling: bool): runner.run(decoded_tokens=[EOS_TOKEN_ID], expected_loaded=(3, 4, 5)) +@pytest.mark.parametrize("async_scheduling", [False, True]) +def test_preempted_request_resumed_in_the_same_step( + request_runner, async_scheduling: bool +): + """A sleep that keeps the tier preempts and resumes in one step. + + `sleep(clear_cache=True)` preempts the running requests, and nothing is + scheduled until the wake-up, so the step that reports the preemption is + also the step that resumes the request. With the offloaded tier kept, that + step already carries the resume load, and the preemption flush must not + take it for one of the stores it is there to settle. + """ + block_size, blocks_per_chunk = 4, 3 + tokens_per_chunk = block_size * blocks_per_chunk + runner = request_runner( + block_size=block_size, + num_gpu_blocks=100, + async_scheduling=async_scheduling, + blocks_per_chunk=blocks_per_chunk, + ) + + runner.new_request(token_ids=[0] * tokens_per_chunk * 2) + runner.manager.prepare_store.side_effect = lambda keys, ctx: generate_store_output( + keys + ) + runner._run([0], complete_transfers=True) + # Retire the store, so the resume step starts with no job registered. + runner._run([0] * 4, complete_transfers=True) + + # What _reset_caches(reset_connector=False) does on the way into sleep. + runner.scheduler.reset_prefix_cache( + reset_running_requests=True, reset_connector=False + ) + assert runner.scheduler.reset_preempted_req_ids + + runner.connector_scheduler._maximal_prefix_lookup = lambda keys, ctx, *_: 2 + runner.manager.prepare_store.side_effect = lambda keys, ctx: generate_store_output( + keys + ) + runner._run([0], complete_transfers=True) + + @pytest.mark.parametrize("async_scheduling", [True, False]) def test_request_preemption(request_runner, async_scheduling: bool): block_size = 4 diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/base.py b/vllm/distributed/kv_transfer/kv_connector/v1/base.py index 072018c039ed..0bf7cd78d5e0 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/base.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/base.py @@ -727,6 +727,10 @@ def build_prom_metrics( """ return None + # Whether the connector tolerates surviving a pause: the resume then lands + # in the step that reports the pause's preemptions, carrying a load. + supports_retained_cache_on_pause: bool = False + def reset_cache(self) -> bool | None: """Reset the connector's internal cache. diff --git a/vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py b/vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py index a6f0d4cafb75..e964ed6b3643 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/multi_connector.py @@ -696,6 +696,10 @@ def build_prom_metrics( prom_metrics, ) + @property + def supports_retained_cache_on_pause(self) -> bool: # type: ignore[override] + return all(c.supports_retained_cache_on_pause for c in self._connectors) + def reset_cache(self) -> bool: results = [c.reset_cache() is not False for c in self._connectors] return all(results) 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 d6c210a6f9bf..ea27bdc1327e 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/offloading/scheduler.py @@ -1786,11 +1786,15 @@ def build_connector_meta( # Flush jobs for preempted requests. for req_id in scheduler_output.preempted_req_ids or (): req_status = self._req_status.get(req_id) - if req_status is None or not req_status.transfer_jobs: + if req_status is None: continue - any_jid = next(iter(req_status.transfer_jobs)) - assert self._jobs[any_jid].is_store - self._current_batch_jobs_to_flush.update(req_status.transfer_jobs) + # A wake-up that kept the offloaded tier resumes the request in the + # step that reports the preemption, so its load is already here; it + # reads no freed block, and only the earlier stores need flushing. + store_jobs = { + jid for jid in req_status.transfer_jobs if self._jobs[jid].is_store + } + self._current_batch_jobs_to_flush.update(store_jobs) # Flush jobs that contain re-allocated blocks. if ( 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 73570881753d..c321935c8cb9 100644 --- a/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py +++ b/vllm/distributed/kv_transfer/kv_connector/v1/offloading_connector.py @@ -232,6 +232,8 @@ def get_required_kvcache_layout(cls, vllm_config: VllmConfig) -> str | None: return "BLHNC" return "LBHNC" + supports_retained_cache_on_pause: bool = True + def reset_cache(self) -> bool | None: assert self.connector_scheduler is not None self.connector_scheduler.reset_cache() diff --git a/vllm/engine/protocol.py b/vllm/engine/protocol.py index 1e6011a2300f..21dbf659639f 100644 --- a/vllm/engine/protocol.py +++ b/vllm/engine/protocol.py @@ -187,7 +187,12 @@ async def reset_prefix_cache( ... @abstractmethod - async def sleep(self, level: int = 1, mode: "PauseMode" = "abort") -> None: + async def sleep( + self, + level: int = 1, + mode: "PauseMode" = "abort", + clear_connector_cache: bool = True, + ) -> None: """Sleep the engine.""" ... @@ -218,6 +223,7 @@ async def pause_generation( mode: "PauseMode" = "abort", wait_for_inflight_requests: bool = False, clear_cache: bool = True, + clear_connector_cache: bool = True, ) -> None: """Pause new generation/encoding requests. @@ -231,6 +237,7 @@ async def pause_generation( wait_for_inflight_requests: DEPRECATED. Use ``mode="wait"`` instead. clear_cache: DEPRECATED. Whether to clear KV and prefix caches after draining. + clear_connector_cache: Evict the KV tier too; unsafe if weights change. """ ... diff --git a/vllm/entrypoints/llm.py b/vllm/entrypoints/llm.py index 999310aa6169..623083267357 100644 --- a/vllm/entrypoints/llm.py +++ b/vllm/entrypoints/llm.py @@ -805,7 +805,12 @@ def reset_prefix_cache( reset_running_requests, reset_connector ) - def sleep(self, level: int = 1, mode: PauseMode = "abort"): + def sleep( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ): """Put the engine to sleep. The engine should not process any requests. The caller should guarantee that no requests are being processed during the sleep period, before `wake_up` is called. @@ -826,9 +831,12 @@ def sleep(self, level: int = 1, mode: PauseMode = "abort"): CPU memory pressure. mode: How to handle any existing requests, can be "abort", "wait", or "keep". + clear_connector_cache: Evict the KV tier too; unsafe if weights change. """ - self.llm_engine.sleep(level=level, mode=mode) + self.llm_engine.sleep( + level=level, mode=mode, clear_connector_cache=clear_connector_cache + ) def release_kv_cache_memory(self) -> None: """Release the GPU physical memory backing the KV cache. diff --git a/vllm/entrypoints/serve/dev/rlhf/api_router.py b/vllm/entrypoints/serve/dev/rlhf/api_router.py index 603fdf785c3d..551bf219cef0 100644 --- a/vllm/entrypoints/serve/dev/rlhf/api_router.py +++ b/vllm/entrypoints/serve/dev/rlhf/api_router.py @@ -32,6 +32,7 @@ async def pause_generation( mode: Annotated[PauseMode, Query()] = "abort", wait_for_inflight_requests: bool = Query(False), clear_cache: Annotated[bool, Query()] = True, + clear_connector_cache: Annotated[bool, Query()] = True, ) -> JSONResponse: """Pause generation requests to allow weight updates. @@ -44,7 +45,8 @@ async def pause_generation( - ``"keep"``: Freeze requests in queue; they resume on /resume. wait_for_inflight_requests: DEPRECATED. Use ``mode="wait"`` instead. clear_cache: DEPRECATED. Whether to clear KV/prefix caches after - draining. Ignored when mode="keep". + draining. Applies to every mode, "keep" included. + clear_connector_cache: Evict the KV tier too; unsafe if weights change. """ engine = engine_client(raw_request) @@ -53,6 +55,7 @@ async def pause_generation( await engine.pause_generation( mode=mode, clear_cache=clear_cache, + clear_connector_cache=clear_connector_cache, wait_for_inflight_requests=wait_for_inflight_requests, ) return JSONResponse( diff --git a/vllm/entrypoints/serve/dev/sleep/api_router.py b/vllm/entrypoints/serve/dev/sleep/api_router.py index 846a487706b7..e99f979e5286 100644 --- a/vllm/entrypoints/serve/dev/sleep/api_router.py +++ b/vllm/entrypoints/serve/dev/sleep/api_router.py @@ -2,7 +2,9 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project -from fastapi import APIRouter, FastAPI, Request +from typing import Annotated + +from fastapi import APIRouter, FastAPI, Query, Request from fastapi.responses import JSONResponse, Response from vllm.engine.protocol import EngineClient @@ -19,11 +21,14 @@ def engine_client(request: Request) -> EngineClient: @router.post("/sleep") -async def sleep(raw_request: Request): +async def sleep( + raw_request: Request, + clear_connector_cache: Annotated[bool, Query()] = True, +): # get POST params level = raw_request.query_params.get("level", "1") mode = raw_request.query_params.get("mode", "abort") - await engine_client(raw_request).sleep(int(level), mode) + await engine_client(raw_request).sleep(int(level), mode, clear_connector_cache) return Response(status_code=200) diff --git a/vllm/v1/engine/async_llm.py b/vllm/v1/engine/async_llm.py index 67b8ac560d52..1aa6afde45eb 100644 --- a/vllm/v1/engine/async_llm.py +++ b/vllm/v1/engine/async_llm.py @@ -906,6 +906,7 @@ async def pause_generation( mode: PauseMode = "abort", wait_for_inflight_requests: bool | None = None, clear_cache: bool = True, + clear_connector_cache: bool = True, ) -> None: """Pause generation to allow model weight updates. @@ -923,6 +924,12 @@ async def pause_generation( wait_for_inflight_requests: DEPRECATED: use mode argument. clear_cache: Whether to clear KV cache and prefix cache after draining. Set to ``False`` to preserve cache for faster resume. + clear_connector_cache: Whether clearing also evicts an external KV + connector tier. Only applies when ``clear_cache`` is set. That + tier survives the pause, so a caller that does not replace the + weights can keep it and resume from the offloaded blocks; + block hashes do not cover the weights, so a caller that does + replace them must leave it at ``True``. """ if wait_for_inflight_requests: @@ -936,7 +943,11 @@ async def pause_generation( mode = "wait" if clear_cache: await self.renderer.clear_mm_cache_async() - await self.engine_core.pause_scheduler_async(mode=mode, clear_cache=clear_cache) + await self.engine_core.pause_scheduler_async( + mode=mode, + clear_cache=clear_cache, + clear_connector_cache=clear_connector_cache, + ) # Small sleep to help ensure that final outputs from any in-flight requests are # returned prior to this method returning. These outputs come out of the engine # prior to the wait-for-idle completion event, but involve additional async @@ -1082,10 +1093,15 @@ async def reset_prefix_cache( async def reset_encoder_cache(self) -> None: await self.engine_core.reset_encoder_cache_async() - async def sleep(self, level: int = 1, mode: PauseMode = "abort") -> None: + async def sleep( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ) -> None: if level >= 1: await self.renderer.clear_mm_cache_async() - await self.engine_core.sleep_async(level, mode) + await self.engine_core.sleep_async(level, mode, clear_connector_cache) if self.logger_manager is not None: self.logger_manager.record_sleep_state(1, level) diff --git a/vllm/v1/engine/core.py b/vllm/v1/engine/core.py index a51278b7617c..457c248584de 100644 --- a/vllm/v1/engine/core.py +++ b/vllm/v1/engine/core.py @@ -875,15 +875,30 @@ def _reset_caches( self.reset_mm_cache() self.reset_encoder_cache() - def _finish_pause(self, clear_cache: bool) -> None: + def _refuse_unsafe_retention(self, clear_connector_cache: bool) -> None: + connector = getattr(self.scheduler, "connector", None) + if clear_connector_cache or connector is None: + return + if not connector.supports_retained_cache_on_pause: + raise ValueError( + f"{type(connector).__name__} cannot keep its cache across a " + "pause: the resume lands in the step that reports the pause's " + "preemptions, which this connector does not expect. Drop " + "clear_connector_cache=False." + ) + + def _finish_pause(self, clear_cache: bool, clear_connector_cache: bool) -> None: # A completed pause promises an idle device: nothing else waits on # the last dummy batch an idle DP rank launches. self.model_executor.collective_rpc("synchronize_device") if clear_cache: - self._reset_caches() + self._reset_caches(reset_connector=clear_connector_cache) def pause_scheduler( - self, mode: PauseMode = "abort", clear_cache: bool = True + self, + mode: PauseMode = "abort", + clear_cache: bool = True, + clear_connector_cache: bool = True, ) -> Future | None: """Pause generation; behavior depends on mode. @@ -897,18 +912,23 @@ def pause_scheduler( optionally clear caches, then complete the returned Future. - ``keep``: Set PAUSED_ALL; return a Future that completes when the output queue is empty. + + ``clear_connector_cache=False`` keeps the external KV tier, which + survives the pause; unsafe if the weights change. """ if mode not in get_args(PauseMode): raise ValueError(f"Invalid pause mode: {mode}") if mode == "wait": raise ValueError("'wait' mode can't be used in inproc-engine mode") + if clear_cache: + self._refuse_unsafe_retention(clear_connector_cache) if mode == "abort": self.scheduler.finish_requests(None, RequestStatus.FINISHED_ABORTED) pause_state = PauseState.PAUSED_ALL if mode == "keep" else PauseState.PAUSED_NEW self.scheduler.set_pause_state(pause_state) - self._finish_pause(clear_cache) + self._finish_pause(clear_cache, clear_connector_cache) return None @@ -920,7 +940,12 @@ def is_scheduler_paused(self) -> bool: """Return whether the scheduler is in any pause state.""" return self.scheduler.pause_state != PauseState.UNPAUSED - def sleep(self, level: int = 1, mode: PauseMode = "abort") -> None | Future: + def sleep( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ) -> None | Future: """Put the engine to sleep at the specified level. Args: @@ -931,11 +956,16 @@ def sleep(self, level: int = 1, mode: PauseMode = "abort") -> None | Future: - Level 2: Discard all GPU memory. mode: Pause mode - how to deal with any existing requests, see documentation of pause_scheduler method. + clear_connector_cache: See pause_scheduler; ignored at level 0. """ # Pause scheduler before sleeping. clear_prefix_cache = level >= 1 - pause_future = self.pause_scheduler(mode=mode, clear_cache=clear_prefix_cache) + pause_future = self.pause_scheduler( + mode=mode, + clear_cache=clear_prefix_cache, + clear_connector_cache=clear_connector_cache, + ) if level < 1: return pause_future @@ -1983,7 +2013,10 @@ def _handle_request_preproc_error(self, request: EngineCoreRequest) -> None: self._send_error_outputs_to_client([request.request_id], request.client_index) def pause_scheduler( - self, mode: PauseMode = "abort", clear_cache: bool = True + self, + mode: PauseMode = "abort", + clear_cache: bool = True, + clear_connector_cache: bool = True, ) -> Future | None: """Pause generation; behavior depends on mode. @@ -1997,13 +2030,18 @@ def pause_scheduler( optionally clear caches, then complete the returned Future. - ``keep``: Set PAUSED_ALL; return a Future that completes when the output queue is empty. + + ``clear_connector_cache=False`` keeps the external KV tier, which + survives the pause; unsafe if the weights change. """ if mode not in get_args(PauseMode): raise ValueError(f"Invalid pause mode: {mode}") + if clear_cache: + self._refuse_unsafe_retention(clear_connector_cache) def engine_idle_callback(engine: "EngineCoreProc", future: Future[Any]) -> None: try: - engine._finish_pause(clear_cache) + engine._finish_pause(clear_cache, clear_connector_cache) except Exception as e: future.set_exception(e) else: @@ -2019,7 +2057,7 @@ def engine_idle_callback(engine: "EngineCoreProc", future: Future[Any]) -> None: self.scheduler.set_pause_state(pause_state) if self._pause_complete(): - self._finish_pause(clear_cache) + self._finish_pause(clear_cache, clear_connector_cache) return None future = Future[Any]() diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index 6ad05e1f8ced..2090e3c6f3f6 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -195,7 +195,12 @@ def reset_prefix_cache( def reset_encoder_cache(self) -> None: raise NotImplementedError - def sleep(self, level: int = 1, mode: PauseMode = "abort") -> None: + def sleep( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ) -> None: raise NotImplementedError def release_kv_cache_memory(self) -> None: @@ -290,7 +295,12 @@ async def reset_prefix_cache_async( async def reset_encoder_cache_async(self) -> None: raise NotImplementedError - async def sleep_async(self, level: int = 1, mode: PauseMode = "abort") -> None: + async def sleep_async( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ) -> None: raise NotImplementedError async def release_kv_cache_memory_async(self) -> None: @@ -402,10 +412,15 @@ def reset_prefix_cache( def reset_encoder_cache(self) -> None: self.engine_core.reset_encoder_cache() - def sleep(self, level: int = 1, mode: PauseMode = "abort") -> None: + def sleep( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ) -> None: if mode == "wait": raise ValueError("'wait' pause mode is not supported in inproc-engine mode") - result = self.engine_core.sleep(level, mode) + result = self.engine_core.sleep(level, mode, clear_connector_cache) assert result is None def release_kv_cache_memory(self) -> None: @@ -1041,8 +1056,15 @@ def list_loras(self) -> set[int]: def pin_lora(self, lora_id: int) -> bool: return self.call_utility("pin_lora", lora_id) - def sleep(self, level: int = 1, mode: PauseMode = "abort") -> None: - self.call_utility("sleep", level, mode) + def sleep( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ) -> None: + # Unconditional, unlike the Rust client: client and engine ship in the + # same package, so they cannot disagree about the signature. + self.call_utility("sleep", level, mode, clear_connector_cache) def release_kv_cache_memory(self) -> None: self.call_utility("release_kv_cache_memory") @@ -1260,9 +1282,14 @@ async def abort_requests_async(self, request_ids: list[str]) -> None: await self._send_input(EngineCoreRequestType.ABORT, request_ids) async def pause_scheduler_async( - self, mode: PauseMode = "abort", clear_cache: bool = True + self, + mode: PauseMode = "abort", + clear_cache: bool = True, + clear_connector_cache: bool = True, ) -> None: - await self.call_utility_async("pause_scheduler", mode, clear_cache) + await self.call_utility_async( + "pause_scheduler", mode, clear_cache, clear_connector_cache + ) async def resume_scheduler_async(self) -> None: await self.call_utility_async("resume_scheduler") @@ -1288,8 +1315,13 @@ async def reset_prefix_cache_async( async def reset_encoder_cache_async(self) -> None: await self.call_utility_async("reset_encoder_cache") - async def sleep_async(self, level: int = 1, mode: PauseMode = "abort") -> None: - await self.call_utility_async("sleep", level, mode) + async def sleep_async( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ) -> None: + await self.call_utility_async("sleep", level, mode, clear_connector_cache) async def release_kv_cache_memory_async(self) -> None: await self.call_utility_async("release_kv_cache_memory") diff --git a/vllm/v1/engine/llm_engine.py b/vllm/v1/engine/llm_engine.py index a2d97b5544fb..0610dc161bff 100644 --- a/vllm/v1/engine/llm_engine.py +++ b/vllm/v1/engine/llm_engine.py @@ -367,10 +367,15 @@ def reset_encoder_cache(self) -> None: """ self.engine_core.reset_encoder_cache() - def sleep(self, level: int = 1, mode: PauseMode = "abort"): + def sleep( + self, + level: int = 1, + mode: PauseMode = "abort", + clear_connector_cache: bool = True, + ): if level >= 1: self.renderer.clear_mm_cache() - self.engine_core.sleep(level, mode) + self.engine_core.sleep(level, mode, clear_connector_cache) if self.logger_manager is not None: self.logger_manager.record_sleep_state(1, level)