Skip to content

Qwen3-Omni: Offline inference

Source https://github.com/vllm-project/vllm-omni/tree/main/examples/offline_inference/qwen3_omni.

Setup

Use --deploy-config for deployment overrides such as stage memory allocation. See the pipeline and deploy configuration documentation.

Run examples

Multiple Prompts

Get into the example folder

cd examples/offline_inference/qwen3_omni
Then run the command below. Note: for processing large volume data, it uses py_generator mode, which will return a python generator from Omni class.
bash run_multiple_prompts.sh

Single Prompt

Get into the example folder

cd examples/offline_inference/qwen3_omni
Then run the command below.
bash run_single_prompt.sh
If you have not enough memory, you can set thinker with tensor parallel. Just run the command below.
bash run_single_prompt_tp.sh

Modality control

If you want to control output modalities, e.g. only output text, you can run the command below:

python end2end.py --output-wav output_audio \
                  --query-type use_audio \
                  --modalities text

Using Local Media Files

The end2end.py script supports local media files (audio, video, image) via command-line arguments:

# Use local video file
python end2end.py --query-type use_video --video-path /path/to/video.mp4

# Use local image file
python end2end.py --query-type use_image --image-path /path/to/image.jpg

# Use local audio file
python end2end.py --query-type use_audio --audio-path /path/to/audio.wav

# Combine multiple local media files
python end2end.py --query-type mixed_modalities \
    --video-path /path/to/video.mp4 \
    --image-path /path/to/image.jpg \
    --audio-path /path/to/audio.wav

If media file paths are not provided, the script will use default assets. Supported query types: - use_video: Video input - use_image: Image input - use_audio: Audio input - text: Text-only query - multi_audios: Multiple audio inputs - mixed_modalities: Combination of video, image, and audio inputs

Async-chunk (offline)

For true stage-level concurrency -- where downstream stages (Talker, Code2Wav) start before the upstream stage (Thinker) finishes -- use the async_chunk example. This requires:

  1. A deploy YAML with async_chunk: true (for example, an overlay based on vllm_omni/deploy/qwen3_omni_moe.yaml).
  2. Hardware that matches the config (e.g. 2x H100 for the default 3-stage config).

The async_chunk example uses AsyncOmni instead of the synchronous Omni class, which enables the async orchestrator to receive stage-0 intermediate outputs and trigger downstream stages early. Chunk data flows directly between stage workers via the in-worker OmniChunkTransferAdapter / connector, not through the orchestrator.

Single prompt

cd examples/offline_inference/qwen3_omni
bash run_single_prompt_async_chunk.sh

Multiple prompts with concurrency control

bash run_multiple_prompts_async_chunk.sh --max-in-flight 4

Text-only output (skip audio generation)

python end2end_async_chunk.py --query-type text --modalities text

Custom deploy config

python end2end_async_chunk.py \
    --query-type use_audio \
    --deploy-config /path/to/your_async_chunk.yaml

Note: The synchronous end2end.py (using Omni) is still the recommended entry point for non-async-chunk workflows. Only use the async_chunk example when you need the stage-level concurrency semantics described in PR #962 / #1151.

Example materials

