From e862971f18f61035ae09d53e0ff8506b1f8f5ecd Mon Sep 17 00:00:00 2001 From: Wang Zupeng Date: Mon, 14 Sep 2026 11:56:50 +0800 Subject: [PATCH] [Feature][Core] Support draining in-process requests before sleep Drive running requests and pending model batches to completion before sleeping, preserving waiting requests and buffered generation results. Keep stop-string processing active and reuse per-request output collectors. Cover sleep levels, output kinds, partial wake, and drain failures. Co-authored-by: Codex Signed-off-by: Wang Zupeng --- docs/features/sleep_mode.md | 16 +++ tests/v1/engine/test_llm_engine.py | 191 ++++++++++++++++++++++++++++- vllm/v1/engine/core_client.py | 30 ++++- vllm/v1/engine/llm_engine.py | 56 ++++++++- 4 files changed, 286 insertions(+), 7 deletions(-) diff --git a/docs/features/sleep_mode.md b/docs/features/sleep_mode.md index dcb6ad8c940c..e7e25669f7e2 100644 --- a/docs/features/sleep_mode.md +++ b/docs/features/sleep_mode.md @@ -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. diff --git a/tests/v1/engine/test_llm_engine.py b/tests/v1/engine/test_llm_engine.py index 65dbb0e48199..00460cd83d80 100644 --- a/tests/v1/engine/test_llm_engine.py +++ b/tests/v1/engine/test_llm_engine.py @@ -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: @@ -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") diff --git a/vllm/v1/engine/core_client.py b/vllm/v1/engine/core_client.py index c354ba699c75..2b6e8f8ddcc9 100644 --- a/vllm/v1/engine/core_client.py +++ b/vllm/v1/engine/core_client.py @@ -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 @@ -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, @@ -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() @@ -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 diff --git a/vllm/v1/engine/llm_engine.py b/vllm/v1/engine/llm_engine.py index 32096079e066..d5e580eb6e2c 100644 --- a/vllm/v1/engine/llm_engine.py +++ b/vllm/v1/engine/llm_engine.py @@ -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 @@ -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) @@ -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) @@ -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) @@ -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() @@ -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)