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
1 change: 1 addition & 0 deletions docs/.nav.yml
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,7 @@ nav:
- Runtime and Stage Execution:
- Execution Modes and Streaming: user_guide/diffusion/execution_modes.md
- Prefill-Decode Disaggregation (experimental): features/pd_disaggregation.md
- Independent Stage Execution (experimental): features/independent_stage_execution.md
- Sleep Mode: features/sleep_mode.md
- Quantization:
- Overview: user_guide/quantization/overview.md
Expand Down
1 change: 1 addition & 0 deletions docs/features/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ an implementation contract has a user-facing workflow.
| --- | --- | --- |
| Choose serial, batched, step-wise, or streaming diffusion execution | [Execution Modes and Streaming](../user_guide/diffusion/execution_modes.md) | [Diffusion Continuous Batching](../design/feature/diffusion_continuous_batching.md), [Async Diffusion Output](../design/feature/async_diffusion_output.md) |
| Reclaim stage memory without restarting the server | [Sleep Mode](sleep_mode.md) | Runtime lifecycle behavior is documented in the user guide |
| Run and route each stage independently (experimental) | [Independent Stage Execution](independent_stage_execution.md) | See user guide |

The [Prefill-Decode Disaggregation design reference](pd_disaggregation.md)
describes the experimental Qwen3-Omni runtime topology and its current
Expand Down
138 changes: 138 additions & 0 deletions docs/features/independent_stage_execution.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,138 @@
# Independent Stage Execution (experimental)

!!! warning
This feature is **experimental**. The `POST /v1/run` request and response
formats may change without notice.

vLLM-Omni supports running stages independently, where the payload is yielded back to the user rather than submitted to the next stage. This provides the groundwork for high-scale distributed serving by decoupling the frontend (API Server & Orchestrator) from independently scalable headless stages, which will eventually allow users to manage their own routing, e.g., based on scoring heuristics.

## Deployment

We will use Qwen3-TTS as an example for independent stage execution, illustrating how to pass the payload from one stage to the next. First, start the frontend, which also brings up the orchestrator in its own thread. For now, this can be accomplished by passing a stage-id of `-1`. Note that you must configure the Omni master server settings for remote replicas.

```bash
# Starts the API server on port 8000
vllm serve Qwen/Qwen3-TTS-12Hz-1.7B-CustomVoice \
--trust-remote-code \
--no-async-chunk \
--stage-id -1 \
--omni-master-address 127.0.0.1 \
--omni-master-port 21212 \
--omni
```

Wait until you see logs indicating that the addresses for remote stages have been pre-allocated.

```text
(APIServer pid=434098) INFO ... Pre-allocated addresses for stages [0, 1] (master=127.0.0.1:21212)
(APIServer pid=434098) INFO ... [OmniMasterServer] Listening on tcp://127.0.0.1:21212
(APIServer pid=434098) INFO ... [DistStageRuntime] OmniMasterServer started for stages [0, 1]
```

In another terminal, start the talker and code2wav stages as shown below:

```bash
# STAGE_ID=0 should be used to start the talker. Use STAGE_ID=1 for code2wav.
STAGE_ID=0

vllm serve Qwen/Qwen3-TTS-12Hz-1.7B-CustomVoice \
--trust-remote-code \
--no-async-chunk \
--stage-id $STAGE_ID \
--headless \
--omni-master-address 127.0.0.1 \
--omni-master-port 21212 \
--omni
```

You should see both stage 0 and stage 1 register as replicas with the `OmniMasterServer` and that the `DistStageRuntime`s have attached successfully. Once the stages have finished initialization, the server will show as ready.

## `POST /v1/run`

In the preliminary implementation, there is no preprocessing or postprocessing done on the request and response objects, so you need to build the stage 0 inputs and unpack the final response directly.

After the first inference call, we can pipe the output directly back into `/v1/run` to run the second stage. Currently, the output is encoded since it is the raw yielded payload, so we provide a helper to unpack the raw message. You can see the full flow below.