end2end.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
This example shows how to use vLLM for running offline inference
with the correct prompt format on Qwen3-Omni (thinker only).
"""

import os
import time
from typing import NamedTuple

import numpy as np
import soundfile as sf
import vllm
from PIL import Image
from vllm import SamplingParams
from vllm.assets.audio import AudioAsset
from vllm.assets.image import ImageAsset
from vllm.assets.video import VideoAsset, video_to_ndarrays
from vllm.multimodal.image import convert_image_mode
from vllm.multimodal.media.audio import load_audio

from vllm_omni.entrypoints.omni import Omni
from vllm_omni.utils.tracking_parser import TrackingArgumentParser

SEED = 42


class QueryResult(NamedTuple):
    inputs: dict
    limit_mm_per_prompt: dict[str, int]


# NOTE: The default `max_num_seqs` and `max_model_len` may result in OOM on
# lower-end GPUs.
# Unless specified, these settings have been tested to work on a single L4.

default_system = (
    "You are Qwen, a virtual human developed by the Qwen Team, Alibaba "
    "Group, capable of perceiving auditory and visual inputs, as well as "
    "generating text and speech."
)


def get_text_query(question: str = None) -> QueryResult:
    if question is None:
        question = "Explain the system architecture for a scalable audio generation pipeline. Answer in 15 words."
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n"
        f"{question}<|im_end|>\n"
        f"<|im_start|>assistant\n"
    )
    return QueryResult(
        inputs={
            "prompt": prompt,
        },
        limit_mm_per_prompt={},
    )


def get_video_query(question: str = None, video_path: str | None = None, num_frames: int = 16) -> QueryResult:
    if question is None:
        question = "Why is this video funny?"
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|vision_start|><|video_pad|><|vision_end|>"
        f"{question}<|im_end|>\n"
        f"<|im_start|>assistant\n"
    )

    if video_path:
        if not os.path.exists(video_path):
            raise FileNotFoundError(f"Video file not found: {video_path}")
        video_frames = video_to_ndarrays(video_path, num_frames=num_frames)
    else:
        video_frames = VideoAsset(name="baby_reading", num_frames=num_frames).np_ndarrays

    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {
                "video": video_frames,
            },
        },
        limit_mm_per_prompt={"video": 1},
    )


def get_image_query(question: str = None, image_path: str | None = None) -> QueryResult:
    if question is None:
        question = "What is the content of this image?"
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>"
        f"{question}<|im_end|>\n"
        f"<|im_start|>assistant\n"
    )

    if image_path:
        if not os.path.exists(image_path):
            raise FileNotFoundError(f"Image file not found: {image_path}")
        pil_image = Image.open(image_path)
        image_data = convert_image_mode(pil_image, "RGB")
    else:
        image_data = convert_image_mode(ImageAsset("cherry_blossom").pil_image, "RGB")

    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {
                "image": image_data,
            },
        },
        limit_mm_per_prompt={"image": 1},
    )


def get_audio_query(question: str = None, audio_path: str | None = None, sampling_rate: int = 16000) -> QueryResult:
    if question is None:
        question = "What is the content of this audio?"
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|audio_start|><|audio_pad|><|audio_end|>"
        f"{question}<|im_end|>\n"
        f"<|im_start|>assistant\n"
    )

    if audio_path:
        if not os.path.exists(audio_path):
            raise FileNotFoundError(f"Audio file not found: {audio_path}")
        audio_signal, sr = load_audio(audio_path, sr=sampling_rate)
        audio_data = (audio_signal.astype(np.float32), sr)
    else:
        audio_data = AudioAsset("mary_had_lamb").audio_and_sample_rate

    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {
                "audio": audio_data,
            },
        },
        limit_mm_per_prompt={"audio": 1},
    )


def get_mixed_modalities_query(
    video_path: str | None = None,
    image_path: str | None = None,
    audio_path: str | None = None,
    num_frames: int = 16,
    sampling_rate: int = 16000,
) -> QueryResult:
    question = "What is recited in the audio? What is the content of this image? Why is this video funny?"
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|audio_start|><|audio_pad|><|audio_end|>"
        "<|vision_start|><|image_pad|><|vision_end|>"
        "<|vision_start|><|video_pad|><|vision_end|>"
        f"{question}<|im_end|>\n"
        f"<|im_start|>assistant\n"
    )

    # Load video
    if video_path:
        if not os.path.exists(video_path):
            raise FileNotFoundError(f"Video file not found: {video_path}")
        video_frames = video_to_ndarrays(video_path, num_frames=num_frames)
    else:
        video_frames = VideoAsset(name="baby_reading", num_frames=num_frames).np_ndarrays

    # Load image
    if image_path:
        if not os.path.exists(image_path):
            raise FileNotFoundError(f"Image file not found: {image_path}")
        pil_image = Image.open(image_path)
        image_data = convert_image_mode(pil_image, "RGB")
    else:
        image_data = convert_image_mode(ImageAsset("cherry_blossom").pil_image, "RGB")

    # Load audio
    if audio_path:
        if not os.path.exists(audio_path):
            raise FileNotFoundError(f"Audio file not found: {audio_path}")
        audio_signal, sr = load_audio(audio_path, sr=sampling_rate)
        audio_data = (audio_signal.astype(np.float32), sr)
    else:
        audio_data = AudioAsset("mary_had_lamb").audio_and_sample_rate

    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {
                "audio": audio_data,
                "image": image_data,
                "video": video_frames,
            },
        },
        limit_mm_per_prompt={"audio": 1, "image": 1, "video": 1},
    )


def get_multi_audios_query() -> QueryResult:
    question = "Are these two audio clips the same?"
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|audio_start|><|audio_pad|><|audio_end|>"
        "<|audio_start|><|audio_pad|><|audio_end|>"
        f"{question}<|im_end|>\n"
        f"<|im_start|>assistant\n"
    )
    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {
                "audio": [
                    AudioAsset("winning_call").audio_and_sample_rate,
                    AudioAsset("mary_had_lamb").audio_and_sample_rate,
                ],
            },
        },
        limit_mm_per_prompt={
            "audio": 2,
        },
    )


def get_use_audio_in_video_query() -> QueryResult:
    question = "Describe the content of the video in details, then convert what the baby say into text."
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|vision_start|><|video_pad|><|vision_end|>"
        f"{question}<|im_end|>\n"
        f"<|im_start|>assistant\n"
    )
    asset = VideoAsset(name="baby_reading", num_frames=16)
    audio = asset.get_audio(sampling_rate=16000)
    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {
                "video": asset.np_ndarrays,
                "audio": audio,
            },
            "mm_processor_kwargs": {
                "use_audio_in_video": True,
            },
        },
        limit_mm_per_prompt={"audio": 1, "video": 1},
    )


query_map = {
    "text": get_text_query,
    "use_audio": get_audio_query,
    "use_image": get_image_query,
    "use_video": get_video_query,
    "use_multi_audios": get_multi_audios_query,
    "use_mixed_modalities": get_mixed_modalities_query,
    "use_audio_in_video": get_use_audio_in_video_query,
}


def main(args):
    model_name = args.model
    print("=" * 20, "\n", f"vllm version: {vllm.__version__}", "\n", "=" * 20)

    # Get paths from args
    video_path = getattr(args, "video_path", None)
    image_path = getattr(args, "image_path", None)
    audio_path = getattr(args, "audio_path", None)

    # Get the query function and call it with appropriate parameters
    query_func = query_map[args.query_type]
    if args.query_type == "use_video":
        query_result = query_func(video_path=video_path, num_frames=getattr(args, "num_frames", 16))
    elif args.query_type == "use_image":
        query_result = query_func(image_path=image_path)
    elif args.query_type == "use_audio":
        query_result = query_func(audio_path=audio_path, sampling_rate=getattr(args, "sampling_rate", 16000))
    elif args.query_type == "mixed_modalities":
        query_result = query_func(
            video_path=video_path,
            image_path=image_path,
            audio_path=audio_path,
            num_frames=getattr(args, "num_frames", 16),
            sampling_rate=getattr(args, "sampling_rate", 16000),
        )
    elif args.query_type == "multi_audios":
        query_result = query_func()
    elif args.query_type == "use_audio_in_video":
        query_result = query_func()
    else:
        query_result = query_func()

    omni_kwargs = vars(args).copy()
    # Override CLI --model with the derived model_name.
    omni_kwargs["model"] = model_name
    omni = Omni(**omni_kwargs)

    thinker_sampling_params = SamplingParams(
        temperature=0.9,
        top_p=0.9,
        top_k=-1,
        max_tokens=1200,
        repetition_penalty=1.05,
        logit_bias={},
        seed=SEED,
    )

    talker_sampling_params = SamplingParams(
        temperature=0.9,
        top_k=50,
        max_tokens=4096,
        seed=SEED,
        detokenize=False,
        repetition_penalty=1.05,
        stop_token_ids=[2150],  # TALKER_CODEC_EOS_TOKEN_ID
    )

    # Sampling parameters for Code2Wav stage (audio generation)
    code2wav_sampling_params = SamplingParams(
        temperature=0.0,
        top_p=1.0,
        top_k=-1,
        max_tokens=4096 * 16,
        seed=SEED,
        detokenize=True,
        repetition_penalty=1.1,
    )

    all_sampling_params = [
        thinker_sampling_params,
        talker_sampling_params,  # code predictor is integrated into talker for Qwen3 Omni
        code2wav_sampling_params,
    ]
    # Match sampling params to the number of configured stages
    num_stages = omni.num_stages
    sampling_params_list = all_sampling_params[:num_stages]

    if args.txt_prompts is None:
        prompts = [query_result.inputs for _ in range(args.num_prompts)]
    else:
        assert args.query_type == "text", "txt-prompts is only supported for text query type"
        with open(args.txt_prompts, encoding="utf-8") as f:
            lines = [ln.strip() for ln in f.readlines()]
            prompts = [get_text_query(ln).inputs for ln in lines if ln != ""]
            print(f"[Info] Loaded {len(prompts)} prompts from {args.txt_prompts}")

    if args.modalities is not None:
        output_modalities = args.modalities.split(",")
        for i, prompt in enumerate(prompts):
            prompt["modalities"] = output_modalities

    profiler_enabled = args.enable_profiler
    if profiler_enabled:
        omni.start_profile(stages=args.profiler_stages)
    omni_generator = omni.generate(prompts, sampling_params_list, py_generator=args.py_generator)
    # Determine output directory: prefer --output-dir; fallback to --output-wav
    output_dir = args.output_dir if getattr(args, "output_dir", None) else args.output_wav
    os.makedirs(output_dir, exist_ok=True)

    total_requests = len(prompts)
    processed_count = 0

    print(f"query type: {args.query_type}")

    for stage_outputs in omni_generator:
        output = stage_outputs
        if stage_outputs.final_output_type == "text":
            request_id = output.request_id
            text_output = output.outputs[0].text
            # Save aligned text file per request
            prompt_text = output.prompt
            out_txt = os.path.join(output_dir, f"{request_id}.txt")
            lines = []
            lines.append("Prompt:\n")
            lines.append(str(prompt_text) + "\n")
            lines.append("vllm_text_output:\n")
            lines.append(str(text_output).strip() + "\n")
            try:
                with open(out_txt, "w", encoding="utf-8") as f:
                    f.writelines(lines)
            except Exception as e:
                print(f"[Warn] Failed writing text file {out_txt}: {e}")
            print(f"Request ID: {request_id}, Text saved to {out_txt}")
        elif stage_outputs.final_output_type == "audio":
            request_id = output.request_id
            audio_tensor = output.outputs[0].multimodal_output["audio"]
            output_wav = os.path.join(output_dir, f"output_{request_id}.wav")

            # Convert to numpy array and ensure correct format
            # In async_chunk mode, audio may arrive as a list of chunks
            if isinstance(audio_tensor, list):
                import torch

                audio_tensor = torch.cat(
                    [(t if isinstance(t, torch.Tensor) else torch.tensor(t)).flatten() for t in audio_tensor]
                )
            audio_numpy = audio_tensor.float().detach().cpu().numpy()

            # Ensure audio is 1D (flatten if needed)
            if audio_numpy.ndim > 1:
                audio_numpy = audio_numpy.flatten()

            # Save audio file with explicit WAV format
            sf.write(output_wav, audio_numpy, samplerate=24000, format="WAV")
            print(f"Request ID: {request_id}, Saved audio to {output_wav}")

        processed_count += 1
        if profiler_enabled and processed_count >= total_requests:
            print(f"[Info] Processed {processed_count}/{total_requests}. Stopping profiler inside active loop...")
            # Stop the profiler while workers are still alive
            omni.stop_profile(stages=args.profiler_stages)

            print("[Info] Waiting 30s for workers to write trace files to disk...")
            time.sleep(30)
            print("[Info] Trace export wait time finished.")
    omni.close()


def parse_args():
    parser = TrackingArgumentParser(description="Demo on using vLLM for offline inference with audio language models")
    parser.add_argument(
        "--model",
        type=str,
        default="Qwen/Qwen3-Omni-30B-A3B-Instruct",
        help="Model name or path.",
    )
    parser.add_argument(
        "--query-type",
        "-q",
        type=str,
        default="use_mixed_modalities",
        choices=query_map.keys(),
        help="Query type.",
    )
    parser.add_argument(
        "--log-stats",
        action="store_true",
        default=False,
        help="Enable writing detailed statistics (default: disabled)",
    )
    parser.add_argument(
        "--stage-init-timeout",
        type=int,
        default=300,
        help="Timeout for initializing a single stage in seconds (default: 300)",
    )
    parser.add_argument(
        "--batch-timeout",
        type=int,
        default=5,
        help="Timeout for batching in seconds (default: 5)",
    )
    parser.add_argument(
        "--init-timeout",
        type=int,
        default=300,
        help="Timeout for initializing stages in seconds (default: 300)",
    )
    parser.add_argument(
        "--shm-threshold-bytes",
        type=int,
        default=65536,
        help="Threshold for using shared memory in bytes (default: 65536)",
    )
    parser.add_argument(
        "--output-wav",
        default="output_audio",
        help="[Deprecated] Output wav directory (use --output-dir).",
    )
    parser.add_argument(
        "--num-prompts",
        type=int,
        default=1,
        help="Number of prompts to generate.",
    )
    parser.add_argument(
        "--txt-prompts",
        type=str,
        default=None,
        help="Path to a .txt file with one prompt per line (preferred).",
    )
    parser.add_argument(
        "--deploy-config",
        type=str,
        default=None,
        help="Path to a deploy config YAML.",
    )
    parser.add_argument(
        "--video-path",
        "-v",
        type=str,
        default=None,
        help="Path to local video file. If not provided, uses default video asset.",
    )
    parser.add_argument(
        "--image-path",
        "-i",
        type=str,
        default=None,
        help="Path to local image file. If not provided, uses default image asset.",
    )
    parser.add_argument(
        "--audio-path",
        "-a",
        type=str,
        default=None,
        help="Path to local audio file. If not provided, uses default audio asset.",
    )
    parser.add_argument(
        "--num-frames",
        type=int,
        default=16,
        help="Number of frames to extract from video (default: 16).",
    )
    parser.add_argument(
        "--sampling-rate",
        type=int,
        default=16000,
        help="Sampling rate for audio loading (default: 16000).",
    )
    parser.add_argument(
        "--log-dir",
        type=str,
        default="logs",
        help="Log directory (default: logs).",
    )
    parser.add_argument(
        "--modalities",
        type=str,
        default=None,
        help="Output modalities to use for the prompts.",
    )
    parser.add_argument(
        "--py-generator",
        action="store_true",
        default=False,
        help="Use py_generator mode. The returned type of Omni.generate() is a Python Generator object.",
    )
    parser.add_argument(
        "--enable-profiler",
        action="store_true",
        default=False,
        help="Enables profiling when set.",
    )
    parser.add_argument(
        "--profiler-stages",
        type=int,
        nargs="*",
        default=None,
        help="List of stage IDs to profile. If not set, profiles all stages.",
    )
    parser.add_argument(
        "--dtype",
        type=str,
        default="auto",
        help="Model dtype (auto, half, float16, bfloat16, float, float32).",
    )

    return parser.parse_args()


if __name__ == "__main__":
    args = parse_args()
    main(args)
end2end_async_chunk.py
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Offline inference with async_chunk enabled via AsyncOmni.

This script uses AsyncOmni (the async orchestrator) to run offline inference
with async_chunk semantics: downstream stages (Talker, Code2Wav) start
*before* upstream stages finish, consuming chunks as they arrive via
the in-worker OmniChunkTransferAdapter / connector.

Compared to the synchronous ``end2end.py`` (which uses ``Omni``), this
entry point achieves true stage-level concurrency -- stage-1/2 are
actively processing while stage-0 is still generating.

Usage
-----
    python end2end_async_chunk.py --query-type use_audio \
        --deploy-config <path-to-deploy-config-yaml>

See ``--help`` for all options.
"""

