diff --git a/docs/configuration/environment_variables.md b/docs/configuration/environment_variables.md index 368baaceb1c..ace89704893 100644 --- a/docs/configuration/environment_variables.md +++ b/docs/configuration/environment_variables.md @@ -83,6 +83,7 @@ depends on the installed kernels and model path. | `SPEAKER_SAMPLES_DIR` | Filesystem path; default `~/.cache/vllm-omni/speakers` | Speech server; read when speaker storage initializes | Environment-only setting. The directory is created; filesystem errors propagate. | Stable | | `SPEAKER_MAX_UPLOADED` | Integer; default `1000` | Speech server; read when speaker storage initializes | Environment-only setting. A non-integer logs a warning and uses `1000`; range is not otherwise validated. | Stable | | `VLLM_OMNI_ASYNC_OUTPUT_TIMEOUT` | Float seconds; default `600` | Diffusion engine async-output wait in `step_streaming`; resolved per call on the request path, not at import | Environment-only setting. A non-float or `<=0` value warns once and uses the default. | Experimental | +| `VLLM_OMNI_EVENT_DRIVEN_ORCH` | `1`, `true`, `yes` or `on` enables; default `0` (off) | Orchestration loop and the serving-side final-output drain; read once when the `Orchestrator` is constructed | Environment-only setting. Values are stripped and case-normalized; any unrecognized value leaves the legacy poll loop selected. | Experimental | | `VLLM_OMNI_INPUT_WAIT_TIMEOUT_S` | Float seconds; default `600`; `<=0` disables | Full-payload input coordinator, not async-chunk transfer; read when the scheduler module imports in each worker | Environment-only setting. A non-float logs a warning and uses `600`. | Stable operational control | | `VLLM_OMNI_ORCH_MONITOR_PATH` | Filesystem path; default `/vllm_omni_orch_monitor_.json` | Orchestrator monitor enabled by `--enable-orch-monitor`; read when the monitor is created | Environment-only path override. Parent directories are created; write errors are logged. | Diagnostic | | `VLLM_OMNI_VIDEO_SYNC_TIMEOUT` | Float seconds; default `600` | Synchronous Videos API; read when the API server module imports | Environment-only setting. A non-float raises `ValueError` during import. | Experimental | diff --git a/docs/contributing/profiling.md b/docs/contributing/profiling.md index b57de441ed1..62fff29c558 100644 --- a/docs/contributing/profiling.md +++ b/docs/contributing/profiling.md @@ -322,6 +322,15 @@ Each 1-second window records: On shutdown the server also logs a short summary (`loop_active_pct`, per-replica queue averages/maxima). +`loop_idle` and `loop_active` count loop iterations, so their absolute scale is +tied to which orchestration loop is running. Under the default poll loop an idle +orchestrator records roughly one iteration per millisecond. Under the +event-driven loop (`VLLM_OMNI_EVENT_DRIVEN_ORCH=1`, see +[Speech API](../serving/speech_api.md#orchestration-loop-experimental)) an idle +orchestrator wakes only once per 0.5 s reconcile timeout, while a busy one still +records one iteration per routed output. Window counts and the `loop_active_pct` +summary are therefore not comparable across the two modes. + ### Relationship to other diagnostics This monitor is intentionally separate from the existing profiling tools: diff --git a/docs/design/module/engine_orchestration.md b/docs/design/module/engine_orchestration.md index 5bb0ea0df44..2e176b35d9c 100644 --- a/docs/design/module/engine_orchestration.md +++ b/docs/design/module/engine_orchestration.md @@ -66,6 +66,12 @@ the boundary affected by the in-flight stage client/process refactor in [#5441](https://github.com/vllm-project/vllm-omni/pull/5441). Names and responsibilities proposed only by that PR are not current contracts. +The orchestration loop also has an opt-in event-driven mode +(`VLLM_OMNI_EVENT_DRIVEN_ORCH=1`, default off) proposed in +[#5221](https://github.com/vllm-project/vllm-omni/pull/5221). It changes poll +cadence only: the routing, ordering, and terminal-state contracts below hold +identically on both loops. + ## Ownership boundary This document owns `AsyncOmniEngine`, `Orchestrator`, request-state creation, diff --git a/docs/serving/speech_api.md b/docs/serving/speech_api.md index 92321e536f9..98fb7f38032 100644 --- a/docs/serving/speech_api.md +++ b/docs/serving/speech_api.md @@ -811,6 +811,51 @@ If you encounter OOM errors: Use `/v1/audio/voices` to list available voices for the loaded model. +## Orchestration Loop (experimental) + +Multi-stage omni deployments route stage outputs through a single orchestrator +loop. By default that loop polls every stage replica on a 1 ms cadence. An +opt-in event-driven mode replaces the poll with one reader task per live stage +replica awaiting its client directly, and switches the serving-side +final-output drain to a condition-variable wakeup at the same time. + +**Configuration (environment variables):** + +| Variable | Default | Description | +|----------|---------|-------------| +| `VLLM_OMNI_EVENT_DRIVEN_ORCH` | `0` (off) | Switches the orchestration loop and the final-output drain from the legacy 1 ms poll to event-driven wakeups. Enabled by `1`, `true`, `yes`, or `on`, matched case-insensitively after surrounding whitespace is stripped; any other value leaves it off. | + +Set it on the process that runs the orchestrator (stage 0 of an omni +deployment) before starting the server: + +```bash +export VLLM_OMNI_EVENT_DRIVEN_ORCH=1 +vllm serve Qwen/Qwen3-TTS-12Hz-1.7B-Base \ + --omni \ + --port 8091 +``` + +The server logs the selected loop mode and its reader/poller counts once at +startup, so you can confirm which loop is live. + +Routing, output ordering, and terminal-state behavior are identical on both +loops; only the poll cadence changes. Leaving the variable unset keeps the +legacy poll loop, which is the supported default. + +**Known limitations:** + +- The measured serving A/B (idle CPU 2.43% to 0.07%; TTFP p99 -32% at + concurrency 8) predates the rebuild on the per-replica fault-isolation work + in [#4285](https://github.com/vllm-project/vllm-omni/pull/4285). That work + changed dead-replica handling and reader/poller lifecycle rather than the + steady-state output path, and the parity suite covers it, but the serving + A/B has not been re-run on the current head. +- The diffusion-poller branch is covered by unit tests only. Deployments whose + stages all run as standard engine cores never exercise it, including GLM-TTS, + which deploys its DiT without `stage_type: diffusion`. +- Concurrency 1 and 32 measured at parity with the legacy loop. At 32 the + latency is admission-bound, which this mode does not address. + ## Development Enable debug logging: diff --git a/tests/config/test_environment_variables.py b/tests/config/test_environment_variables.py index 3369fd42934..00e0dffdb2e 100644 --- a/tests/config/test_environment_variables.py +++ b/tests/config/test_environment_variables.py @@ -164,7 +164,7 @@ def test_inventory_matches_reviewed_snapshot_counts(): """Make an inventory expansion an explicit review decision.""" category_counts = Counter(item.category for item in ENVIRONMENT_VARIABLE_INVENTORY.values()) assert category_counts == { - EnvironmentVariableCategory.PUBLIC_OMNI: 23, + EnvironmentVariableCategory.PUBLIC_OMNI: 24, EnvironmentVariableCategory.INHERITED_VLLM: 20, EnvironmentVariableCategory.PLATFORM_EXTERNAL: 27, EnvironmentVariableCategory.MODEL_SPECIFIC: 56, diff --git a/tests/diffusion/test_diffusion_engine_metrics.py b/tests/diffusion/test_diffusion_engine_metrics.py index 8caf4a8dcfb..dc909a2b920 100644 --- a/tests/diffusion/test_diffusion_engine_metrics.py +++ b/tests/diffusion/test_diffusion_engine_metrics.py @@ -172,12 +172,29 @@ def test_abort_keeps_output_consumers_alive_for_terminal_snapshot(self) -> None: assert ".cancel(" not in abort_branch def test_orchestrator_consumes_metrics_only_output_without_routing(self) -> None: + """A metrics-only sentinel contributes its queue depth and is not routed. + + Absorption lives in ``_absorb_diffusion_metrics`` rather than inline in a + loop, so both orchestration loops share one implementation. The ordering + is asserted where it now lives: the snapshot inside the helper, and the + absorb-before-route guard inside every loop that polls diffusion output. + """ source = _read_source(_ORCHESTRATOR_PATH) - loop_source = _get_function_source(source, "Orchestrator", "_orchestration_loop") - snapshot_pos = loop_source.index("_update_stage_replica_waiting(") - sentinel_pos = loop_source.index("diffusion_output.request_id == DIFFUSION_METRICS_ONLY_REQUEST_ID") - route_pos = loop_source.index("pool.record_output_timestamps([diffusion_output])") - assert snapshot_pos < sentinel_pos < route_pos + + # Snapshot before the verdict, so a sentinel still reports its waiting + # depth on the way out instead of being dropped unaccounted. + absorb_source = _get_function_source(source, "Orchestrator", "_absorb_diffusion_metrics") + snapshot_pos = absorb_source.index("_update_stage_replica_waiting(") + sentinel_pos = absorb_source.index("diffusion_output.request_id == DIFFUSION_METRICS_ONLY_REQUEST_ID") + assert snapshot_pos < sentinel_pos + + # Absorb before routing, or a sentinel reaches the downstream consumer + # as if it were a real request output. + for loop_name in ("_orchestration_loop", "_orchestration_loop_event_driven"): + loop_source = _get_function_source(source, "Orchestrator", loop_name) + absorb_pos = loop_source.index("self._absorb_diffusion_metrics(") + route_pos = loop_source.index("record_output_timestamps(") + assert absorb_pos < route_pos, loop_name class TestVaeDecodeEmit: diff --git a/tests/engine/test_orchestrator_event_driven.py b/tests/engine/test_orchestrator_event_driven.py new file mode 100644 index 00000000000..284dde55d1c --- /dev/null +++ b/tests/engine/test_orchestrator_event_driven.py @@ -0,0 +1,247 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Event-driven orchestration loop (``VLLM_OMNI_EVENT_DRIVEN_ORCH=1``) tests. + +Parity suite: re-runs the legacy orchestration scenarios from +``test_orchestrator.py`` / ``test_orchestrator_error_handling.py`` with the +event-driven loop selected, so both loops are held to the same behavior. Plus +event-driven-specific coverage: reader reconcile on client swap, the blocking +final-output drain, and flag parsing. +""" + +from __future__ import annotations + +import asyncio +from types import SimpleNamespace + +import janus +import pytest + +from vllm_omni.engine.async_omni_engine import AsyncOmniEngine +from vllm_omni.engine.messages import ShutdownRequestMessage +from vllm_omni.engine.orchestrator import _event_driven_orch_enabled + +from . import test_orchestrator as legacy +from . import test_orchestrator_error_handling as legacy_errors +from .test_orchestrator import ( + FakeOutputProcessor, + FakeStageClient, + OrchestratorFixture, + _build_harness, + _build_request_output, + _engine_core_outputs, + _enqueue_add_request, + _get_output_message, + _sampling_params, + _shutdown_orchestrator, + _wait_for, +) + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + + +@pytest.fixture +def orchestrator_factory(monkeypatch): + """Flag-setting clone of the legacy harness fixture. + + Sets ``VLLM_OMNI_EVENT_DRIVEN_ORCH=1`` before any Orchestrator is + constructed and asserts the flag actually took effect, so the parity tests + cannot silently exercise the legacy poll loop. + """ + monkeypatch.setenv("VLLM_OMNI_EVENT_DRIVEN_ORCH", "1") + fixtures: list[OrchestratorFixture] = [] + + def _factory(*args, **kwargs) -> OrchestratorFixture: + fixture = _build_harness(*args, **kwargs) + assert fixture.orchestrator._event_driven_orch is True + fixtures.append(fixture) + return fixture + + yield _factory + + for fixture in fixtures: + if fixture.thread.is_alive(): + fixture.request_sync_q.put_nowait(ShutdownRequestMessage()) + fixture.thread.join(timeout=5) + for q in fixture.queues: + q.close() + + +# --------------------------------------------------------------------------- +# Parity: the legacy scenario matrix, re-run through the event-driven loop +# --------------------------------------------------------------------------- + +_PARITY_TESTS = [ + legacy.test_run_two_stage_llm, + legacy.test_run_single_stage_diffusion, + legacy.test_run_single_stage_diffusion_streaming_forwards_intermediate_chunks, + legacy.test_run_llm_to_diffusion, + legacy.test_run_async_chunk, + legacy.test_run_shutdown, + legacy.test_run_abort, + legacy.test_multi_replica_round_robin_distribution, + legacy.test_multi_replica_abort_broadcasts_to_all_replicas, + legacy.test_multi_replica_shutdown_all_replicas, + legacy.test_multi_replica_cfg_companion_inherits_parent_affinity, + # Stats plumbing. The scheduler-stats cases matter most here: a batch with + # no request outputs still carries SchedulerStats on throttled ticks, and a + # reader that drops every output-less batch silently stops reporting + # KV/queue gauges under the event-driven loop. + legacy.test_orchestrator_records_iteration_stats_without_scheduler_stats, + legacy.test_orchestrator_records_scheduler_stats_without_outputs, + legacy.test_orchestrator_does_not_build_iteration_stats_for_finished_only_batch, + legacy.test_orchestrator_does_not_build_iteration_stats_without_stat_logger, + # Per-replica fault isolation (#4285): a dead replica must be evicted and + # the server kept up, on both the LLM reader path and the diffusion poller. + legacy_errors.test_engine_dead_error_evicts_replica_and_keeps_running, + legacy_errors.test_engine_dead_error_fails_only_dead_replica_requests, + legacy_errors.test_forward_to_dead_downstream_stage_fails_request_not_server, + legacy_errors.test_add_request_to_dead_stage_fails_request_not_server, + legacy_errors.test_diffusion_replica_death_on_poll_keeps_server, + legacy_errors.test_diffusion_error_output_routed_as_finished, + legacy_errors.test_diffusion_client_error_output_propagates_status_code, +] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("legacy_test", _PARITY_TESTS, ids=lambda f: f.__name__) +async def test_event_driven_parity(legacy_test, orchestrator_factory) -> None: + await legacy_test(orchestrator_factory) + + +# --------------------------------------------------------------------------- +# Event-driven-specific behavior +# --------------------------------------------------------------------------- + + +def test_flag_parsing(monkeypatch) -> None: + monkeypatch.delenv("VLLM_OMNI_EVENT_DRIVEN_ORCH", raising=False) + assert _event_driven_orch_enabled() is False + for value in ("1", "true", "True", "YES", "on"): + monkeypatch.setenv("VLLM_OMNI_EVENT_DRIVEN_ORCH", value) + assert _event_driven_orch_enabled() is True + for value in ("0", "false", "off", ""): + monkeypatch.setenv("VLLM_OMNI_EVENT_DRIVEN_ORCH", value) + assert _event_driven_orch_enabled() is False + + +def test_default_is_legacy_loop(monkeypatch) -> None: + """Without the env flag, the harness runs the legacy poll loop.""" + monkeypatch.delenv("VLLM_OMNI_EVENT_DRIVEN_ORCH", raising=False) + fixture = _build_harness([FakeStageClient(stage_type="llm", final_output=True)]) + try: + assert fixture.orchestrator._event_driven_orch is False + finally: + fixture.request_sync_q.put_nowait(ShutdownRequestMessage()) + fixture.thread.join(timeout=5) + for q in fixture.queues: + q.close() + + +@pytest.mark.asyncio +async def test_reader_reconcile_picks_up_swapped_client(orchestrator_factory) -> None: + """Outputs from a replica whose client object was replaced still flow. + + The event-driven loop binds one reader task per client object; the + periodic reconcile must respawn the reader when ``pool.clients[replica]`` + is swapped (replica replacement), otherwise the new client's outputs + would never be drained. + """ + stage0 = FakeStageClient(stage_type="llm", final_output=True) + processor = FakeOutputProcessor(request_outputs=[_build_request_output("req-swap", token_ids=[3], finished=True)]) + orchestrator_fixture = orchestrator_factory([stage0], output_processors=[processor]) + request = SimpleNamespace(request_id="req-swap", prompt_token_ids=[1, 2]) + + try: + await _enqueue_add_request( + orchestrator_fixture, + request_id="req-swap", + prompt=request, + original_prompt={"prompt": "swap"}, + sampling_params_list=[_sampling_params()], + final_stage_id=0, + ) + await _wait_for(lambda: len(stage0.add_request_calls) == 1) + + # Swap in a fresh client for the same replica slot; keep the pool's + # other wiring intact. The reconcile tick (0.5 s) must respawn the + # reader bound to the new client object. + pool = orchestrator_fixture.orchestrator.stage_pools[0] + replacement = FakeStageClient(stage_type="llm", final_output=True) + replacement.stage_id = stage0.stage_id + replacement.replica_id = stage0.replica_id + pool.clients[0] = replacement + + replacement.push_engine_core_outputs(_engine_core_outputs("swapped-raw", 1.0)) + + output_msg = await _get_output_message(orchestrator_fixture, timeout=5.0) + assert output_msg.request_id == "req-swap" + assert output_msg.finished is True + finally: + await _shutdown_orchestrator(orchestrator_fixture) + + +# --------------------------------------------------------------------------- +# Blocking final-output drain (AsyncOmniEngine.get_output_blocking_async) +# --------------------------------------------------------------------------- + + +def _drain_engine(alive: bool = True) -> AsyncOmniEngine: + engine = object.__new__(AsyncOmniEngine) + engine.output_queue = janus.Queue() + engine.orchestrator_thread = SimpleNamespace(is_alive=lambda: alive) + return engine + + +def _drain_cleanup(engine: AsyncOmniEngine) -> None: + if engine._output_drain_executor is not None: + engine._output_drain_executor.shutdown(wait=False) + engine._output_drain_executor = None + engine.output_queue.close() + + +@pytest.mark.asyncio +async def test_blocking_drain_returns_queued_message() -> None: + engine = _drain_engine() + try: + engine.output_queue.sync_q.put_nowait("msg-1") + assert await engine.get_output_blocking_async(timeout=1.0) == "msg-1" + finally: + _drain_cleanup(engine) + + +@pytest.mark.asyncio +async def test_blocking_drain_wakes_on_late_message() -> None: + """A message put after the wait starts wakes the drain, no polling.""" + engine = _drain_engine() + try: + + async def _delayed_put() -> None: + await asyncio.sleep(0.05) + engine.output_queue.sync_q.put_nowait("late-msg") + + put_task = asyncio.create_task(_delayed_put()) + msg = await engine.get_output_blocking_async(timeout=5.0) + await put_task + assert msg == "late-msg" + finally: + _drain_cleanup(engine) + + +@pytest.mark.asyncio +async def test_blocking_drain_timeout_returns_none_when_alive() -> None: + engine = _drain_engine(alive=True) + try: + assert await engine.get_output_blocking_async(timeout=0.05) is None + finally: + _drain_cleanup(engine) + + +@pytest.mark.asyncio +async def test_blocking_drain_raises_when_orchestrator_dead() -> None: + engine = _drain_engine(alive=False) + try: + with pytest.raises(RuntimeError, match="Orchestrator died"): + await engine.get_output_blocking_async(timeout=0.05) + finally: + _drain_cleanup(engine) diff --git a/vllm_omni/config/environment_variable_inventory.py b/vllm_omni/config/environment_variable_inventory.py index b5e3ca599de..2a7385c1d8d 100644 --- a/vllm_omni/config/environment_variable_inventory.py +++ b/vllm_omni/config/environment_variable_inventory.py @@ -64,6 +64,7 @@ def is_public_omni(self) -> bool: "SPEAKER_MAX_UPLOADED", "SPEAKER_SAMPLES_DIR", "VLLM_OMNI_ASYNC_OUTPUT_TIMEOUT", + "VLLM_OMNI_EVENT_DRIVEN_ORCH", "VLLM_OMNI_INPUT_WAIT_TIMEOUT_S", "VLLM_OMNI_ORCH_MONITOR_PATH", "VLLM_OMNI_SERVER_STORAGE__FILE_CONCURRENCY", diff --git a/vllm_omni/engine/async_omni_engine.py b/vllm_omni/engine/async_omni_engine.py index 12c3413f185..a2aa4e7aeba 100644 --- a/vllm_omni/engine/async_omni_engine.py +++ b/vllm_omni/engine/async_omni_engine.py @@ -109,6 +109,8 @@ class AsyncOmniEngine: _transfer_emitter: Any = None _prom_metrics: Any = None _enable_orch_monitor: bool = False + # Lazily created by get_output_blocking_async(). + _output_drain_executor: concurrent.futures.ThreadPoolExecutor | None = None def __init__( self, @@ -1739,6 +1741,42 @@ async def try_get_output_async(self) -> EngineQueueMessage | None: raise RuntimeError("Orchestrator died unexpectedly. See logs above.") return None + async def get_output_blocking_async(self, timeout: float = 1.0) -> EngineQueueMessage | None: + """Blocking-wait read from the Orchestrator output queue. + + Waits up to ``timeout`` seconds in a dedicated drain thread for the + next message (condition-variable wakeup instead of a poll cadence); + returns ``None`` on timeout so the caller keeps its liveness check, + mirroring ``try_get_output_async``'s contract. Used by the serving + final-output drain when ``VLLM_OMNI_EVENT_DRIVEN_ORCH`` is on. + """ + executor = self._output_drain_executor + if executor is None: + executor = concurrent.futures.ThreadPoolExecutor( + max_workers=1, + thread_name_prefix="omni-output-drain", + ) + self._output_drain_executor = executor + + sync_q = self.output_queue.sync_q + + def _drain_get() -> EngineQueueMessage | None: + # Exceptions are swallowed to a None sentinel: the queue may be + # closed mid-shutdown, and an exception left on an executor future + # after task cancellation would warn as never-retrieved. + try: + return sync_q.get(timeout=timeout) + except queue.Empty: + return None + except Exception: + return None + + loop = asyncio.get_running_loop() + msg = await loop.run_in_executor(executor, _drain_get) + if msg is None and not self.is_alive(): + raise RuntimeError("Orchestrator died unexpectedly. See logs above.") + return msg + def get_stage_metadata(self, stage_id: int) -> StageRuntimeInfo: """Get cached metadata for a stage.""" return self.stage_metadata[stage_id] @@ -1928,6 +1966,12 @@ def shutdown(self) -> None: except Exception: pass + if self._output_drain_executor is not None: + # Any in-flight blocking get bails out within its ≤1 s timeout + # (or immediately via the queue close above), so don't wait. + self._output_drain_executor.shutdown(wait=False) + self._output_drain_executor = None + if hasattr(self, "_runtime") and self._runtime is not None and orchestrator_stopped: try: self._runtime.shutdown() diff --git a/vllm_omni/engine/orchestrator.py b/vllm_omni/engine/orchestrator.py index e1306fe3e18..ab052b49409 100644 --- a/vllm_omni/engine/orchestrator.py +++ b/vllm_omni/engine/orchestrator.py @@ -15,6 +15,7 @@ from __future__ import annotations import asyncio +import os import time as _time from collections.abc import Awaitable, Callable from dataclasses import dataclass, field @@ -65,6 +66,23 @@ logger = init_logger(__name__) +# VLLM_OMNI_EVENT_DRIVEN_ORCH=1 switches the orchestration loop (and the +# serving-side final-output drain in entrypoints/async_omni.py) from the legacy +# 1 ms poll cadence to event-driven wakeups: one reader task per live LLM stage +# replica awaits `client.get_output_async()` directly — the same pattern vLLM's +# own AsyncLLM output handler uses — and feeds a single serial dispatch queue. +# Default off; the legacy poll loop remains the fallback. +_EVENT_DRIVEN_ORCH_ENV = "VLLM_OMNI_EVENT_DRIVEN_ORCH" + +# How often the event-driven loop reconciles its reader-task set against +# `available_replica_ids()` (elastic membership, replica eviction) while idle. +_ORCH_READER_RECONCILE_INTERVAL_S = 0.5 + + +def _event_driven_orch_enabled() -> bool: + return os.environ.get(_EVENT_DRIVEN_ORCH_ENV, "0").strip().lower() in ("1", "true", "yes", "on") + + if TYPE_CHECKING: from vllm_omni.experimental.fullduplex.engine.contracts import ( DuplexControlPlanePort, @@ -449,6 +467,7 @@ def __init__( self._shutdown_event = asyncio.Event() self._stages_shutdown = False + self._event_driven_orch = _event_driven_orch_enabled() # Distributed membership (optional, injected by DistStageRuntime) self._membership = membership_controller @@ -964,11 +983,105 @@ def _sample_replica_metrics(self) -> dict[str, tuple[int, int]]: async def _orchestration_output_handler(self) -> None: """Poll all stages, handle transfers, send final outputs to main.""" try: - await self._orchestration_loop() + if self._event_driven_orch: + await self._orchestration_loop_event_driven() + else: + await self._orchestration_loop() except asyncio.CancelledError: logger.debug("[Orchestrator] _orchestration_output_handler cancelled") return + async def _process_llm_stage_outputs( + self, + stage_id: int, + replica_id: int, + raw_outputs: Any, + raw_terminal_request_ids: set[str], + ) -> list[Any]: + """Process one raw LLM poll result; returns processed request outputs. + + Shared by the legacy poll loop and the event-driven dispatcher so both + run the exact same per-output handling. Callers own the + ``EngineDeadError`` catch, since eviction needs the replica id, and + supply ``raw_terminal_request_ids`` — the per-poll accumulator this + fills for the caller to drain through + ``_finish_raw_terminal_requests`` once routing is done. + """ + pool = self.stage_pools[stage_id] + await self._handle_kv_ready_raw_outputs(stage_id, raw_outputs) + for eco in raw_outputs.outputs: + # Emit kv_wait_s before _handle_kv_ready_raw_outputs' + # async_chunk early-return so it lands in all modes. + kv_params = getattr(eco, "kv_transfer_params", None) + if ( + self._prom_metrics is not None + and isinstance(kv_params, dict) + and (kv_wait_s := kv_params.get("kv_wait_s")) is not None + ): + self._prom_metrics.observe_kv_wait( + kv_params.get("connector_type") or "unknown", + float(kv_wait_s), + ) + req_state = self.request_states.get(getattr(eco, "request_id", None)) + if req_state is None or not req_state.streaming.enabled: + continue + req_state.streaming.segment_finished = bool(getattr(eco, "is_segment_finished", False)) + req_state.streaming.segment_token_ids = ( + self._coerce_int_list(getattr(eco, "new_token_ids", None)) + if req_state.streaming.segment_finished + else [] + ) + raw_mm = self._completion_multimodal_output(eco, None) + req_state.streaming.segment_output_metadata = ( + dict(raw_mm) if req_state.streaming.segment_finished and isinstance(raw_mm, dict) else {} + ) + req_state.streaming.new_prompt_len_snapshot = getattr( + eco, + "new_prompt_len_snapshot", + None, + ) + if await self._apply_raw_terminal_stage_finish(stage_id, eco, req_state): + raw_terminal_request_ids.add(req_state.request_id) + iteration_stats = IterationStats() if (self._stat_logger is not None and raw_outputs.outputs) else None + processed = await pool.process_llm_raw_outputs( + replica_id, + raw_outputs, + iteration_stats=iteration_stats, + ) + if self._stat_logger is not None and (raw_outputs.scheduler_stats is not None or iteration_stats is not None): + self._stat_logger.record( + raw_outputs.scheduler_stats, + iteration_stats, + engine_idx=self._stage_replica_to_engine_idx[(stage_id, replica_id)], + ) + _sched_stats = raw_outputs.scheduler_stats + if ( + self._prom_metrics is not None + and _sched_stats is not None + and getattr(_sched_stats, "num_waiting_reqs", None) is not None + ): + self._update_stage_replica_waiting( + stage_id, + replica_id, + int(_sched_stats.num_waiting_reqs), + ) + return processed + + def _absorb_diffusion_metrics(self, stage_id: int, replica_id: int, diffusion_output: Any) -> bool: + """Drain a diffusion output's piggybacked metrics. + + Returns ``True`` when the output carried nothing but metrics and must + not be routed. Shared by both orchestration loops: the event-driven + poller would otherwise hand a metrics-only sentinel to + ``_handle_processed_outputs`` as if it were a real request output. + """ + output_metrics = getattr(diffusion_output, "metrics", None) + if isinstance(output_metrics, dict): + n_waiting = output_metrics.pop(metric_defs.DIFFUSION_SCHEDULER_WAITING_KEY, None) + if n_waiting is not None: + self._update_stage_replica_waiting(stage_id, replica_id, int(n_waiting)) + return diffusion_output.request_id == DIFFUSION_METRICS_ONLY_REQUEST_ID + async def _orchestration_loop(self) -> None: """Poll stage pools and route logical outputs.""" while not self._shutdown_event.is_set(): @@ -988,13 +1101,7 @@ async def _orchestration_loop(self) -> None: if diffusion_output is None: continue - output_metrics = getattr(diffusion_output, "metrics", None) - if isinstance(output_metrics, dict): - n_waiting = output_metrics.pop(metric_defs.DIFFUSION_SCHEDULER_WAITING_KEY, None) - if n_waiting is not None: - self._update_stage_replica_waiting(stage_id, replica_id, int(n_waiting)) - - if diffusion_output.request_id == DIFFUSION_METRICS_ONLY_REQUEST_ID: + if self._absorb_diffusion_metrics(stage_id, replica_id, diffusion_output): idle = False continue @@ -1005,69 +1112,9 @@ async def _orchestration_loop(self) -> None: if raw_outputs is None: continue - await self._handle_kv_ready_raw_outputs(stage_id, raw_outputs) - for eco in raw_outputs.outputs: - # Emit kv_wait_s before _handle_kv_ready_raw_outputs' - # async_chunk early-return so it lands in all modes. - kv_params = getattr(eco, "kv_transfer_params", None) - if ( - self._prom_metrics is not None - and isinstance(kv_params, dict) - and (kv_wait_s := kv_params.get("kv_wait_s")) is not None - ): - self._prom_metrics.observe_kv_wait( - kv_params.get("connector_type") or "unknown", - float(kv_wait_s), - ) - req_state = self.request_states.get(getattr(eco, "request_id", None)) - if req_state is None or not req_state.streaming.enabled: - continue - req_state.streaming.segment_finished = bool(getattr(eco, "is_segment_finished", False)) - req_state.streaming.segment_token_ids = ( - self._coerce_int_list(getattr(eco, "new_token_ids", None)) - if req_state.streaming.segment_finished - else [] - ) - raw_mm = self._completion_multimodal_output(eco, None) - req_state.streaming.segment_output_metadata = ( - dict(raw_mm) - if req_state.streaming.segment_finished and isinstance(raw_mm, dict) - else {} - ) - req_state.streaming.new_prompt_len_snapshot = getattr( - eco, - "new_prompt_len_snapshot", - None, - ) - if await self._apply_raw_terminal_stage_finish(stage_id, eco, req_state): - raw_terminal_request_ids.add(req_state.request_id) - iteration_stats = ( - IterationStats() if (self._stat_logger is not None and raw_outputs.outputs) else None + processed = await self._process_llm_stage_outputs( + stage_id, replica_id, raw_outputs, raw_terminal_request_ids ) - processed = await pool.process_llm_raw_outputs( - replica_id, - raw_outputs, - iteration_stats=iteration_stats, - ) - if self._stat_logger is not None and ( - raw_outputs.scheduler_stats is not None or iteration_stats is not None - ): - self._stat_logger.record( - raw_outputs.scheduler_stats, - iteration_stats, - engine_idx=self._stage_replica_to_engine_idx[(stage_id, replica_id)], - ) - _sched_stats = raw_outputs.scheduler_stats - if ( - self._prom_metrics is not None - and _sched_stats is not None - and getattr(_sched_stats, "num_waiting_reqs", None) is not None - ): - self._update_stage_replica_waiting( - stage_id, - replica_id, - int(_sched_stats.num_waiting_reqs), - ) except asyncio.CancelledError: raise except EngineDeadError as e: @@ -1093,6 +1140,258 @@ async def _orchestration_loop(self) -> None: else: await asyncio.sleep(0) + async def _orchestration_loop_event_driven(self) -> None: + """Event-driven variant of ``_orchestration_loop``. + + Selected by ``VLLM_OMNI_EVENT_DRIVEN_ORCH=1``. One reader task per + available LLM stage replica awaits ``client.get_output_async()`` + directly — the same pattern vLLM's own ``AsyncLLM`` output handler uses + — and feeds a single dispatch queue. This coroutine consumes that queue + serially, so routing/handling semantics are identical to the legacy + loop; only the 1 ms poll cadence (and its per-tick ``asyncio.wait_for`` + task churn) is removed. Diffusion stages keep their nowait-poll contract + via a per-pool poller task feeding the same queue. + """ + ready_q: asyncio.Queue[tuple[str, int, int, Any]] = asyncio.Queue() + readers: dict[tuple[int, int], tuple[asyncio.Task, Any]] = {} + pollers: dict[int, asyncio.Task] = {} + reaped: list[asyncio.Task] = [] + + async def _llm_replica_reader(stage_id: int, replica_id: int, client: Any) -> None: + try: + while not self._shutdown_event.is_set(): + raw_outputs = await client.get_output_async() + # Same keep/drop rule as StagePool._poll_stage_raw: a batch + # with no request outputs still carries SchedulerStats on + # throttled ticks, and dropping it loses the KV/queue gauges + # for that interval. Only a fully empty batch is dropped. A + # poll-style client returns one instead of blocking, so keep + # the legacy 1 ms cadence for those; a blocking client only + # lands here rarely, where 1 ms is irrelevant. + if ( + not raw_outputs.outputs + and raw_outputs.scheduler_stats is None + and not raw_outputs.finished_requests + ): + await asyncio.sleep(0.001) + continue + await ready_q.put(("llm", stage_id, replica_id, raw_outputs)) + except asyncio.CancelledError: + raise + except BaseException as e: # noqa: BLE001 - routed to the dispatcher + await ready_q.put(("error", stage_id, replica_id, e)) + + async def _diffusion_poller(stage_id: int, pool: StagePool) -> None: + try: + while not self._shutdown_event.is_set(): + got = False + for replica_id in pool.available_replica_ids(): + try: + output = pool.poll_diffusion_output(replica_id) + except EngineDeadError as e: + await ready_q.put(("error", stage_id, replica_id, e)) + continue + if output is None: + continue + await ready_q.put(("diffusion", stage_id, replica_id, output)) + got = True + await asyncio.sleep(0 if got else 0.001) + except asyncio.CancelledError: + raise + except BaseException as e: # noqa: BLE001 - routed to the dispatcher + await ready_q.put(("error", stage_id, -1, e)) + + def _reconcile_readers() -> None: + """Track available replicas: spawn/reap LLM readers and diffusion pollers. + + Walks ``available_replica_ids()`` for parity with the legacy loop, + so a replica evicted by ``_handle_dead_replica`` (or marked + unavailable by membership) stops being read here too. + + Diffusion pollers are reconciled here rather than spawned once at + startup: a pool with no live replica yet reports ``stage_type is + None``, so a stage whose first replica registers at runtime would + otherwise never get a poller and its outputs would never drain. + """ + live: set[tuple[int, int]] = set() + live_pools: set[int] = set() + for stage_id in range(self.num_stages): + pool = self.stage_pools[stage_id] + stage_type = pool.stage_type + if stage_type is None: + # No live replica yet; nothing to attach to. A later + # reconcile picks the stage up once one registers. + continue + if stage_type == "diffusion": + live_pools.add(stage_id) + existing_poller = pollers.get(stage_id) + if existing_poller is None or existing_poller.done(): + pollers[stage_id] = asyncio.create_task( + _diffusion_poller(stage_id, pool), + name=f"orch-diffusion-poller-s{stage_id}", + ) + continue + for replica_id in pool.available_replica_ids(): + client = pool.clients[replica_id] + if client is None: + continue + key = (stage_id, replica_id) + live.add(key) + existing = readers.get(key) + if existing is None: + readers[key] = ( + asyncio.create_task( + _llm_replica_reader(stage_id, replica_id, client), + name=f"orch-reader-s{stage_id}r{replica_id}", + ), + client, + ) + continue + if existing[0].done(): + # The reader exited, which means it already queued its + # failure. Leave the slot alone until the dispatcher + # handles that error and evicts the replica; respawning + # now just re-raises the same error against the same + # dead client and queues a duplicate eviction. + continue + if existing[1] is client: + continue + # Client swapped underneath us: retire the old reader. + existing[0].cancel() + reaped.append(existing[0]) + readers[key] = ( + asyncio.create_task( + _llm_replica_reader(stage_id, replica_id, client), + name=f"orch-reader-s{stage_id}r{replica_id}", + ), + client, + ) + for key in list(readers): + if key not in live: + task, _ = readers.pop(key) + task.cancel() + reaped.append(task) + for stage_id in list(pollers): + if stage_id not in live_pools: + task = pollers.pop(stage_id) + task.cancel() + reaped.append(task) + + _reconcile_readers() + logger.info( + "[Orchestrator] Event-driven orchestration loop enabled (%s): %d LLM replica readers, %d diffusion pollers", + _EVENT_DRIVEN_ORCH_ENV, + len(readers), + len(pollers), + ) + + shutdown_task = asyncio.create_task(self._shutdown_event.wait(), name="orch-shutdown-wait") + pending_get: asyncio.Task | None = None + next_reconcile = _time.monotonic() + _ORCH_READER_RECONCILE_INTERVAL_S + try: + while not self._shutdown_event.is_set(): + if pending_get is None: + pending_get = asyncio.create_task(ready_q.get(), name="orch-dispatch-get") + done, _ = await asyncio.wait( + {pending_get, shutdown_task}, + timeout=max(0.0, next_reconcile - _time.monotonic()), + return_when=asyncio.FIRST_COMPLETED, + ) + if shutdown_task in done: + return + # Reconcile on a wall-clock schedule, not only when the queue + # goes idle: under continuous traffic `asyncio.wait` returns on + # every output and a timeout-only reconcile would never fire, + # so a replica that registers at runtime would never get a + # reader and its outputs would never drain. + if _time.monotonic() >= next_reconcile: + next_reconcile = _time.monotonic() + _ORCH_READER_RECONCILE_INTERVAL_S + _reconcile_readers() + if not done: + self._orch_monitor.note_loop(idle=True) + continue + kind, stage_id, replica_id, payload = pending_get.result() + pending_get = None + + if kind == "error": + # replica_id < 0 means the failure was raised by a poller + # task itself rather than attributed to one replica, so + # there is nothing to evict -- and evict_replica() rejects + # an out-of-range id, which would turn a survivable death + # into a ValueError out of the dispatcher. + if isinstance(payload, EngineDeadError) and replica_id >= 0: + # Same per-replica isolation as the legacy loop (#4285): + # evict and keep serving. Reconcile immediately so the + # dead replica's reader is reaped now rather than at the + # next scheduled tick. Skip a replica already evicted by + # an earlier error from the same slot -- evict_replica() + # would shut a None client down a second time. + if replica_id in self.stage_pools[stage_id].live_replica_ids(): + await self._handle_dead_replica(stage_id, replica_id, payload) + _reconcile_readers() + continue + if self._shutdown_event.is_set(): + return + logger.error( + "[Orchestrator] Stage-%s replica-%s reader failed: %r", + stage_id, + replica_id, + payload, + ) + raise payload + + raw_terminal_request_ids: set[str] = set() + try: + if kind == "diffusion": + pool = self.stage_pools[stage_id] + if self._absorb_diffusion_metrics(stage_id, replica_id, payload): + self._orch_monitor.note_loop(idle=False) + continue + pool.record_output_timestamps([payload]) + processed = [payload] + else: + processed = await self._process_llm_stage_outputs( + stage_id, replica_id, payload, raw_terminal_request_ids + ) + except asyncio.CancelledError: + raise + except EngineDeadError as e: + if replica_id in self.stage_pools[stage_id].live_replica_ids(): + await self._handle_dead_replica(stage_id, replica_id, e) + _reconcile_readers() + continue + except Exception: + if self._shutdown_event.is_set(): + return + logger.exception( + "[Orchestrator] Stage-%s replica-%s processing failed", + stage_id, + replica_id, + ) + raise + + await self._handle_processed_outputs(stage_id, replica_id, processed) + await self._finish_raw_terminal_requests(stage_id, replica_id, raw_terminal_request_ids) + # Mirrors the legacy loop, which sets idle=False only after a + # poll produced outputs that were routed -- an evicted replica + # `continue`s above without marking the tick active. + self._orch_monitor.note_loop(idle=False) + finally: + shutdown_task.cancel() + if pending_get is not None: + pending_get.cancel() + reader_tasks = [task for task, _ in readers.values()] + poller_tasks = list(pollers.values()) + for task in (*reader_tasks, *poller_tasks): + task.cancel() + # `reaped` holds readers/pollers retired mid-run (client swap, + # eviction, stage teardown). They were cancelled but never awaited, + # so gather them here too rather than leaving dangling tasks. + cleanup = [shutdown_task, *reader_tasks, *poller_tasks, *reaped] + if pending_get is not None: + cleanup.append(pending_get) + await asyncio.gather(*cleanup, return_exceptions=True) + async def _handle_processed_outputs(self, stage_id: int, replica_id: int, outputs: list[Any]) -> None: """Route processed stage outputs produced by one stage poll.""" pool = self.stage_pools[stage_id] diff --git a/vllm_omni/entrypoints/async_omni.py b/vllm_omni/entrypoints/async_omni.py index 2ab7b2a25e9..af710cba23f 100644 --- a/vllm_omni/entrypoints/async_omni.py +++ b/vllm_omni/entrypoints/async_omni.py @@ -53,6 +53,11 @@ logger = init_logger(__name__) _FINAL_OUTPUT_IDLE_SLEEP_S = 0.001 +# Blocking-wait interval for the event-driven final-output drain +# (VLLM_OMNI_EVENT_DRIVEN_ORCH=1): a message wakes the drain immediately via +# the janus queue's condition variable; this timeout only bounds how often the +# orchestrator liveness check runs while the pipeline is idle. +_FINAL_OUTPUT_BLOCKING_WAIT_S = 1.0 class AsyncEventResolver: @@ -880,14 +885,29 @@ def _final_output_handler(self) -> None: engine = self.engine + # Event-driven drain (VLLM_OMNI_EVENT_DRIVEN_ORCH=1): block on the + # queue's condition variable in a dedicated thread instead of the + # get_nowait + 1 ms sleep cadence. Same flag as the orchestrator-side + # event-driven loop (vllm_omni/engine/orchestrator.py). + from vllm_omni.engine.orchestrator import _event_driven_orch_enabled + + event_driven_drain = _event_driven_orch_enabled() and hasattr(engine, "get_output_blocking_async") + async def _final_output_loop(): """Background coroutine that dispatches final outputs to request queues.""" try: while True: - msg = await engine.try_get_output_async() - if msg is None: - await asyncio.sleep(_FINAL_OUTPUT_IDLE_SLEEP_S) - continue + if event_driven_drain: + msg = await engine.get_output_blocking_async(timeout=_FINAL_OUTPUT_BLOCKING_WAIT_S) + if msg is None: + # Timed out with the orchestrator alive; loop for + # the periodic liveness check. + continue + else: + msg = await engine.try_get_output_async() + if msg is None: + await asyncio.sleep(_FINAL_OUTPUT_IDLE_SLEEP_S) + continue if isinstance(msg, dict) and msg.get("type") == "ack": ack_data = msg.get("ack")