```python
import argparse
import json

import httpx
import soundfile as sf
from huggingface_hub import hf_hub_download
from transformers import AutoTokenizer

from vllm_omni.entrypoints.openai.protocol.audio import OpenAICreateSpeechRequest
from vllm_omni.entrypoints.openai.serving_run import decode_output
from vllm_omni.entrypoints.openai.serving_speech import OmniOpenAIServingSpeech
from vllm_omni.entrypoints.openai.tts_adapters.base import conditioning_cache_salt
from vllm_omni.model_executor.models.qwen3_tts.prompt_embeds_builder import Qwen3TTSPromptEmbedsBuilder


def build_stage0_input(model: str, text: str, speaker: str) -> dict:
"""Build the talker's input the way the speech endpoint does (CustomVoice, built-in speaker)."""
tts_params = {"text": [text], "task_type": ["CustomVoice"], "language": ["Auto"], "speaker": [speaker]}
tokenizer = AutoTokenizer.from_pretrained(model, trust_remote_code=True, padding_side="left")
with open(hf_hub_download(model, "config.json")) as f:
talker_config = json.load(f)["talker_config"]
prompt_len = Qwen3TTSPromptEmbedsBuilder.estimate_prompt_len_from_additional_information(
additional_information=tts_params,
task_type="CustomVoice",
tokenize_prompt=lambda t: tokenizer(t, padding=False)["input_ids"],
codec_language_id=talker_config.get("codec_language_id"),
spk_is_dialect=talker_config.get("spk_is_dialect"),
)
cache_salt = conditioning_cache_salt(OpenAICreateSpeechRequest(input=text, voice=speaker), tts_params)
return {"prompt_token_ids": [1] * prompt_len, "additional_information": tts_params, "cache_salt": cache_salt}


def main() -> None:
parser = argparse.ArgumentParser()
# NOTE: the port here is your API server port, not the omni master port
parser.add_argument("--url", default="http://localhost:8000")
parser.add_argument("--model", default="Qwen/Qwen3-TTS-12Hz-1.7B-CustomVoice")
parser.add_argument("--text", default="Hello world")
parser.add_argument("--speaker", default="vivian")
parser.add_argument("--out", default="out.wav")
args = parser.parse_args()

url = f"{args.url}/v1/run"

# Stage 0 (talker) returns the request for stage 1.
stage1_request = httpx.post(
url,
json={
"stage_input": build_stage0_input(args.model, args.text, args.speaker)
},
timeout=600,
)
stage1_request.raise_for_status()

# Stage 1 (code2wav) returns the final output.
final = httpx.post(
url,
content=stage1_request.content,
headers={"Content-Type": "application/json"},
timeout=600,
)
final.raise_for_status()

# For now, since we don't have postprocessing, we call a helper to unpack the message for us
decoded_output = decode_output(final.json()["output"])
# Then extract the speech and save it
audio_output, audio_key = OmniOpenAIServingSpeech._extract_audio_output(decoded_output)
sf.write(args.out, audio_output[audio_key].float().numpy().squeeze(), int(audio_output["sr"]))
print(f"wrote {args.out}")


if __name__ == "__main__":
main()
```

## Limitations

This feature is early in its development and currently has the following constraints, many of which are being actively worked on.