import asyncio
import logging
import os
import time
import uuid
from typing import NamedTuple

import numpy as np
import soundfile as sf
import torch

os.environ["VLLM_WORKER_MULTIPROC_METHOD"] = "spawn"

from PIL import Image
from vllm import SamplingParams
from vllm.assets.audio import AudioAsset
from vllm.assets.image import ImageAsset
from vllm.assets.video import VideoAsset, video_to_ndarrays
from vllm.multimodal.image import convert_image_mode
from vllm.multimodal.media.audio import load_audio

from vllm_omni.entrypoints.async_omni import AsyncOmni
from vllm_omni.utils.tracking_parser import TrackingArgumentParser

logger = logging.getLogger(__name__)

# ---------------------------------------------------------------------------
# Query builders (reuse the patterns from end2end.py)
# ---------------------------------------------------------------------------

default_system = (
    "You are Qwen, a virtual human developed by the Qwen Team, Alibaba "
    "Group, capable of perceiving auditory and visual inputs, as well as "
    "generating text and speech."
)


class QueryResult(NamedTuple):
    inputs: dict
    limit_mm_per_prompt: dict[str, int]


def get_text_query(question: str = None) -> QueryResult:
    if question is None:
        question = "Explain the system architecture for a scalable audio generation pipeline. Answer in 15 words."
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n"
        f"{question}<|im_end|>\n"
        "<|im_start|>assistant\n"
    )
    return QueryResult(inputs={"prompt": prompt}, limit_mm_per_prompt={})


