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
16 changes: 16 additions & 0 deletions docs/features/sleep_mode.md
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,22 @@ llm.collective_rpc("reload_weights")
llm.wake_up(tags=["kv_cache"])
```

#### Draining running requests

Use `llm.sleep(level=1, mode="wait")` to finish running requests before
releasing GPU memory. Requests still in the waiting queue remain queued until
`wake_up()`. The `wait` mode is also available with the in-process engine
(`VLLM_ENABLE_V1_MULTIPROCESSING=0`) for a single data-parallel replica,
including tensor-parallel execution.

The in-process engine continues applying stop strings and collecting token
and logprob outputs during the drain. Buffered results are returned by the next
`LLMEngine.step()`, or by `llm.wait_for_completion()` after waking the engine.
Streaming deltas produced during the drain are combined per request.

Wake a sleeping engine before requesting another drain. In-process data
parallelism (`data_parallel_size > 1`) does not yet support `mode="wait"`.

#### RLHF weight updates

During RLHF training, vLLM allows you to selectively wake up only the model weights or the KV cache using the tags argument in wake_up(). This fine-grained control is especially useful when updating model weights: by waking up just the weights (e.g., llm.wake_up(tags=["weights"])), you avoid allocating memory for the KV cache until after the weight update is complete. This approach helps prevent GPU out-of-memory (OOM) errors, particularly with large models, by minimizing peak memory usage during weight synchronization and update operations.
Expand Down
191 changes: 190 additions & 1 deletion tests/v1/engine/test_llm_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,13 @@

import pytest

from tests.utils import create_new_process_for_each_test
from vllm import LLM
from vllm.sampling_params import SamplingParams, StructuredOutputsParams
from vllm.sampling_params import (
RequestOutputKind,
SamplingParams,
StructuredOutputsParams,
)
from vllm.v1.metrics.reader import Counter, Gauge, Histogram, Metric, Vector

if TYPE_CHECKING:
Expand Down Expand Up @@ -236,3 +241,187 @@ def test_skip_tokenizer_initialization(model: str):
assert len(completions) > 0
assert completions[0].text == ""
assert completions[0].token_ids


@pytest.mark.parametrize("output_kind", list(RequestOutputKind))
@pytest.mark.parametrize("level", [0, 1, 2])
@create_new_process_for_each_test()
def test_inproc_sleep_wait_preserves_requests(
vllm_runner, monkeypatch, output_kind, level
):
"""Drain running generations without losing outputs or admitting waiters."""
from vllm.v1.engine.output_processor import RequestOutputCollector

monkeypatch.setenv("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
with vllm_runner(
MODEL,
dtype=DTYPE,
max_model_len=128,
max_num_seqs=1,
enforce_eager=level != 1,
enable_sleep_mode=True,
enable_prefix_caching=False,
async_scheduling=level == 1,
compilation_config={
"mode": 0,
"cudagraph_mode": "FULL",
"cudagraph_capture_sizes": [1],
},
gpu_memory_utilization=0.2,
) as model:
llm = model.llm
engine = llm.llm_engine
prompt = "The capital of France is"
params = SamplingParams(
temperature=0,
max_tokens=16,
ignore_eos=True,
logprobs=3,
output_kind=RequestOutputKind.FINAL_ONLY,
)
expected = llm.generate([prompt], params, use_tqdm=False)[0].outputs[0]
if output_kind == RequestOutputKind.DELTA:
params = SamplingParams(
temperature=0,
max_tokens=16,
ignore_eos=True,
logprobs=3,
output_kind=RequestOutputKind.FINAL_ONLY,
stop=[llm.get_tokenizer().decode(expected.token_ids[:4])],
)
expected = llm.generate([prompt], params, use_tqdm=False)[0].outputs[0]
assert expected.finish_reason == "stop"
assert len(expected.token_ids) < params.max_tokens
params.output_kind = output_kind
collector = RequestOutputCollector(output_kind, "running")
engine.add_request("running", prompt, params)
core = engine.engine_core.engine_core
while not core.scheduler.running:
for output in engine.step():
collector.put(output)
engine.add_request("waiting", prompt, params)

llm.sleep(level=level, mode="wait")

assert not core.scheduler.running
assert len(core.scheduler.waiting) == 1
assert engine.is_sleeping()
assert engine.has_unfinished_requests()
assert engine.get_num_unfinished_requests() == 2
drained = engine.step()
assert [output.request_id for output in drained] == ["running"]
for output in drained:
collector.put(output)
result = collector.get_nowait()
assert result is not None and result.finished
actual = result.outputs[0]
assert actual.token_ids == expected.token_ids
assert actual.text == expected.text
assert actual.finish_reason == expected.finish_reason
assert len(actual.logprobs) == len(expected.logprobs)
for actual_step, expected_step in zip(actual.logprobs, expected.logprobs):
assert actual_step.keys() == expected_step.keys()
for token_id in actual_step:
assert actual_step[token_id].logprob == pytest.approx(
expected_step[token_id].logprob, abs=1e-4
)
assert len(core.scheduler.waiting) == 1

assert engine.get_num_unfinished_requests() == 1
if level >= 1:
llm.wake_up(tags=["weights"])
assert engine.is_sleeping()
assert engine.step() == []
if level == 2:
llm.collective_rpc("reload_weights")
llm.wake_up(tags=["kv_cache"])
else:
llm.wake_up(tags=["scheduling"])
waiting = RequestOutputCollector(output_kind, "waiting")
while engine.has_unfinished_requests():
for output in engine.step():
assert output.request_id == "waiting"
waiting.put(output)
result = waiting.get_nowait()
assert result is not None and result.finished
assert result.outputs[0].token_ids == expected.token_ids


def _inproc_client_for_drain():
from collections import deque
from unittest.mock import MagicMock

from vllm.v1.engine.core_client import InprocClient

client = object.__new__(InprocClient)
client.engine_core = MagicMock()
client.engine_core.vllm_config.parallel_config.data_parallel_size = 1
client.engine_core.model_executor.is_sleeping = False
client.engine_core.sleep.return_value = None
client.engine_core.batch_queue = []
client._pending_outputs = deque()
return client


def test_inproc_wait_delivers_raw_outputs_after_batch_cleanup():
"""A drained client retains every output and flushes queued model batches."""
from vllm.v1.engine import EngineCoreOutput, EngineCoreOutputs

client = _inproc_client_for_drain()
core = client.engine_core
core.scheduler.has_requests.side_effect = [True, False, False]
core.batch_queue = [object()]
count = 0

def step():
nonlocal count
count += 1
if count == 2:
core.batch_queue.clear()
return {0: EngineCoreOutputs(outputs=[EngineCoreOutput("r", [count])])}, True

core.step_fn.side_effect = step
client.sleep(level=1, mode="wait")
core.sleep.assert_called_once_with(1, "keep")
assert core.post_step.call_count == 2
assert client.get_output().outputs[0].new_token_ids == [1]
assert client.get_output().outputs[0].new_token_ids == [2]
assert core.step_fn.call_count == 2


@pytest.mark.parametrize("dp_size,asleep", [(2, False), (1, True)])
def test_inproc_drain_rejects_unsafe_state_before_scheduling(dp_size, asleep):
"""Never run a local-only drain across DP ranks or unmapped memory."""
from unittest.mock import MagicMock

client = _inproc_client_for_drain()
core = client.engine_core
core.vllm_config.parallel_config.data_parallel_size = dp_size
core.model_executor.is_sleeping = asleep
step = MagicMock()
with pytest.raises(ValueError, match="data parallelism|Wake the engine"):
client.drain_requests(step)
core.scheduler.set_pause_state.assert_not_called()
step.assert_not_called()
core.sleep.assert_not_called()


def test_inproc_wait_does_not_sleep_after_step_failure():
"""A failed drain must propagate the error before releasing GPU memory."""
client = _inproc_client_for_drain()
core = client.engine_core
core.scheduler.has_requests.return_value = True
core.step_fn.side_effect = RuntimeError("model execution failed")
with pytest.raises(RuntimeError, match="model execution failed"):
client.sleep(level=1, mode="wait")
core.sleep.assert_not_called()


def test_inproc_wait_with_no_running_requests_does_not_step():
"""Waiting-only and idle engines can sleep without admitting a new request."""
client = _inproc_client_for_drain()
core = client.engine_core
core.scheduler.has_requests.return_value = False
client.sleep(level=0, mode="wait")
core.step_fn.assert_not_called()
core.sleep.assert_called_once_with(0, "keep")
30 changes: 28 additions & 2 deletions vllm/v1/engine/core_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import uuid
import weakref
from abc import ABC, abstractmethod
from collections import Counter, defaultdict
from collections import Counter, defaultdict, deque
from collections.abc import Awaitable, Callable, Sequence
from concurrent.futures import Future
from dataclasses import dataclass
Expand Down Expand Up @@ -35,6 +35,7 @@
get_open_zmq_inproc_path,
make_zmq_socket,
)
from vllm.v1.core.sched.interface import PauseState
from vllm.v1.engine import (
EEP_NOTIFICATION_CALL_ID,
FT_STATUS_CALL_ID,
Expand Down Expand Up @@ -357,8 +358,14 @@ def __init__(
log_stats,
executor_fail_callback=executor_fail_callback,
)
self._pending_outputs: deque[EngineCoreOutputs] = deque()

def get_output(self) -> EngineCoreOutputs:
if self._pending_outputs:
return self._pending_outputs.popleft()
return self._step()

def _step(self) -> EngineCoreOutputs:
outputs, model_executed = self.engine_core.step_fn()
self.engine_core.post_step(model_executed=model_executed)
return outputs and outputs.get(0) or EngineCoreOutputs()
Expand Down Expand Up @@ -393,9 +400,28 @@ def reset_prefix_cache(
def reset_encoder_cache(self) -> None:
self.engine_core.reset_encoder_cache()

def drain_requests(self, step: Callable[[], Any]) -> None:
"""Drive running requests to completion, leaving new requests queued.

The caller supplies the step function so frontend stop conditions and
output delivery continue to run while the in-process engine drains.
"""
core = self.engine_core
if core.vllm_config.parallel_config.data_parallel_size > 1:
raise ValueError(
"'wait' mode is not supported with in-process data parallelism"
)
if core.model_executor.is_sleeping:
raise ValueError("Wake the engine before draining requests")

core.scheduler.set_pause_state(PauseState.PAUSED_NEW)
while core.scheduler.has_requests() or core.batch_queue:
step()

def sleep(self, level: int = 1, mode: PauseMode = "abort") -> None:
if mode == "wait":
raise ValueError("'wait' pause mode is not supported in inproc-engine mode")
self.drain_requests(lambda: self._pending_outputs.append(self._step()))
mode = "keep"
result = self.engine_core.sleep(level, mode)
assert result is None

Expand Down
56 changes: 52 additions & 4 deletions vllm/v1/engine/llm_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,9 +29,9 @@
from vllm.tracing import init_tracer
from vllm.usage.usage_lib import UsageContext
from vllm.v1.engine import EngineCoreRequest, PauseMode
from vllm.v1.engine.core_client import EngineCoreClient
from vllm.v1.engine.core_client import EngineCoreClient, InprocClient
from vllm.v1.engine.input_processor import InputProcessor
from vllm.v1.engine.output_processor import OutputProcessor
from vllm.v1.engine.output_processor import OutputProcessor, RequestOutputCollector
from vllm.v1.engine.parallel_sampling import ParentRequest
from vllm.v1.executor import Executor
from vllm.v1.metrics.loggers import StatLoggerFactory, StatLoggerManager
Expand Down Expand Up @@ -87,6 +87,7 @@ def __init__(
else:
self.dp_group = None
self.should_execute_dummy_batch = False
self._drained_outputs: dict[str, RequestOutputCollector] = {}

self.renderer = renderer = renderer_from_config(self.vllm_config)

Expand Down Expand Up @@ -191,10 +192,17 @@ def from_engine_args(
)

def get_num_unfinished_requests(self) -> int:
return self.output_processor.get_num_unfinished_requests()
undelivered = sum(
isinstance(collector.output, (RequestOutput, PoolingRequestOutput))
and collector.output.finished
for collector in self._drained_outputs.values()
)
return self.output_processor.get_num_unfinished_requests() + undelivered

def has_unfinished_requests(self) -> bool:
has_unfinished = self.output_processor.has_unfinished_requests()
has_unfinished = bool(
self._drained_outputs or self.output_processor.has_unfinished_requests()
)
if self.dp_group is None:
return has_unfinished or self.engine_core.dp_engines_running()
return self.has_unfinished_requests_dp(has_unfinished)
Expand All @@ -217,6 +225,8 @@ def get_supported_tasks(self) -> tuple[SupportedTask, ...]:
def abort_request(self, request_ids: list[str], internal: bool = False) -> None:
"""Remove request_ids from EngineCore and Detokenizer."""

for request_id in request_ids:
self._drained_outputs.pop(request_id, None)
request_ids = self.output_processor.abort_requests(request_ids, internal)
self.engine_core.abort_requests(request_ids)

Expand Down Expand Up @@ -302,6 +312,17 @@ def add_request(
return req_id

def step(self) -> list[RequestOutput | PoolingRequestOutput]:
if self._drained_outputs:
outputs = [
output
for collector in self._drained_outputs.values()
if (output := collector.get_nowait()) is not None
]
self._drained_outputs.clear()
return outputs
return self._step()

def _step(self) -> list[RequestOutput | PoolingRequestOutput]:
if self.should_execute_dummy_batch:
self.should_execute_dummy_batch = False
self.engine_core.execute_dummy_batch()
Expand Down Expand Up @@ -369,7 +390,34 @@ def reset_encoder_cache(self) -> None:
"""
self.engine_core.reset_encoder_cache()

def _drain_inproc_requests(self, client: InprocClient) -> None:
states = list(self.output_processor.request_states.values())
for state in states:
request_id = (
state.parent_req.external_req_id
if state.parent_req is not None
else state.external_req_id
)
collector = self._drained_outputs.get(request_id)
if collector is None:
collector = RequestOutputCollector(state.output_kind, request_id)
self._drained_outputs[request_id] = collector
state.queue = collector
try:
client.drain_requests(self._step)
finally:
for state in states:
state.queue = None
self._drained_outputs = {
request_id: collector
for request_id, collector in self._drained_outputs.items()
if collector.output is not None
}

def sleep(self, level: int = 1, mode: PauseMode = "abort"):
if mode == "wait" and isinstance(self.engine_core, InprocClient):
self._drain_inproc_requests(self.engine_core)
mode = "keep"
if level >= 1:
self.renderer.clear_mm_cache()
self.engine_core.sleep(level, mode)
Expand Down
Loading