- Support for emulating different entrypoints' preprocessing and postprocessing through `/v1/run`
- API for passing the `/v1/run` request to a specific replica
- Support for async chunking
1 change: 1 addition & 0 deletions docs/serving/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -175,6 +175,7 @@ the other protocol.
| TTS voice management | `GET/POST /v1/audio/voices`, `DELETE /v1/audio/voices/{name}` | [Voices](speech_api.md#voices-endpoint) |
| Video job lifecycle | `GET /v1/videos`, `GET/DELETE /v1/videos/{video_id}`, `GET /v1/videos/{video_id}/content` | [Video endpoints](videos_api.md#endpoints) |
| Release and restore stage memory | `POST /v1/omni/sleep`, `POST /v1/omni/wakeup` | [Sleep Mode](../features/sleep_mode.md) |
| Run one stage per call (experimental) | `POST /v1/run` | [Independent Stage Execution](../features/independent_stage_execution.md) |

## Standalone Experimental Servers

Expand Down
5 changes: 3 additions & 2 deletions tests/core/sched/test_generation_batch_coalescing.py
Original file line number Diff line number Diff line change
Expand Up @@ -97,6 +97,7 @@ def scheduler(monkeypatch, initial=(), delayed=()):
requests_with_ready_chunks=set(), finished_requests=set(), input_terminal_req_ids=set()
)
s._latest_omni_connector_output = None
s._outputs_awaiting_stage_payload = {}
s._omni_connector_output_inbox = Inbox(clock, initial, delayed)
return s, clock

Expand Down Expand Up @@ -167,7 +168,7 @@ def test_deadline_is_new_for_each_batch_not_extended_by_arrivals(monkeypatch):

def test_cancelled_notification_is_filtered_before_coordinator(monkeypatch, mocker):
event = notification("0", "cancelled")
event.request_metadata = {"0": {"code_predictor_codes": [7]}, "cancelled": {"code_predictor_codes": [9]}}
event.request_metadata = {"0": {"codes": {"audio": [7]}}, "cancelled": {"codes": {"audio": [9]}}}
s, _ = scheduler(monkeypatch, [event])
s._generation_max_wait_s = 0
s.waiting = []
Expand All @@ -177,7 +178,7 @@ def test_cancelled_notification_is_filtered_before_coordinator(monkeypatch, mock
s.input_coordinator.process_pending_chunks = mocker.Mock()
s._consume_pending_connector_output("generation")
s.input_coordinator.update_request_metadata.assert_called_once_with(
s.requests, {"0": {"code_predictor_codes": [7]}}, model_mode="generation"
s.requests, {"0": {"codes": {"audio": [7]}}}, model_mode="generation"
)
s.input_coordinator.process_pending_chunks.assert_called_once()
waiting, running, ready, finished = s.input_coordinator.process_pending_chunks.call_args.args
Expand Down
2 changes: 2 additions & 0 deletions tests/core/sched/test_omni_ar_scheduler_logprobs.py
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,7 @@ def __init__(self, request_id: str) -> None:
self.stop_reason = None
self.trace_headers = None
self.num_nans_in_logits = 0
self.additional_information = None

def is_finished(self) -> bool:
return RequestStatus.is_finished(self.status)
Expand Down Expand Up @@ -148,6 +149,7 @@ def _make_scheduler_stub(requests: list[_Request]) -> SimpleNamespace:
waiting_for_transfer_free=set(),
_new_prompt_len_snapshot={},
_pooling_output_decoder=None,
_outputs_awaiting_stage_payload={},
finished_req_ids=set(),
finished_req_ids_dict=defaultdict(set),
kv_cache_manager=SimpleNamespace(take_events=lambda: None),
Expand Down
1 change: 1 addition & 0 deletions tests/core/sched/test_omni_ar_scheduler_streaming.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@ def _make_scheduler(*, stage_id: int = 0, session_mode: str = "turn") -> OmniARS
sched._free_request_blocks = MagicMock()
sched.encoder_cache_manager = MagicMock()
sched._inflight_prefills = set()
sched._outputs_awaiting_stage_payload = {}
return sched


Expand Down
22 changes: 22 additions & 0 deletions tests/core/sched/test_omni_scheduler_finish_requests_purge.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
import vllm_omni.core.sched.omni_generation_scheduler as gen_sched_mod
from vllm_omni.core.sched.omni_ar_scheduler import OmniARScheduler
from vllm_omni.core.sched.omni_generation_scheduler import OmniGenerationScheduler
from vllm_omni.engine import OmniEngineCoreOutput

pytestmark = [pytest.mark.core_model, pytest.mark.cpu]

Expand Down Expand Up @@ -57,6 +58,7 @@ def _make_scheduler(scheduler_cls, *, requests, running, waiting):
scheduler.waiting = waiting
scheduler.kv_holding_waiting = []
scheduler.deferred_waiting = set()
scheduler._outputs_awaiting_stage_payload = {}
return scheduler


Expand Down Expand Up @@ -230,3 +232,23 @@ def test_finish_requests_does_not_reopen_off_queue_deferred_free_terminal(schedu
assert scheduler_cls.finish_requests(scheduler, [terminal.request_id], RequestStatus.FINISHED_ABORTED) == []
assert scheduler.requests == {terminal.request_id: terminal}
assert terminal.status == RequestStatus.FINISHED_STOPPED


# None finishes every request (abort-all pause, fault recovery).
@pytest.mark.parametrize("request_ids", [["run"], None])
@pytest.mark.parametrize(("scheduler_cls", "scheduler_mod"), _SCHEDULER_PARAMS)
def test_finish_requests_drops_the_held_output_of_a_run_request(
monkeypatch: pytest.MonkeyPatch,
scheduler_cls,
scheduler_mod,
request_ids: list[str] | None,
) -> None:
"""Ensure aborting a run request whose finished output waits for its stage payload drops that output."""
# The run request already finished and was freed; only its output is held.
scheduler = _make_scheduler(scheduler_cls, requests={}, running=[], waiting=[])
scheduler._outputs_awaiting_stage_payload = {"run": (0, OmniEngineCoreOutput(request_id="run", new_token_ids=[]))}
monkeypatch.setattr(scheduler_mod.VLLMScheduler, "finish_requests", lambda self, request_ids, finished_status: [])

scheduler_cls.finish_requests(scheduler, request_ids, RequestStatus.FINISHED_ABORTED)

assert scheduler._outputs_awaiting_stage_payload == {}
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@ def test_finish_requests_cleans_input_coordinator_for_finished_ids(
scheduler.waiting = []
scheduler.kv_holding_waiting = []
scheduler.deferred_waiting = set()
scheduler._outputs_awaiting_stage_payload = {}

def fake_finish_requests(self, request_ids, finished_status):
assert request_ids == ["req-a", "req-b"]
Expand Down
40 changes: 40 additions & 0 deletions tests/core/sched/test_omni_scheduler_mixin_shared.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
from vllm_omni.core.sched import omni_scheduler_mixin
from vllm_omni.core.sched.omni_scheduler_mixin import OmniSchedulerMixin
from vllm_omni.core.sched.output import OmniChunkRecvHandle
from vllm_omni.data_entry_keys import RETURN_STAGE_PAYLOAD_KEY
from vllm_omni.outputs import OmniConnectorOutput

pytestmark = [pytest.mark.core_model, pytest.mark.cpu]

Expand Down Expand Up @@ -182,6 +184,7 @@ def test_finished_request_attachment_keeps_ar_abort_policy_explicit(
):
scheduler = _Scheduler()
scheduler.finished_req_ids_dict = defaultdict(set, {2: {"req-finished"}})
scheduler._outputs_awaiting_stage_payload = {}
outputs: dict = {}

scheduler._attach_finished_request_sets(
Expand Down Expand Up @@ -237,3 +240,40 @@ def test_native_downstream_sender_stage_does_not_wait_for_chunks(role, coordinat
assert scheduler._native_data_plane
assert (scheduler.input_coordinator is not None) is coordinated
assert scheduler._async_chunk_transport_enabled() is coordinated


def test_only_run_requests_wait_for_their_stage_payload():
"""Ensure a run request's finished output is held, not reported as aborted, then sent once with its payload."""
scheduler = _Scheduler()
# NOTE: run requests don't currently support async chunk
scheduler.vllm_config = SimpleNamespace(model_config=SimpleNamespace(stage_id=0, async_chunk=False))
scheduler._init_omni_io_scheduling_state()
scheduler.finished_req_ids_dict = defaultdict(set, {0: {"run"}})
outputs: dict = defaultdict(list)

for request_id, tags in (("normal", None), ("run", {RETURN_STAGE_PAYLOAD_KEY: True})):
request = SimpleNamespace(
request_id=request_id,
client_index=0,
trace_headers=None,
take_events=lambda: [],
additional_information=tags,
)
scheduler._append_request_output(outputs, request, new_token_ids=[], finish_reason=FinishReason.STOP)

# Ensure that attaching the finished request sets emits nothing for the held run request
# while the normal request's output is already in this step's outputs.
held_step: dict = {}
scheduler._attach_finished_request_sets(held_step, synthesize_abort_outputs=True)
assert [output.request_id for output in outputs[0]] == ["normal"]
assert held_step[0].outputs == []

# Ensure that once the worker's stage payload arrives, the held
# run request's finished is emitted once with the payload,
scheduler.enqueue_omni_connector_output(OmniConnectorOutput(stage_payloads={"run": b"payload"}))
scheduler._consume_pending_connector_output("ar")
released_step: dict = {}
scheduler._attach_finished_request_sets(released_step, synthesize_abort_outputs=True)

[output] = released_step[0].outputs
assert (output.request_id, output.finish_reason, output.stage_payload) == ("run", FinishReason.STOP, b"payload")
2 changes: 2 additions & 0 deletions tests/core/sched/test_omni_scheduler_mixin_timeouts.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ def __init__(self, requests, coordinator):
self.requests = requests
self.input_coordinator = coordinator
self.finish_calls = []
self._outputs_awaiting_stage_payload = {}

def finish_requests(self, req_ids, status):
self.finish_calls.append((set(req_ids), status))
Expand Down Expand Up @@ -127,6 +128,7 @@ def __init__(self, requests, adapter):
self.requests = requests
self.chunk_transfer_adapter = adapter
self.finish_calls = []
self._outputs_awaiting_stage_payload = {}

def finish_requests(self, req_ids, status):
self.finish_calls.append((set(req_ids), status))
Expand Down
3 changes: 3 additions & 0 deletions tests/core/sched/test_omni_scheduler_ready_inbox.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ def test_ready_inbox_coalesces_live_events_and_drops_cancelled(mocker):
_async_chunk=True, update_request_metadata=mocker.Mock(), process_pending_chunks=mocker.Mock()
)
scheduler.input_coordinator = coordinator
scheduler._outputs_awaiting_stage_payload = {}
scheduler._init_omni_connector_output_inbox()
scheduler.enqueue_omni_connector_output(OmniConnectorOutput(chunk_ready_req_ids={"r"}))
scheduler.enqueue_omni_connector_output(
Expand Down Expand Up @@ -54,6 +55,7 @@ def test_native_input_gate_parks_and_restores_kv_holders(policy, async_chunk):
scheduler.waiting = create_request_queue(policy)
scheduler.kv_holding_waiting = create_request_queue(policy)
scheduler.deferred_waiting = set()
scheduler._outputs_awaiting_stage_payload = {}
scheduler.running = []
scheduler.chunk_transfer_adapter = None
scheduler.input_coordinator = OmniSchedulingCoordinator(
Expand Down Expand Up @@ -110,6 +112,7 @@ def test_native_input_gate_preserves_upstream_blocked_waits(async_chunk, status)
scheduler.waiting = create_request_queue(SchedulingPolicy.FCFS)
scheduler.kv_holding_waiting = create_request_queue(SchedulingPolicy.FCFS)
scheduler.deferred_waiting = set()
scheduler._outputs_awaiting_stage_payload = {}
scheduler.running = []
request = Request("blocked", [1, 2, 3], SamplingParams(max_tokens=1), pooling_params=None)
request.status = status
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -62,6 +62,7 @@ def _make_scheduler(scheduler_cls, *, requests, running, waiting, counter, skipp
scheduler.kv_holding_waiting = skipped if skipped is not None else []
scheduler.deferred_waiting = set()
scheduler.num_waiting_for_streaming_input = counter
scheduler._outputs_awaiting_stage_payload = {}
return scheduler


Expand Down
Loading
Loading