def get_audio_query(
    question: str = None,
    audio_path: str | None = None,
    sampling_rate: int = 16000,
) -> QueryResult:
    if question is None:
        question = "What is the content of this audio?"
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|audio_start|><|audio_pad|><|audio_end|>"
        f"{question}<|im_end|>\n"
        "<|im_start|>assistant\n"
    )
    if audio_path:
        if not os.path.exists(audio_path):
            raise FileNotFoundError(f"Audio file not found: {audio_path}")
        audio_signal, sr = load_audio(audio_path, sr=sampling_rate)
        audio_data = (audio_signal.astype(np.float32), sr)
    else:
        audio_data = AudioAsset("mary_had_lamb").audio_and_sample_rate
    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {"audio": audio_data},
        },
        limit_mm_per_prompt={"audio": 1},
    )


def get_image_query(question: str = None, image_path: str | None = None) -> QueryResult:
    if question is None:
        question = "What is the content of this image?"
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|vision_start|><|image_pad|><|vision_end|>"
        f"{question}<|im_end|>\n"
        "<|im_start|>assistant\n"
    )
    if image_path:
        if not os.path.exists(image_path):
            raise FileNotFoundError(f"Image file not found: {image_path}")
        pil_image = Image.open(image_path)
        image_data = convert_image_mode(pil_image, "RGB")
    else:
        image_data = convert_image_mode(ImageAsset("cherry_blossom").pil_image, "RGB")
    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {"image": image_data},
        },
        limit_mm_per_prompt={"image": 1},
    )


