Skip to content
Open
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
4 changes: 4 additions & 0 deletions docs/features/sleep_mode.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 4 additions & 0 deletions docs/training/async_rl.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
2 changes: 2 additions & 0 deletions rust/proto/control.proto
Original file line number Diff line number Diff line change
Expand Up @@ -126,6 +126,7 @@ enum PauseMode {
message PauseGenerationRequest {
PauseMode mode = 1;
optional bool clear_cache = 2;
optional bool clear_connector_cache = 3;
}
message PauseGenerationResponse {}

Expand All @@ -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 {}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
29 changes: 25 additions & 4 deletions rust/src/engine-core-client/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(())
}

Expand All @@ -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(())
}

Expand Down
6 changes: 4 additions & 2 deletions rust/src/server/src/grpc/control.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {}))
Expand Down Expand Up @@ -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 {}))
Expand Down
12 changes: 11 additions & 1 deletion rust/src/server/src/routes/pause.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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)]
Expand All @@ -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"))
}
Expand All @@ -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))?;

Expand Down
8 changes: 7 additions & 1 deletion rust/src/server/src/routes/sleep.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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"))
}
Expand All @@ -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))?;

Expand Down
84 changes: 84 additions & 0 deletions rust/src/server/src/routes/tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Expand Down Expand Up @@ -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() {
Expand Down
42 changes: 42 additions & 0 deletions tests/v1/core/test_scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@
KVCacheGroupSpec,
MambaSpec,
)
from vllm.v1.metrics.stats import PrefixCacheStats
from vllm.v1.outputs import (
DraftTokenIds,
ECConnectorOutput,
Expand Down Expand Up @@ -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.
Expand Down
42 changes: 41 additions & 1 deletion tests/v1/engine/test_engine_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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:
Expand Down
Loading
Loading