def get_video_query(
    question: str = None,
    video_path: str | None = None,
    num_frames: int = 16,
) -> QueryResult:
    if question is None:
        question = "Why is this video funny?"
    prompt = (
        f"<|im_start|>system\n{default_system}<|im_end|>\n"
        "<|im_start|>user\n<|vision_start|><|video_pad|><|vision_end|>"
        f"{question}<|im_end|>\n"
        "<|im_start|>assistant\n"
    )
    if video_path:
        if not os.path.exists(video_path):
            raise FileNotFoundError(f"Video file not found: {video_path}")
        video_frames = video_to_ndarrays(video_path, num_frames=num_frames)
    else:
        video_frames = VideoAsset(name="baby_reading", num_frames=num_frames).np_ndarrays
    return QueryResult(
        inputs={
            "prompt": prompt,
            "multi_modal_data": {"video": video_frames},
        },
        limit_mm_per_prompt={"video": 1},
    )


query_map = {
    "text": get_text_query,
    "use_audio": get_audio_query,
    "use_image": get_image_query,
    "use_video": get_video_query,
}

# ---------------------------------------------------------------------------
# Core async routine
# ---------------------------------------------------------------------------


def clone_prompt_for_request(template: dict) -> dict:
    """Shallow-clone prompt dict so concurrent requests own independent containers."""
    cloned = dict(template)
    for key in ("multi_modal_data", "mm_processor_kwargs", "additional_information"):
        value = template.get(key)
        if isinstance(value, dict):
            cloned[key] = dict(value)
        elif isinstance(value, list):
            cloned[key] = list(value)
    return cloned


def _default_deploy_config_path() -> str | None:
    """Best-effort default deploy config for running Qwen3-Omni with async_chunk.

    The default ``vllm_omni/deploy/qwen3_omni_moe.yaml`` ships with
    ``async_chunk: true`` at the top level, so loading it is enough to
    enable async-chunk semantics. To disable it, copy the YAML and set
    ``async_chunk: false`` (or pass ``--deploy-config`` to a YAML that
    overrides the flag).

    When this example is executed from within the repository, we resolve
    the default YAML path relative to this file. When installed elsewhere,
    the file may not exist and callers should pass ``--deploy-config``
    explicitly.
    """
    repo_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "../../.."))
    candidate = os.path.join(
        repo_root,
        "vllm_omni",
        "deploy",
        "qwen3_omni_moe.yaml",
    )
    return candidate if os.path.exists(candidate) else None


async def run_single_request(
    async_omni: AsyncOmni,
    prompt: dict,
    request_id: str,
    sampling_params_list: list[SamplingParams] | None,
    output_dir: str,
    output_modalities: list[str] | None = None,
    stream_audio_to_disk: bool = False,
) -> dict:
    """Run one request through AsyncOmni and collect outputs.

    Returns a dict with timing information and saved file paths.
    """
    t_start = time.perf_counter()
    text_parts: list[str] = []
    audio_chunks: list[torch.Tensor] = []
    audio_sr: int | None = None
    first_audio_ts: float | None = None
    audio_list_consumed: int = 0
    audio_last_tensor: torch.Tensor | None = None
    stage_0_first_output_ts: float | None = None

    samplerate = 24000
    wav_file = os.path.join(output_dir, f"output_{request_id}.wav")
    sf_writer: sf.SoundFile | None = None
    audio_samples_written: int = 0

    try:
        async for omni_output in async_omni.generate(
            prompt=prompt,
            request_id=request_id,
            sampling_params_list=sampling_params_list,
            output_modalities=output_modalities,
        ):
            output = omni_output
            if omni_output.final_output_type == "text":
                if stage_0_first_output_ts is None:
                    stage_0_first_output_ts = time.perf_counter()
                text_output = output.outputs[0].text
                text_parts.append(text_output)
            elif omni_output.final_output_type == "audio":
                mm_out = output.outputs[0].multimodal_output
                if mm_out and "audio" in mm_out:
                    if first_audio_ts is None:
                        first_audio_ts = time.perf_counter()
                    if audio_sr is None and "sr" in mm_out:
                        sr_val = mm_out["sr"]
                        audio_sr = sr_val.item() if hasattr(sr_val, "item") else int(sr_val)
                        samplerate = audio_sr
                    audio_data = mm_out["audio"]
                    if isinstance(audio_data, list):
                        new_chunks = audio_data[audio_list_consumed:]
                        audio_list_consumed = len(audio_data)
                    elif isinstance(audio_data, torch.Tensor):
                        new_chunks = [audio_data]
                        audio_last_tensor = audio_data
                    else:
                        new_chunks = []

                    if stream_audio_to_disk and new_chunks:
                        if sf_writer is None:
                            sf_writer = sf.SoundFile(
                                wav_file,
                                mode="w",
                                samplerate=samplerate,
                                channels=1,
                                subtype="FLOAT",
                            )
                        for chunk in new_chunks:
                            chunk_np = chunk.float().detach().cpu().numpy().flatten()
                            sf_writer.write(chunk_np)
                            audio_samples_written += len(chunk_np)
                    else:
                        audio_chunks.extend(new_chunks)
    finally:
        if sf_writer is not None:
            sf_writer.close()

    t_end = time.perf_counter()
    result = {
        "request_id": request_id,
        "e2e_latency_s": t_end - t_start,
        "saved_files": [],
    }

    # Save text output
    if text_parts:
        text_file = os.path.join(output_dir, f"{request_id}.txt")
        with open(text_file, "w", encoding="utf-8") as f:
            f.write("".join(text_parts))
        result["saved_files"].append(text_file)
        print(
            f"[Request {request_id}] Text saved to {text_file} "
            f"(stage-0 first output at {stage_0_first_output_ts - t_start:.3f}s)"
        )

    # Save audio output
    if stream_audio_to_disk and audio_samples_written > 0:
        result["saved_files"].append(wav_file)
        result["audio_duration_s"] = audio_samples_written / samplerate
        result["num_audio_chunks"] = audio_list_consumed
        ttfa = (first_audio_ts - t_start) if first_audio_ts else None
        result["time_to_first_audio_s"] = ttfa
        ttfa_str = f"{ttfa:.3f}s" if ttfa is not None else "N/A"
        print(
            f"[Request {request_id}] Audio streamed to {wav_file} "
            f"(duration={result['audio_duration_s']:.2f}s, "
            f"TTFA={ttfa_str}, "
            f"e2e={result['e2e_latency_s']:.3f}s)"
        )
    elif audio_chunks or audio_last_tensor is not None:
        if audio_chunks:
            if len(audio_chunks) > 1:
                audio_tensor = torch.cat(audio_chunks, dim=-1)
            else:
                audio_tensor = audio_chunks[0]
        else:
            audio_tensor = audio_last_tensor
        audio_numpy = audio_tensor.float().detach().cpu().numpy()
        if audio_numpy.ndim > 1:
            audio_numpy = audio_numpy.flatten()
        sf.write(wav_file, audio_numpy, samplerate=samplerate, format="WAV")
        result["saved_files"].append(wav_file)
        result["audio_duration_s"] = len(audio_numpy) / samplerate
        result["num_audio_chunks"] = len(audio_chunks)
        ttfa = (first_audio_ts - t_start) if first_audio_ts else None
        result["time_to_first_audio_s"] = ttfa
        ttfa_str = f"{ttfa:.3f}s" if ttfa is not None else "N/A"
        print(
            f"[Request {request_id}] Audio saved to {wav_file} "
            f"({len(audio_chunks)} chunks, "
            f"duration={result['audio_duration_s']:.2f}s, "
            f"TTFA={ttfa_str}, "
            f"e2e={result['e2e_latency_s']:.3f}s)"
        )

    return result


async def run_all(args):
    """Main async entry: build prompts, create AsyncOmni, run requests."""
    # Build query
    query_func = query_map[args.query_type]
    if args.query_type == "use_video":
        query_result = query_func(
            video_path=getattr(args, "video_path", None),
            num_frames=getattr(args, "num_frames", 16),
        )
    elif args.query_type == "use_image":
        query_result = query_func(image_path=getattr(args, "image_path", None))
    elif args.query_type == "use_audio":
        query_result = query_func(
            audio_path=getattr(args, "audio_path", None),
            sampling_rate=getattr(args, "sampling_rate", 16000),
        )
    else:
        query_result = query_func()

    # Build prompt list
    if args.txt_prompts is not None:
        assert args.query_type == "text", "txt-prompts is only supported for text query type"
        with open(args.txt_prompts, encoding="utf-8") as f:
            lines = [ln.strip() for ln in f if ln.strip()]
        prompts = [get_text_query(ln).inputs for ln in lines]
        print(f"[Info] Loaded {len(prompts)} prompts from {args.txt_prompts}")
    else:
        prompts = [clone_prompt_for_request(query_result.inputs) for _ in range(args.num_prompts)]

    # Inject output modalities if specified
    output_modalities = None
    if args.modalities is not None:
        output_modalities = args.modalities.split(",")
        for prompt in prompts:
            prompt["modalities"] = output_modalities

    # Create AsyncOmni
    print(f"[Info] Creating AsyncOmni with deploy_config={args.deploy_config}")
    async_omni = None
    try:
        # ``from_cli_args`` forwards only explicitly-passed CLI args so
        # argparse defaults do not silently override deploy YAML values.
        async_omni = AsyncOmni.from_cli_args(args, model=args.model)

        # Use default sampling params from the resolved pipeline and deploy
        # config.
        #
        # NOTE: Since we do not set the sampling params directly, .generate in
        # will automatically set the output kind to delta, since this is what
        # makes sense for most multimodal use-cases.
        sampling_params_list = None

        output_dir = args.output_dir
        os.makedirs(output_dir, exist_ok=True)

        # Run requests with concurrency control
        semaphore = asyncio.Semaphore(args.max_in_flight)
        request_timeout = getattr(args, "request_timeout_s", None)
        stream_audio = getattr(args, "stream_audio_to_disk", False)

        async def _run_one(idx: int, prompt: dict):
            async with semaphore:
                request_id = f"req_{idx}_{uuid.uuid4().hex[:8]}"
                coro = run_single_request(
                    async_omni=async_omni,
                    prompt=prompt,
                    request_id=request_id,
                    sampling_params_list=sampling_params_list,
                    output_dir=output_dir,
                    output_modalities=output_modalities,
                    stream_audio_to_disk=stream_audio,
                )
                if request_timeout and request_timeout > 0:
                    return await asyncio.wait_for(coro, timeout=request_timeout)
                return await coro

        wall_start = time.perf_counter()
        tasks = [_run_one(i, p) for i, p in enumerate(prompts)]
        all_results = await asyncio.gather(*tasks, return_exceptions=True)
        wall_end = time.perf_counter()

        # Print summary
        print("\n" + "=" * 60)
        print("Summary")
        print("=" * 60)
        success_count = 0
        total_audio_dur = 0.0
        for r in all_results:
            if isinstance(r, Exception):
                print(f"  [ERROR] {type(r).__name__}: {r}")
            else:
                success_count += 1
                total_audio_dur += r.get("audio_duration_s", 0.0)
                print(f"  [{r['request_id']}] e2e={r['e2e_latency_s']:.3f}s  files={r['saved_files']}")
        wall_time = wall_end - wall_start
        print(f"\nTotal: {success_count}/{len(prompts)} succeeded")
        print(f"Wall time: {wall_time:.3f}s")
        if total_audio_dur > 0:
            print(f"Total audio duration: {total_audio_dur:.2f}s")
            print(f"Real-time factor: {total_audio_dur / wall_time:.2f}x")
        print("=" * 60)
    finally:
        if async_omni is not None:
            async_omni.shutdown()


# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------


def parse_args():
    parser = TrackingArgumentParser(
        description=(
            "Offline inference with async_chunk enabled via AsyncOmni. "
            "Downstream stages start before upstream stages finish, "
            "achieving true stage-level concurrency."
        )
    )
    parser.add_argument(
        "--model",
        type=str,
        default="Qwen/Qwen3-Omni-30B-A3B-Instruct",
        help="Model name or path.",
    )
    parser.add_argument(
        "--query-type",
        "-q",
        type=str,
        default="use_audio",
        choices=query_map.keys(),
        help="Query type.",
    )
    parser.add_argument(
        "--deploy-config",
        type=str,
        default=_default_deploy_config_path(),
        help=(
            "Path to a deploy config YAML. "
            "If not set, uses the model's default config "
            "(make sure it has async_chunk: true)."
        ),
    )
    parser.add_argument(
        "--log-stats",
        action="store_true",
        default=False,
        help="Enable writing detailed statistics.",
    )
    parser.add_argument(
        "--stage-init-timeout",
        type=int,
        default=300,
        help="Timeout for initializing a single stage (seconds).",
    )
    parser.add_argument(
        "--output-dir",
        type=str,
        default="output_audio_async_chunk",
        help="Directory to save output files.",
    )
    parser.add_argument(
        "--num-prompts",
        type=int,
        default=1,
        help="Number of prompts to generate (duplicated from query).",
    )
    parser.add_argument(
        "--txt-prompts",
        type=str,
        default=None,
        help="Path to a .txt file with one prompt per line.",
    )
    parser.add_argument(
        "--max-in-flight",
        type=int,
        default=1,
        help="Maximum concurrent requests (default: 1).",
    )
    parser.add_argument(
        "--request-timeout-s",
        type=float,
        default=None,
        help=(
            "Per-request timeout in seconds. When set, a request that "
            "exceeds this duration is cancelled and reported as an error. "
            "Default: None (no timeout)."
        ),
    )
    parser.add_argument(
        "--batch-timeout-s",
        type=float,
        default=None,
        help=(
            "Global timeout for the entire batch in seconds. When set, "
            "the whole run_all() is cancelled if it exceeds this duration. "
            "Default: None (no global timeout)."
        ),
    )
    parser.add_argument(
        "--stream-audio-to-disk",
        action="store_true",
        default=False,
        help=(
            "Write audio chunks to WAV incrementally instead of "
            "accumulating in memory. Useful for very long audio or "
            "high --max-in-flight to reduce memory footprint."
        ),
    )
    parser.add_argument(
        "--modalities",
        type=str,
        default=None,
        help="Comma-separated output modalities filter (e.g. 'text', 'audio', 'text,audio').",
    )
    parser.add_argument(
        "--audio-path",
        "-a",
        type=str,
        default=None,
        help="Path to local audio file.",
    )
    parser.add_argument(
        "--image-path",
        "-i",
        type=str,
        default=None,
        help="Path to local image file.",
    )
    parser.add_argument(
        "--video-path",
        "-v",
        type=str,
        default=None,
        help="Path to local video file.",
    )
    parser.add_argument(
        "--num-frames",
        type=int,
        default=16,
        help="Number of frames to extract from video.",
    )
    parser.add_argument(
        "--sampling-rate",
        type=int,
        default=16000,
        help="Sampling rate for audio loading.",
    )
    return parser.parse_args()


if __name__ == "__main__":
    args = parse_args()

    async def _main():
        batch_timeout = getattr(args, "batch_timeout_s", None)
        if batch_timeout and batch_timeout > 0:
            await asyncio.wait_for(run_all(args), timeout=batch_timeout)
        else:
            await run_all(args)

    try:
        asyncio.run(_main())
    except asyncio.TimeoutError:
        print(
            f"\n[TIMEOUT] Batch exceeded --batch-timeout-s="
            f"{args.batch_timeout_s}s. AsyncOmni shutdown was handled "
            f"by finally block."
        )
    except KeyboardInterrupt:
        print("\nInterrupted by user. AsyncOmni shutdown was handled by finally block.")
run_multiple_prompts.sh
python end2end.py --output-wav output_audio \
                  --query-type text \
                  --txt-prompts text_prompts_10.txt \
                  --py-generator
run_multiple_prompts_async_chunk.sh
#!/bin/bash
# Run multiple Qwen3-Omni requests with async_chunk enabled.
#
# Uses AsyncOmni with --max-in-flight to control request-level
# concurrency (each request still gets true stage-level concurrency
# via async_chunk).
#
# Usage:
#   bash run_multiple_prompts_async_chunk.sh
#   bash run_multiple_prompts_async_chunk.sh --max-in-flight 4

set -euo pipefail

SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../.." && pwd)"

python "${SCRIPT_DIR}/end2end_async_chunk.py" \
    --query-type text \
    --txt-prompts "${SCRIPT_DIR}/text_prompts_10.txt" \
    --deploy-config "${REPO_ROOT}/vllm_omni/deploy/qwen3_omni_moe.yaml" \
    --output-dir output_audio_async_chunk \
    --max-in-flight 2 \
    "$@"
run_single_prompt.sh
python end2end.py --output-wav output_audio \
                  --query-type use_audio
run_single_prompt_async_chunk.sh
#!/bin/bash
# Run a single Qwen3-Omni request with async_chunk enabled.
#
# This uses AsyncOmni (async orchestrator) so that downstream stages
# (Talker, Code2Wav) start *before* stage-0 (Thinker) finishes,
# achieving true stage-level concurrency via chunk-level streaming.
#
# Prerequisites:
#   - A deploy config YAML (e.g. qwen3_omni_moe.yaml)
#   - Hardware matching the config (e.g. 2x H100 for the default 3-stage config)
#
# Usage:
#   bash run_single_prompt_async_chunk.sh
#   bash run_single_prompt_async_chunk.sh --query-type text --modalities text
#   bash run_single_prompt_async_chunk.sh --deploy-config /path/to/custom.yaml

set -euo pipefail

SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "${SCRIPT_DIR}/../../.." && pwd)"

python "${SCRIPT_DIR}/end2end_async_chunk.py" \
    --query-type use_audio \
    --deploy-config "${REPO_ROOT}/vllm_omni/deploy/qwen3_omni_moe.yaml" \
    --output-dir output_audio_async_chunk \
    "$@"
run_single_prompt_tp.sh
python end2end.py --output-wav output_audio \
                  --query-type use_audio \
                  --stage-init-timeout 300

# stage-init-timeout sets the maximum wait to avoid two vLLM stages initializing at the same time on the same card.
text_prompts_10.txt
What is the capital of France?
How many planets are in our solar system?
What is the largest ocean on Earth?
Who wrote the novel "1984"?
What is the chemical symbol for water?
What year did World War II end?
What is the tallest mountain in the world?
What is the speed of light in vacuum?
Who painted the Mona Lisa?
What is the smallest prime number?