diff --git a/docs/.nav.yml b/docs/.nav.yml index 1e3907a6e81..2e7097efb93 100644 --- a/docs/.nav.yml +++ b/docs/.nav.yml @@ -72,6 +72,7 @@ nav: - CPU Offloading: user_guide/diffusion/cpu_offload_diffusion.md - LoRA: user_guide/diffusion/lora.md - Custom Pipeline: features/custom_pipeline.md + - Request-Level Batching: user_guide/diffusion/request_batching.md - Step Execution: user_guide/diffusion/step_execution.md - Quantization: - Overview: user_guide/quantization/overview.md @@ -118,6 +119,7 @@ nav: - design/feature/async_chunk.md - design/feature/vae_parallel.md - design/feature/diffusion_step_execution.md + - design/feature/diffusion_request_level_batching.md - design/feature/diffusion_continuous_batching.md - Module Design: - design/module/ar_module.md diff --git a/docs/configuration/stage_configs.md b/docs/configuration/stage_configs.md index 3a17f1c4013..0b0ee9237c9 100644 --- a/docs/configuration/stage_configs.md +++ b/docs/configuration/stage_configs.md @@ -379,6 +379,20 @@ The maximum number of sequences for concurrent processing in this stage. For LLM Default: `1` +#### `engine_args.request_batch_max_wait_ms` + +The maximum time, in milliseconds, that a diffusion request-mode stage may wait +before the first `schedule()` of a new scheduler wave so compatible requests can +accumulate for request-level batching. This only applies to diffusion pipelines +that support request-level batching with `step_execution` disabled. + +Use this together with `max_num_seqs > 1` for bursty serving traffic. `0` +disables admission waiting and preserves the lowest first-request latency. +For diffusion request-level batching tuning, see +[Request-Level Batching](../user_guide/diffusion/request_batching.md). + +Default: `0.0` + ### `engine_args` Engine arguments for configuring the LLM engine, diffusion engine, or other engine types used by this stage. diff --git a/docs/contributing/model/adding_diffusion_model.md b/docs/contributing/model/adding_diffusion_model.md index 35ff6dae202..7ff13fd1b6e 100644 --- a/docs/contributing/model/adding_diffusion_model.md +++ b/docs/contributing/model/adding_diffusion_model.md @@ -349,12 +349,14 @@ class YourModelPipeline(nn.Module): - def __call__( + def forward( self, -+ req: OmniDiffusionRequest, # ← Add request parameter here ++ req: DiffusionRequestBatch, # ← Add request-batch parameter here - ): -+ ) -> DiffusionOutput: # ← Add return type ++ ) -> list[DiffusionOutput]: # ← Add return type ``` -[`OmniDiffusionRequest`](https://docs.vllm.ai/projects/vllm-omni/en/latest/api/vllm_omni/diffusion/request/#vllm_omni.diffusion.request.OmniDiffusionRequest) is a dataclass that contains the **prompts** and **sampling parameters** [`OmniDiffusionSamplingParams`](https://docs.vllm.ai/projects/vllm-omni/en/latest/api/vllm_omni/inputs/data/#vllm_omni.inputs.data.OmniDiffusionSamplingParams) for the diffusion pipeline execution. It also contains a request_id for other components to trace this request and its outputs. +[`OmniDiffusionRequest`](https://docs.vllm.ai/projects/vllm-omni/en/latest/api/vllm_omni/diffusion/request/#vllm_omni.diffusion.request.OmniDiffusionRequest) is a dataclass that contains one **prompt** and the **sampling parameters** [`OmniDiffusionSamplingParams`](https://docs.vllm.ai/projects/vllm-omni/en/latest/api/vllm_omni/inputs/data/#vllm_omni.inputs.data.OmniDiffusionSamplingParams) for one logical diffusion request. It also contains a request_id for other components to trace this request and its outputs. Before pipeline execution, the runner wraps one or more independent requests into `DiffusionRequestBatch`. + +[`DiffusionRequestBatch`](https://docs.vllm.ai/projects/vllm-omni/en/latest/api/vllm_omni/diffusion/worker/request_batch/#vllm_omni.diffusion.worker.request_batch.DiffusionRequestBatch) exposes compatibility properties such as `prompts`, `sampling_params`, and `request_id`. Pipelines that can execute the whole request batch in one forward pass should set `supports_request_batch = True`; other pipelines still receive a single-request batch and return a one-element output list. See some parameters in `OmniDiffusionSamplingParams` as follows: @@ -367,19 +369,18 @@ See some parameters in `OmniDiffusionSamplingParams` as follows: **Extract parameters from request:** ```python -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.data import DiffusionOutput +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch def forward( self, - req: OmniDiffusionRequest, -) -> DiffusionOutput: - # Extract prompts from request - if req.prompts is not None: - prompt = [ - p if isinstance(p, str) else (p.get("prompt") or "") - for p in req.prompts - ] + req: DiffusionRequestBatch, +) -> list[DiffusionOutput]: + # Extract prompts from the request batch + prompts = [ + p if isinstance(p, str) else (p.get("prompt") or "") + for p in req.prompts + ] # Extract sampling parameters sampling_params = req.sampling_params @@ -388,14 +389,16 @@ def forward( height = sampling_params.height or (self.default_sample_size * self.vae_scale_factor) width = sampling_params.width or (self.default_sample_size * self.vae_scale_factor) - # For image editing pipelines, extract images from multi_modal_data - if hasattr(req, 'multi_modal_data') and req.multi_modal_data: - input_images = req.multi_modal_data.get('image', []) + # For image editing pipelines, extract media from each prompt dict + input_images = [] + for p in req.prompts: + multi_modal_data = p.get("multi_modal_data", {}) if isinstance(p, dict) else {} + input_images.append(multi_modal_data.get("image")) # ... rest of generation logic ``` -For an image editing model, an example `OmniDiffusionRequest` is like: +For an image editing model, the request `prompt` can be a dict like: ```python { "prompt": "turn this cat to a dog", @@ -472,12 +475,12 @@ def get_your_model_pre_process_func( def pre_process_func( request: OmniDiffusionRequest, ): - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - # image pre-processing - # after pre-processing, update the request attributes - ... + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + # image pre-processing + # after pre-processing, update the request attributes + ... return request return pre_process_func @@ -923,11 +926,11 @@ When implementing a new pipeline, avoid putting all logic inside a single functi For example: ``` -def forward(self, req: OmniDiffusionRequest): +def forward(self, req: DiffusionRequestBatch) -> list[DiffusionOutput]: prompt_embeds = self.encode_prompt(req) latents = self.diffuse(prompt_embeds, req) images = self.vae.decode(latents) - return DiffusionOutput(output=images) + return [DiffusionOutput(output=images)] ``` This allows the timing utility to measure each stage (e.g., encode_prompt, diffuse, vae.decode) separately and helps identify performance bottlenecks more easily. diff --git a/docs/design/feature/diffusion_continuous_batching.md b/docs/design/feature/diffusion_continuous_batching.md index c443dea78fa..f90969369ed 100644 --- a/docs/design/feature/diffusion_continuous_batching.md +++ b/docs/design/feature/diffusion_continuous_batching.md @@ -38,9 +38,13 @@ With continuous batching enabled: The current implementation is conservative: - only compatible requests are batched together -- request-mode diffusion still runs with `max_num_seqs=1` - per-request progress and completion remain independent +Here, "continuous batching" means the step-wise path enabled by +`step_execution=True`. Request-mode `DiffusionRequestBatch` is static +request-level batching for one full pipeline `forward()` call; it does not +admit or remove requests between denoise steps. + ## Enablement Use `--step-execution` as the feature gate, then increase `--max-num-seqs` @@ -72,7 +76,7 @@ which is built from shape-sensitive and CFG-sensitive sampling fields. This is the core correctness rule for batching: requests are only co-batched when they share the same denoise tensor contract. -There are two important details: +There are three important details: - `num_inference_steps` is not part of the key, so requests with different total step counts can still share a batch @@ -88,8 +92,10 @@ key also covers LoRA identity (`lora_int_id`, `lora_scale`), so requests targeting different adapters or scales run in separate batches and the worker can activate exactly one adapter per step. -The current batching unit is one `OmniDiffusionRequest`. Requests with -multiple prompts do not participate in batching today. +The scheduler batching unit is one logical `OmniDiffusionRequest`. In the +step-wise path, runtime tensor batching is represented as `StepInputBatch`. For +request-mode prompt semantics, see +[Request-Level Batching](../../user_guide/diffusion/request_batching.md). ## Runner @@ -126,17 +132,15 @@ request-local scheduler state and outputs. the background loop and async add-request path needed for multiple requests to accumulate in the scheduler. -This is supporting infrastructure, not the main design point. The batching -behavior is defined by scheduler-side compatibility gating and runner-side -batch packing. +When `step_execution=True`, the engine routes work through the step-wise +executor path. The continuous batching behavior is defined by scheduler-side +compatibility gating and runner-side `StepInputBatch` packing. ## Current Limitations - Experimental feature; use `max_num_seqs=1` for the older conservative path. - Only native pipelines that already support `step_execution=True`. -- Request-mode diffusion still clamps `max_num_seqs` back to `1`. - Only homogeneous batches keyed by `SamplingParamsKey` are supported. -- Multi-prompt requests are not batched. - `cache_backend`, KV transfer, and other request-mode extras are not wired into the batched step-wise path yet. - Future work can relax the current same-shape restriction with richer diff --git a/docs/design/feature/diffusion_request_level_batching.md b/docs/design/feature/diffusion_request_level_batching.md new file mode 100644 index 00000000000..57287ada49d --- /dev/null +++ b/docs/design/feature/diffusion_request_level_batching.md @@ -0,0 +1,175 @@ +# Request-Level Batching for Diffusion + +This document describes the request-mode batching path for diffusion pipelines. +For end-user enablement and tuning, see +[Request-Level Batching](../../user_guide/diffusion/request_batching.md). + +This is separate from +[Continuous Batching for Step-Wise Diffusion](diffusion_continuous_batching.md). +Request-level batching runs one full pipeline `forward()` over a static batch of +compatible requests. Step-wise continuous batching admits work between denoise +steps when `step_execution=True`. + +## Why It Helps + +The request-level design avoids coupling several logical requests to one request +object. This keeps request identity, abort/error handling, and per-request +metadata unambiguous while still allowing one fused pipeline forward pass for +bursty or concurrent traffic. + +## Overview + +With request-level batching enabled: + +- each `OmniDiffusionRequest` contains one `prompt` and one `request_id` +- the scheduler groups compatible waiting requests into one scheduler wave +- `DiffusionRequestBatch` wraps the scheduled requests for pipeline `forward()` +- batch-capable pipelines return `list[DiffusionOutput]`, one output per + request +- `BatchRunnerOutput` maps each result back to its original `request_id` + +Pipelines opt in with `supports_request_batch = True` and a `forward()` method +that accepts `DiffusionRequestBatch` and returns `list[DiffusionOutput]`. +Pipelines that do not opt in keep the existing per-request execution path. + +## Enablement + +Request-level batching is the request-mode path, so `step_execution` must remain +disabled. Increase `max_num_seqs` above `1` to let the scheduler keep multiple +compatible requests active: + +```bash +vllm serve Qwen/Qwen-Image --omni \ + --port 8091 \ + --max-num-seqs 4 +``` + +For bursty online ingress, `request_batch_max_wait_ms` can add a bounded +admission wait before the first `schedule()` of a scheduler wave: + +```bash +vllm serve Qwen/Qwen-Image --omni \ + --port 8091 \ + --max-num-seqs 4 \ + --request-batch-max-wait-ms 20 +``` + +`request_batch_max_wait_ms=0` disables this wait and is the default. + +## Request Contract + +`OmniDiffusionRequest` represents one logical request. It owns one prompt, +sampling parameters, request id, and request-local metadata. Runtime batches are +formed by the scheduler and represented separately from the request payload. + +Runtime batching is represented by: + +- [`DiffusionSchedulerOutput`](gh-file:vllm_omni/diffusion/sched/interface.py) + for scheduled request ids and request payloads +- [`DiffusionRequestBatch`](gh-file:vllm_omni/diffusion/worker/request_batch.py) + for the pipeline-facing request batch +- [`BatchRunnerOutput`](gh-file:vllm_omni/diffusion/worker/utils.py) for + per-request results + +`DiffusionRequestBatch` intentionally exposes compatibility properties such as +`prompts`, `sampling_params`, `request_id`, and `kv_sender_info` so migrated +pipelines can stay close to upstream code while using a batch-aware contract. + +## Scheduler + +The scheduler derives its capacity from `max_num_seqs` through +`max_num_running_reqs`. It exposes waiting/running queue counters so the engine +can decide whether admission wait is useful before scheduling a new wave. + +Batch compatibility is controlled by +[`SamplingParamsKey`](gh-file:vllm_omni/diffusion/sched/interface.py). The key +contains shape-sensitive and guidance-sensitive fields, including output count +and LoRA identity. Requests with incompatible shapes, CFG settings, output +counts, LoRA adapters, or LoRA scales are kept in separate batches. + +Admission is conservative: + +- the scheduler only batches compatible requests +- FIFO ordering is preserved +- an incompatible request at the head of the waiting queue blocks later + compatible requests + +## Engine + +[`DiffusionEngine`](gh-file:vllm_omni/diffusion/diffusion_engine.py) resolves +request-batch capability during initialization from the configured pipeline +class, including custom pipeline classes. + +The capability check uses the pipeline class attribute +`supports_request_batch = True`. Pipelines that set this attribute must implement +a request-batch-compatible `forward()` contract and return one +`DiffusionOutput` per request; the runner validates that return shape at runtime. + +When the selected pipeline is batch-capable and `step_execution=False`, request +mode routes scheduler waves through the batch executor path. Otherwise it keeps +the per-request executor path. + +The optional admission wait runs only when: + +- request batching is supported +- `step_execution=False` +- `request_batch_max_wait_ms > 0` +- no requests are currently running + +The wait exits early when the waiting queue reaches capacity, when the queue is +stable for a short window, when the deadline expires, or when the engine stops. + +## Executor And Runner + +The executor exposes two request-mode entries: + +- `execute_request`: one worker call per scheduled request +- `execute_batch`: one worker call for the whole `DiffusionSchedulerOutput` + +On the batch path, the worker builds a `DiffusionRequestBatch` and runs the +pipeline once. Request-local setup remains per request: + +- KV transfer metadata +- random generator and seed handling +- request output/error/abort mapping + +Shared batch setup happens once per batch when possible: + +- cache refresh +- LoRA activation for the homogeneous adapter key +- pipeline `forward(req_batch)` + +Large tensor IPC still uses the shared-memory packing path. The packer traverses +both normal `RunnerOutput.result` wrappers and nested batch results so batched +outputs do not fall back to pickle IPC for tensor payloads. + +## Current Limitations + +- Only pipelines that declare the request-batch contract use fused batch + execution. +- Batches are homogeneous under `SamplingParamsKey`; heterogeneous resolution or + incompatible guidance settings do not co-batch yet. +- FIFO scheduling can reduce batching opportunities when an incompatible + request is at the front of the queue. +- `request_batch_max_wait_ms` improves burst coalescing but can add latency to + the first request in a scheduler wave. Keep it small for latency-sensitive + serving. +- Step-wise continuous batching is documented separately and only applies when + `step_execution=True`. + +## Related Files + +- Request object and request batch: + [`vllm_omni/diffusion/request.py`](gh-file:vllm_omni/diffusion/request.py) +- Scheduler interface: + [`vllm_omni/diffusion/sched/interface.py`](gh-file:vllm_omni/diffusion/sched/interface.py) +- Scheduler base: + [`vllm_omni/diffusion/sched/base_scheduler.py`](gh-file:vllm_omni/diffusion/sched/base_scheduler.py) +- Engine: + [`vllm_omni/diffusion/diffusion_engine.py`](gh-file:vllm_omni/diffusion/diffusion_engine.py) +- Worker runner: + [`vllm_omni/diffusion/worker/diffusion_model_runner.py`](gh-file:vllm_omni/diffusion/worker/diffusion_model_runner.py) +- Executor interface: + [`vllm_omni/diffusion/executor/abstract.py`](gh-file:vllm_omni/diffusion/executor/abstract.py) +- Tests: + [`tests/diffusion/test_diffusion_engine.py`](gh-file:tests/diffusion/test_diffusion_engine.py) diff --git a/docs/design/index.md b/docs/design/index.md index 8789f2e1b22..cfd177d65da 100644 --- a/docs/design/index.md +++ b/docs/design/index.md @@ -11,6 +11,7 @@ This section contains design documents and architecture specifications for vLLM- - [Disaggregated Inference](feature/disaggregated_inference.md) - [Ray-based Execution](feature/ray_based_execution.md) - [Adding Step Execution Support for Diffusion Pipelines](feature/diffusion_step_execution.md) +- [Request-Level Batching for Diffusion](feature/diffusion_request_level_batching.md) - [Continuous Batching for Step-Wise Diffusion](feature/diffusion_continuous_batching.md) ## Infrastructure Design Documents diff --git a/docs/design/module/dit_module.md b/docs/design/module/dit_module.md index 50890c59742..953e11f7fea 100644 --- a/docs/design/module/dit_module.md +++ b/docs/design/module/dit_module.md @@ -199,7 +199,7 @@ class _BaseScheduler(SchedulerInterface): - **Shared cleanup logic**: Duplicate-request checks, finish handling, and state removal are centralized instead of duplicated in each policy. -- **Current constraint boundary**: `_BaseScheduler` derives `max_num_running_reqs` from `max_num_seqs`, but request-mode diffusion is still clamped back to `1` by the engine. The step-wise path can keep this above `1` for compatible-request batching. +- **Current constraint boundary**: `_BaseScheduler` derives `max_num_running_reqs` from `max_num_seqs`. Request-mode diffusion can use that capacity for compatible independent requests when the configured pipeline declares `supports_request_batch = True`; step-wise diffusion uses the same scheduler capacity for compatible step batches. #### 2.4 Current `RequestScheduler` Policy @@ -215,7 +215,7 @@ class RequestScheduler(_BaseScheduler): - **FIFO request scheduling**: Waiting requests are promoted in queue order. -- **Single-request admission**: `RequestScheduler` still admits one active request at a time because request-mode execution completes a whole request per dispatch. +- **Compatible request admission**: `RequestScheduler` admits waiting requests while capacity remains and the request's `SamplingParamsKey` is compatible with the active batch. Request-mode execution keeps each logical request independent, while batch-capable pipelines receive the scheduled requests as a runner-side `DiffusionRequestBatch`. - **Executor result feedback**: `update_from_output()` converts executor output into `FINISHED_COMPLETED` or `FINISHED_ERROR` and returns finished request ids. @@ -235,7 +235,12 @@ while True: - **No scheduler-owned IPC**: Scheduler no longer talks to workers directly. -- **Split concurrency model**: Request-mode diffusion remains single-active-request, while the step-wise path can keep multiple compatible requests running and advance them independently between denoise steps. +- **Split concurrency model**: Request-mode diffusion can schedule a static + batch of compatible independent requests for one full pipeline forward when + the pipeline supports request-level batching; the step-wise path can also + admit or remove compatible requests between denoise steps. See + [Request-Level Batching](../feature/diffusion_request_level_batching.md) and + [Continuous Batching for Step-Wise Diffusion](../feature/diffusion_continuous_batching.md). --- diff --git a/docs/getting_started/quickstart.md b/docs/getting_started/quickstart.md index 652c512a245..1e7a87b9004 100644 --- a/docs/getting_started/quickstart.md +++ b/docs/getting_started/quickstart.md @@ -31,7 +31,6 @@ uv pip install -e . For additional installation methods — please see the [installation guide](installation/README.md). - !!! note It is important to install the same major & minor version of vLLM and vLLM Omni, otherwise things may not work as expected. If the versions are misaligned, you will see a warning when you import vLLM Omni. @@ -52,14 +51,14 @@ if __name__ == "__main__": images[0].save("coffee.png") ``` -You can pass a list of prompts and wait for them to process altogether, shown below. +You can pass a list of prompts and wait for the independent requests to finish, +as shown below. !!! info - However, it is not currently recommended to do so - because not all models support batch inference, - and batch requesting mostly does not provide significant performance improvement (despite the impression that it does). - This feature is primarily for the sake of interface compatibility with vLLM and to allow for future improvements. + For diffusion pipelines, each prompt becomes a separate logical request. + The runtime may automatically batch compatible in-flight requests through + the scheduler and runner. ```python from vllm_omni.entrypoints.omni import Omni @@ -88,9 +87,9 @@ if __name__ == "__main__": !!! info - For diffusion pipelines, the stage config field `stage_args.[].engine_args.max_num_seqs` is 1 by default, and the input - list is sliced into single-item requests before feeding into the diffusion pipeline. For models that do internally support - batched inputs, you can [modify this configuration](../configuration/stage_configs.md) to let the model accept a longer batch of prompts. + For diffusion request-level batching controls such as `max_num_seqs` and + `request_batch_max_wait_ms`, see + [Request-Level Batching](../user_guide/diffusion/request_batching.md). For more usages, please refer to [offline inference](../user_guide/examples/offline_inference/qwen2_5_omni.md) diff --git a/docs/user_guide/diffusion/request_batching.md b/docs/user_guide/diffusion/request_batching.md new file mode 100644 index 00000000000..5cd24802bcc --- /dev/null +++ b/docs/user_guide/diffusion/request_batching.md @@ -0,0 +1,110 @@ +# Request-Level Batching + +Request-level batching lets diffusion serving combine multiple compatible +logical requests into one pipeline `forward()` call. Each prompt remains a +separate request with its own `request_id`, sampling parameters, seed, output, +error, and abort state. The scheduler decides which requests can run together. + +!!! warning "Prompt List Semantics" + + Diffusion request-level batching does not support a top-level packed + list-prompt request. Submit multiple prompts as independent requests and let + the scheduler batch compatible in-flight requests. Multimodal payloads stay + inside a single prompt dict, for example + `{"prompt": "...", "multi_modal_data": {"image": image}}`. + +## Enablement + +Increase `max_num_seqs` above `1` to allow the request scheduler to keep more +than one compatible request active: + +```bash +vllm serve Qwen/Qwen-Image --omni \ + --port 8091 \ + --max-num-seqs 4 +``` + +For bursty online traffic, you can also set a small admission wait window. This +lets the engine wait briefly before the first `schedule()` of a new wave so +nearby compatible requests can arrive and share the same fused forward pass: + +```bash +vllm serve Qwen/Qwen-Image --omni \ + --port 8091 \ + --max-num-seqs 4 \ + --request-batch-max-wait-ms 20 +``` + +`--request-batch-max-wait-ms 0` is the default and disables admission waiting, +so there is no added wait latency. + +For deploy YAMLs, configure the diffusion stage engine args: + +```yaml +stage_args: + - stage_id: 0 + stage_type: diffusion + engine_args: + max_num_seqs: 4 + request_batch_max_wait_ms: 20 +``` + +## Compatibility + +Only pipelines that declare request-batch support use the fused request-batch +path. The engine validates that the pipeline `forward()` uses the request-batch +contract and returns `list[DiffusionOutput]`. Pipelines that do not support this +contract do not use fused `pipeline.forward(batch)`; scheduled requests are +executed through per-request worker calls. + +The scheduler batches only compatible requests. Compatibility is based on +shape-sensitive and guidance-sensitive sampling fields, including resolution, +frame count, CFG settings, output count, and LoRA identity. Requests with +different LoRA adapters or scales are kept in separate batches so the worker +activates one adapter per batch. + +Request-level batching applies when `step_execution=False`. For the separate +step-wise runtime, see [Step Execution](step_execution.md). + +## Tuning + +- `max_num_seqs` caps the number of active compatible requests in one scheduler + wave. +- `request_batch_max_wait_ms` is an upper bound on extra admission wait before a + new wave starts. Keep it small for latency-sensitive serving; values such as + `10` to `50` ms are a practical starting range for bursty HTTP ingress. +- `0` disables admission waiting and preserves the lowest first-request latency. +- FIFO ordering is conservative: an incompatible request at the front of the + waiting queue can block later compatible requests from joining the current + batch. + +## Python API + +When constructing `Omni`, pass the same engine arguments: + +```python +from vllm_omni.entrypoints.omni import Omni + +omni = Omni( + model="Qwen/Qwen-Image", + max_num_seqs=4, + request_batch_max_wait_ms=20.0, +) + +outputs = omni.generate( + [ + "a cup of coffee on a table", + "a toy dinosaur on a sandy beach", + "a fox waking up in bed and yawning", + ] +) +``` + +`Omni.generate([...])` submits each list item as its own logical diffusion +request. The runtime may batch those requests internally when their sampling +parameters are compatible. + +## For Contributors + +For implementation details and model-author guidance, see +[Request-Level Batching for Diffusion](../../design/feature/diffusion_request_level_batching.md). diff --git a/docs/user_guide/diffusion_features.md b/docs/user_guide/diffusion_features.md index 3e1bff799b1..3ac70a7540a 100644 --- a/docs/user_guide/diffusion_features.md +++ b/docs/user_guide/diffusion_features.md @@ -75,13 +75,17 @@ Extension methods add specialized capabilities to diffusion models beyond standa ### Execution Modes -Execution modes control how the diffusion pipeline processes denoise steps. +Execution modes control how the diffusion pipeline processes requests and +denoise steps. | Method | Description | Best For | |--------|-------------|----------| +| **[Request-Level Batching](diffusion/request_batching.md)** | Scheduler batches compatible independent diffusion requests into one pipeline forward pass | Bursty online serving and multi-request throughput | | **[Step Execution](diffusion/step_execution.md)** | Per-step denoise execution with mid-request abort support | Request cancellation between denoise steps, fine-grained execution control | -**Note:** Step execution is currently supported by QwenImagePipeline only. See [Supported Models](#supported-models) for details. +**Note:** Request-level batching is available for pipelines that declare the +request-batch forward contract. Step execution is currently supported by +QwenImagePipeline only. See [Supported Models](#supported-models) for details. ### Quantization Methods diff --git a/docs/user_guide/examples/offline_inference/text_to_image.md b/docs/user_guide/examples/offline_inference/text_to_image.md index 3a97ffbf74b..7cd508d297a 100644 --- a/docs/user_guide/examples/offline_inference/text_to_image.md +++ b/docs/user_guide/examples/offline_inference/text_to_image.md @@ -160,9 +160,11 @@ python examples/offline_inference/text_to_image/text_to_image.py \ --output flux2-dev.png ``` -### Batch Requests (Multiple Prompts) +### Multiple Prompts -You can pass multiple prompts in a single `generate` call. +You can pass multiple prompts in a single `generate` call. For diffusion +pipelines, each prompt is submitted as a separate logical request; compatible +requests may be automatically batched by the scheduler and runner. ```python from vllm_omni.entrypoints.omni import Omni @@ -181,17 +183,8 @@ if __name__ == "__main__": !!! info - Not all models support batch inference, and batch requesting mostly does not provide significant - performance improvement. This feature is primarily for interface compatibility with vLLM and to - allow for future improvements. - -!!! info - - For diffusion pipelines, the stage config field `stage_args.[].runtime.max_batch_size` is 1 by - default, and the input list is sliced into single-item requests before feeding into the diffusion - pipeline. For models that do internally support batched inputs, you can - [modify this configuration](https://github.com/vllm-project/vllm-omni/tree/main/configuration/stage_configs.md) to let the model accept a - longer batch of prompts. + For diffusion request-level batching controls such as `max_num_seqs`, see + [Request-Level Batching](../../diffusion/request_batching.md). ### Negative Prompts diff --git a/examples/offline_inference/text_to_image/README.md b/examples/offline_inference/text_to_image/README.md index 02538b40fac..60ca0fb8eed 100644 --- a/examples/offline_inference/text_to_image/README.md +++ b/examples/offline_inference/text_to_image/README.md @@ -179,7 +179,6 @@ python examples/offline_inference/text_to_image/text_to_image.py \ --output flux2-dev.png ``` - ### HiDream-I1-Full Models The `--auxiliary-text-encoder` parameter is required when running HiDream‑I1‑Full: @@ -199,7 +198,9 @@ python examples/offline_inference/text_to_image/text_to_image.py \ ### Batch Requests (Multiple Prompts) -You can pass multiple prompts in a single `generate` call. +You can pass multiple prompts in a single `Omni.generate` call. `Omni` +submits each prompt as an independent request and returns one output per +prompt. ```python from vllm_omni.entrypoints.omni import Omni @@ -224,11 +225,10 @@ if __name__ == "__main__": !!! info - For diffusion pipelines, the stage config field `stage_args.[].runtime.max_batch_size` is 1 by - default, and the input list is sliced into single-item requests before feeding into the diffusion - pipeline. For models that do internally support batched inputs, you can - [modify this configuration](../../../configuration/stage_configs.md) to let the model accept a - longer batch of prompts. + For diffusion pipelines, the input list is sliced into single-item requests + before feeding into the diffusion pipeline. For request-level batching + controls such as `max_num_seqs`, see + [Request-Level Batching](../../../docs/user_guide/diffusion/request_batching.md). ### Negative Prompts diff --git a/pyproject.toml b/pyproject.toml index 962f58a8ffd..ace2f3bf21c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -271,6 +271,7 @@ markers = [ "skipif_npu: Skip if the num of NPU cards is less than the required", "skipif_musa: Skip if the num of MUSA cards is less than the required", # more detailed markers + "sp: Sequence parallelism tests (multi-GPU)", "slow: Slow tests (may skip in quick CI)", "benchmark: Benchmark tests", ] diff --git a/tests/diffusion/batching/test_diffusion_batching.py b/tests/diffusion/batching/test_diffusion_batching.py index 01638a16360..fa96b54a45c 100644 --- a/tests/diffusion/batching/test_diffusion_batching.py +++ b/tests/diffusion/batching/test_diffusion_batching.py @@ -220,76 +220,11 @@ async def run_batch( return elapsed -# ------------------------------------------------------------------ -# Explicit batch — single generate() call with list of prompts -# ------------------------------------------------------------------ - - -async def run_batch_explicit( - omni: AsyncOmni, - prompts: list[dict[str, str]], - label: str = "batch_explicit", -) -> float: - """Send all prompts as a single batch via generate(prompt=[...]). - - Passes a *list* of prompts so that all are processed in **one** - ``DiffusionEngine.step()`` call. The orchestrator detects - ``isinstance(prompt, list)`` and routes to the batch path. - - A single ``OmniRequestOutput`` is yielded containing ALL generated - images combined. - """ - print(f"⚡ Running {label.upper()} mode – {len(prompts)} prompts in ONE engine call ...") - sp = _default_sampling_params() - request_id = f"{label}-{uuid.uuid4().hex[:8]}" - start = time.perf_counter() - - result: OmniRequestOutput | None = None - async for output in omni.generate( - prompt=prompts, - request_id=request_id, - sampling_params_list=[sp], - ): - result = output - - elapsed = time.perf_counter() - start - if result is not None: - images = _extract_images(result) - print(f" Got {len(images)} images total from batch, request_id={result.request_id}") - else: - print(" ⚠️ No output received from batch generate()") - - print(f" ✅ Total {label} mode: {elapsed:.2f}s\n") - return elapsed - - # ------------------------------------------------------------------ # Async validation helpers # ------------------------------------------------------------------ -async def validate_batch_explicit(omni: AsyncOmni, prompts: list[dict[str, str]]) -> None: - """Validate generate(prompt=[...]) returns a single result with all images.""" - print(f"🔍 Validating batch generate() correctness with {len(prompts)} prompts ...") - sp = _default_sampling_params() - request_id = f"validate-batch-{uuid.uuid4().hex[:8]}" - - result: OmniRequestOutput | None = None - async for output in omni.generate( - prompt=prompts, - request_id=request_id, - sampling_params_list=[sp], - ): - result = output - - assert result is not None, "No output received from batch generate()" - images = _extract_images(result) - # Batch mode returns ONE output with ALL images combined - assert len(images) == len(prompts), f"Expected {len(prompts)} images (one per prompt), got {len(images)}" - assert result.request_id == request_id, f"Expected request_id={request_id}, got {result.request_id}" - print(f" ✅ Batch returned {len(images)} images with correct request_id.\n") - - async def validate_concurrent(omni: AsyncOmni, prompts: list[dict[str, str]]) -> None: """Validate that every concurrent request receives a distinct result with its own request_id.""" @@ -330,17 +265,14 @@ async def compare_single_vs_parallel( await warmup(omni, WARMUP_PROMPTS) single_time = await run_single(omni, prompts) parallel_time = await run_batch(omni, prompts, label="parallel") - explicit_time = await run_batch_explicit(omni, prompts, label="batch_explicit") finally: omni.shutdown() speedup_parallel = single_time / parallel_time if parallel_time > 0 else float("inf") - speedup_explicit = single_time / explicit_time if explicit_time > 0 else float("inf") print("=" * 60) print(f"📊 Summary ({len(prompts)} prompts)") print(f" Sequential : {single_time:.2f}s") print(f" Parallel (gather) : {parallel_time:.2f}s ({speedup_parallel:.2f}x)") - print(f" Explicit batch : {explicit_time:.2f}s ({speedup_explicit:.2f}x)") print("=" * 60) @@ -362,12 +294,8 @@ async def main(model: str, num_prompts: int, mode: str, batch_size: int = 1) -> if mode == "validate": await validate_concurrent(omni, prompts) - elif mode == "validate_batch": - await validate_batch_explicit(omni, prompts) elif mode == "batch": await run_batch(omni, prompts, label="measurement") - elif mode == "batch_explicit": - await run_batch_explicit(omni, prompts) elif mode == "single": await run_single(omni, prompts) else: @@ -495,13 +423,10 @@ async def _inner(): @pytest.mark.diffusion @hardware_test(res={"cuda": "L4", "rocm": "MI325", "xpu": "B60"}) @pytest.mark.parametrize("model_name", models) -def test_diffusion_batching_async_explicit_batch(model_name: str): - """Test that AsyncOmni batch mode (generate(prompt=[...])) dispatches - all prompts in a single engine call and returns a single combined result. - - The list-prompt path routes through the orchestrator's - ``add_batch_request_async`` → ``AsyncOmni.generate_batch`` - and yields ONE ``OmniRequestOutput`` with ALL images combined. +def test_diffusion_batching_list_prompt_rejected(model_name: str): + """Test that list-prompt batch requests are rejected at the diffusion + stage boundary. Users should submit multiple independent requests to + leverage scheduler batching instead. """ async def _inner(): @@ -511,25 +436,14 @@ async def _inner(): sp = _default_sampling_params() request_id = f"explicit-batch-{uuid.uuid4().hex[:8]}" - # Batch mode: pass list of prompts → single request_id - result: OmniRequestOutput | None = None - async for output in omni.generate( - prompt=prompts, - request_id=request_id, - sampling_params_list=[sp], - ): - result = output - - assert result is not None, "No output received from batch generate()" - - images = _extract_images(result) - # One image per prompt, all in a single output - assert len(images) == len(prompts), f"Expected {len(prompts)} images in combined output, got {len(images)}" - assert result.request_id == request_id, f"Expected request_id={request_id}, got {result.request_id}" - for i, img in enumerate(images): - assert img.width == 256, f"Image {i} width mismatch" - assert img.height == 256, f"Image {i} height mismatch" - print(f" ✅ Batch returned {len(images)} images, request_id={result.request_id}") + with pytest.raises(ValueError, match="Diffusion stages accept only a single prompt per request"): + async for _output in omni.generate( + prompt=prompts, + request_id=request_id, + sampling_params_list=[sp], + ): + pass + print(" ✅ List-prompt batch correctly rejected") finally: omni.shutdown() @@ -613,12 +527,11 @@ def test_diffusion_batching_distinct_results(model_name: str): parser.add_argument("--batch-size", type=int, default=1, help="Diffusion batch size (1 = no batching)") parser.add_argument( "--mode", - choices=["batch", "batch_explicit", "single", "compare", "validate", "validate_batch"], + choices=["batch", "single", "compare", "validate"], default="compare", help=( - "Run mode: 'batch' (parallel gather), 'batch_explicit' (list-prompt batch API), " - "'single' (sequential), 'compare' (all three), " - "'validate' (concurrent correctness), 'validate_batch' (list-prompt correctness)" + "Run mode: 'batch' (parallel gather), 'single' (sequential), " + "'compare' (single vs parallel), 'validate' (concurrent correctness)" ), ) args = parser.parse_args() diff --git a/tests/diffusion/diffusion_backend/test_diffusers_backend.py b/tests/diffusion/diffusion_backend/test_diffusers_backend.py index b64d6f738c2..b9b4a72b74a 100644 --- a/tests/diffusion/diffusion_backend/test_diffusers_backend.py +++ b/tests/diffusion/diffusion_backend/test_diffusers_backend.py @@ -18,6 +18,7 @@ ) from vllm_omni.diffusion.models.diffusers_adapter import DiffusersAdapterPipeline from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniDiffusionSamplingParams pytestmark = [pytest.mark.diffusion] @@ -48,7 +49,7 @@ def _make_request(**overrides) -> OmniDiffusionRequest: prompt_obj["negative_prompt"] = negative_prompt defaults = { - "prompts": [prompt_obj], + "prompt": prompt_obj, "sampling_params": OmniDiffusionSamplingParams( num_inference_steps=20, guidance_scale=7.5, @@ -66,6 +67,11 @@ def _make_request(**overrides) -> OmniDiffusionRequest: return OmniDiffusionRequest(**defaults) +def _make_batch(**overrides) -> DiffusionRequestBatch: + """Wrap a single request in a DiffusionRequestBatch, matching the forward contract.""" + return DiffusionRequestBatch(requests=[_make_request(**overrides)]) + + @pytest.mark.core_model @pytest.mark.cpu class TestPipelineArgumentsHandling: @@ -84,7 +90,7 @@ def test_adapter_forward_returns_output(self, mocker): "__call__", return_value=MockPipelineOutput(image=stub_image), ) - output = adapter.forward(request) + output = adapter.forward(DiffusionRequestBatch(requests=[request])) assert isinstance(output, DiffusionOutput) assert isinstance(output.output, MockPipelineOutput) @@ -145,7 +151,7 @@ def test_adapter_guard_unknown_output_type(self, mocker): "__call__", return_value=raw_output, ) - output = adapter.forward(_make_request()) + output = adapter.forward(_make_batch()) assert isinstance(output, DiffusionOutput) assert output.output == raw_output @@ -205,7 +211,7 @@ def to(self, device): ), ) - kwargs = adapter._build_call_kwargs(req) + kwargs = adapter._build_call_kwargs(DiffusionRequestBatch(requests=[req])) assert kwargs["prompt"] == "a cat on mars" assert kwargs["negative_prompt"] == "low quality" @@ -403,7 +409,7 @@ def to(self, device): ), ) with pytest.raises(ValueError): - pipeline.forward(problematic_request) + pipeline.forward(DiffusionRequestBatch(requests=[problematic_request])) @pytest.mark.advanced_model diff --git a/tests/diffusion/models/cosmos3/test_cosmos3_pipeline.py b/tests/diffusion/models/cosmos3/test_cosmos3_pipeline.py index 6696bfcc8bf..c4eab98126d 100644 --- a/tests/diffusion/models/cosmos3/test_cosmos3_pipeline.py +++ b/tests/diffusion/models/cosmos3/test_cosmos3_pipeline.py @@ -15,6 +15,8 @@ from PIL import Image from torch import nn +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch + pytestmark = [pytest.mark.core_model, pytest.mark.cpu, pytest.mark.diffusion] @@ -243,6 +245,31 @@ def make_sampling_params(**overrides: Any) -> SimpleNamespace: return SimpleNamespace(**values) +def make_request_batch(prompt: Any, sampling_params: SimpleNamespace) -> DiffusionRequestBatch: + if isinstance(prompt, list): + return DiffusionRequestBatch( + requests=[ + SimpleNamespace( + prompt=item, + request_id=f"cosmos3-test-{idx}", + sampling_params=sampling_params, + kv_sender_info=None, + ) + for idx, item in enumerate(prompt) + ] + ) + return DiffusionRequestBatch( + requests=[ + SimpleNamespace( + prompt=prompt, + request_id="cosmos3-test", + sampling_params=sampling_params, + kv_sender_info=None, + ) + ] + ) + + def _ids(value: int) -> torch.Tensor: return torch.tensor([[value]], dtype=torch.long) @@ -466,34 +493,34 @@ def test_preprocess_i2v_image_and_action_video_inputs() -> None: preprocess = get_cosmos3_pre_process_func(SimpleNamespace()) i2v = SimpleNamespace( - prompts=[{"prompt": "A slow camera push.", "multi_modal_data": {"image": Image.new("RGB", (320, 160))}}], - sampling_params=SimpleNamespace(height=None, width=None, extra_args={}), + prompt={"prompt": "A slow camera push.", "multi_modal_data": {"image": Image.new("RGB", (320, 160))}}, + sampling_params=make_sampling_params(height=None, width=None, extra_args={}), ) result = preprocess(i2v) assert (result.sampling_params.height, result.sampling_params.width) == (672, 1344) - assert tuple(result.prompts[0]["additional_information"]["preprocessed_image"].shape[-2:]) == (672, 1344) + assert tuple(result.prompt["additional_information"]["preprocessed_image"].shape[-2:]) == (672, 1344) frames = [Image.new("RGB", (8, 4), color) for color in ("red", "green", "blue")] action = SimpleNamespace( - prompts=[{"prompt": "Move.", "multi_modal_data": {"video": frames}}], - sampling_params=SimpleNamespace(height=16, width=32, extra_args={"action_mode": "forward_dynamics"}), + prompt={"prompt": "Move.", "multi_modal_data": {"video": frames}}, + sampling_params=make_sampling_params(height=16, width=32, extra_args={"action_mode": "forward_dynamics"}), ) - additional = preprocess(action).prompts[0]["additional_information"] + additional = preprocess(action).prompt["additional_information"] assert tuple(additional["preprocessed_image"].shape) == (1, 3, 16, 32) assert tuple(additional["preprocessed_video"].shape) == (1, 3, 3, 16, 32) frames = [Image.new("RGB", (8, 4), color) for color in ("red", "green", "blue", "yellow", "purple", "black")] v2v = SimpleNamespace( - prompts=[{"prompt": "Continue.", "multi_modal_data": {"video": frames}}], - sampling_params=SimpleNamespace( + prompt={"prompt": "Continue.", "multi_modal_data": {"video": frames}}, + sampling_params=make_sampling_params( height=16, width=32, extra_args={"condition_frame_indexes_vision": [0, 1], "condition_video_keep": "last"}, ), ) - additional = preprocess(v2v).prompts[0]["additional_information"] + additional = preprocess(v2v).prompt["additional_information"] assert tuple(additional["preprocessed_video"].shape) == (1, 3, 5, 16, 32) assert additional["condition_frame_indexes_vision"] == [0, 1] @@ -565,15 +592,16 @@ class FramesWithFps(list): fps = 12.5 frames = FramesWithFps(Image.new("RGB", (8, 4), color) for color in ("red", "green", "blue", "yellow", "black")) + prompt = {"prompt": "transfer", "multi_modal_data": {"video": frames}} request = SimpleNamespace( - prompts=[{"prompt": "transfer", "multi_modal_data": {"video": frames}}], + prompt=prompt, sampling_params=SimpleNamespace( height=16, width=32, extra_args={"edge": True, "max_frames": 4, "resolution": "256"}, ), ) - additional = preprocess(request).prompts[0]["additional_information"] + additional = preprocess(request).prompt["additional_information"] assert (request.sampling_params.height, request.sampling_params.width) == (192, 320) assert tuple(additional["preprocessed_transfer_video"].shape) == (1, 3, 4, 192, 320) assert additional["transfer_input_fps"] == 12.5 @@ -1332,7 +1360,7 @@ def test_forward_defaults_and_mode_selection( pipeline = make_cosmos3_pipeline() captured = self._install_forward_stubs(pipeline) - output = pipeline.forward(SimpleNamespace(prompts=[prompt], sampling_params=sampling_params)) + output = pipeline.forward(make_request_batch(prompt, sampling_params)) assert expected["key"] in output.output assert captured["format"]["is_t2i"] is expected["is_t2i"] @@ -1353,15 +1381,13 @@ def test_forward_i2v_sound_and_action_routes(self, make_cosmos3_pipeline) -> Non torch.zeros(1, 2, 1, 1, 1), ) pipeline.forward( - SimpleNamespace( - prompts=[ - { - "prompt": "move", - "modalities": ["video"], - "additional_information": {"preprocessed_image": image_tensor}, - } - ], - sampling_params=make_sampling_params(height=16, width=16, num_frames=5), + make_request_batch( + { + "prompt": "move", + "modalities": ["video"], + "additional_information": {"preprocessed_image": image_tensor}, + }, + make_sampling_params(height=16, width=16, num_frames=5), ) ) assert captured["diffuse_calls"][-1]["shared_kwargs"]["noisy_frame_mask"] is velocity_mask @@ -1375,18 +1401,16 @@ def test_forward_i2v_sound_and_action_routes(self, make_cosmos3_pipeline) -> Non v2v_condition, ) pipeline.forward( - SimpleNamespace( - prompts=[ - { - "prompt": "continue", - "modalities": ["video"], - "additional_information": { - "preprocessed_video": video_tensor, - "condition_frame_indexes_vision": [0], - }, - } - ], - sampling_params=make_sampling_params(height=16, width=16, num_frames=5), + make_request_batch( + { + "prompt": "continue", + "modalities": ["video"], + "additional_information": { + "preprocessed_video": video_tensor, + "condition_frame_indexes_vision": [0], + }, + }, + make_sampling_params(height=16, width=16, num_frames=5), ) ) assert captured["flow_shifts"][-1] == 10.0 @@ -1401,9 +1425,9 @@ def test_forward_i2v_sound_and_action_routes(self, make_cosmos3_pipeline) -> Non pipeline._prepare_sound_latents = lambda *args: (sound_latents, 4) pipeline._decode_sound_latents = lambda *args: torch.ones(1, 2, 20) output = pipeline.forward( - SimpleNamespace( - prompts=[{"prompt": "A robot", "modalities": ["video"], "generate_sound": True}], - sampling_params=make_sampling_params(num_frames=9, frame_rate=3.0), + make_request_batch( + {"prompt": "A robot", "modalities": ["video"], "generate_sound": True}, + make_sampling_params(num_frames=9, frame_rate=3.0), ) ) assert captured["diffuse_calls"][-1]["sound_latents"] is sound_latents @@ -1411,15 +1435,13 @@ def test_forward_i2v_sound_and_action_routes(self, make_cosmos3_pipeline) -> Non pipeline.transformer = pipeline.transformer.__class__(latent_channel_size=2, action_gen=True, action_dim=4) output = pipeline.forward( - SimpleNamespace( - prompts=[ - { - "prompt": "Pick the block.", - "modalities": ["video"], - "additional_information": {"preprocessed_image": image_tensor}, - } - ], - sampling_params=make_sampling_params( + make_request_batch( + { + "prompt": "Pick the block.", + "modalities": ["video"], + "additional_information": {"preprocessed_image": image_tensor}, + }, + make_sampling_params( height=16, width=16, extra_args={ @@ -1498,7 +1520,7 @@ def fake_prepare_action_video(*args, **kwargs): AssertionError("RoboLab should not decode video") ) - output = pipeline.forward(SimpleNamespace(prompts=["ignored"], sampling_params=make_sampling_params())) + output = pipeline.forward(make_request_batch("ignored", make_sampling_params())) assert captured["format"] == { "prompt": "Pick the cube.", @@ -1528,24 +1550,24 @@ def fake_prepare_action_video(*args, **kwargs): ("prompt", "sampling_params", "message"), [ (["one", "two"], make_sampling_params(), "single prompt"), - ([{"prompt": "one", "modalities": ["image", "video"]}], make_sampling_params(), "both image and video"), + ({"prompt": "one", "modalities": ["image", "video"]}, make_sampling_params(), "both image and video"), ( - [{"prompt": "x", "modalities": ["image"], "generate_sound": True}], + {"prompt": "x", "modalities": ["image"], "generate_sound": True}, make_sampling_params(), "only for video", ), ( - [{"prompt": "x", "modalities": ["image"]}], + {"prompt": "x", "modalities": ["image"]}, make_sampling_params(extra_args={"edge": {"control_path": "/tmp/control.mp4"}}), "transfer inference is supported only for video outputs", ), ( - [{"prompt": "x", "modalities": ["video"], "generate_sound": True}], + {"prompt": "x", "modalities": ["video"], "generate_sound": True}, make_sampling_params(extra_args={"edge": {"control_path": "/tmp/control.mp4"}}), "cannot be combined with sound generation", ), ( - [{"prompt": "x", "modalities": ["video"]}], + {"prompt": "x", "modalities": ["video"]}, make_sampling_params( extra_args={ "edge": {"control_path": "/tmp/control.mp4"}, @@ -1567,4 +1589,4 @@ def test_forward_rejects_invalid_public_requests( pipeline.transformer = pipeline.transformer.__class__(latent_channel_size=2, sound_gen=True, sound_dim=3) with pytest.raises(ValueError, match=message): - pipeline.forward(SimpleNamespace(prompts=prompt, sampling_params=sampling_params)) + pipeline.forward(make_request_batch(prompt, sampling_params)) diff --git a/tests/diffusion/models/dmd2/test_dmd2_request_sanitization.py b/tests/diffusion/models/dmd2/test_dmd2_request_sanitization.py index 1365a3477c9..21679b5b7ff 100644 --- a/tests/diffusion/models/dmd2/test_dmd2_request_sanitization.py +++ b/tests/diffusion/models/dmd2/test_dmd2_request_sanitization.py @@ -43,10 +43,10 @@ def _mock_base_init(self, *a, **kw): return pipeline -def _make_request(prompts=None, **sp_kwargs) -> OmniDiffusionRequest: +def _make_request(prompt=None, **sp_kwargs) -> OmniDiffusionRequest: sp = OmniDiffusionSamplingParams(**sp_kwargs) return OmniDiffusionRequest( - prompts=prompts or [{"prompt": "a cat dancing"}], + prompt=prompt if prompt is not None else {"prompt": "a cat dancing"}, sampling_params=sp, request_id="dmd2-sanitize", ) @@ -133,35 +133,23 @@ def test_is_cfg_negative_forced_false(pipeline): def test_negative_prompt_stripped_from_prompt_dict(pipeline): - req = _make_request(prompts=[{"prompt": "a cat", "negative_prompt": "blurry"}]) + req = _make_request(prompt={"prompt": "a cat", "negative_prompt": "blurry"}) pipeline._sanitize_dmd2_request(req) - assert "negative_prompt" not in req.prompts[0] - assert req.prompts[0]["prompt"] == "a cat" + assert "negative_prompt" not in req.prompt + assert req.prompt["prompt"] == "a cat" def test_no_negative_prompt_unchanged(pipeline): - req = _make_request(prompts=[{"prompt": "a cat"}]) + req = _make_request(prompt={"prompt": "a cat"}) pipeline._sanitize_dmd2_request(req) - assert req.prompts[0] == {"prompt": "a cat"} + assert req.prompt == {"prompt": "a cat"} def test_string_prompt_not_mutated(pipeline): """String prompts (not dicts) must pass through unchanged.""" - req = _make_request(prompts=["a cat dancing"]) + req = _make_request(prompt="a cat dancing") pipeline._sanitize_dmd2_request(req) - assert req.prompts == ["a cat dancing"] - - -def test_multiple_prompts_all_sanitized(pipeline): - req = _make_request( - prompts=[ - {"prompt": "a cat", "negative_prompt": "blurry"}, - {"prompt": "a dog", "negative_prompt": "ugly"}, - ] - ) - pipeline._sanitize_dmd2_request(req) - for p in req.prompts: - assert "negative_prompt" not in p + assert req.prompt == "a cat dancing" # --------------------------------------------------------------------------- diff --git a/tests/diffusion/models/dmd2/test_dmd2_scheduler.py b/tests/diffusion/models/dmd2/test_dmd2_scheduler.py index e5a1b6d8361..0b6ca3a9e51 100644 --- a/tests/diffusion/models/dmd2/test_dmd2_scheduler.py +++ b/tests/diffusion/models/dmd2/test_dmd2_scheduler.py @@ -48,7 +48,7 @@ def _mock_base_init(self, *a, **kw): def _make_request(**sp_kwargs) -> OmniDiffusionRequest: sp = OmniDiffusionSamplingParams(**sp_kwargs) - return OmniDiffusionRequest(prompts=[{"prompt": "a cat"}], sampling_params=sp, request_id="dmd2-scheduler") + return OmniDiffusionRequest(prompt={"prompt": "a cat"}, sampling_params=sp, request_id="dmd2-scheduler") @pytest.fixture( diff --git a/tests/diffusion/models/flux/test_flux_pipeline.py b/tests/diffusion/models/flux/test_flux_pipeline.py new file mode 100644 index 00000000000..fd5d7a2acdb --- /dev/null +++ b/tests/diffusion/models/flux/test_flux_pipeline.py @@ -0,0 +1,182 @@ +from types import SimpleNamespace + +import pytest +import torch +from torch import nn + +from vllm_omni.diffusion.models.flux.pipeline_flux import FluxPipeline +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + + +def _make_flux_sampling(**overrides): + values = { + "height": 32, + "width": 32, + "num_inference_steps": 2, + "sigmas": None, + "guidance_scale": 3.5, + "generator": None, + "true_cfg_scale": 4.0, + "num_outputs_per_prompt": 0, + "latents": None, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def _make_flux_pipeline(): + pipeline = object.__new__(FluxPipeline) + nn.Module.__init__(pipeline) + pipeline.default_sample_size = 128 + pipeline.vae_scale_factor = 8 + pipeline.device = torch.device("cpu") + pipeline.text_encoder = None + pipeline.transformer = SimpleNamespace( + in_channels=4, + guidance_embeds=False, + dtype=torch.float32, + ) + return pipeline + + +def test_forward_collates_request_prompt_tensors_for_flux(monkeypatch): + monkeypatch.setattr( + "vllm_omni.diffusion.models.flux.pipeline_flux.get_classifier_free_guidance_world_size", + lambda: 1, + ) + pipeline = _make_flux_pipeline() + encode_calls = [] + prepare_latents_call = {} + diffuse_call = {} + + def _fake_encode_prompt(**kwargs): + encode_calls.append(kwargs) + return ( + kwargs["prompt_embeds"], + kwargs["pooled_prompt_embeds"], + torch.zeros(kwargs["prompt_embeds"].shape[1], 3), + ) + + def _fake_prepare_latents( + batch_size, + num_channels_latents, + height, + width, + dtype, + device, + generator, + latents, + ): + prepare_latents_call.update( + { + "batch_size": batch_size, + "num_channels_latents": num_channels_latents, + "height": height, + "width": width, + "dtype": dtype, + "device": device, + "generator": generator, + "latents": latents, + } + ) + return torch.zeros(batch_size, 1, 1), torch.zeros(1, 3) + + def _fake_diffuse( + prompt_embeds, + pooled_prompt_embeds, + negative_prompt_embeds, + negative_pooled_prompt_embeds, + *args, + **kwargs, + ): + diffuse_call.update( + { + "prompt_embeds": prompt_embeds, + "pooled_prompt_embeds": pooled_prompt_embeds, + "negative_prompt_embeds": negative_prompt_embeds, + "negative_pooled_prompt_embeds": negative_pooled_prompt_embeds, + } + ) + return torch.arange(2, dtype=torch.float32).view(2, 1) + + pipeline.encode_prompt = _fake_encode_prompt + pipeline.prepare_latents = _fake_prepare_latents + pipeline.prepare_timesteps = lambda *args, **kwargs: (torch.tensor([1.0]), 1) + pipeline.diffuse = _fake_diffuse + + prompt_embeds_a = torch.zeros(2, 3) + prompt_embeds_b = torch.ones(2, 3) + pooled_prompt_embeds_a = torch.full((4,), 2.0) + pooled_prompt_embeds_b = torch.full((4,), 3.0) + negative_prompt_embeds_a = torch.full((2, 3), 4.0) + negative_prompt_embeds_b = torch.full((2, 3), 5.0) + negative_pooled_prompt_embeds_a = torch.full((4,), 6.0) + negative_pooled_prompt_embeds_b = torch.full((4,), 7.0) + latents_a = torch.zeros(1, 1, 1) + latents_b = torch.ones(1, 1, 1) + gen_a = torch.Generator(device="cpu").manual_seed(1) + gen_b = torch.Generator(device="cpu").manual_seed(2) + + batch = DiffusionRequestBatch( + requests=[ + SimpleNamespace( + request_id="flux-prompt-a", + prompt={ + "prompt": "prompt-a", + "negative_prompt": "negative-a", + "prompt_embeds": prompt_embeds_a, + "pooled_prompt_embeds": pooled_prompt_embeds_a, + "negative_prompt_embeds": negative_prompt_embeds_a, + "negative_pooled_prompt_embeds": negative_pooled_prompt_embeds_a, + }, + sampling_params=_make_flux_sampling(generator=gen_a, latents=latents_a), + ), + SimpleNamespace( + request_id="flux-prompt-b", + prompt={ + "prompt": "prompt-b", + "negative_prompt": "negative-b", + "additional_information": { + "prompt_embeds": [prompt_embeds_b], + "pooled_prompt_embeds": [pooled_prompt_embeds_b], + "negative_prompt_embeds": [negative_prompt_embeds_b], + "negative_pooled_prompt_embeds": [negative_pooled_prompt_embeds_b], + }, + }, + sampling_params=_make_flux_sampling(generator=gen_b, latents=latents_b), + ), + ] + ) + + outputs = pipeline.forward(batch, output_type="latent") + + assert encode_calls[0]["prompt"] is None + assert encode_calls[0]["prompt_2"] is None + torch.testing.assert_close( + encode_calls[0]["prompt_embeds"], + torch.stack([prompt_embeds_a, prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + encode_calls[0]["pooled_prompt_embeds"], + torch.stack([pooled_prompt_embeds_a, pooled_prompt_embeds_b], dim=0), + ) + assert encode_calls[1]["prompt"] is None + assert encode_calls[1]["prompt_2"] is None + torch.testing.assert_close( + encode_calls[1]["prompt_embeds"], + torch.stack([negative_prompt_embeds_a, negative_prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + encode_calls[1]["pooled_prompt_embeds"], + torch.stack([negative_pooled_prompt_embeds_a, negative_pooled_prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + diffuse_call["negative_pooled_prompt_embeds"], + torch.stack([negative_pooled_prompt_embeds_a, negative_pooled_prompt_embeds_b], dim=0), + ) + assert prepare_latents_call["generator"] == [gen_a, gen_b] + torch.testing.assert_close(prepare_latents_call["latents"], torch.cat([latents_a, latents_b], dim=0)) + torch.testing.assert_close(outputs[0].output, torch.tensor([[0.0]])) + torch.testing.assert_close(outputs[1].output, torch.tensor([[1.0]])) diff --git a/tests/diffusion/models/flux2/test_flux2_klein_num_inference_steps.py b/tests/diffusion/models/flux2/test_flux2_klein_num_inference_steps.py index 3c6b5bada06..5afcc4e7830 100644 --- a/tests/diffusion/models/flux2/test_flux2_klein_num_inference_steps.py +++ b/tests/diffusion/models/flux2/test_flux2_klein_num_inference_steps.py @@ -39,7 +39,7 @@ def _make_minimal_request( ) req = MagicMock() req.sampling_params = params - req.prompts = [prompt] + req.prompt = prompt req.multi_modal_data = {} return req diff --git a/tests/diffusion/models/gr00t/test_pipeline.py b/tests/diffusion/models/gr00t/test_pipeline.py index 72118c3f93d..3221121074f 100644 --- a/tests/diffusion/models/gr00t/test_pipeline.py +++ b/tests/diffusion/models/gr00t/test_pipeline.py @@ -84,7 +84,7 @@ def test_pipeline_initializes_local_policy(): def test_forward_returns_dict_actions_in_output(): pipeline = _pipeline() req = OmniDiffusionRequest( - prompts=["pick"], + prompt="pick", request_id="req", sampling_params=OmniDiffusionSamplingParams( extra_args={ @@ -118,7 +118,7 @@ def test_forward_returns_dict_actions_in_output(): def test_dummy_warmup_returns_shape_correct_zero_actions(): pipeline = _pipeline() req = OmniDiffusionRequest( - prompts=["dummy run"], + prompt="dummy run", request_id="dummy_req_id", sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), ) diff --git a/tests/diffusion/models/hunyuan_image3/test_hunyuan_image3_step_execution.py b/tests/diffusion/models/hunyuan_image3/test_hunyuan_image3_step_execution.py index 7b6c94c62f2..b4f1e7493f2 100644 --- a/tests/diffusion/models/hunyuan_image3/test_hunyuan_image3_step_execution.py +++ b/tests/diffusion/models/hunyuan_image3/test_hunyuan_image3_step_execution.py @@ -19,6 +19,7 @@ HunyuanImage3Pipeline, ) from vllm_omni.diffusion.worker.input_batch import InputBatch +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.diffusion.worker.utils import DiffusionRequestState pytestmark = [pytest.mark.core_model, pytest.mark.diffusion, pytest.mark.cpu] @@ -41,7 +42,7 @@ def _state(request_id: str, step_index: int) -> DiffusionRequestState: state = DiffusionRequestState( request_id=request_id, sampling=SimpleNamespace(), - prompts=["prompt"], + prompt="prompt", ) state.step_index = step_index state.timesteps = torch.tensor([1.0, 0.5, 0.25, 0.0]) @@ -135,7 +136,7 @@ def fake_prepare_model_inputs(**kwargs): state = DiffusionRequestState( request_id="req-bot-task", sampling=sampling, - prompts=[prompt_item], + prompt=prompt_item, ) with pytest.raises(RuntimeError, match="stop after prepare_model_inputs"): @@ -160,10 +161,14 @@ def fake_prepare_model_inputs(**kwargs): monkeypatch.setattr(hy3_module, "get_system_prompt", fake_get_system_prompt) pipeline.prepare_model_inputs = fake_prepare_model_inputs - req = SimpleNamespace( - request_id="req-forward-bot-task", - sampling_params=_sampling_params(bot_task="think_recaption", use_system_prompt="dynamic"), - prompts=[{"prompt": "prompt", "bot_task": "vanilla"}], + req = DiffusionRequestBatch( + requests=[ + SimpleNamespace( + request_id="req-forward-bot-task", + sampling_params=_sampling_params(bot_task="think_recaption", use_system_prompt="dynamic"), + prompt={"prompt": "prompt", "bot_task": "vanilla"}, + ) + ] ) with pytest.raises(RuntimeError, match="stop after prepare_model_inputs"): diff --git a/tests/diffusion/models/hunyuan_video/test_hunyuan_video_quant_config_propagation.py b/tests/diffusion/models/hunyuan_video/test_hunyuan_video_quant_config_propagation.py index 524a9b93b33..d4bbb0f8ef0 100644 --- a/tests/diffusion/models/hunyuan_video/test_hunyuan_video_quant_config_propagation.py +++ b/tests/diffusion/models/hunyuan_video/test_hunyuan_video_quant_config_propagation.py @@ -10,6 +10,7 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch +import PIL.Image import pytest import torch @@ -19,6 +20,21 @@ pytestmark = [pytest.mark.core_model, pytest.mark.cpu, pytest.mark.diffusion] +def test_i2v_pre_process_uses_single_request_prompt(): + preprocess = i2v_module.get_hunyuan_video_15_i2v_pre_process_func(SimpleNamespace()) + image = PIL.Image.new("RGB", (640, 320)) + request = SimpleNamespace( + prompt={"prompt": "turn this into a video", "multi_modal_data": {"image": image}}, + sampling_params=SimpleNamespace(height=None, width=None), + ) + + result = preprocess(request) + + assert result is request + assert request.sampling_params.height == 448 + assert request.sampling_params.width == 896 + + class TestHunyuanVideoQuantConfigPropagation: """Verify quant_config is propagated to the transformer model in HunyuanVideo-1.5 pipelines.""" diff --git a/tests/diffusion/models/ltx2/test_ltx2_3_pipeline.py b/tests/diffusion/models/ltx2/test_ltx2_3_pipeline.py index cb39c0dc817..96051416b9d 100644 --- a/tests/diffusion/models/ltx2/test_ltx2_3_pipeline.py +++ b/tests/diffusion/models/ltx2/test_ltx2_3_pipeline.py @@ -465,6 +465,99 @@ def fake_predict_noise(**kwargs): class TestCFGParallelForwardPath: """Test the LTX-2.3 CFG-parallel denoising path without loading model weights.""" + def test_forward_collates_request_prompt_embeds_and_mask_aliases(self, monkeypatch): + from vllm_omni.diffusion.models.ltx2 import pipeline_ltx2_3 as ltx23 + from vllm_omni.diffusion.request import OmniDiffusionRequest + from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch + from vllm_omni.inputs.data import OmniDiffusionSamplingParams + + pipe = object.__new__(ltx23.LTX23Pipeline) + torch.nn.Module.__init__(pipe) + pipe.device = torch.device("cpu") + pipe.tokenizer_max_length = 4 + monkeypatch.setattr(ltx23, "get_classifier_free_guidance_world_size", lambda: 1) + + class StopAtEncodePromptError(Exception): + pass + + captured = {} + + def fake_encode_prompt(**kwargs): + captured.update(kwargs) + raise StopAtEncodePromptError + + object.__setattr__(pipe, "encode_prompt", fake_encode_prompt) + + prompt_embeds_a = torch.zeros(2, 3) + prompt_embeds_b = torch.ones(2, 3) + negative_prompt_embeds_a = torch.full((2, 3), 2.0) + negative_prompt_embeds_b = torch.full((2, 3), 3.0) + prompt_attention_mask_a = torch.tensor([True, True]) + prompt_attention_mask_b = torch.tensor([True, False]) + negative_attention_mask_a = torch.tensor([False, True]) + negative_attention_mask_b = torch.tensor([False, False]) + + requests = [ + OmniDiffusionRequest( + prompt={ + "prompt": "prompt-a", + "negative_prompt": "negative-a", + "prompt_embeds": prompt_embeds_a, + "negative_prompt_embeds": negative_prompt_embeds_a, + "prompt_attention_mask": prompt_attention_mask_a, + "negative_prompt_attention_mask": negative_attention_mask_a, + }, + sampling_params=OmniDiffusionSamplingParams( + height=32, + width=32, + num_frames=1, + frame_rate=1.0, + num_inference_steps=2, + ), + request_id="ltx23-prompt-local-a", + ), + OmniDiffusionRequest( + prompt={ + "prompt": "prompt-b", + "negative_prompt": "negative-b", + "prompt_embeds": prompt_embeds_b, + "negative_prompt_embeds": negative_prompt_embeds_b, + "attention_mask": prompt_attention_mask_b, + "negative_attention_mask": negative_attention_mask_b, + }, + sampling_params=OmniDiffusionSamplingParams( + height=32, + width=32, + num_frames=1, + frame_rate=1.0, + num_inference_steps=2, + ), + request_id="ltx23-prompt-local-b", + ), + ] + + with pytest.raises(StopAtEncodePromptError): + pipe.forward(DiffusionRequestBatch(requests=requests)) + + assert captured["prompt"] is None + assert captured["negative_prompt"] is None + torch.testing.assert_close( + captured["prompt_embeds"], + torch.stack([prompt_embeds_a, prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + captured["negative_prompt_embeds"], + torch.stack([negative_prompt_embeds_a, negative_prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + captured["prompt_attention_mask"], + torch.stack([prompt_attention_mask_a, prompt_attention_mask_b], dim=0), + ) + torch.testing.assert_close( + captured["negative_prompt_attention_mask"], + torch.stack([negative_attention_mask_a, negative_attention_mask_b], dim=0), + ) + @pytest.mark.parametrize(("cfg_rank", "expected_prompt_value"), [(0, 1.0), (1, 0.0)]) @pytest.mark.parametrize( ("frame_rate_input", "audio_sampling_rate", "expected_frame_rate"), @@ -481,6 +574,7 @@ def test_forward_cfg_parallel_steps_video_and_audio_scheduler( ): from vllm_omni.diffusion.models.ltx2 import pipeline_ltx2_3 as ltx23 from vllm_omni.diffusion.request import OmniDiffusionRequest + from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniDiffusionSamplingParams pipe = object.__new__(ltx23.LTX23Pipeline) @@ -616,7 +710,7 @@ def fake_encode_prompt(**kwargs): video_latents = torch.tensor([[[1.0, -2.0]]]) audio_latents = torch.tensor([[[0.5, 3.0]]]) req = OmniDiffusionRequest( - prompts=[{"prompt": "prompt", "negative_prompt": "negative"}], + prompt={"prompt": "prompt", "negative_prompt": "negative"}, sampling_params=OmniDiffusionSamplingParams( height=32, width=32, @@ -631,7 +725,7 @@ def fake_encode_prompt(**kwargs): request_id="ltx23-cfg-parallel-forward-test", ) - output = pipe.forward(req) + output = pipe.forward(DiffusionRequestBatch(requests=[req]))[0] expected_video_noise = ltx23.LTX23Pipeline._combine_x0_space_cfg( video_latents, diff --git a/tests/diffusion/models/ovis_image/test_ovis_image.py b/tests/diffusion/models/ovis_image/test_ovis_image.py index 8e357e6c636..05092ade234 100644 --- a/tests/diffusion/models/ovis_image/test_ovis_image.py +++ b/tests/diffusion/models/ovis_image/test_ovis_image.py @@ -24,6 +24,7 @@ from vllm_omni.diffusion.distributed.utils import get_local_device from vllm_omni.diffusion.models.ovis_image.pipeline_ovis_image import OvisImagePipeline from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniDiffusionSamplingParams @@ -166,7 +167,7 @@ def test_basic_generation(ovis_pipeline): """Test the forward pass logic.""" # Setup request req = OmniDiffusionRequest( - prompts=["A photo of a cat"], + prompt="A photo of a cat", request_id="ovis-basic", sampling_params=OmniDiffusionSamplingParams( height=256, @@ -176,7 +177,7 @@ def test_basic_generation(ovis_pipeline): ), ) - output = ovis_pipeline(req) + output = ovis_pipeline(DiffusionRequestBatch(requests=[req])) assert output is not None assert output.output is not None @@ -201,12 +202,10 @@ def test_guidance_scale(ovis_pipeline, monkeypatch: pytest.MonkeyPatch): monkeypatch.setattr(f"{_cfg_parallel}.get_classifier_free_guidance_rank", lambda: 0) req = OmniDiffusionRequest( - prompts=[ - { - "prompt": "A photo of a cat", - "negative_prompt": "bad quality", - } - ], + prompt={ + "prompt": "A photo of a cat", + "negative_prompt": "bad quality", + }, request_id="ovis-guidance", sampling_params=OmniDiffusionSamplingParams( height=256, @@ -216,7 +215,7 @@ def test_guidance_scale(ovis_pipeline, monkeypatch: pytest.MonkeyPatch): ), ) - ovis_pipeline(req) + ovis_pipeline(DiffusionRequestBatch(requests=[req])) assert ovis_pipeline.transformer.call_count >= 2 @@ -226,7 +225,7 @@ def test_resolution_check(ovis_pipeline): """Test resolution divisible validation logic if present.""" # Pass odd resolution req = OmniDiffusionRequest( - prompts=["test"], + prompt="test", request_id="ovis-resolution", sampling_params=OmniDiffusionSamplingParams( height=250, # Not divisible by 16 (8*2) @@ -237,7 +236,7 @@ def test_resolution_check(ovis_pipeline): # Should warn but proceed (as per code I read earlier) or resize? # The code had `logger.warning(...)` - output = ovis_pipeline(req) + output = ovis_pipeline(DiffusionRequestBatch(requests=[req])) assert output is not None diff --git a/tests/diffusion/models/qwen_image/test_qwen_image_edit_plus.py b/tests/diffusion/models/qwen_image/test_qwen_image_edit_plus.py index 873b52bf7a6..6f09817ee7b 100644 --- a/tests/diffusion/models/qwen_image/test_qwen_image_edit_plus.py +++ b/tests/diffusion/models/qwen_image/test_qwen_image_edit_plus.py @@ -25,12 +25,10 @@ def test_qwen_image_edit_plus_rejects_too_many_input_images(tmp_path: Path): pre_process = get_qwen_image_edit_plus_pre_process_func(SimpleNamespace(model=str(tmp_path))) image = Image.fromarray(np.zeros((32, 32, 3), dtype=np.uint8)) request = SimpleNamespace( - prompts=[ - { - "prompt": "combine", - "multi_modal_data": {"image": [image, image, image, image, image]}, - } - ], + prompt={ + "prompt": "combine", + "multi_modal_data": {"image": [image, image, image, image, image]}, + }, sampling_params=SimpleNamespace(height=None, width=None), ) diff --git a/tests/diffusion/models/qwen_image/test_qwen_image_max_sequence_length.py b/tests/diffusion/models/qwen_image/test_qwen_image_max_sequence_length.py index f5676a0056f..3f85a6b17ff 100644 --- a/tests/diffusion/models/qwen_image/test_qwen_image_max_sequence_length.py +++ b/tests/diffusion/models/qwen_image/test_qwen_image_max_sequence_length.py @@ -17,6 +17,7 @@ from vllm_omni.diffusion.models.qwen_image.pipeline_qwen_image_layered import ( QwenImageLayeredPipeline, ) +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch pytestmark = [pytest.mark.core_model, pytest.mark.cpu] @@ -160,7 +161,7 @@ def _fake_prepare_generation_context(**kwargs): pipeline._prepare_generation_context = _fake_prepare_generation_context state = SimpleNamespace( - prompts=["prompt"], + prompt="prompt", sampling=SimpleNamespace( height=None, width=None, @@ -179,6 +180,104 @@ def _fake_prepare_generation_context(**kwargs): assert captured["max_sequence_length"] == 1024 +def _make_request_batch_prompt_sampling(**overrides): + values = { + "height": 32, + "width": 32, + "num_inference_steps": 2, + "sigmas": None, + "max_sequence_length": None, + "num_outputs_per_prompt": 0, + "generator": None, + "latents": None, + "true_cfg_scale": None, + "guidance_scale_provided": False, + "guidance_scale": 1.0, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def test_forward_collates_request_prompt_tensors_for_qwen_image(): + pipeline = object.__new__(QwenImagePipeline) + nn.Module.__init__(pipeline) + pipeline.vae_scale_factor = 8 + pipeline.default_sample_size = 128 + + class StopAfterPrepareContextError(Exception): + pass + + captured = {} + + def _fake_prepare_generation_context(**kwargs): + captured.update(kwargs) + raise StopAfterPrepareContextError + + pipeline._prepare_generation_context = _fake_prepare_generation_context + + prompt_embeds_a = torch.zeros(2, 3) + prompt_embeds_b = torch.ones(2, 3) + prompt_embeds_mask_a = torch.tensor([True, True]) + prompt_embeds_mask_b = torch.tensor([True, False]) + negative_prompt_embeds_a = torch.full((2, 3), 2.0) + negative_prompt_embeds_b = torch.full((2, 3), 3.0) + negative_prompt_embeds_mask_a = torch.tensor([False, True]) + negative_prompt_embeds_mask_b = torch.tensor([False, False]) + + batch = DiffusionRequestBatch( + requests=[ + SimpleNamespace( + request_id="qwen-prompt-a", + prompt={ + "prompt": "prompt-a", + "negative_prompt": "negative-a", + "prompt_embeds": prompt_embeds_a, + "prompt_embeds_mask": prompt_embeds_mask_a, + "negative_prompt_embeds": negative_prompt_embeds_a, + "negative_prompt_embeds_mask": negative_prompt_embeds_mask_a, + }, + sampling_params=_make_request_batch_prompt_sampling(), + ), + SimpleNamespace( + request_id="qwen-prompt-b", + prompt={ + "prompt": "prompt-b", + "negative_prompt": "negative-b", + "additional_information": { + "prompt_embeds": [prompt_embeds_b], + "prompt_embeds_mask": [prompt_embeds_mask_b], + "negative_prompt_embeds": [negative_prompt_embeds_b], + "negative_prompt_embeds_mask": [negative_prompt_embeds_mask_b], + }, + }, + sampling_params=_make_request_batch_prompt_sampling(), + ), + ] + ) + + with pytest.raises(StopAfterPrepareContextError): + pipeline.forward(batch) + + assert captured["prompt"] is None + assert captured["negative_prompt"] is None + torch.testing.assert_close( + captured["prompt_embeds"], + torch.stack([prompt_embeds_a, prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + captured["prompt_embeds_mask"], + torch.stack([prompt_embeds_mask_a, prompt_embeds_mask_b], dim=0), + ) + torch.testing.assert_close( + captured["negative_prompt_embeds"], + torch.stack([negative_prompt_embeds_a, negative_prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + captured["negative_prompt_embeds_mask"], + torch.stack([negative_prompt_embeds_mask_a, negative_prompt_embeds_mask_b], dim=0), + ) + + @pytest.mark.parametrize( ("pipeline_class", "drop_idx"), [ diff --git a/tests/diffusion/models/sd3/test_sd3_pipeline.py b/tests/diffusion/models/sd3/test_sd3_pipeline.py new file mode 100644 index 00000000000..3164e5dda81 --- /dev/null +++ b/tests/diffusion/models/sd3/test_sd3_pipeline.py @@ -0,0 +1,154 @@ +from types import SimpleNamespace + +import pytest +import torch +from torch import nn + +from vllm_omni.diffusion.models.sd3.pipeline_sd3 import StableDiffusion3Pipeline +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch + +pytestmark = [pytest.mark.core_model, pytest.mark.cpu] + + +def _make_sd3_sampling(**overrides): + values = { + "height": 32, + "width": 32, + "num_inference_steps": 2, + "sigmas": None, + "max_sequence_length": None, + "num_outputs_per_prompt": 0, + "generator": None, + "latents": None, + "guidance_scale": 4.0, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def _make_sd3_pipeline(): + pipeline = object.__new__(StableDiffusion3Pipeline) + nn.Module.__init__(pipeline) + pipeline.vae_scale_factor = 8 + pipeline.patch_size = 2 + pipeline.default_sample_size = 128 + pipeline.transformer = SimpleNamespace(in_channels=1) + return pipeline + + +def test_forward_collates_request_prompt_tensors_for_sd3(): + pipeline = _make_sd3_pipeline() + + class StopAfterDiffuseError(Exception): + pass + + encode_calls = [] + diffuse_call = {} + + def _fake_encode_prompt(**kwargs): + encode_calls.append(kwargs) + prompt_embeds = kwargs["prompt_embeds"] + if prompt_embeds is None: + prompt_embeds = torch.empty(2, 2, 3) + return prompt_embeds, kwargs.get("pooled_prompt_embeds") + + def _fake_diffuse(**kwargs): + diffuse_call.update(kwargs) + raise StopAfterDiffuseError + + pipeline.encode_prompt = _fake_encode_prompt + pipeline.prepare_latents = lambda *args, **kwargs: torch.zeros(2, 1, 1, 1) + pipeline.prepare_timesteps = lambda *args, **kwargs: (torch.tensor([1.0]), 1) + pipeline.diffuse = _fake_diffuse + + prompt_embeds_a = torch.zeros(2, 3) + prompt_embeds_b = torch.ones(2, 3) + pooled_prompt_embeds_a = torch.full((4,), 2.0) + pooled_prompt_embeds_b = torch.full((4,), 3.0) + negative_prompt_embeds_a = torch.full((2, 3), 4.0) + negative_prompt_embeds_b = torch.full((2, 3), 5.0) + negative_pooled_prompt_embeds_a = torch.full((4,), 6.0) + negative_pooled_prompt_embeds_b = torch.full((4,), 7.0) + + batch = DiffusionRequestBatch( + requests=[ + SimpleNamespace( + request_id="sd3-prompt-a", + prompt={ + "prompt": "prompt-a", + "negative_prompt": "negative-a", + "prompt_embeds": prompt_embeds_a, + "pooled_prompt_embeds": pooled_prompt_embeds_a, + "negative_prompt_embeds": negative_prompt_embeds_a, + "negative_pooled_prompt_embeds": negative_pooled_prompt_embeds_a, + }, + sampling_params=_make_sd3_sampling(), + ), + SimpleNamespace( + request_id="sd3-prompt-b", + prompt={ + "prompt": "prompt-b", + "negative_prompt": "negative-b", + "additional_information": { + "prompt_embeds": [prompt_embeds_b], + "pooled_prompt_embeds": [pooled_prompt_embeds_b], + "negative_prompt_embeds": [negative_prompt_embeds_b], + "negative_pooled_prompt_embeds": [negative_pooled_prompt_embeds_b], + }, + }, + sampling_params=_make_sd3_sampling(), + ), + ] + ) + + with pytest.raises(StopAfterDiffuseError): + pipeline.forward(batch) + + assert encode_calls[0]["prompt"] is None + assert encode_calls[0]["prompt_2"] is None + assert encode_calls[0]["prompt_3"] is None + torch.testing.assert_close( + encode_calls[0]["prompt_embeds"], + torch.stack([prompt_embeds_a, prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + encode_calls[0]["pooled_prompt_embeds"], + torch.stack([pooled_prompt_embeds_a, pooled_prompt_embeds_b], dim=0), + ) + + assert encode_calls[1]["prompt"] is None + assert encode_calls[1]["prompt_2"] is None + assert encode_calls[1]["prompt_3"] is None + torch.testing.assert_close( + encode_calls[1]["prompt_embeds"], + torch.stack([negative_prompt_embeds_a, negative_prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + encode_calls[1]["pooled_prompt_embeds"], + torch.stack([negative_pooled_prompt_embeds_a, negative_pooled_prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + diffuse_call["pooled_prompt_embeds"], + torch.stack([pooled_prompt_embeds_a, pooled_prompt_embeds_b], dim=0), + ) + torch.testing.assert_close( + diffuse_call["negative_pooled_prompt_embeds"], + torch.stack([negative_pooled_prompt_embeds_a, negative_pooled_prompt_embeds_b], dim=0), + ) + + +def test_encode_prompt_preserves_direct_pooled_prompt_embeds(): + pipeline = _make_sd3_pipeline() + prompt_embeds = torch.zeros(1, 2, 3) + pooled_prompt_embeds = torch.ones(1, 4) + + actual_prompt_embeds, actual_pooled_prompt_embeds = pipeline.encode_prompt( + prompt=None, + prompt_2=None, + prompt_3=None, + prompt_embeds=prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, + ) + + assert actual_prompt_embeds is prompt_embeds + assert actual_pooled_prompt_embeds is pooled_prompt_embeds diff --git a/tests/diffusion/models/wan2_2/test_wan22_i2v_pipeline.py b/tests/diffusion/models/wan2_2/test_wan22_i2v_pipeline.py index 576678e2cf0..9d5b11d3ba8 100644 --- a/tests/diffusion/models/wan2_2/test_wan22_i2v_pipeline.py +++ b/tests/diffusion/models/wan2_2/test_wan22_i2v_pipeline.py @@ -34,19 +34,19 @@ def _make_i2v_pipeline(*, expand_timesteps: bool) -> Wan22I2VPipeline: def test_i2v_preprocess_requires_image_and_resizes_to_480p_aspect() -> None: preprocess = get_wan22_i2v_pre_process_func(SimpleNamespace()) request = SimpleNamespace( - prompts=[{"prompt": "p", "multi_modal_data": {"image": Image.new("RGB", (320, 160), "red")}}], + prompt={"prompt": "p", "multi_modal_data": {"image": Image.new("RGB", (320, 160), "red")}}, sampling_params=SimpleNamespace(height=None, width=None), ) result = preprocess(request) - prompt = result.prompts[0] + prompt = result.prompt assert result.sampling_params.height == 432 assert result.sampling_params.width == 880 assert prompt["multi_modal_data"]["image"].size == (880, 432) missing_image = SimpleNamespace( - prompts=[{"prompt": "p", "multi_modal_data": {}}], + prompt={"prompt": "p", "multi_modal_data": {}}, sampling_params=SimpleNamespace(height=None, width=None), ) with pytest.raises(ValueError, match="No image is provided"): @@ -100,7 +100,10 @@ def fake_predict_noise_maybe_with_cfg(**kwargs): timestep_dtype = calls[0]["timestep_values"].dtype torch.testing.assert_close(calls[0]["timestep_values"][0, :4], torch.zeros(4, dtype=timestep_dtype)) torch.testing.assert_close(calls[0]["timestep_values"][0, 4:], torch.full((4,), 900, dtype=timestep_dtype)) - torch.testing.assert_close(calls[0]["hidden_states"][:, :, 0], torch.ones(1, 4, 4, 4)) + torch.testing.assert_close( + calls[0]["hidden_states"][:, :, 0], + torch.ones_like(calls[0]["hidden_states"][:, :, 0]), + ) torch.testing.assert_close(result, torch.full_like(latents, 2.0)) diff --git a/tests/diffusion/models/wan2_2/test_wan22_pipeline_diffuse.py b/tests/diffusion/models/wan2_2/test_wan22_pipeline_diffuse.py index 54bb672ef81..d695766f64e 100644 --- a/tests/diffusion/models/wan2_2/test_wan22_pipeline_diffuse.py +++ b/tests/diffusion/models/wan2_2/test_wan22_pipeline_diffuse.py @@ -9,6 +9,7 @@ from torch import nn from vllm_omni.diffusion.models.wan2_2.pipeline_wan2_2 import Wan22Pipeline +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch pytestmark = [pytest.mark.core_model, pytest.mark.cpu, pytest.mark.diffusion] @@ -77,8 +78,9 @@ def _fake_diffuse(**kwargs): pipeline.diffuse = _fake_diffuse # type: ignore[method-assign] - req = SimpleNamespace( - prompts=["prompt"], + mock_req = SimpleNamespace( + prompt="prompt", + request_id="test-req", sampling_params=SimpleNamespace( height=None, width=None, @@ -96,8 +98,9 @@ def _fake_diffuse(**kwargs): extra_args={}, ), ) + batch = DiffusionRequestBatch(requests=[mock_req]) - output = pipeline.forward(req, prompt_embeds=prompt_embeds, output_type="latent", guidance_scale=1.0) + output = pipeline.forward(batch, prompt_embeds=prompt_embeds, output_type="latent", guidance_scale=1.0) assert torch.equal(output.output, torch.ones((1, 4, 1, 8, 8))) assert torch.equal(captured["timesteps"], pipeline.scheduler.timesteps) diff --git a/tests/diffusion/models/wan2_2/test_wan22_vace_pipeline.py b/tests/diffusion/models/wan2_2/test_wan22_vace_pipeline.py index 97a000df1e0..714b79d5066 100644 --- a/tests/diffusion/models/wan2_2/test_wan22_vace_pipeline.py +++ b/tests/diffusion/models/wan2_2/test_wan22_vace_pipeline.py @@ -38,21 +38,19 @@ def test_vace_preprocess_collects_reference_video_and_mask_inputs() -> None: frame = Image.new("RGB", (64, 64), "black") mask = Image.new("L", (64, 64), 255) request = SimpleNamespace( - prompts=[ - { - "prompt": "p", - "multi_modal_data": { - "image": ref, - "video": [frame], - "mask": mask, - }, - } - ], + prompt={ + "prompt": "p", + "multi_modal_data": { + "image": ref, + "video": [frame], + "mask": mask, + }, + }, sampling_params=SimpleNamespace(height=None, width=None), ) result = preprocess(request) - additional_info = result.prompts[0]["additional_information"] + additional_info = result.prompt["additional_information"] assert result.sampling_params.height == 432 assert result.sampling_params.width == 880 diff --git a/tests/diffusion/test_diffusion_engine.py b/tests/diffusion/test_diffusion_engine.py index eb862a91ae0..5f9a5a5c0dc 100644 --- a/tests/diffusion/test_diffusion_engine.py +++ b/tests/diffusion/test_diffusion_engine.py @@ -11,8 +11,20 @@ import pytest import torch - -from vllm_omni.diffusion.diffusion_engine import _move_tensor_tree_to_cpu +from pytest_mock import MockerFixture + +import vllm_omni.diffusion.diffusion_engine as diffusion_engine_module +from vllm_omni.diffusion.data import DiffusionOutput, OmniDiffusionConfig +from vllm_omni.diffusion.diffusion_engine import DiffusionEngine, _move_tensor_tree_to_cpu +from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.sched.interface import ( + CachedRequestData, + NewRequestData, +) +from vllm_omni.diffusion.sched.interface import ( + DiffusionSchedulerOutput as RealDiffusionSchedulerOutput, +) +from vllm_omni.inputs.data import OmniDiffusionSamplingParams @dataclass @@ -66,6 +78,324 @@ def update_from_output(self, sched_output, runner_output): return [req.request_id for req in sched_output.scheduled_new_reqs] +class _BatchCapablePipeline: + supports_request_batch = True + + +class _SingleRequestPipeline: + pass + + +class _SingleRequestOverridePipeline(_BatchCapablePipeline): + def forward(self, req, prompt_ids=None): + return DiffusionOutput(output=None) + + +def _make_request_mode_sched_output(*request_ids: str) -> RealDiffusionSchedulerOutput: + new_reqs = [ + NewRequestData( + request_id=request_id, + req=OmniDiffusionRequest( + prompt=f"prompt_{request_id}", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), + request_id=request_id, + ), + ) + for request_id in request_ids + ] + return RealDiffusionSchedulerOutput( + step_id=0, + scheduled_new_reqs=new_reqs, + scheduled_cached_reqs=CachedRequestData.make_empty(), + finished_req_ids=set(), + num_running_reqs=len(new_reqs), + num_waiting_reqs=0, + ) + + +class TestRequestBatchCapability: + pytestmark = [pytest.mark.core_model, pytest.mark.diffusion, pytest.mark.cpu] + + def test_supports_request_batch_uses_registered_model_class(self, monkeypatch: pytest.MonkeyPatch) -> None: + od_config = SimpleNamespace(model_class_name="BatchPipeline", custom_pipeline_args=None) + + monkeypatch.setattr( + diffusion_engine_module.DiffusionModelRegistry, + "_try_load_model_cls", + lambda model_class_name: _BatchCapablePipeline if model_class_name == "BatchPipeline" else None, + ) + + assert diffusion_engine_module.supports_request_batch(od_config) is True + + def test_supports_request_batch_uses_custom_pipeline_class(self, monkeypatch: pytest.MonkeyPatch) -> None: + od_config = SimpleNamespace( + model_class_name="SinglePipeline", + custom_pipeline_args={"pipeline_class": _BatchCapablePipeline}, + ) + + monkeypatch.setattr( + diffusion_engine_module.DiffusionModelRegistry, + "_try_load_model_cls", + lambda model_class_name: _SingleRequestPipeline, + ) + + assert diffusion_engine_module.supports_request_batch(od_config) is True + + def test_supports_request_batch_uses_custom_pipeline_class_name(self, monkeypatch: pytest.MonkeyPatch) -> None: + od_config = SimpleNamespace( + model_class_name="SinglePipeline", + custom_pipeline_args={"pipeline_class": "test.module.BatchPipeline"}, + ) + + monkeypatch.setattr( + diffusion_engine_module, + "resolve_obj_by_qualname", + lambda qualname: _BatchCapablePipeline if qualname == "test.module.BatchPipeline" else None, + ) + monkeypatch.setattr( + diffusion_engine_module.DiffusionModelRegistry, + "_try_load_model_cls", + lambda model_class_name: _SingleRequestPipeline, + ) + + assert diffusion_engine_module.supports_request_batch(od_config) is True + + def test_supports_request_batch_uses_only_explicit_pipeline_attribute_for_custom_override( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + od_config = SimpleNamespace( + model_class_name="BatchPipeline", + custom_pipeline_args={"pipeline_class": _SingleRequestOverridePipeline}, + ) + monkeypatch.setattr( + diffusion_engine_module.DiffusionModelRegistry, + "_try_load_model_cls", + lambda model_class_name: None, + ) + + assert diffusion_engine_module.supports_request_batch(od_config) is True + + def test_supports_request_batch_honors_explicit_false_on_custom_override( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + class _ExplicitlyUnsupportedOverride(_BatchCapablePipeline): + supports_request_batch = False + + def forward(self, req, prompt_ids=None): + return DiffusionOutput(output=None) + + od_config = SimpleNamespace( + model_class_name="BatchPipeline", + custom_pipeline_args={"pipeline_class": _ExplicitlyUnsupportedOverride}, + ) + monkeypatch.setattr( + diffusion_engine_module.DiffusionModelRegistry, + "_try_load_model_cls", + lambda model_class_name: None, + ) + + assert diffusion_engine_module.supports_request_batch(od_config) is False + + def test_supports_request_batch_rejects_invalid_custom_pipeline_class_name( + self, + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, + ) -> None: + od_config = SimpleNamespace( + model_class_name="BatchPipeline", + custom_pipeline_args={"pipeline_class": "test.module.MissingPipeline"}, + ) + + def fail_resolve(qualname): + raise ImportError(qualname) + + monkeypatch.setattr(diffusion_engine_module, "resolve_obj_by_qualname", fail_resolve) + registry_load = mocker.Mock(return_value=_BatchCapablePipeline) + monkeypatch.setattr( + diffusion_engine_module.DiffusionModelRegistry, + "_try_load_model_cls", + registry_load, + ) + + with pytest.raises(ValueError, match="Failed to resolve custom diffusion pipeline class"): + diffusion_engine_module.supports_request_batch(od_config) + registry_load.assert_not_called() + + def test_engine_disables_batch_dispatch_for_single_request_pipeline( + self, + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, + ) -> None: + od_config = SimpleNamespace( + model_class_name="SinglePipeline", + custom_pipeline_args=None, + streaming_output=False, + ) + fake_executor = SimpleNamespace( + execute_request=mocker.Mock(return_value="per-request"), + execute_batch=mocker.Mock(return_value="batch"), + execute_step=mocker.Mock(return_value="step"), + ) + fake_executor_cls = mocker.Mock(return_value=fake_executor) + + monkeypatch.setattr( + "vllm_omni.diffusion.diffusion_engine.get_diffusion_post_process_func", + lambda *args, **kwargs: None, + ) + monkeypatch.setattr( + "vllm_omni.diffusion.diffusion_engine.get_diffusion_pre_process_func", + lambda *args, **kwargs: None, + ) + monkeypatch.setattr( + "vllm_omni.diffusion.diffusion_engine.DiffusionExecutor.get_class", + lambda *args, **kwargs: fake_executor_cls, + ) + monkeypatch.setattr( + diffusion_engine_module.DiffusionModelRegistry, + "_try_load_model_cls", + lambda model_class_name: _SingleRequestPipeline, + ) + monkeypatch.setattr(DiffusionEngine, "_dummy_run", lambda self: None) + + engine = DiffusionEngine(od_config) + output = engine.execute_fn(_make_request_mode_sched_output("req-a", "req-b")) + + assert engine.supports_request_batch is False + assert output == "per-request" + fake_executor.execute_request.assert_called_once() + fake_executor.execute_batch.assert_not_called() + + @pytest.mark.parametrize("request_ids", [("req-a",), ("req-a", "req-b")]) + def test_engine_enables_batch_dispatch_for_request_batch_pipeline( + self, + request_ids: tuple[str, ...], + monkeypatch: pytest.MonkeyPatch, + mocker: MockerFixture, + ) -> None: + od_config = SimpleNamespace( + model_class_name="BatchPipeline", + custom_pipeline_args=None, + streaming_output=False, + ) + fake_executor = SimpleNamespace( + execute_request=mocker.Mock(return_value="per-request"), + execute_batch=mocker.Mock(return_value="batch"), + execute_step=mocker.Mock(return_value="step"), + ) + fake_executor_cls = mocker.Mock(return_value=fake_executor) + + monkeypatch.setattr( + "vllm_omni.diffusion.diffusion_engine.get_diffusion_post_process_func", + lambda *args, **kwargs: None, + ) + monkeypatch.setattr( + "vllm_omni.diffusion.diffusion_engine.get_diffusion_pre_process_func", + lambda *args, **kwargs: None, + ) + monkeypatch.setattr( + "vllm_omni.diffusion.diffusion_engine.DiffusionExecutor.get_class", + lambda *args, **kwargs: fake_executor_cls, + ) + monkeypatch.setattr( + diffusion_engine_module.DiffusionModelRegistry, + "_try_load_model_cls", + lambda model_class_name: _BatchCapablePipeline, + ) + monkeypatch.setattr(DiffusionEngine, "_dummy_run", lambda self: None) + + engine = DiffusionEngine(od_config) + output = engine.execute_fn(_make_request_mode_sched_output(*request_ids)) + + assert engine.supports_request_batch is True + assert output == "batch" + fake_executor.execute_batch.assert_called_once() + fake_executor.execute_request.assert_not_called() + + +class TestRequestBatchAdmission: + pytestmark = [pytest.mark.core_model, pytest.mark.diffusion, pytest.mark.cpu] + + def test_config_rejects_negative_request_batch_max_wait_ms(self) -> None: + with pytest.raises(ValueError, match="request_batch_max_wait_ms"): + OmniDiffusionConfig(model="test", request_batch_max_wait_ms=-1.0) + + def test_config_normalizes_request_batch_max_wait_ms_to_float(self) -> None: + config = OmniDiffusionConfig(model="test", request_batch_max_wait_ms=5) + + assert config.request_batch_max_wait_ms == 5.0 + assert isinstance(config.request_batch_max_wait_ms, float) + + def test_scheduler_exposes_waiting_and_running_counts(self) -> None: + from vllm_omni.diffusion.sched import RequestScheduler + + od_config = SimpleNamespace(max_num_seqs=4) + scheduler = RequestScheduler() + scheduler.initialize(od_config) + + assert scheduler.num_waiting_requests() == 0 + assert scheduler.num_running_requests() == 0 + + scheduler.add_request( + OmniDiffusionRequest( + prompt="prompt_a", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), + request_id="req-a", + ) + ) + scheduler.add_request( + OmniDiffusionRequest( + prompt="prompt_b", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), + request_id="req-b", + ) + ) + assert scheduler.num_waiting_requests() == 2 + assert scheduler.num_running_requests() == 0 + + scheduler.schedule() + assert scheduler.num_waiting_requests() == 0 + assert scheduler.num_running_requests() == 2 + + def test_request_batch_admission_exits_early_when_waiting_queue_stable(self) -> None: + from vllm_omni.diffusion.sched import RequestScheduler + + od_config = SimpleNamespace( + max_num_seqs=32, + request_batch_max_wait_ms=1000.0, + step_execution=False, + ) + scheduler = RequestScheduler() + scheduler.initialize(od_config) + for idx in range(2): + scheduler.add_request( + OmniDiffusionRequest( + prompt=f"prompt_{idx}", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), + request_id=f"req-{idx}", + ) + ) + + engine = object.__new__(DiffusionEngine) + engine.od_config = od_config + engine.scheduler = scheduler + engine.step_execution = False + engine.supports_request_batch = True + engine.stop_event = threading.Event() + engine._rpc_lock = threading.RLock() + engine._cv = threading.Condition(engine._rpc_lock) + + start = time.monotonic() + with engine._cv: + engine._wait_for_request_batch_admission_locked() + waited_s = time.monotonic() - start + + # Stable-window exit (~50ms), not the full 1000ms deadline. + assert waited_s < 0.5 + assert waited_s >= 0.04 + assert scheduler.num_waiting_requests() == 2 + assert scheduler.num_running_requests() == 0 + + @pytest.mark.core_model @pytest.mark.diffusion @pytest.mark.cpu @@ -142,8 +472,6 @@ def test_move_tensor_tree_moves_nested_cuda_tensors_to_cpu() -> None: @pytest.mark.asyncio async def test_async_add_req_and_wait_for_response(): - from vllm_omni.diffusion.diffusion_engine import DiffusionEngine - engine = object.__new__(DiffusionEngine) engine.scheduler = MockScheduler() engine._out_queue = {} @@ -153,8 +481,10 @@ async def test_async_add_req_and_wait_for_response(): engine._cv = threading.Condition(engine._rpc_lock) engine._init_lock = asyncio.Lock() engine._closed = False + engine.od_config = SimpleNamespace(streaming_output=False) engine._loop_started = False engine.main_loop = None + engine.supports_request_batch = False engine._finalize_finished_request = lambda rid, out, err: out.result diff --git a/tests/diffusion/test_diffusion_engine_cleanup.py b/tests/diffusion/test_diffusion_engine_cleanup.py index f940c11d62e..3cc3d48f7a5 100644 --- a/tests/diffusion/test_diffusion_engine_cleanup.py +++ b/tests/diffusion/test_diffusion_engine_cleanup.py @@ -20,7 +20,7 @@ def _make_request(request_id: str) -> OmniDiffusionRequest: return OmniDiffusionRequest( - prompts=[f"prompt_{request_id}"], + prompt=f"prompt_{request_id}", sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), request_id=request_id, ) diff --git a/tests/diffusion/test_diffusion_engine_rpc_routing.py b/tests/diffusion/test_diffusion_engine_rpc_routing.py index 852f6dc2ee2..4b7d0efbb1b 100644 --- a/tests/diffusion/test_diffusion_engine_rpc_routing.py +++ b/tests/diffusion/test_diffusion_engine_rpc_routing.py @@ -32,12 +32,12 @@ from vllm_omni.diffusion.data import DiffusionOutput from vllm_omni.diffusion.diffusion_engine import DiffusionEngine, _RpcTask from vllm_omni.diffusion.sched import RequestScheduler -from vllm_omni.diffusion.sched.interface import SamplingParamsKey +from vllm_omni.diffusion.sched.interface import RequestBatchSamplingParamsKey from vllm_omni.diffusion.worker.utils import RunnerOutput # Default values for every batch-key field, so SimpleNamespace-based # sampling_params satisfy ``get_sampling_params_key``'s attribute lookups. -_SAMPLING_KEY_DEFAULTS = {f.name: f.default for f in _dc_fields(SamplingParamsKey)} +_SAMPLING_KEY_DEFAULTS = {f.name: f.default for f in _dc_fields(RequestBatchSamplingParamsKey)} pytestmark = [pytest.mark.diffusion, pytest.mark.cpu] @@ -59,6 +59,7 @@ def __init__(self, rpc_delay: float = 0.0): self.rpc_delay = rpc_delay self.is_failed = False self._closed = False + self.od_config = SimpleNamespace() def collective_rpc( self, @@ -101,14 +102,15 @@ def execute_request(self, scheduler_output) -> RunnerOutput: # Mimic the real MultiprocDiffusionExecutor.execute_request: it # forwards a single request through collective_rpc. new_req = scheduler_output.scheduled_new_reqs[0] + req = new_req.req result = self.collective_rpc( "execute_model", - args=(new_req.req,), + args=(req, self.od_config), unique_reply_rank=0, exec_all_ranks=True, ) return RunnerOutput( - request_id=new_req.request_id, + request_id=req.request_id, step_index=None, finished=True, result=result, @@ -119,10 +121,12 @@ def shutdown(self) -> None: def _make_request(tag: str): + sampling_params = dict(_SAMPLING_KEY_DEFAULTS) + sampling_params["num_inference_steps"] = 1 return SimpleNamespace( request_id=tag, - prompts=[f"prompt_{tag}"], - sampling_params=SimpleNamespace(num_inference_steps=1, **_SAMPLING_KEY_DEFAULTS), + prompt=f"prompt_{tag}", + sampling_params=SimpleNamespace(**sampling_params), ) @@ -137,17 +141,20 @@ def _make_engine_with_loop( """ engine = DiffusionEngine.__new__(DiffusionEngine) engine._closed = False + engine.od_config = SimpleNamespace(streaming_output=False) engine.executor = _ConcurrencyTrackingExecutor(rpc_delay=rpc_delay) sched = RequestScheduler() sched.initialize(SimpleNamespace(max_num_seqs=1)) engine.scheduler = sched engine.step_execution = False + engine.supports_request_batch = False engine.execute_fn = engine.executor.execute_request engine._rpc_lock = threading.RLock() engine._cv = threading.Condition(engine._rpc_lock) engine._out_queue = {} + engine._out_queue_streaming = {} engine._closed = False engine.abort_queue = queue.Queue() engine._rpc_queue = queue.Queue() diff --git a/tests/diffusion/test_diffusion_ipc.py b/tests/diffusion/test_diffusion_ipc.py index 57c653ef104..f7cc39d9430 100644 --- a/tests/diffusion/test_diffusion_ipc.py +++ b/tests/diffusion/test_diffusion_ipc.py @@ -15,6 +15,7 @@ pack_diffusion_output_shm, unpack_diffusion_output_shm, ) +from vllm_omni.diffusion.worker.utils import BatchRunnerOutput, RunnerOutput pytestmark = [pytest.mark.core_model, pytest.mark.diffusion, pytest.mark.cpu] @@ -107,6 +108,30 @@ def test_rpc_result_envelope_diffusion_output_round_trips_through_shm() -> None: assert unpacked["rank_statuses"] == [{"rank": 0, "ok": True}] +def test_batch_runner_output_round_trips_nested_results_through_shm() -> None: + first = torch.arange(_large_numel(torch.float32), dtype=torch.float32) + second = torch.arange(_large_numel(torch.float32), dtype=torch.float32) + 1 + output = BatchRunnerOutput.from_list( + [ + RunnerOutput(request_id="req-0", finished=True, result=DiffusionOutput(output=first)), + RunnerOutput(request_id="req-1", finished=True, result=DiffusionOutput(output={"image": second})), + RunnerOutput(request_id="req-error", finished=True, result=DiffusionOutput(error="boom")), + ] + ) + + pack_diffusion_output_shm(output) + + assert output.runner_outputs[0].result.output["__tensor_shm__"] is True + assert output.runner_outputs[1].result.output["image"]["__tensor_shm__"] is True + assert output.runner_outputs[2].result.error == "boom" + + unpack_diffusion_output_shm(output) + + torch.testing.assert_close(output["req-0"].result.output, first) + torch.testing.assert_close(output["req-1"].result.output["image"], second) + assert output["req-error"].result.error == "boom" + + def test_pack_value_keeps_tensor_at_threshold_inline() -> None: tensor = torch.arange( _SHM_TENSOR_THRESHOLD // torch.empty((), dtype=torch.float32).element_size(), diff --git a/tests/diffusion/test_diffusion_model_runner.py b/tests/diffusion/test_diffusion_model_runner.py index 4e5cb0689c3..55584609dd3 100644 --- a/tests/diffusion/test_diffusion_model_runner.py +++ b/tests/diffusion/test_diffusion_model_runner.py @@ -11,6 +11,7 @@ from tests.helpers.mark import hardware_test from vllm_omni.diffusion.data import DiffusionOutput from vllm_omni.diffusion.worker.diffusion_model_runner import DiffusionModelRunner +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch, split_diffusion_output_by_request pytestmark = [pytest.mark.diffusion] @@ -22,14 +23,39 @@ def _noop_forward_context(*args, **kwargs): class _DummyPipeline: + supports_request_batch = True + def __init__(self, output): self._output = output self.forward_calls = 0 + self.last_req = None def forward(self, req): - del req + self.last_req = req self.forward_calls += 1 - return self._output + return [self._output] + + +class _SingleRequestBatchPipeline: + supports_request_batch = False + + def __init__(self): + self.last_req = None + + def forward(self, req): + self.last_req = req + return [DiffusionOutput(output=req.prompts[0])] + + +class _SingleRequestDiffusionOutputPipeline: + supports_request_batch = False + + def __init__(self): + self.last_req = None + + def forward(self, req): + self.last_req = req + return DiffusionOutput(output=req.prompts[0]) class _ChunkStepPipeline: @@ -81,9 +107,28 @@ def _make_request(skip_cache_refresh: bool = True): ) return SimpleNamespace( request_id="req-test", - prompts=["a prompt"], + prompt="a prompt", sampling_params=sampling_params, skip_cache_refresh=skip_cache_refresh, + kv_sender_info=None, + ) + + +def _make_request_with_params(req_id: str, sampling_params): + return SimpleNamespace( + request_id=req_id, + prompt=f"prompt-{req_id}", + prompts=[f"prompt-{req_id}"], + sampling_params=sampling_params, + skip_cache_refresh=True, + ) + + +def _fake_platform_for_peak_memory(): + return SimpleNamespace( + reset_peak_memory_stats=lambda: None, + max_memory_reserved=lambda: 0, + max_memory_allocated=lambda: 0, ) @@ -91,7 +136,7 @@ def _make_runner(cache_backend, cache_backend_name: str, enable_cache_dit_summar runner = object.__new__(DiffusionModelRunner) runner.vllm_config = object() runner.device = torch.device("cpu") - runner.pipeline = _DummyPipeline(output=SimpleNamespace(output="ok")) + runner.pipeline = _DummyPipeline(output=DiffusionOutput(output="ok")) runner.cache_backend = cache_backend runner.offload_backend = None runner.state_cache = {} @@ -238,6 +283,232 @@ def is_enabled(self): assert cache_summary_calls == [(runner.pipeline, True)] +@pytest.mark.core_model +@pytest.mark.cpu +def test_execute_model_passes_single_request_batch_to_non_admission_pipeline(monkeypatch): + runner = _make_runner(cache_backend=None, cache_backend_name="none") + runner.pipeline = _SingleRequestBatchPipeline() + req = _make_request(skip_cache_refresh=True) + + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + + output = DiffusionModelRunner.execute_model(runner, req) + + assert output.output == "a prompt" + assert isinstance(runner.pipeline.last_req, DiffusionRequestBatch) + assert runner.pipeline.last_req.num_reqs == 1 + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_execute_model_accepts_bare_diffusion_output_from_single_request_pipeline(monkeypatch): + runner = _make_runner(cache_backend=None, cache_backend_name="none") + runner.pipeline = _SingleRequestDiffusionOutputPipeline() + req = _make_request(skip_cache_refresh=True) + + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + + output = DiffusionModelRunner.execute_model(runner, req) + + assert output.output == "a prompt" + assert isinstance(runner.pipeline.last_req, DiffusionRequestBatch) + assert runner.pipeline.last_req.num_reqs == 1 + + +class _BatchPipeline: + """Pipeline returning a configurable list of outputs from forward().""" + + supports_request_batch = True + + def __init__(self, outputs): + self._outputs = outputs + self.last_batch = None + + def forward(self, batch): + self.last_batch = batch + return list(self._outputs) + + +class _SingleRequestPipeline: + def forward(self, batch): + del batch + return [DiffusionOutput(output="single")] + + +class _BatchSingleOutputPipeline: + supports_request_batch = True + + def forward(self, batch): + return DiffusionOutput(output=batch.prompts[0]) + + +def _make_batch_runner(pipeline): + runner = object.__new__(DiffusionModelRunner) + runner.vllm_config = object() + runner.device = torch.device("cpu") + runner.pipeline = pipeline + runner.cache_backend = None + runner.offload_backend = None + runner.od_config = SimpleNamespace( + cache_backend="none", + enable_cache_dit_summary=False, + parallel_config=SimpleNamespace(use_hsdp=False), + ) + runner.kv_transfer_manager = SimpleNamespace( + receive_multi_kv_cache_distributed=lambda req, cfg_kv_collect_func=None, target_device=None: None, + ) + return runner + + +def _make_scheduler_output(num_reqs: int): + reqs = [_make_request() for _ in range(num_reqs)] + for i, req in enumerate(reqs): + req.request_id = f"req-{i}" + return SimpleNamespace(scheduled_new_reqs=[SimpleNamespace(req=req) for req in reqs]) + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_execute_model_batch_rejects_output_count_mismatch(monkeypatch): + """A pipeline returning the wrong number of outputs must fail loudly, + not silently drop requests or IndexError on the per-request mapping.""" + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + monkeypatch.setattr(model_runner_module, "current_omni_platform", _fake_platform_for_peak_memory()) + # forward returns 1 output for 2 scheduled requests + runner = _make_batch_runner(_BatchPipeline(outputs=[DiffusionOutput(output="only-one")])) + sched = _make_scheduler_output(num_reqs=2) + + with pytest.raises(RuntimeError, match="returned 1 outputs for 2 requests"): + DiffusionModelRunner.execute_model_batch(runner, sched, runner.od_config) + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_execute_model_batch_routes_one_output_per_request(monkeypatch): + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + monkeypatch.setattr(model_runner_module, "current_omni_platform", _fake_platform_for_peak_memory()) + outs = [DiffusionOutput(output="a"), DiffusionOutput(output="b")] + runner = _make_batch_runner(_BatchPipeline(outputs=outs)) + sched = _make_scheduler_output(num_reqs=2) + + result = DiffusionModelRunner.execute_model_batch(runner, sched, runner.od_config) + + assert len(result.runner_outputs) == 2 + assert result.runner_outputs[0].request_id == "req-0" + assert result.runner_outputs[0].result.output == "a" + assert result.runner_outputs[1].request_id == "req-1" + assert result.runner_outputs[1].result.output == "b" + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_execute_model_batch_preserves_per_request_sampling_and_seeds_generators(monkeypatch): + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + monkeypatch.setattr(model_runner_module, "current_omni_platform", _fake_platform_for_peak_memory()) + outputs = [DiffusionOutput(output="a"), DiffusionOutput(output="b")] + pipeline = _BatchPipeline(outputs=outputs) + runner = _make_batch_runner(pipeline) + sched = _make_scheduler_output(num_reqs=2) + sched.scheduled_new_reqs[0].req.sampling_params.seed = 111 + sched.scheduled_new_reqs[1].req.sampling_params.seed = 222 + + DiffusionModelRunner.execute_model_batch(runner, sched, runner.od_config) + + assert pipeline.last_batch is not None + assert [sp.seed for sp in pipeline.last_batch.sampling_params_list] == [111, 222] + assert [sp.generator.initial_seed() for sp in pipeline.last_batch.sampling_params_list] == [111, 222] + with pytest.raises(AssertionError, match="multiple requests"): + _ = pipeline.last_batch.sampling_params + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_execute_model_batch_rejects_pipeline_without_request_batch_support(monkeypatch): + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + monkeypatch.setattr(model_runner_module, "current_omni_platform", _fake_platform_for_peak_memory()) + runner = _make_batch_runner(_SingleRequestPipeline()) + sched = _make_scheduler_output(num_reqs=2) + + with pytest.raises(RuntimeError, match="does not support request-batch forward"): + DiffusionModelRunner.execute_model_batch(runner, sched, runner.od_config) + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_execute_model_batch_rejects_single_diffusion_output(monkeypatch): + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + monkeypatch.setattr(model_runner_module, "current_omni_platform", _fake_platform_for_peak_memory()) + runner = _make_batch_runner(_BatchSingleOutputPipeline()) + sched = _make_scheduler_output(num_reqs=1) + + with pytest.raises(RuntimeError, match="request-batch forward must return list\\[DiffusionOutput\\]"): + DiffusionModelRunner.execute_model_batch(runner, sched, runner.od_config) + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_execute_model_batch_uses_runner_output_helper(monkeypatch): + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + monkeypatch.setattr(model_runner_module, "current_omni_platform", _fake_platform_for_peak_memory()) + outs = [DiffusionOutput(output="a"), DiffusionOutput(output="b")] + runner = _make_batch_runner(_BatchPipeline(outputs=outs)) + sched = _make_scheduler_output(num_reqs=2) + helper_calls = [] + original_from_outputs = DiffusionModelRunner._runner_output_from_outputs + + def _recording_from_outputs(self, reqs, outputs): + helper_calls.append(([req.request_id for req in reqs], [output.output for output in outputs])) + return original_from_outputs(self, reqs, outputs) + + monkeypatch.setattr(DiffusionModelRunner, "_runner_output_from_outputs", _recording_from_outputs) + + result = DiffusionModelRunner.execute_model_batch(runner, sched, runner.od_config) + + assert [runner_output.result.output for runner_output in result.runner_outputs] == ["a", "b"] + assert helper_calls == [(["req-0", "req-1"], ["a", "b"])] + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_split_diffusion_output_by_request_slices_single_and_multi_request_outputs(): + reqs = [_make_request(), _make_request()] + reqs[0].request_id = "req-0" + reqs[1].request_id = "req-1" + batch = DiffusionRequestBatch(requests=reqs) + result = DiffusionOutput(output=["img-0a", "img-0b", "img-1a", "img-1b"], stage_durations={"decode": 1.0}) + + outputs = split_diffusion_output_by_request(result, batch, num_outputs_per_prompt=2) + + assert [output.output for output in outputs] == [["img-0a", "img-0b"], ["img-1a", "img-1b"]] + assert [output.stage_durations for output in outputs] == [{"decode": 1.0}, {"decode": 1.0}] + + single = split_diffusion_output_by_request( + result, DiffusionRequestBatch(requests=reqs[:1]), num_outputs_per_prompt=2 + ) + + assert single[0].output == ["img-0a", "img-0b"] + + +@pytest.mark.core_model +@pytest.mark.cpu +def test_split_diffusion_output_by_request_slices_tuple_outputs(): + reqs = [_make_request(), _make_request()] + batch = DiffusionRequestBatch(requests=reqs) + result = DiffusionOutput( + output=( + ["video-0", "video-1"], + torch.tensor([10, 20]), + ) + ) + + outputs = split_diffusion_output_by_request(result, batch, num_outputs_per_prompt=1) + + assert outputs[0].output[0] == ["video-0"] + assert torch.equal(outputs[0].output[1], torch.tensor([10])) + assert outputs[1].output[0] == ["video-1"] + assert torch.equal(outputs[1].output[1], torch.tensor([20])) + + @pytest.mark.core_model @pytest.mark.cpu def test_load_model_clears_cache_backend_for_unsupported_pipeline(monkeypatch): diff --git a/tests/diffusion/test_diffusion_output_formatter.py b/tests/diffusion/test_diffusion_output_formatter.py index 49c07d98f40..185d51b54e3 100644 --- a/tests/diffusion/test_diffusion_output_formatter.py +++ b/tests/diffusion/test_diffusion_output_formatter.py @@ -23,12 +23,12 @@ def _request( - prompts: list[str] | None = None, + prompt: str | dict | None = None, *, num_outputs_per_prompt: int = 1, ) -> OmniDiffusionRequest: return OmniDiffusionRequest( - prompts=prompts or ["prompt"], + prompt=prompt or "prompt", request_id="req-1", sampling_params=OmniDiffusionSamplingParams( num_inference_steps=1, @@ -69,7 +69,7 @@ def test_formatter_preserves_single_video_audio_actions_and_custom_output( ) results = format_diffusion_outputs( - request=_request(["prompt-0"]), + request=_request("prompt-0"), od_config=_config(), diffusion_output=DiffusionOutput( output=None, @@ -118,7 +118,7 @@ def test_formatter_preserves_text_custom_output(monkeypatch: pytest.MonkeyPatch) ) [result] = format_diffusion_outputs( - request=_request(["describe this"]), + request=_request("describe this"), od_config=_config(), diffusion_output=DiffusionOutput(output="caption"), output_data="caption", @@ -147,7 +147,7 @@ class AudioModel: postprocess_output = normalize_diffusion_postprocess_output(["waveform"], {}) [result] = format_diffusion_outputs( - request=_request(["speak"]), + request=_request("speak"), od_config=_config("audio_model"), diffusion_output=DiffusionOutput(output=["waveform"]), output_data=["waveform"], @@ -180,7 +180,7 @@ def test_formatter_preserves_audio_model_video_audio_and_actions( ) [result] = format_diffusion_outputs( - request=_request(["watch and listen"]), + request=_request("watch and listen"), od_config=_config("audio_video_model"), diffusion_output=DiffusionOutput(output=None), output_data={"raw": "output"}, @@ -212,7 +212,7 @@ def test_formatter_preserves_audio_only_postprocess_dict( ) [result] = format_diffusion_outputs( - request=_request(["speak"]), + request=_request("speak"), od_config=_config("audio_model"), diffusion_output=DiffusionOutput(output=None), output_data={"raw": "output"}, @@ -241,7 +241,7 @@ def test_formatter_preserves_single_prompt_multiple_audio_outputs( postprocess_output = normalize_diffusion_postprocess_output(["waveform-0", "waveform-1"], {}) [result] = format_diffusion_outputs( - request=_request(["speak"], num_outputs_per_prompt=2), + request=_request("speak", num_outputs_per_prompt=2), od_config=_config("audio_model"), diffusion_output=DiffusionOutput(output=["waveform-0", "waveform-1"]), output_data=["waveform-0", "waveform-1"], @@ -255,7 +255,7 @@ def test_formatter_preserves_single_prompt_multiple_audio_outputs( assert result.multimodal_output == {"audio": ["waveform-0", "waveform-1"]} -def test_formatter_preserves_multi_prompt_audio_and_action_slicing( +def test_formatter_preserves_single_prompt_audio_and_action_payloads( monkeypatch: pytest.MonkeyPatch, ) -> None: monkeypatch.setattr(output_formatter, "supports_audio_output", lambda _: False) @@ -271,7 +271,7 @@ def test_formatter_preserves_multi_prompt_audio_and_action_slicing( ) results = format_diffusion_outputs( - request=_request(["prompt-0", "prompt-1"]), + request=_request("prompt-0"), od_config=_config(), diffusion_output=DiffusionOutput(output=None), output_data={"raw": "output"}, @@ -279,29 +279,22 @@ def test_formatter_preserves_multi_prompt_audio_and_action_slicing( timings=_timings(), ) - assert len(results) == 2 - assert results[0].images == ["frame-0"] - assert results[1].images == ["frame-1"] + assert len(results) == 1 + assert results[0].images == ["frame-0", "frame-1"] + assert results[0].prompt == "prompt-0" assert results[0].custom_output == {"shared": True} - assert results[1].custom_output == {"shared": True} - assert results[0].multimodal_output["audio"] == "audio-0" - assert results[1].multimodal_output["audio"] == "audio-1" + assert results[0].multimodal_output["audio"] == ["audio-0", "audio-1"] assert results[0].multimodal_output["fps"] == 12.5 - assert results[1].multimodal_output["fps"] == 12.5 torch.testing.assert_close( results[0].multimodal_output["actions"], - torch.tensor([1.0, 2.0]), - ) - torch.testing.assert_close( - results[1].multimodal_output["actions"], - torch.tensor([3.0, 4.0]), + torch.tensor([[1.0, 2.0], [3.0, 4.0]]), ) def test_format_empty_diffusion_outputs_preserves_empty_response_shape() -> None: - results = format_empty_diffusion_outputs(_request(["prompt-0", "prompt-1"])) + results = format_empty_diffusion_outputs(_request("prompt-0")) - assert len(results) == 2 - assert [result.prompt for result in results] == ["prompt-0", "prompt-1"] - assert [result.images for result in results] == [[], []] - assert [result.metrics for result in results] == [{}, {}] + assert len(results) == 1 + assert [result.prompt for result in results] == ["prompt-0"] + assert [result.images for result in results] == [[]] + assert [result.metrics for result in results] == [{}] diff --git a/tests/diffusion/test_diffusion_request.py b/tests/diffusion/test_diffusion_request.py index 9087a6e2f1c..ffaab3c530f 100644 --- a/tests/diffusion/test_diffusion_request.py +++ b/tests/diffusion/test_diffusion_request.py @@ -13,7 +13,7 @@ def _make_request() -> OmniDiffusionRequest: return OmniDiffusionRequest( - prompts=[{"prompt": "a cup of coffee on a table"}], + prompt={"prompt": "a cup of coffee on a table"}, sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), request_id="request-test", ) @@ -22,7 +22,7 @@ def _make_request() -> OmniDiffusionRequest: def test_request_id_is_required(): with pytest.raises(TypeError, match="request_id"): OmniDiffusionRequest( - prompts=[{"prompt": "a cup of coffee on a table"}], + prompt={"prompt": "a cup of coffee on a table"}, sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), ) @@ -37,7 +37,7 @@ def test_request_ids_identity_list_is_removed(): def test_request_id_must_be_non_empty(): with pytest.raises(ValueError, match="request_id must be a non-empty string"): OmniDiffusionRequest( - prompts=[{"prompt": "a cup of coffee on a table"}], + prompt={"prompt": "a cup of coffee on a table"}, sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), request_id="", ) @@ -45,7 +45,7 @@ def test_request_id_must_be_non_empty(): def test_dummy_run_request_is_identified_by_reserved_request_id(): req = OmniDiffusionRequest( - prompts=[{"prompt": "dummy run"}], + prompt={"prompt": "dummy run"}, sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), request_id=DUMMY_DIFFUSION_REQUEST_ID, ) diff --git a/tests/diffusion/test_diffusion_scheduler.py b/tests/diffusion/test_diffusion_scheduler.py index 280b5642038..7d6d2d11bff 100644 --- a/tests/diffusion/test_diffusion_scheduler.py +++ b/tests/diffusion/test_diffusion_scheduler.py @@ -29,7 +29,7 @@ def _make_request(req_id: str) -> OmniDiffusionRequest: return OmniDiffusionRequest( - prompts=[f"prompt_{req_id}"], + prompt=f"prompt_{req_id}", sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), request_id=req_id, ) @@ -67,7 +67,7 @@ def _make_step_request( sampling_params: OmniDiffusionSamplingParams | None = None, ) -> OmniDiffusionRequest: return OmniDiffusionRequest( - prompts=[f"prompt_{req_id}"], + prompt=f"prompt_{req_id}", sampling_params=sampling_params or OmniDiffusionSamplingParams( num_inference_steps=num_inference_steps, @@ -128,6 +128,12 @@ def update_from_output(self, sched_output, output) -> set[str]: def has_requests(self) -> bool: return not self._scheduled + def num_waiting_requests(self) -> int: + return 0 if self._scheduled else 1 + + def num_running_requests(self) -> int: + return 1 if self._scheduled else 0 + def get_request_state(self, request_id: str): del request_id return self._state @@ -164,7 +170,7 @@ def _make(lora_int_id: int | None = None, lora_scale: float = 1.0) -> OmniDiffus ) sp.lora_scale = lora_scale return OmniDiffusionRequest( - prompts=["prompt"], + prompt="prompt", sampling_params=sp, request_id=f"req-{lora_int_id}-{lora_scale}", ) @@ -196,6 +202,41 @@ def test_equal_for_same_lora_identity(self) -> None: assert a == b +class TestGetRequestBatchSamplingParamsKey: + """Pure-function tests for the request-batch compatibility key builder.""" + + @staticmethod + def _make( + *, + num_inference_steps: int = 2, + seed: int | None = 123, + generator: torch.Generator | None = None, + ) -> OmniDiffusionRequest: + sp = OmniDiffusionSamplingParams( + num_inference_steps=num_inference_steps, + seed=seed, + generator=generator, + ) + return OmniDiffusionRequest(prompt="prompt", sampling_params=sp, request_id=f"req-{num_inference_steps}") + + def test_distinguishes_num_inference_steps(self) -> None: + from vllm_omni.diffusion.sched.base_scheduler import get_request_batch_sampling_params_key + + assert get_request_batch_sampling_params_key( + self._make(num_inference_steps=2) + ) != get_request_batch_sampling_params_key(self._make(num_inference_steps=4)) + + def test_ignores_seed_and_generator(self) -> None: + from vllm_omni.diffusion.sched.base_scheduler import get_request_batch_sampling_params_key + + gen_a = torch.Generator(device="cpu").manual_seed(1) + gen_b = torch.Generator(device="cpu").manual_seed(2) + + assert get_request_batch_sampling_params_key( + self._make(seed=1, generator=gen_a) + ) == get_request_batch_sampling_params_key(self._make(seed=2, generator=gen_b)) + + class TestRequestScheduler: def setup_method(self) -> None: self.scheduler: RequestScheduler = RequestScheduler() @@ -297,8 +338,18 @@ def test_batches_compatible_requests_up_to_max_num_seqs(self) -> None: scheduler = RequestScheduler() scheduler.initialize(SimpleNamespace(max_num_seqs=2)) - req_id_a = scheduler.add_request(_make_request("a")) - req_id_b = scheduler.add_request(_make_request("b")) + req_id_a = scheduler.add_request( + _make_step_request( + "a", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1, seed=123), + ) + ) + req_id_b = scheduler.add_request( + _make_step_request( + "b", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1, seed=123), + ) + ) sched_output = scheduler.schedule() @@ -306,6 +357,50 @@ def test_batches_compatible_requests_up_to_max_num_seqs(self) -> None: assert sched_output.num_running_reqs == 2 assert sched_output.num_waiting_reqs == 0 + def test_batches_incompatible_request_sampling_params_separately(self) -> None: + scheduler = RequestScheduler() + scheduler.initialize(SimpleNamespace(max_num_seqs=2)) + + req_id_a = scheduler.add_request( + _make_step_request( + "a", num_inference_steps=2, sampling_params=OmniDiffusionSamplingParams(num_inference_steps=2, seed=123) + ) + ) + scheduler.add_request( + _make_step_request( + "b", num_inference_steps=4, sampling_params=OmniDiffusionSamplingParams(num_inference_steps=4, seed=123) + ) + ) + + first = scheduler.schedule() + + assert _new_ids(first) == [req_id_a] + assert first.num_running_reqs == 1 + assert first.num_waiting_reqs == 1 + + def test_batches_different_request_local_seed_together(self) -> None: + scheduler = RequestScheduler() + scheduler.initialize(SimpleNamespace(max_num_seqs=2)) + + req_id_a = scheduler.add_request( + _make_step_request( + "a", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=2, seed=123), + ) + ) + req_id_b = scheduler.add_request( + _make_step_request( + "b", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=2, seed=456), + ) + ) + + first = scheduler.schedule() + + assert _new_ids(first) == [req_id_a, req_id_b] + assert first.num_running_reqs == 2 + assert first.num_waiting_reqs == 0 + def test_incompatible_waiting_head_blocks_later_compatible_request(self) -> None: scheduler = RequestScheduler() scheduler.initialize(SimpleNamespace(max_num_seqs=3)) @@ -313,7 +408,7 @@ def test_incompatible_waiting_head_blocks_later_compatible_request(self) -> None req_id_a = scheduler.add_request(_make_request("a")) req_id_b = scheduler.add_request( OmniDiffusionRequest( - prompts=["prompt_b"], + prompt="prompt_b", sampling_params=OmniDiffusionSamplingParams(width=768), request_id="b", ) @@ -374,7 +469,7 @@ def test_has_requests_state_transition(self) -> None: def test_request_id_is_scheduler_key(self) -> None: request = OmniDiffusionRequest( - prompts=["prompt_map_a", "prompt_map_b"], + prompt="prompt_map_a", sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), request_id="map-parent", ) @@ -618,124 +713,6 @@ def test_dummy_run_raises_on_output_error(self, mocker: MockerFixture) -> None: with pytest.raises(RuntimeError, match="Dummy run failed: boom"): engine._dummy_run() - @pytest.mark.asyncio - async def test_step_multi_request_reuses_multimodal_slice_logic(self, mocker: MockerFixture) -> None: - engine = DiffusionEngine.__new__(DiffusionEngine) - engine.od_config = SimpleNamespace( - model_class_name="mock_model", - enable_cpu_offload=False, - ) - engine.pre_process_func = None - engine.post_process_func = None - engine._check_and_start_background_loop = mocker.AsyncMock() - engine.async_add_req_and_wait_for_response = mocker.AsyncMock( - return_value=DiffusionOutput( - output={ - "video": ["frame-0", "frame-1"], - "audio": ["audio-0", "audio-1"], - "actions": torch.tensor([[1.0, 2.0], [3.0, 4.0]]), - } - ) - ) - - request = OmniDiffusionRequest( - prompts=["prompt-0", "prompt-1"], - sampling_params=OmniDiffusionSamplingParams( - num_inference_steps=1, - num_outputs_per_prompt=1, - ), - request_id="req-batch", - ) - - mocker.patch("vllm_omni.diffusion.output_formatter.supports_audio_output", return_value=False) - outputs = await engine.step(request) - - assert len(outputs) == 2 - assert outputs[0].images == ["frame-0"] - assert outputs[1].images == ["frame-1"] - assert outputs[0].multimodal_output["audio"] == "audio-0" - assert outputs[1].multimodal_output["audio"] == "audio-1" - torch.testing.assert_close( - outputs[0].multimodal_output["actions"], - torch.tensor([1.0, 2.0]), - ) - torch.testing.assert_close( - outputs[1].multimodal_output["actions"], - torch.tensor([3.0, 4.0]), - ) - - @pytest.mark.asyncio - async def test_step_empty_dict_output_still_runs_postprocess(self, mocker: MockerFixture) -> None: - engine = DiffusionEngine.__new__(DiffusionEngine) - engine.od_config = SimpleNamespace( - model_class_name="mock_model", - enable_cpu_offload=False, - ) - engine.pre_process_func = None - engine.post_process_func = mocker.Mock(return_value={"video": ["processed"]}) - engine._post_process_accepts_sampling_params = False - engine._check_and_start_background_loop = mocker.AsyncMock() - engine.async_add_req_and_wait_for_response = mocker.AsyncMock( - return_value=DiffusionOutput( - output={}, - custom_output={"actions": torch.tensor([[1.0, 2.0]])}, - ) - ) - - request = OmniDiffusionRequest( - prompts=["prompt"], - sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), - request_id="req-action", - ) - - mocker.patch("vllm_omni.diffusion.diffusion_engine.supports_audio_output", return_value=False) - outputs = await engine.step(request) - - engine.post_process_func.assert_called_once_with({}) - assert outputs[0].images == ["processed"] - torch.testing.assert_close(outputs[0].multimodal_output["actions"], torch.tensor([[1.0, 2.0]])) - - @pytest.mark.asyncio - async def test_step_action_only_flag_skips_postprocess(self, mocker: MockerFixture) -> None: - engine = DiffusionEngine.__new__(DiffusionEngine) - engine.od_config = SimpleNamespace( - model_class_name="mock_model", - enable_cpu_offload=False, - ) - engine.pre_process_func = None - engine.post_process_func = mocker.Mock(side_effect=AssertionError("postprocess should be skipped")) - engine.action_post_process_func = mocker.Mock(return_value=torch.tensor([[3.0, 4.0]])) - engine._post_process_accepts_sampling_params = False - engine._action_post_process_accepts_custom_output = True - engine._action_post_process_accepts_sampling_params = False - engine._check_and_start_background_loop = mocker.AsyncMock() - raw_action = torch.tensor([[1.0, 2.0]]) - engine.async_add_req_and_wait_for_response = mocker.AsyncMock( - return_value=DiffusionOutput( - output={}, - custom_output={ - "action": raw_action, - "action_only_output": True, - }, - ) - ) - - request = OmniDiffusionRequest( - prompts=["prompt"], - sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), - request_id="req-action", - ) - - mocker.patch("vllm_omni.diffusion.diffusion_engine.supports_audio_output", return_value=False) - outputs = await engine.step(request) - - engine.post_process_func.assert_not_called() - engine.action_post_process_func.assert_called_once() - assert engine.action_post_process_func.call_args.args[0] is raw_action - assert "custom_output" in engine.action_post_process_func.call_args.kwargs - assert outputs[0].images == [] - torch.testing.assert_close(outputs[0].multimodal_output["actions"], torch.tensor([[3.0, 4.0]])) - class TestStepScheduler: def setup_method(self) -> None: diff --git a/tests/diffusion/test_diffusion_step_pipeline.py b/tests/diffusion/test_diffusion_step_pipeline.py index 6d26d073d0a..bd2289283e5 100644 --- a/tests/diffusion/test_diffusion_step_pipeline.py +++ b/tests/diffusion/test_diffusion_step_pipeline.py @@ -114,6 +114,23 @@ def post_decode(self, state, **kwargs): raise AssertionError("post_decode should not run after interrupt") +class _FakePeakMemoryPlatform: + def __init__(self, reserved_mb: list[float]): + self._reserved_mb = reserved_mb + self.reset_calls = 0 + + def reset_peak_memory_stats(self): + self.reset_calls += 1 + + def max_memory_reserved(self): + index = min(self.reset_calls - 1, len(self._reserved_mb) - 1) + return int(self._reserved_mb[index] * 1024**2) + + def max_memory_allocated(self): + index = min(self.reset_calls - 1, len(self._reserved_mb) - 1) + return int((self._reserved_mb[index] - 100) * 1024**2) + + class _IdentityNoiseTransformer(torch.nn.Module): def forward(self, x: torch.Tensor, **kwargs): del kwargs @@ -205,15 +222,10 @@ def post_decode(self, state, **kwargs): def _make_step_request(num_inference_steps: int = 2): - return SimpleNamespace( - prompts=["a prompt"], + return OmniDiffusionRequest( + prompt="a prompt", request_id="req-1", - sampling_params=SimpleNamespace( - generator=None, - seed=None, - generator_device=None, - num_inference_steps=num_inference_steps, - ), + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=num_inference_steps), ) @@ -226,7 +238,7 @@ def _assert_aborted_output(output: DiffusionOutput, request_id: str) -> None: def _make_engine_request(req_id: str = "req-1", num_inference_steps: int = 2) -> OmniDiffusionRequest: return OmniDiffusionRequest( - prompts=[f"prompt-{req_id}"], + prompt=f"prompt-{req_id}", sampling_params=OmniDiffusionSamplingParams(num_inference_steps=num_inference_steps), request_id=req_id, ) @@ -282,6 +294,7 @@ def _make_distributed_runner(mode: str, device: torch.device): def _make_scheduler_output(req, request_id="req-1", step_id=0, finished_req_ids=None): + req.request_id = request_id return DiffusionSchedulerOutput( step_id=step_id, scheduled_new_reqs=[NewRequestData(request_id=request_id, req=req)], @@ -309,7 +322,7 @@ def _make_input_batch_state(request_id: str, latent_value: float) -> DiffusionRe state = DiffusionRequestState( request_id=request_id, sampling=SimpleNamespace(), - prompts=None, + prompt=None, ) state.latents = torch.tensor([[latent_value]]) state.timesteps = torch.tensor([1.0]) @@ -456,6 +469,24 @@ def test_completes_request_and_clears_state(self, monkeypatch): assert runner.pipeline.scheduler_calls == 2 assert runner.pipeline.decode_calls == 1 + def test_carries_peak_memory_across_stepwise_request_lifecycle(self, monkeypatch): + runner = _make_runner() + req = _make_step_request() + fake_platform = _FakePeakMemoryPlatform([1500.0, 1200.0]) + monkeypatch.setattr(model_runner_module, "set_forward_context", _noop_forward_context) + monkeypatch.setattr(model_runner_module, "current_omni_platform", fake_platform) + + first = DiffusionModelRunner.execute_stepwise(runner, _make_scheduler_output(req, step_id=0)) + first_output = first.get_request_output("req-1") + assert first_output.finished is False + assert first_output.result is None + + second = DiffusionModelRunner.execute_stepwise(runner, _make_cached_scheduler_output(step_id=1)) + second_output = second.get_request_output("req-1") + assert second_output.finished is True + assert second_output.result is not None + assert second_output.result.peak_memory_mb == pytest.approx(1500.0) + def test_rejects_multi_request_step_batch(self): runner = _make_runner() req_1 = _make_step_request() diff --git a/tests/diffusion/test_multiproc_engine_concurrency.py b/tests/diffusion/test_multiproc_engine_concurrency.py index 5bfb8d73888..070d899e6ee 100644 --- a/tests/diffusion/test_multiproc_engine_concurrency.py +++ b/tests/diffusion/test_multiproc_engine_concurrency.py @@ -18,10 +18,16 @@ from vllm_omni.diffusion.diffusion_engine import DiffusionEngine from vllm_omni.diffusion.executor.multiproc_executor import MultiprocDiffusionExecutor from vllm_omni.diffusion.ipc import DIFFUSION_RPC_RESULT_ENVELOPE +from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.sched import RequestScheduler +from vllm_omni.diffusion.sched.interface import ( + CachedRequestData, + DiffusionSchedulerOutput, + NewRequestData, +) from vllm_omni.diffusion.stage_diffusion_proc import StageDiffusionProc from vllm_omni.diffusion.worker.diffusion_worker import WorkerProc -from vllm_omni.diffusion.worker.utils import RunnerOutput +from vllm_omni.diffusion.worker.utils import BatchRunnerOutput, RunnerOutput from vllm_omni.inputs.data import OmniDiffusionSamplingParams from vllm_omni.outputs import OmniRequestOutput @@ -40,7 +46,7 @@ def _mock_request(tag: str): """Return a lightweight request object identifiable by *tag*.""" return SimpleNamespace( request_id=tag, - prompts=[f"prompt_{tag}"], + prompt=f"prompt_{tag}", sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), ) @@ -101,13 +107,24 @@ def _run(): req = req_q.get(timeout=10) method = req.get("method", "") args = req.get("args", ()) - if method in {"generate", "execute_model"} and args and hasattr(args[0], "request_id"): + if method == "execute_model_batch" and args and isinstance(args[0], DiffusionSchedulerOutput): + sched_output = args[0] + runner_outputs = [] + for nr in sched_output.scheduled_new_reqs: + tag = f"result_for_{nr.request_id}" + runner_outputs.append( + RunnerOutput(request_id=nr.request_id, finished=True, result=_tagged_output(tag)) + ) + res_q.put(BatchRunnerOutput.from_list(runner_outputs)) + elif method in {"generate", "execute_model"} and args and hasattr(args[0], "request_id"): tag = f"result_for_{args[0].request_id}" + res_q.put(_tagged_output(tag)) elif args: tag = f"result_for_{args[0]}" + res_q.put(_tagged_output(tag)) else: tag = f"result_for_{method}" - res_q.put(_tagged_output(tag)) + res_q.put(_tagged_output(tag)) t = threading.Thread(target=_run, daemon=True) t.start() @@ -172,6 +189,64 @@ def _b(): assert results["B"].error == "result_for_B" +# ───────────────── request-mode dispatch (per-request vs batch) ───────────── + + +def _make_sched_output(*request_ids: str) -> DiffusionSchedulerOutput: + """Build a request-mode scheduler output with the given new requests.""" + new_reqs = [ + NewRequestData( + request_id=rid, + req=OmniDiffusionRequest( + prompt=f"prompt_{rid}", + sampling_params=OmniDiffusionSamplingParams(num_inference_steps=1), + request_id=rid, + ), + ) + for rid in request_ids + ] + return DiffusionSchedulerOutput( + step_id=0, + scheduled_new_reqs=new_reqs, + scheduled_cached_reqs=CachedRequestData.make_empty(), + finished_req_ids=set(), + num_running_reqs=len(new_reqs), + num_waiting_reqs=0, + ) + + +class TestRequestModeDispatch: + """Request-batch-capable dispatch uses ``execute_batch`` for request-mode cycles.""" + + @pytest.mark.parametrize("request_ids", [("solo",), ("A", "B", "C")]) + def test_request_batch_capable_pipeline_uses_execute_batch(self, request_ids): + engine, executor, _, _ = _make_engine() + executor.execute_request = Mock(return_value="per-request") + executor.execute_batch = Mock(return_value="batch") + engine.execute_fn = executor.execute_batch + + out = engine.execute_fn(_make_sched_output(*request_ids)) + + executor.execute_batch.assert_called_once() + assert out == "batch" + executor.execute_request.assert_not_called() + + @pytest.mark.parametrize("request_ids", [("solo",), ("A", "B")]) + def test_batch_path_routes_results_through_worker(self, request_ids): + """End-to-end: a request-batch cycle goes out as one ``execute_model_batch`` + RPC and comes back as a per-request-routed ``BatchRunnerOutput``.""" + engine, executor, req_q, res_q = _make_engine() + engine.execute_fn = executor.execute_batch + wt = _start_worker(req_q, res_q, count=1) + + out = engine.execute_fn(_make_sched_output(*request_ids)) + wt.join(5) + + assert isinstance(out, BatchRunnerOutput) + results = {ro.request_id: ro.result.error for ro in out.runner_outputs} + assert results == {request_id: f"result_for_{request_id}" for request_id in request_ids} + + # ───────────────── concurrent collective RPC ───────────────── diff --git a/tests/diffusion/test_stage_diffusion_proc.py b/tests/diffusion/test_stage_diffusion_proc.py index 3038d4e90cf..c96aa8ba118 100644 --- a/tests/diffusion/test_stage_diffusion_proc.py +++ b/tests/diffusion/test_stage_diffusion_proc.py @@ -3,9 +3,7 @@ import asyncio import time -from concurrent.futures import ThreadPoolExecutor from dataclasses import asdict, dataclass -from types import SimpleNamespace import pytest @@ -16,68 +14,6 @@ pytestmark = [pytest.mark.core_model, pytest.mark.diffusion, pytest.mark.cpu] -def test_process_batch_request_preserves_parent_request_id_and_kv_sender_info(): - async def run_test(): - captured = {} - - async def step(request): - captured["request"] = request - return [ - SimpleNamespace( - images=["img-1"], - _multimodal_output={}, - _custom_output={}, - metrics={}, - stage_durations={}, - peak_memory_mb=0.0, - latents=None, - trajectory_latents=None, - trajectory_timesteps=None, - trajectory_log_probs=None, - trajectory_decoded=None, - final_output_type="image", - ), - SimpleNamespace( - images=["img-2"], - _multimodal_output={}, - _custom_output={}, - metrics={}, - stage_durations={}, - peak_memory_mb=0.0, - latents=None, - trajectory_latents=None, - trajectory_timesteps=None, - trajectory_log_probs=None, - trajectory_decoded=None, - final_output_type="image", - ), - ] - - proc = object.__new__(StageDiffusionProc) - proc._engine = SimpleNamespace(step=step) - proc._od_config = SimpleNamespace(streaming_output=False) - proc._executor = ThreadPoolExecutor(max_workers=1) - - try: - result = await proc._process_batch_request( - request_id="req-parent", - prompts=["hello", "world"], - sampling_params_dict=asdict(OmniDiffusionSamplingParams()), - kv_sender_info={0: {"host": "10.0.0.2", "zmq_port": 50151}}, - ) - finally: - proc._executor.shutdown(wait=True) - - request = captured["request"] - assert request.request_id == "req-parent" - assert not hasattr(request, "request_ids") - assert request.kv_sender_info == {0: {"host": "10.0.0.2", "zmq_port": 50151}} - assert result.request_id == "req-parent" - assert result.images == ["img-1", "img-2"] - - asyncio.run(run_test()) - - @dataclass class MockOmniRequestOutput: request_id: str = "" diff --git a/tests/distributed/omni_connectors/test_kv_async_prefetch.py b/tests/distributed/omni_connectors/test_kv_async_prefetch.py index b6feba26d76..6ef5a3f8cc3 100644 --- a/tests/distributed/omni_connectors/test_kv_async_prefetch.py +++ b/tests/distributed/omni_connectors/test_kv_async_prefetch.py @@ -349,7 +349,7 @@ def test_consume_then_apply_attaches_payload(): data, _ = receiver.consume_prefetched_kv(_req("rid-apply")) assert data is not None - req = OmniDiffusionRequest(prompts=["p"], sampling_params=OmniDiffusionSamplingParams(), request_id="rid-apply") + req = OmniDiffusionRequest(prompt="p", sampling_params=OmniDiffusionSamplingParams(), request_id="rid-apply") # Mirror consume_and_distribute_kv_cache's LOCAL apply path (CPU: record_stream no-op). receiver._record_stream_for_prefetched(data) receiver.apply_kv_cache_to_request(req, data) diff --git a/tests/distributed/omni_connectors/test_kv_flow.py b/tests/distributed/omni_connectors/test_kv_flow.py index 8bf2b5d710e..83206fefbda 100644 --- a/tests/distributed/omni_connectors/test_kv_flow.py +++ b/tests/distributed/omni_connectors/test_kv_flow.py @@ -370,7 +370,7 @@ def test_manager_reception(kv_config, mock_connector, common_constants): mock_connector.store[store_key] = data_to_receive req = OmniDiffusionRequest( - prompts=["test_recv"], + prompt="test_recv", sampling_params=OmniDiffusionSamplingParams(), request_id=req_id, ) @@ -414,7 +414,7 @@ def test_manager_reception_prefers_parent_request_id_for_batched_request(kv_conf mock_connector.store[store_key] = data_to_receive req = OmniDiffusionRequest( - prompts=["prompt-a", "prompt-b"], + prompt="prompt-a", sampling_params=OmniDiffusionSamplingParams(), request_id=parent_req_id, ) @@ -439,7 +439,7 @@ def collect_cfg(request_id, cfg_request_ids, kv_transfer_manager, target_device) return {"cfg_text_kv_metadata": {"ok": True}} req = OmniDiffusionRequest( - prompts=["prompt-a", "prompt-b"], + prompt="prompt-a", sampling_params=OmniDiffusionSamplingParams(), request_id="req-parent", ) @@ -500,7 +500,7 @@ def test_integration_flow(common_constants): receiver_manager._connector = connector req = OmniDiffusionRequest( - prompts=["test_integ"], + prompt="test_integ", sampling_params=OmniDiffusionSamplingParams(), request_id=req_id, ) diff --git a/tests/e2e/offline_inference/custom_pipeline/qwen_image_pipeline_with_logprob.py b/tests/e2e/offline_inference/custom_pipeline/qwen_image_pipeline_with_logprob.py index 709c6655565..97efab0fe63 100644 --- a/tests/e2e/offline_inference/custom_pipeline/qwen_image_pipeline_with_logprob.py +++ b/tests/e2e/offline_inference/custom_pipeline/qwen_image_pipeline_with_logprob.py @@ -21,7 +21,7 @@ from vllm_omni.diffusion.data import DiffusionOutput, OmniDiffusionConfig from vllm_omni.diffusion.distributed.utils import get_local_device from vllm_omni.diffusion.models.qwen_image import QwenImagePipeline -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch def _maybe_to_cpu(v): @@ -33,6 +33,8 @@ def _maybe_to_cpu(v): # Custom pipeline class for QwenImage that returns log probabilities during the diffusion process. # This is for test class QwenImagePipelineWithLogProbForTest(QwenImagePipeline): + supports_request_batch = False + def __init__(self, *, od_config: OmniDiffusionConfig, prefix: str = ""): super().__init__(od_config=od_config, prefix=prefix) self.device = get_local_device() @@ -211,7 +213,7 @@ def diffuse( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt_ids: torch.Tensor | list[int] | None = None, prompt_mask: torch.Tensor | None = None, negative_prompt_ids: torch.Tensor | list[int] | None = None, @@ -239,8 +241,12 @@ def forward( sde_type: Literal["sde", "cps"] = "sde", logprobs: bool = True, ) -> DiffusionOutput: - # Extract prompt data from OmniCustomPrompt in req.prompts[0] - custom_prompt = req.prompts[0] if req.prompts else {} + if req.num_reqs != 1: + raise ValueError("QwenImagePipelineWithLogProbForTest supports only single-request forward.") + request = req.requests[0] + + # Extract prompt data from OmniCustomPrompt in req.prompt + custom_prompt = request.prompt if request.prompt is not None else {} if isinstance(custom_prompt, dict): prompt_ids = custom_prompt.get("prompt_ids", prompt_ids) prompt_mask = custom_prompt.get("prompt_mask", prompt_mask) @@ -248,7 +254,7 @@ def forward( negative_prompt_mask = custom_prompt.get("negative_prompt_mask", negative_prompt_mask) # Read sampling params from req.sampling_params - sp = req.sampling_params + sp = request.sampling_params height = sp.height or self.default_sample_size * self.vae_scale_factor width = sp.width or self.default_sample_size * self.vae_scale_factor num_inference_steps = sp.num_inference_steps or num_inference_steps diff --git a/tests/entrypoints/test_async_omni_diffusion_config.py b/tests/entrypoints/test_async_omni_diffusion_config.py index eba1878ec85..61a420a4cfa 100644 --- a/tests/entrypoints/test_async_omni_diffusion_config.py +++ b/tests/entrypoints/test_async_omni_diffusion_config.py @@ -284,6 +284,29 @@ def test_serve_cli_accepts_diffusion_attention_backend(): assert diffusion_attention_config.default.backend == "FLASH_ATTN" +def test_serve_cli_accepts_request_batch_max_wait_ms(): + """Ensure diffusion serve CLI forwards request-batch admission wait to stage config.""" + parser = TrackingArgumentParser() + subparsers = parser.add_subparsers(dest="command") + OmniServeCommand().subparser_init(subparsers) + + args = parser.parse_args( + [ + "serve", + "Qwen/Qwen-Image", + "--omni", + "--request-batch-max-wait-ms", + "250", + ] + ) + + explicit_kwargs = args.get_explicit_kwargs_dict() + stage_cfg = AsyncOmniEngine._create_default_diffusion_stage_cfg(explicit_kwargs)[0] + + assert args.request_batch_max_wait_ms == 250.0 + assert stage_cfg["engine_args"]["request_batch_max_wait_ms"] == 250.0 + + def test_serve_cli_accepts_additional_config(): """Ensure diffusion serve CLI exposes additional_config and forwards it to stage config.""" parser = TrackingArgumentParser() diff --git a/tests/model_executor/stage_input_processors/test_glm_image.py b/tests/model_executor/stage_input_processors/test_glm_image.py index 029ba99c29e..8dc333fb6f2 100644 --- a/tests/model_executor/stage_input_processors/test_glm_image.py +++ b/tests/model_executor/stage_input_processors/test_glm_image.py @@ -284,12 +284,11 @@ def test_basic_t2i(self): prompt = {"prompt": "a cat", "mm_processor_kwargs": {"target_h": 1024, "target_w": 1024}} - result = ar2diffusion(source_outputs, prompt=[prompt]) - assert len(result) == 1 - assert result[0]["prompt"] == "a cat" - assert result[0]["height"] == 1024 - assert result[0]["width"] == 1024 - assert "prior_token_ids" in result[0]["extra"] + result = ar2diffusion(source_outputs, prompt=prompt) + assert result["prompt"] == "a cat" + assert result["height"] == 1024 + assert result["width"] == 1024 + assert "prior_token_ids" in result["extra"] def test_i2i_with_mm_output(self): """Test image-to-image with prior_token_image_ids from AR model.""" @@ -306,9 +305,8 @@ def test_i2i_with_mm_output(self): "multi_modal_data": {"image": img}, } - result = ar2diffusion(source_outputs, prompt=[prompt]) - assert len(result) == 1 - assert result[0]["extra"]["prior_token_image_ids"] is not None + result = ar2diffusion(source_outputs, prompt=prompt) + assert result["extra"]["prior_token_image_ids"] is not None def test_i2i_detected_via_modalities(self): """Test i2i mode detected via modalities field.""" @@ -321,11 +319,11 @@ def test_i2i_detected_via_modalities(self): "modalities": ["img2img"], } - result = ar2diffusion(source_outputs, prompt=[prompt]) - assert len(result) == 1 + result = ar2diffusion(source_outputs, prompt=prompt) + assert result["prompt"] == "edit this" - def test_empty_source_outputs_returns_empty_list(self): - assert ar2diffusion([], prompt={}) == [] + def test_empty_source_outputs_returns_none(self): + assert ar2diffusion([], prompt=None) is None def test_default_dimensions(self): """When no height/width in prompt, defaults to 1024x1024.""" @@ -333,9 +331,9 @@ def test_default_dimensions(self): source_outputs = [_source_output(token_ids)] prompt = {"prompt": "test"} - result = ar2diffusion(source_outputs, prompt=[prompt]) - assert result[0]["height"] == 1024 - assert result[0]["width"] == 1024 + result = ar2diffusion(source_outputs, prompt=prompt) + assert result["height"] == 1024 + assert result["width"] == 1024 def test_requires_multimodal_data_with_pil_image(self): """Test that pil_image is included when requires_multimodal_data=True.""" @@ -350,8 +348,8 @@ def test_requires_multimodal_data_with_pil_image(self): "multi_modal_data": {"image": img}, } - result = ar2diffusion(source_outputs, prompt=[prompt], requires_multimodal_data=True) - assert result[0]["pil_image"] is img + result = ar2diffusion(source_outputs, prompt=prompt, requires_multimodal_data=True) + assert result["pil_image"] is img def test_extra_params_passed_through(self): """Test that seed, num_inference_steps, guidance_scale, negative_prompt are passed.""" @@ -366,24 +364,22 @@ def test_extra_params_passed_through(self): "negative_prompt": "blurry", } - result = ar2diffusion(source_outputs, prompt=[prompt]) - assert result[0]["seed"] == 42 - assert result[0]["num_inference_steps"] == 50 - assert result[0]["guidance_scale"] == 7.5 - assert result[0]["negative_prompt"] == "blurry" + result = ar2diffusion(source_outputs, prompt=prompt) + assert result["prompt"] == "test" + assert result["seed"] == 42 + assert result["num_inference_steps"] == 50 + assert result["guidance_scale"] == 7.5 + assert result["negative_prompt"] == "blurry" - def test_batch_requests(self): - """Test processing multiple requests in a batch.""" + def test_multiple_source_outputs_uses_first_payload_only(self): + """Test the GLM bridge keeps a single diffusion payload for one request.""" tokens1 = list(range(256)) + list(range(1024)) + [16385] - tokens2 = list(range(256)) + list(range(1024)) + [16385] + tokens2 = [1, 2, 3] source_outputs = [_source_output(tokens1), _source_output(tokens2)] - prompts = [ - {"prompt": "first", "mm_processor_kwargs": {"target_h": 1024, "target_w": 1024}}, - {"prompt": "second", "mm_processor_kwargs": {"target_h": 512, "target_w": 512}}, - ] + prompt = {"prompt": "first", "mm_processor_kwargs": {"target_h": 1024, "target_w": 1024}} - result = ar2diffusion(source_outputs, prompt=prompts) - assert len(result) == 2 - assert result[0]["prompt"] == "first" - assert result[1]["prompt"] == "second" + result = ar2diffusion(source_outputs, prompt=prompt) + assert result["prompt"] == "first" + assert result["height"] == 1024 + assert result["width"] == 1024 diff --git a/vllm_omni/diffusion/cache/teacache/coefficient_estimator.py b/vllm_omni/diffusion/cache/teacache/coefficient_estimator.py index 60638cca876..49ba806d55b 100644 --- a/vllm_omni/diffusion/cache/teacache/coefficient_estimator.py +++ b/vllm_omni/diffusion/cache/teacache/coefficient_estimator.py @@ -201,15 +201,17 @@ def __init__( def collect_from_prompt(self, prompt: str, **generate_kwargs): self.hook.start_collection() req = OmniDiffusionRequest( - prompts=[prompt], + prompt=prompt, request_id="teacache-coefficient-estimator", sampling_params=OmniDiffusionSamplingParams( num_inference_steps=generate_kwargs.get("num_inference_steps", 20), seed=generate_kwargs.get("seed", 42), ), ) + from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch + with torch.no_grad(): - self.pipeline.forward(req) + self.pipeline.forward(DiffusionRequestBatch(requests=[req])) trajectory = self.hook.stop_collection() if trajectory: self.collected_data.append(trajectory) diff --git a/vllm_omni/diffusion/data.py b/vllm_omni/diffusion/data.py index 4c031bc1557..f5fe4c30aeb 100644 --- a/vllm_omni/diffusion/data.py +++ b/vllm_omni/diffusion/data.py @@ -770,6 +770,12 @@ class OmniDiffusionConfig: # Maximum number of sequences to generate in a batch max_num_seqs: int = 1 + # Request-mode batch admission: wait briefly for compatible requests to + # accumulate in the scheduler waiting queue before the first schedule() of + # a wave. Improves fused forward batch sizes under bursty HTTP ingress. + # 0 disables admission (default; no added latency). + request_batch_max_wait_ms: float = 0.0 + # Supplementary model specific parameters extras: dict[str, Any] = Field(default_factory=dict) @@ -842,6 +848,9 @@ def settle_port(self, port: int, port_inc: int = 42, max_attempts: int = 100) -> def __post_init__(self): self.master_port = self._resolve_master_port() + self.request_batch_max_wait_ms = float(self.request_batch_max_wait_ms or 0.0) + if self.request_batch_max_wait_ms < 0: + raise ValueError(f"request_batch_max_wait_ms must be non-negative, got {self.request_batch_max_wait_ms}.") if isinstance(self.profiler_config, dict): from vllm.config import ProfilerConfig diff --git a/vllm_omni/diffusion/diffusion_engine.py b/vllm_omni/diffusion/diffusion_engine.py index 1d8753682bc..2554099e760 100644 --- a/vllm_omni/diffusion/diffusion_engine.py +++ b/vllm_omni/diffusion/diffusion_engine.py @@ -17,6 +17,7 @@ import PIL.Image import torch from vllm.logger import init_logger +from vllm.utils.import_utils import resolve_obj_by_qualname from vllm.v1.engine.exceptions import EngineDeadError from vllm_omni.diffusion.data import ( @@ -38,6 +39,7 @@ normalize_diffusion_postprocess_output, ) from vllm_omni.diffusion.registry import ( + DiffusionModelRegistry, get_diffusion_action_post_process_func, get_diffusion_post_process_func, get_diffusion_pre_process_func, @@ -74,6 +76,37 @@ def _func_accepts_parameter(func: object | None, parameter_name: str) -> bool: ) +def _resolve_custom_pipeline_cls(custom_pipeline_args: dict[str, Any] | None) -> type | None: + if custom_pipeline_args is None: + return None + + try: + pipeline_cls = custom_pipeline_args["pipeline_class"] + except KeyError as exc: + raise ValueError("custom_pipeline_args must include 'pipeline_class'.") from exc + + if isinstance(pipeline_cls, type): + return pipeline_cls + if isinstance(pipeline_cls, str): + try: + return resolve_obj_by_qualname(pipeline_cls) + except (AttributeError, ImportError, ValueError) as exc: + raise ValueError(f"Failed to resolve custom diffusion pipeline class {pipeline_cls!r}.") from exc + raise TypeError( + f"custom_pipeline_args['pipeline_class'] must be a qualified name string or a class, " + f"got {type(pipeline_cls).__name__}" + ) + + +def supports_request_batch(od_config: OmniDiffusionConfig) -> bool: + model_cls = _resolve_custom_pipeline_cls(getattr(od_config, "custom_pipeline_args", None)) + if model_cls is None: + model_cls = DiffusionModelRegistry._try_load_model_cls(getattr(od_config, "model_class_name", None)) + if model_cls is None: + return False + return bool(getattr(model_cls, "supports_request_batch", False)) + + def _move_tensor_tree_to_cpu(value: object) -> object: if isinstance(value, torch.Tensor): return value.cpu() if value.device.type != "cpu" else value @@ -138,10 +171,7 @@ def __init__( StepScheduler() if self.step_execution else RequestScheduler() ) self.scheduler.initialize(od_config) - if self.scheduler.max_num_running_reqs > 1 and not self.step_execution: - max_num_seqs = self.scheduler.max_num_running_reqs - self.scheduler.max_num_running_reqs = 1 - logger.warning(f"Non-stepwise-execution does not support max-num-seqs={max_num_seqs}, set it to 1.") + self.supports_request_batch = False if self.step_execution else supports_request_batch(od_config) self.main_loop: asyncio.AbstractEventLoop | None = None self.stop_event: threading.Event | None = None self.worker_thread: threading.Thread | None = None @@ -161,9 +191,18 @@ def __init__( self._rpc_queue: queue.Queue[_RpcTask] = queue.Queue() if self.step_execution: self.execute_fn = self.executor.execute_step + elif self.supports_request_batch: + self.execute_fn = self.executor.execute_batch else: self.execute_fn = self.executor.execute_request + if self.supports_request_batch: + logger.info( + "[RequestBatch] engine init max_num_seqs=%s max_wait_ms=%s", + getattr(od_config, "max_num_seqs", None), + getattr(od_config, "request_batch_max_wait_ms", None), + ) + try: self._dummy_run() except Exception as e: @@ -348,6 +387,9 @@ def _busy_loop(self): # Only RPC / abort work pending; loop back to drain it. continue + if self.supports_request_batch: + self._wait_for_request_batch_admission_locked() + sched_output = self.scheduler.schedule() if sched_output.is_empty: @@ -390,6 +432,61 @@ def _busy_loop(self): # Engine is stopping: fail any RPCs still queued so callers don't hang. self._fail_pending_rpcs(RuntimeError("DiffusionEngine is shutting down.")) + def _wait_for_request_batch_admission_locked(self) -> None: + """Wait for compatible requests to accumulate before scheduling a wave. + + Caller must hold ``self._cv``. + """ + if self.step_execution or not self.supports_request_batch: + return + + max_wait_s = self.od_config.request_batch_max_wait_ms / 1000.0 + if max_wait_s == 0: + return + + max_batch = self.scheduler.max_num_running_reqs + waiting = self.scheduler.num_waiting_requests() + running = self.scheduler.num_running_requests() + + if running > 0: + return + + start = time.monotonic() + deadline = start + max_wait_s + last_waiting = -1 + stable_since = start + # Require a short idle period with no queue growth so bursty HTTP + # ingress can land before the first schedule() of a wave. + stable_window_s = min(0.05, max_wait_s / 5.0) + + while not self.stop_event.is_set(): + waiting = self.scheduler.num_waiting_requests() + now = time.monotonic() + + if waiting >= max_batch: + break + if waiting > 0 and (now - stable_since) >= stable_window_s: + break + if now >= deadline: + break + + if waiting > last_waiting: + stable_since = now + last_waiting = waiting + + remaining = deadline - now + self._cv.wait(timeout=min(remaining, 0.002)) + + waited_ms = (time.monotonic() - start) * 1000.0 + final_waiting = self.scheduler.num_waiting_requests() + if final_waiting > 0: + logger.info( + "[RequestBatch] admission wait done waiting=%d max_batch=%d waited_ms=%.1f", + final_waiting, + max_batch, + waited_ms, + ) + def _process_rpc_queue(self) -> None: """Execute pending collective_rpc tasks from the busy-loop thread. @@ -680,7 +777,7 @@ def _dummy_run(self): logger.info("Skipping dummy warmup run (num_frames=0)") return req = OmniDiffusionRequest( - prompts=[prompt], + prompt=prompt, request_id=DUMMY_DIFFUSION_REQUEST_ID, sampling_params=OmniDiffusionSamplingParams( height=height, @@ -930,9 +1027,12 @@ def _finalize_finished_request( raise RuntimeError(f"Diffusion scheduler lost state for request {request_id}.") if state.status == DiffusionRequestStatus.FINISHED_ABORTED: + # Preserve runner-provided abort details when available. + if runner_output is not None and runner_output.result is not None and runner_output.result.aborted: + return runner_output.result return DiffusionOutput( aborted=True, - abort_message=f"Request {state.req.request_id} aborted.", + abort_message=f"Request {request_id} aborted.", ) if runner_output is not None and runner_output.result is not None: diff --git a/vllm_omni/diffusion/executor/abstract.py b/vllm_omni/diffusion/executor/abstract.py index 370010d1b21..8bfb5923aa7 100644 --- a/vllm_omni/diffusion/executor/abstract.py +++ b/vllm_omni/diffusion/executor/abstract.py @@ -73,6 +73,11 @@ def execute_request(self, scheduler_output: DiffusionSchedulerOutput) -> BaseRun """Execute request-mode work from a scheduler output.""" pass + @abstractmethod + def execute_batch(self, scheduler_output: DiffusionSchedulerOutput) -> BaseRunnerOutput: + """Execute request-mode work as a single batched RPC.""" + pass + @abstractmethod def execute_step(self, scheduler_output: DiffusionSchedulerOutput) -> BaseRunnerOutput: """Execute step-mode work from a scheduler output.""" diff --git a/vllm_omni/diffusion/executor/multiproc_executor.py b/vllm_omni/diffusion/executor/multiproc_executor.py index cd5e026430d..2e491ff6971 100644 --- a/vllm_omni/diffusion/executor/multiproc_executor.py +++ b/vllm_omni/diffusion/executor/multiproc_executor.py @@ -308,36 +308,68 @@ def register_failure_callback( self._failure_callbacks.append(callback) def execute_request(self, scheduler_output: DiffusionSchedulerOutput) -> BaseRunnerOutput: - """Adapt request-mode scheduler output to worker execute_model RPC.""" - from vllm_omni.diffusion.worker.utils import RunnerOutput + """Adapt request-mode scheduler output to worker execute_model RPCs. + + Returns a BatchRunnerOutput with one RunnerOutput per scheduled request. + """ + from vllm_omni.diffusion.worker.utils import BatchRunnerOutput, RunnerOutput self._ensure_open() - if scheduler_output.num_scheduled_reqs != 1: - raise ValueError( - f"Request mode currently supports batch_size=1, " - f"but got {scheduler_output.num_scheduled_reqs} scheduled requests." - ) + runner_outputs: list[RunnerOutput] = [] + + for new_req in scheduler_output.scheduled_new_reqs: + req = new_req.req + try: + result = self.collective_rpc( + "execute_model", + args=(req, self.od_config, scheduler_output.kv_prefetch_jobs), + unique_reply_rank=0, + exec_all_ranks=True, + ) + if not isinstance(result, DiffusionOutput): + raise RuntimeError(f"Unexpected response type: {type(result)!r}") + runner_outputs.append( + RunnerOutput( + request_id=new_req.request_id, + step_index=None, + finished=True, + result=result, + ) + ) + except Exception as exc: + runner_outputs.append( + RunnerOutput( + request_id=new_req.request_id, + step_index=None, + finished=True, + result=DiffusionOutput(error=str(exc)), + ) + ) - new_req = scheduler_output.scheduled_new_reqs[0] + return BatchRunnerOutput.from_list(runner_outputs) + + def execute_batch(self, scheduler_output: DiffusionSchedulerOutput) -> BaseRunnerOutput: + """Execute request-mode work through a single batched worker RPC. + + The worker builds DiffusionRequestBatch from scheduler output and returns + BatchRunnerOutput with one RunnerOutput per scheduled request. + """ + from vllm_omni.diffusion.worker.utils import BatchRunnerOutput + + self._ensure_open() result = self.collective_rpc( - "execute_model", - args=(new_req.req, self.od_config, getattr(scheduler_output, "kv_prefetch_jobs", None)), + "execute_model_batch", + args=(scheduler_output, self.od_config), unique_reply_rank=0, exec_all_ranks=True, ) - if not isinstance(result, DiffusionOutput): - raise RuntimeError(f"Unexpected response type for execute_request: {type(result)!r}") - - return RunnerOutput( - request_id=new_req.request_id, - step_index=None, - finished=True, - result=result, - ) + if not isinstance(result, BatchRunnerOutput): + raise RuntimeError(f"Unexpected response type for execute_batch: {type(result)!r}") + return result def execute_step(self, scheduler_output: DiffusionSchedulerOutput) -> BaseRunnerOutput: """Forward step-mode scheduler output to worker execute_stepwise RPC.""" - from vllm_omni.diffusion.worker.utils import BaseRunnerOutput, RunnerOutput + from vllm_omni.diffusion.worker.utils import BaseRunnerOutput self._ensure_open() result = self.collective_rpc( @@ -349,18 +381,7 @@ def execute_step(self, scheduler_output: DiffusionSchedulerOutput) -> BaseRunner if isinstance(result, BaseRunnerOutput): return result - # TODO: Remove this fallback; DiffusionOutput cannot faithfully represent - # failed multi-request step batches. - if isinstance(result, DiffusionOutput): - request_id = scheduler_output.scheduled_request_ids[0] if scheduler_output.scheduled_request_ids else "" - return RunnerOutput( - request_id=request_id, - step_index=None, - finished=True, - result=result, - ) - else: - raise RuntimeError(f"Unexpected response type for execute_step: {type(result)!r}") + raise RuntimeError(f"Unexpected response type for execute_step: {type(result)!r}") def collective_rpc( self, diff --git a/vllm_omni/diffusion/inline_stage_diffusion_client.py b/vllm_omni/diffusion/inline_stage_diffusion_client.py index a54ce8b18ff..8b3d90f0097 100644 --- a/vllm_omni/diffusion/inline_stage_diffusion_client.py +++ b/vllm_omni/diffusion/inline_stage_diffusion_client.py @@ -12,8 +12,6 @@ from concurrent.futures import ThreadPoolExecutor from typing import TYPE_CHECKING, Any -import torch -from PIL import Image from vllm.logger import init_logger from vllm.v1.engine.exceptions import EngineDeadError @@ -128,7 +126,7 @@ async def _dispatch_request( ) -> None: try: request = OmniDiffusionRequest( - prompts=[prompt], + prompt=prompt, sampling_params=sampling_params, request_id=request_id, kv_sender_info=kv_sender_info, @@ -161,113 +159,6 @@ async def _dispatch_request( finally: self._tasks.pop(request_id, None) - async def add_batch_request_async( - self, - request_id: str, - prompts: list[OmniPromptType], - sampling_params: OmniDiffusionSamplingParams, - kv_sender_info: dict[int, dict[str, Any]] | None = None, - ) -> None: - logger.debug( - "[InlineStageDiffusionClient] stage-%s [rep-%s] add batch request: %s (%d prompts)", - self.stage_id, - self.replica_id, - request_id, - len(prompts), - ) - task = asyncio.create_task( - self._dispatch_batch( - request_id, - prompts, - sampling_params, - kv_sender_info, - ) - ) - self._tasks[request_id] = task - - async def _dispatch_batch( - self, - request_id: str, - prompts: list[Any], - sampling_params: OmniDiffusionSamplingParams, - kv_sender_info: dict[str, Any] | None = None, - ) -> None: - try: - request = OmniDiffusionRequest( - prompts=prompts, - sampling_params=sampling_params, - request_id=request_id, - kv_sender_info=kv_sender_info, - ) - - results = await self._engine.step(request) - - all_images: list = [] - merged_mm: dict[str, Any] = {} - merged_metrics: dict[str, Any] = {} - merged_durations: dict[str, float] = {} - merged_custom: dict[str, Any] = {} - peak_mem = 0.0 - latents = None - trajectory_latents: list[torch.Tensor] | None = None - trajectory_timesteps: list[torch.Tensor] | None = None - trajectory_log_probs: torch.Tensor | None = None - trajectory_decoded: list[Image.Image] | None = None - final_output_type = "image" - - for r in results: - all_images.extend(r.images) - merged_mm.update(r._multimodal_output) - merged_metrics.update(r.metrics) - merged_durations.update(r.stage_durations) - merged_custom.update(r._custom_output) - peak_mem = max(peak_mem, r.peak_memory_mb) - if latents is None and r.latents is not None: - latents = r.latents - if trajectory_latents is None: - trajectory_latents = r.trajectory_latents - if trajectory_timesteps is None: - trajectory_timesteps = r.trajectory_timesteps - if trajectory_log_probs is None: - trajectory_log_probs = r.trajectory_log_probs - if trajectory_decoded is None: - trajectory_decoded = r.trajectory_decoded - if r.final_output_type != "image": - final_output_type = r.final_output_type - - result = OmniRequestOutput.from_diffusion( - request_id=request_id, - images=all_images, - prompt=prompts[0] if len(prompts) == 1 else None, - metrics=merged_metrics, - latents=latents, - trajectory_latents=trajectory_latents, - trajectory_timesteps=trajectory_timesteps, - trajectory_log_probs=trajectory_log_probs, - trajectory_decoded=trajectory_decoded, - custom_output=merged_custom or None, - multimodal_output=merged_mm or None, - final_output_type=final_output_type, - stage_durations=merged_durations, - peak_memory_mb=peak_mem, - ) - - self._output_queue.put_nowait(result) - except DiffusionRequestAbortedError as e: - logger.info("request_id: %s aborted: %s", request_id, str(e)) - except Exception as e: - logger.exception("Batch diffusion request %s failed: %s", request_id, e) - status_code, error_type = client_error_metadata(e) - error_output = OmniRequestOutput.from_error( - request_id=request_id, - error_message=str(e), - status_code=status_code, - error_type=error_type, - ) - self._output_queue.put_nowait(error_output) - finally: - self._tasks.pop(request_id, None) - def get_diffusion_output_nowait(self) -> OmniRequestOutput | None: try: return self._output_queue.get_nowait() diff --git a/vllm_omni/diffusion/ipc.py b/vllm_omni/diffusion/ipc.py index 2268b45ab1f..cd83e718c19 100644 --- a/vllm_omni/diffusion/ipc.py +++ b/vllm_omni/diffusion/ipc.py @@ -138,9 +138,10 @@ def _is_rpc_result_envelope(output: object) -> bool: def pack_diffusion_output_shm(output: object) -> object: """Replace large tensors in diffusion worker outputs with SHM handles. - Supports either a bare ``DiffusionOutput`` or a wrapper object carrying one - in ``.result`` (for example ``RunnerOutput``), or an RPC result envelope - carrying the diffusion output in ``["result"]``. + Supports a bare ``DiffusionOutput``, a wrapper object carrying one in + ``.result`` (for example ``RunnerOutput``), an RPC result envelope carrying + the diffusion output in ``["result"]``, or a batch wrapper carrying + ``RunnerOutput`` objects in ``.runner_outputs``. """ if isinstance(output, DiffusionOutput): return _pack_diffusion_fields(output) @@ -154,6 +155,11 @@ def pack_diffusion_output_shm(output: object) -> object: result = getattr(output, "result", None) if isinstance(result, DiffusionOutput): output.result = _pack_diffusion_fields(result) + + runner_outputs = getattr(output, "runner_outputs", None) + if isinstance(runner_outputs, list): + for runner_output in runner_outputs: + pack_diffusion_output_shm(runner_output) return output @@ -179,4 +185,9 @@ def unpack_diffusion_output_shm(output: object) -> object: result = getattr(output, "result", None) if isinstance(result, DiffusionOutput): output.result = _unpack_diffusion_fields(result) + + runner_outputs = getattr(output, "runner_outputs", None) + if isinstance(runner_outputs, list): + for runner_output in runner_outputs: + unpack_diffusion_output_shm(runner_output) return output diff --git a/vllm_omni/diffusion/models/audiox/pipeline_audiox.py b/vllm_omni/diffusion/models/audiox/pipeline_audiox.py index b2035b7a8d0..02639f1afcd 100644 --- a/vllm_omni/diffusion/models/audiox/pipeline_audiox.py +++ b/vllm_omni/diffusion/models/audiox/pipeline_audiox.py @@ -27,7 +27,7 @@ from vllm_omni.diffusion.models.audiox.audiox_transformer import MMDiffusionTransformer from vllm_omni.diffusion.models.interface import SupportAudioOutput from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.transformers_utils.processors import audiox as _audiox_transforms _VIDEO_ONLY_TASKS = _audiox_transforms.VIDEO_ONLY_TASKS @@ -359,6 +359,7 @@ def __call__(self, sigma: torch.Tensor, sigma_next: torch.Tensor) -> torch.Tenso class AudioXPipeline(nn.Module, SupportAudioOutput, DiffusionPipelineProfilerMixin): + supports_request_batch = False support_audio_output: ClassVar[bool] = True audio_sample_rate: ClassVar[int] = 44100 audio_channels: ClassVar[int] = 2 @@ -811,7 +812,7 @@ def _build_conditioning_batch( ) return batch - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: if req.prompts is None or len(req.prompts) == 0: raise ValueError("AudioXPipeline requires at least one prompt.") normalized_prompts = _normalize_prompts(list(req.prompts)) @@ -913,10 +914,10 @@ def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: target_samples = int(seconds_total * self._sample_rate) audio = audio[..., :target_samples] + stage_durations = self.stage_durations if getattr(self, "enable_diffusion_pipeline_profiler", False) else None + custom_output = {"audiox_task": task_norm} return DiffusionOutput( output=audio, - custom_output={"audiox_task": task_norm}, - stage_durations=self.stage_durations - if getattr(self, "enable_diffusion_pipeline_profiler", False) - else None, + custom_output=custom_output, + stage_durations=stage_durations, ) diff --git a/vllm_omni/diffusion/models/bagel/pipeline_bagel.py b/vllm_omni/diffusion/models/bagel/pipeline_bagel.py index 7a52b90a09c..4536227658b 100644 --- a/vllm_omni/diffusion/models/bagel/pipeline_bagel.py +++ b/vllm_omni/diffusion/models/bagel/pipeline_bagel.py @@ -29,7 +29,7 @@ from vllm_omni.diffusion.model_loader.diffusers_loader import DiffusersPipelineLoader from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import download_weights_from_hf_specific from .autoencoder import AutoEncoder, AutoEncoderParams, DistributedAutoEncoder @@ -341,7 +341,7 @@ def _regen_init_noise_on_device(self, gen_input: dict, seed: int | None) -> None ) @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: if len(req.prompts) > 1: logger.warning( """This model only supports a single prompt, not a batched request.""", diff --git a/vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py b/vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py index 694ae4e8e42..890b76d6b94 100644 --- a/vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py +++ b/vllm_omni/diffusion/models/cosmos3/pipeline_cosmos3.py @@ -61,6 +61,7 @@ from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin, _is_rank_zero from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.entrypoints.openai.video_api_utils import positive_float from vllm_omni.inputs.data import OmniDiffusionSamplingParams @@ -371,115 +372,114 @@ def _preprocess_condition_video( def pre_process_func(request: OmniDiffusionRequest) -> OmniDiffusionRequest: action_mode = _request_action_mode(request) + prompt = request.prompt if is_guardrails_enabled(od_config, request.sampling_params): - for prompt in request.prompts: - text = prompt if isinstance(prompt, str) else prompt.get("prompt", "") - check_text_safety(text) + text = prompt if isinstance(prompt, str) else prompt.get("prompt", "") + check_text_safety(text) + + if isinstance(prompt, str): + return request + multi_modal_data = prompt.get("multi_modal_data", {}) or {} + raw_image = multi_modal_data.get("image") + raw_video = multi_modal_data.get("video") + if raw_image is None and raw_video is None: + return request + if raw_image is not None and raw_video is not None and action_mode is None: + raise ValueError("Cosmos3 non-action generation accepts either image or video input, not both.") - for i, prompt in enumerate(request.prompts): - if isinstance(prompt, str): - continue - multi_modal_data = prompt.get("multi_modal_data", {}) or {} - raw_image = multi_modal_data.get("image") - raw_video = multi_modal_data.get("video") - if raw_image is None and raw_video is None: - continue - if raw_image is not None and raw_video is not None and action_mode is None: - raise ValueError("Cosmos3 non-action generation accepts either image or video input, not both.") - - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - raw_video_frames: list[Any] | None = None - transfer_input_fps: float | None = None - if raw_video is not None: - transfer_input_fps = _video_payload_fps(raw_video) - raw_video_frames = _video_payload_to_frames(raw_video) - if not raw_video_frames: - raise TypeError("Cosmos3 video input must be a non-empty list of PIL images or image paths.") - - if raw_image is None: - assert raw_video_frames is not None # raw_image and raw_video can't both be None here - image = _pil_to_rgb(raw_video_frames[0]) - else: - image = _pil_to_rgb(raw_image) - extra = _extra_args(request) - transfer_requested = action_mode is None and has_transfer_hints(extra) + if "additional_information" not in prompt: + prompt["additional_information"] = {} - # Auto-calculate H/W from aspect ratio (720p max area) - if transfer_requested: - _set_transfer_size_from_image(request, image) - elif request.sampling_params.height is None or request.sampling_params.width is None: - if action_mode is not None: - _set_action_size_from_image(request, image) - else: - max_area = 720 * 1280 - aspect_ratio = image.height / image.width - mod_value = 16 - height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value - width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value - if request.sampling_params.height is None: - request.sampling_params.height = height - if request.sampling_params.width is None: - request.sampling_params.width = width - - target_w = request.sampling_params.width - target_h = request.sampling_params.height + raw_video_frames: list[Any] | None = None + transfer_input_fps: float | None = None + if raw_video is not None: + transfer_input_fps = _video_payload_fps(raw_video) + raw_video_frames = _video_payload_to_frames(raw_video) + if not raw_video_frames: + raise TypeError("Cosmos3 video input must be a non-empty list of PIL images or image paths.") + + if raw_image is None: + assert raw_video_frames is not None # raw_image and raw_video can't both be None here + image = _pil_to_rgb(raw_video_frames[0]) + else: + image = _pil_to_rgb(raw_image) + extra = _extra_args(request) + transfer_requested = action_mode is None and has_transfer_hints(extra) + + # Auto-calculate H/W from aspect ratio (720p max area) + if transfer_requested: + _set_transfer_size_from_image(request, image) + elif request.sampling_params.height is None or request.sampling_params.width is None: if action_mode is not None: - prompt["additional_information"]["preprocessed_image"] = _preprocess_action_image( - image, - int(target_h), - int(target_w), + _set_action_size_from_image(request, image) + else: + max_area = 720 * 1280 + aspect_ratio = image.height / image.width + mod_value = 16 + height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value + width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value + if request.sampling_params.height is None: + request.sampling_params.height = height + if request.sampling_params.width is None: + request.sampling_params.width = width + + target_w = request.sampling_params.width + target_h = request.sampling_params.height + if action_mode is not None: + prompt["additional_information"]["preprocessed_image"] = _preprocess_action_image( + image, + int(target_h), + int(target_w), + ) + elif raw_video is None: + prompt["additional_information"]["preprocessed_image"] = _preprocess_condition_image( + image, + int(target_h), + int(target_w), + ) + else: + assert raw_video_frames is not None + if transfer_requested: + if transfer_input_fps is not None: + prompt["additional_information"]["transfer_input_fps"] = transfer_input_fps + transfer_frames = media_to_uint8_cthw( + raw_video_frames, + height=int(target_h), + width=int(target_w), + max_frames=transfer_max_frames_from_extra_args(extra), ) - elif raw_video is None: - prompt["additional_information"]["preprocessed_image"] = _preprocess_condition_image( - image, - int(target_h), - int(target_w), + prompt["additional_information"]["preprocessed_transfer_video"] = uint8_cthw_to_normalized_5d( + transfer_frames, + dtype=torch.float32, ) else: - assert raw_video_frames is not None - if transfer_requested: - if transfer_input_fps is not None: - prompt["additional_information"]["transfer_input_fps"] = transfer_input_fps - transfer_frames = media_to_uint8_cthw( - raw_video_frames, - height=int(target_h), - width=int(target_w), - max_frames=transfer_max_frames_from_extra_args(extra), - ) - prompt["additional_information"]["preprocessed_transfer_video"] = uint8_cthw_to_normalized_5d( - transfer_frames, - dtype=torch.float32, - ) - else: - condition_frame_indexes_vision = normalize_condition_frame_indexes_vision( - extra.get( - "condition_frame_indexes_vision", - prompt.get("condition_frame_indexes_vision"), - ) - ) - keep = normalize_condition_video_keep( - extra.get("condition_video_keep", prompt.get("condition_video_keep")) + condition_frame_indexes_vision = normalize_condition_frame_indexes_vision( + extra.get( + "condition_frame_indexes_vision", + prompt.get("condition_frame_indexes_vision"), ) - max_frames = condition_pixel_frame_count(condition_frame_indexes_vision) - prompt["additional_information"]["preprocessed_video"] = _preprocess_condition_video( - raw_video_frames, - int(target_h), - int(target_w), - max_frames, - keep, - ) - prompt["additional_information"]["condition_frame_indexes_vision"] = list( - condition_frame_indexes_vision - ) - if action_mode is not None and raw_video_frames is not None: - prompt["additional_information"]["preprocessed_video"] = _preprocess_action_video( + ) + keep = normalize_condition_video_keep( + extra.get("condition_video_keep", prompt.get("condition_video_keep")) + ) + max_frames = condition_pixel_frame_count(condition_frame_indexes_vision) + prompt["additional_information"]["preprocessed_video"] = _preprocess_condition_video( raw_video_frames, int(target_h), int(target_w), + max_frames, + keep, ) - request.prompts[i] = prompt + prompt["additional_information"]["condition_frame_indexes_vision"] = list( + condition_frame_indexes_vision + ) + if action_mode is not None and raw_video_frames is not None: + prompt["additional_information"]["preprocessed_video"] = _preprocess_action_video( + raw_video_frames, + int(target_h), + int(target_w), + ) + request.prompt = prompt return request @@ -1428,7 +1428,7 @@ def _get_sound_tokenizer(self): return self._sound_tokenizer @staticmethod - def _is_t2i_request(req: OmniDiffusionRequest) -> bool: + def _is_t2i_request(req: DiffusionRequestBatch) -> bool: """Return whether request-level modalities select image output. Only ``"image"`` switches Cosmos3 into T2I. ``"video"`` and omitted @@ -2828,16 +2828,16 @@ def _forward_transfer( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, ) -> DiffusionOutput: pipeline_start = time.time() # --- Parse request --- + prompt_data = req.prompts[0] if req.prompts else "" if len(req.prompts) > 1: raise ValueError("Cosmos3OmniDiffusersPipeline currently supports a single prompt per request.") sp = req.sampling_params - prompt_data = req.prompts[0] robolab_inputs = self._build_robolab_policy_inputs(sp, prompt_data, getattr(req, "request_id", None)) if robolab_inputs is not None: return self._forward_robolab_policy(sp, robolab_inputs, pipeline_start) @@ -3217,12 +3217,14 @@ def _run_diffusion(start_latents): if action_latents is None or raw_action_dim is None or domain_id is None: raise ValueError("Cosmos3 action generation finished without action latents.") action = action_latents[:, :, :raw_action_dim].detach().cpu() - custom_action_output: dict[str, Any] = { - "action": action, - "raw_action_dim": raw_action_dim, - "action_mode": action_mode, - "domain_id": domain_id, - } - return DiffusionOutput(output={"video": video}, custom_output=custom_action_output) + return DiffusionOutput( + output={"video": video}, + custom_output={ + "action": action, + "raw_action_dim": raw_action_dim, + "action_mode": action_mode, + "domain_id": domain_id, + }, + ) return DiffusionOutput(output={"image": video} if is_t2i else {"video": video}) diff --git a/vllm_omni/diffusion/models/diffusers_adapter/pipeline_diffusers_adapter.py b/vllm_omni/diffusion/models/diffusers_adapter/pipeline_diffusers_adapter.py index 8163f785436..4d48053a603 100644 --- a/vllm_omni/diffusion/models/diffusers_adapter/pipeline_diffusers_adapter.py +++ b/vllm_omni/diffusion/models/diffusers_adapter/pipeline_diffusers_adapter.py @@ -25,7 +25,7 @@ from vllm_omni.diffusion.data import DiffusionOutput, OmniDiffusionConfig from vllm_omni.diffusion.models.diffusers_adapter.pipeline_utils import BasePipelineUtils, get_pipeline_utils from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniPromptType, OmniTextPrompt from vllm_omni.platforms import current_omni_platform @@ -46,6 +46,7 @@ class DiffusersAdapterPipeline(nn.Module, DiffusionPipelineProfilerMixin): batching mode. """ + supports_request_batch = False supports_step_execution: bool = False def __init__(self, *, od_config: OmniDiffusionConfig, device: torch.device | None = None): @@ -151,9 +152,8 @@ def post_decode(self, **_: Any) -> Any: # Forward pass # ------------------------------------------------------------------ - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: """Full delegation to diffusers ``pipeline.__call__()``.""" - kwargs = self._build_call_kwargs(req) logger.debug(f"Calling diffusers pipeline with kwargs: {kwargs}") @@ -280,8 +280,8 @@ def _set_attention_backend(self) -> None: f"{dict(zip(attention_backend_attempts, attempt_errors))}" ) - def _build_call_kwargs(self, req: OmniDiffusionRequest) -> dict[str, Any]: - """Translate ``OmniDiffusionRequest`` into diffusers ``__call__`` kwargs.""" + def _build_call_kwargs(self, req: DiffusionRequestBatch) -> dict[str, Any]: + """Translate a ``DiffusionRequestBatch`` into diffusers ``__call__`` kwargs.""" sampling = req.sampling_params input_kwargs = self._extract_input(req.prompts) diff --git a/vllm_omni/diffusion/models/dmd2/mixin.py b/vllm_omni/diffusion/models/dmd2/mixin.py index 48e5719259c..873120fc183 100644 --- a/vllm_omni/diffusion/models/dmd2/mixin.py +++ b/vllm_omni/diffusion/models/dmd2/mixin.py @@ -10,7 +10,7 @@ from vllm_omni.diffusion.models.dmd2.config import DMD2Config from vllm_omni.diffusion.models.schedulers import DMD2EulerScheduler from vllm_omni.diffusion.models.utils import _load_json -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch logger = logging.getLogger(__name__) @@ -35,8 +35,8 @@ def __init_dmd2__(self) -> None: stochastic_sampling=(self.dmd2_config.solver == "sde"), ) - def _sanitize_dmd2_request(self, req: OmniDiffusionRequest) -> None: - """Sanitize CFG-related fields in-place. Mutates req.sampling_params and req.prompts.""" + def _sanitize_dmd2_request(self, req) -> None: + """Sanitize CFG-related fields in-place. Works with both OmniDiffusionRequest and DiffusionRequestBatch.""" sp = req.sampling_params if sp.num_inference_steps and sp.num_inference_steps != self.dmd2_config.num_inference_steps: @@ -75,15 +75,15 @@ def _sanitize_dmd2_request(self, req: OmniDiffusionRequest) -> None: logger.warning("DMD2: ignoring extra_args.%s.", key) extra_args.pop(key) - fixed = [] - for p in req.prompts: + # Strip negative_prompt from each request's prompt in-place. + requests = req.requests if hasattr(req, "requests") else [req] + for request in requests: + p = request.prompt if isinstance(p, dict) and "negative_prompt" in p: logger.warning("DMD2: ignoring negative_prompt.") - p = {k: v for k, v in p.items() if k != "negative_prompt"} - fixed.append(p) - req.prompts = fixed + request.prompt = {k: v for k, v in p.items() if k != "negative_prompt"} - def forward(self, req: OmniDiffusionRequest, **kwargs) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch, **kwargs) -> list[DiffusionOutput]: self._sanitize_dmd2_request(req) kwargs.pop("guidance_scale", None) kwargs.pop("num_inference_steps", None) diff --git a/vllm_omni/diffusion/models/dreamid_omni/pipeline_dreamid_omni.py b/vllm_omni/diffusion/models/dreamid_omni/pipeline_dreamid_omni.py index 1b764e4a5d5..5555e370814 100644 --- a/vllm_omni/diffusion/models/dreamid_omni/pipeline_dreamid_omni.py +++ b/vllm_omni/diffusion/models/dreamid_omni/pipeline_dreamid_omni.py @@ -25,7 +25,7 @@ from vllm_omni.diffusion.distributed.utils import get_local_device from vllm_omni.diffusion.model_loader.diffusers_loader import DiffusersPipelineLoader from vllm_omni.diffusion.models.interface import SupportAudioInput, SupportImageInput, SupportsComponentDiscovery -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch try: from dreamid_omni.utils.divisible_crop import DivisibleCrop @@ -482,7 +482,7 @@ def encode_prompt( def forward( self, - request: OmniDiffusionRequest, + request: DiffusionRequestBatch, **kwargs, ) -> DiffusionOutput: """Main forward pass for DreamID-Omni pipeline for R2AV task.""" diff --git a/vllm_omni/diffusion/models/dreamzero/pipeline_dreamzero.py b/vllm_omni/diffusion/models/dreamzero/pipeline_dreamzero.py index bbff1237364..cf2c0867efc 100644 --- a/vllm_omni/diffusion/models/dreamzero/pipeline_dreamzero.py +++ b/vllm_omni/diffusion/models/dreamzero/pipeline_dreamzero.py @@ -52,7 +52,7 @@ DEFAULT_SIGMA_SHIFT, ) from vllm_omni.diffusion.models.schedulers.scheduling_flow_unipc_multistep import FlowUniPCMultistepScheduler -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.platforms import current_omni_platform logger = logging.getLogger(__name__) @@ -898,7 +898,7 @@ def _transform_robot_obs(self, robot_obs: dict): return transform, transform.transform_input(robot_obs) @torch.no_grad() - def forward(self, req: OmniDiffusionRequest, **kwargs) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch, **kwargs) -> DiffusionOutput: """Full inference step. Called by DiffusionEngine.step().""" extra_args = req.sampling_params.extra_args or {} robot_obs = extra_args.get("robot_obs") diff --git a/vllm_omni/diffusion/models/ernie_image/pipeline_ernie_image.py b/vllm_omni/diffusion/models/ernie_image/pipeline_ernie_image.py index cd60cc93b27..82f534c312d 100644 --- a/vllm_omni/diffusion/models/ernie_image/pipeline_ernie_image.py +++ b/vllm_omni/diffusion/models/ernie_image/pipeline_ernie_image.py @@ -27,6 +27,7 @@ from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import download_weights_from_hf_specific logger = init_logger(__name__) @@ -392,7 +393,7 @@ def check_inputs( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = "", height: int = 1024, diff --git a/vllm_omni/diffusion/models/flux/pipeline_flux.py b/vllm_omni/diffusion/models/flux/pipeline_flux.py index ecbc0536d48..bec8d3a14d1 100644 --- a/vllm_omni/diffusion/models/flux/pipeline_flux.py +++ b/vllm_omni/diffusion/models/flux/pipeline_flux.py @@ -31,8 +31,8 @@ from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery from vllm_omni.diffusion.models.t5_encoder import T5EncoderModel from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import download_weights_from_hf_specific logger = logging.getLogger(__name__) @@ -70,6 +70,8 @@ def post_process_func(images: torch.Tensor): class FluxPipeline( nn.Module, FluxPipelineMixin, CFGParallelMixin, DiffusionPipelineProfilerMixin, SupportsComponentDiscovery ): + supports_request_batch = True + _dit_modules: ClassVar[list[str]] = ["transformer"] _encoder_modules: ClassVar[list[str]] = ["text_encoder", "text_encoder_2"] _vae_modules: ClassVar[list[str]] = ["vae"] @@ -502,7 +504,7 @@ def check_cfg_parallel_validity(self, true_cfg_scale: float, has_neg_prompt: boo def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, prompt_2: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, @@ -525,30 +527,54 @@ def forward( joint_attention_kwargs: dict[str, Any] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 512, - ): + ) -> list[DiffusionOutput]: """Forward pass for flux.""" # TODO: In online mode, sometimes it receives [{"negative_prompt": None}, {...}], so cannot use .get("...", "") # TODO: May be some data formatting operations on the API side. Hack for now. + sampling_params_list = req.sampling_params_list + common_sampling_params = sampling_params_list[0] prompt = [p if isinstance(p, str) else (p.get("prompt") or "") for p in req.prompts] or prompt if all(isinstance(p, str) or p.get("negative_prompt") is None for p in req.prompts): negative_prompt = None elif req.prompts: negative_prompt = ["" if isinstance(p, str) else (p.get("negative_prompt") or "") for p in req.prompts] - height = req.sampling_params.height or self.default_sample_size * self.vae_scale_factor - width = req.sampling_params.width or self.default_sample_size * self.vae_scale_factor - num_inference_steps = req.sampling_params.num_inference_steps or num_inference_steps - sigmas = req.sampling_params.sigmas or sigmas + height = common_sampling_params.height or self.default_sample_size * self.vae_scale_factor + width = common_sampling_params.width or self.default_sample_size * self.vae_scale_factor + num_inference_steps = common_sampling_params.num_inference_steps or num_inference_steps + sigmas = common_sampling_params.sigmas or sigmas guidance_scale = ( - req.sampling_params.guidance_scale if req.sampling_params.guidance_scale is not None else guidance_scale + common_sampling_params.guidance_scale + if common_sampling_params.guidance_scale is not None + else guidance_scale ) - generator = req.sampling_params.generator or generator - true_cfg_scale = req.sampling_params.true_cfg_scale or true_cfg_scale + true_cfg_scale = common_sampling_params.true_cfg_scale or true_cfg_scale num_images_per_prompt = ( - req.sampling_params.num_outputs_per_prompt - if req.sampling_params.num_outputs_per_prompt > 0 + common_sampling_params.num_outputs_per_prompt + if common_sampling_params.num_outputs_per_prompt > 0 else num_images_per_prompt ) + generator = req.collate_request_generators(num_images_per_prompt, generator) + latents = req.collate_request_tensors("latents", latents) + prompt_fields = DiffusionRequestBatch.collate_prompt_field_map( + req.prompts, + { + "prompt_embeds": prompt_embeds, + "negative_prompt_embeds": negative_prompt_embeds, + "pooled_prompt_embeds": pooled_prompt_embeds, + "negative_pooled_prompt_embeds": negative_pooled_prompt_embeds, + }, + ) + prompt_embeds = prompt_fields["prompt_embeds"] + negative_prompt_embeds = prompt_fields["negative_prompt_embeds"] + pooled_prompt_embeds = prompt_fields["pooled_prompt_embeds"] + negative_pooled_prompt_embeds = prompt_fields["negative_pooled_prompt_embeds"] + if prompt_embeds is not None: + prompt = None + prompt_2 = None + if negative_prompt_embeds is not None: + negative_prompt = None + negative_prompt_2 = None # 1. Check inputs. Raise error if not correct self.check_inputs( @@ -666,9 +692,17 @@ def forward( latents = (latents / self.vae.config.scaling_factor) + self.vae.config.shift_factor image = self.vae.decode(latents, return_dict=False)[0] - return DiffusionOutput( - output=image, stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None - ) + stage_durations = self.stage_durations if hasattr(self, "stage_durations") else None + n = num_images_per_prompt + if req.num_reqs == 1: + return [DiffusionOutput(output=image, stage_durations=stage_durations)] + return [ + DiffusionOutput( + output=image[i * n : (i + 1) * n] if isinstance(image, list) else image[i * n : (i + 1) * n], + stage_durations=stage_durations, + ) + for i in range(req.num_reqs) + ] def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) @@ -681,3 +715,6 @@ class FluxDMD2Pipeline(DMD2PipelineMixin, FluxPipeline): def __init__(self, *, od_config: OmniDiffusionConfig, prefix: str = ""): super().__init__(od_config=od_config, prefix=prefix) self.__init_dmd2__() + + def forward(self, req: DiffusionRequestBatch, **kwargs) -> list[DiffusionOutput]: + return super().forward(req, **kwargs) diff --git a/vllm_omni/diffusion/models/flux/pipeline_flux_kontext.py b/vllm_omni/diffusion/models/flux/pipeline_flux_kontext.py index d5790983721..0b3b80cbd2a 100644 --- a/vllm_omni/diffusion/models/flux/pipeline_flux_kontext.py +++ b/vllm_omni/diffusion/models/flux/pipeline_flux_kontext.py @@ -33,8 +33,8 @@ from vllm_omni.diffusion.models.interface import SupportImageInput, SupportsComponentDiscovery from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.logger import init_logger logger = init_logger(__name__) @@ -468,7 +468,7 @@ def num_timesteps(self): @torch.no_grad() def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, image: PIL.Image.Image | list[PIL.Image.Image] | None = None, prompt: str | list[str] | None = None, prompt_2: str | list[str] | None = None, diff --git a/vllm_omni/diffusion/models/flux2/pipeline_flux2.py b/vllm_omni/diffusion/models/flux2/pipeline_flux2.py index a92d7e55c43..f88e582e3f8 100644 --- a/vllm_omni/diffusion/models/flux2/pipeline_flux2.py +++ b/vllm_omni/diffusion/models/flux2/pipeline_flux2.py @@ -34,8 +34,8 @@ from vllm_omni.diffusion.models.mistral_encoder import MistralEncoderModel from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import download_weights_from_hf_specific logger = logging.getLogger(__name__) @@ -859,7 +859,7 @@ def check_cfg_parallel_validity(self, true_cfg_scale: float, has_neg_prompt: boo def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, image: PIL.Image.Image | list[PIL.Image.Image] | None = None, prompt: str | list[str] | None = None, height: int | None = None, diff --git a/vllm_omni/diffusion/models/flux2_klein/pipeline_flux2_klein.py b/vllm_omni/diffusion/models/flux2_klein/pipeline_flux2_klein.py index 74598319fae..81a6f50f107 100644 --- a/vllm_omni/diffusion/models/flux2_klein/pipeline_flux2_klein.py +++ b/vllm_omni/diffusion/models/flux2_klein/pipeline_flux2_klein.py @@ -45,8 +45,8 @@ ) from vllm_omni.diffusion.models.interface import SupportImageInput, SupportsComponentDiscovery from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import download_weights_from_hf_specific logger = init_logger(__name__) @@ -752,7 +752,7 @@ def interrupt(self): def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, image: PIL.Image.Image | list[PIL.Image.Image] | None = None, reference_image: PIL.Image.Image | list[PIL.Image.Image] | None = None, mask_image: PIL.Image.Image | list[PIL.Image.Image] | None = None, diff --git a/vllm_omni/diffusion/models/glm_image/pipeline_glm_image.py b/vllm_omni/diffusion/models/glm_image/pipeline_glm_image.py index b5dcb324d8f..9f3d2940090 100644 --- a/vllm_omni/diffusion/models/glm_image/pipeline_glm_image.py +++ b/vllm_omni/diffusion/models/glm_image/pipeline_glm_image.py @@ -49,6 +49,7 @@ from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, @@ -81,58 +82,54 @@ def get_glm_image_pre_process_func(od_config: OmniDiffusionConfig): def pre_process_func(request: OmniDiffusionRequest): """Pre-process condition images for Image Edit mode.""" - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - if raw_image is None: - # Text-to-image mode, no preprocessing needed - continue - - if not isinstance(raw_image, list): - raw_image = [raw_image] - images = [ - PIL.Image.open(im) if isinstance(im, str) else cast(PIL.Image.Image | np.ndarray | torch.Tensor, im) - for im in raw_image - ] - - preprocessed = [] - height, width = None, None - - for img in images: - if isinstance(img, PIL.Image.Image): - img_h, img_w = img.size[::-1] # PIL is (width, height) - else: - img_h, img_w = img.shape[:2] - - # Align to multiple of vae_scale_factor * patch_size - multiple_of = vae_scale_factor * patch_size - img_h = (img_h // multiple_of) * multiple_of - img_w = (img_w // multiple_of) * multiple_of - - processed = image_processor.preprocess(img, height=img_h, width=img_w) - preprocessed.append(processed) - - # Use first image dimensions as default - if height is None: - height, width = img_h, img_w - - # Store in request - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt, additional_information={}) - elif "additional_information" not in prompt: - prompt["additional_information"] = {} - prompt["additional_information"]["preprocessed_image"] = processed # type: ignore - prompt["additional_information"]["prompt_image"] = images # type: ignore - request.prompts[i] = prompt - if request.sampling_params.height is None: - request.sampling_params.height = height - if request.sampling_params.width is None: - request.sampling_params.width = width + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + if raw_image is None: + # Text-to-image mode, no preprocessing needed + return request + + if not isinstance(raw_image, list): + raw_image = [raw_image] + images = [ + PIL.Image.open(im) if isinstance(im, str) else cast(PIL.Image.Image | np.ndarray | torch.Tensor, im) + for im in raw_image + ] + + preprocessed = [] + height, width = None, None + + for img in images: + if isinstance(img, PIL.Image.Image): + img_h, img_w = img.size[::-1] # PIL is (width, height) + else: + img_h, img_w = img.shape[:2] + + # Align to multiple of vae_scale_factor * patch_size + multiple_of = vae_scale_factor * patch_size + img_h = (img_h // multiple_of) * multiple_of + img_w = (img_w // multiple_of) * multiple_of + + processed = image_processor.preprocess(img, height=img_h, width=img_w) + preprocessed.append(processed) + + # Use first image dimensions as default + if height is None: + height, width = img_h, img_w + + # Store in request + prompt["additional_information"]["preprocessed_image"] = processed # type: ignore + prompt["additional_information"]["prompt_image"] = images # type: ignore + request.prompt = prompt + if request.sampling_params.height is None: + request.sampling_params.height = height + if request.sampling_params.width is None: + request.sampling_params.width = width return request @@ -674,7 +671,7 @@ def _prepare_condition_image_kv_cache( return kv_caches @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: """ Main generation forward pass. diff --git a/vllm_omni/diffusion/models/gr00t/pipeline_gr00t.py b/vllm_omni/diffusion/models/gr00t/pipeline_gr00t.py index 98565a93faf..1154dcf7145 100644 --- a/vllm_omni/diffusion/models/gr00t/pipeline_gr00t.py +++ b/vllm_omni/diffusion/models/gr00t/pipeline_gr00t.py @@ -11,7 +11,8 @@ from vllm_omni.diffusion.data import DiffusionOutput, OmniDiffusionConfig from vllm_omni.diffusion.models.gr00t.policy import Gr00tPolicy -from vllm_omni.diffusion.request import DUMMY_DIFFUSION_REQUEST_ID, OmniDiffusionRequest +from vllm_omni.diffusion.request import DUMMY_DIFFUSION_REQUEST_ID +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch logger = init_logger(__name__) @@ -107,7 +108,7 @@ def _dummy_actions(self) -> dict[str, np.ndarray]: return actions @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest, **kwargs) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch, **kwargs) -> DiffusionOutput: del kwargs extra_args = req.sampling_params.extra_args or {} robot_obs = extra_args.get("robot_obs") diff --git a/vllm_omni/diffusion/models/helios/pipeline_helios.py b/vllm_omni/diffusion/models/helios/pipeline_helios.py index 9acd13b53b8..8ee7c8cf1ed 100644 --- a/vllm_omni/diffusion/models/helios/pipeline_helios.py +++ b/vllm_omni/diffusion/models/helios/pipeline_helios.py @@ -30,6 +30,7 @@ from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.platforms import current_omni_platform if TYPE_CHECKING: @@ -284,7 +285,7 @@ def prepare_encode( """Initialize Helios request state for chunk-wise step execution.""" del kwargs req = OmniDiffusionRequest( - prompts=state.prompts or [], + prompt=state.prompt, sampling_params=state.sampling, request_id=state.request_id, ) @@ -905,7 +906,7 @@ def post_decode( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | None = None, negative_prompt: str | None = None, height: int = 384, diff --git a/vllm_omni/diffusion/models/hidream_image/pipeline_hidream_image.py b/vllm_omni/diffusion/models/hidream_image/pipeline_hidream_image.py index 135e21e1a73..af156b34bf8 100644 --- a/vllm_omni/diffusion/models/hidream_image/pipeline_hidream_image.py +++ b/vllm_omni/diffusion/models/hidream_image/pipeline_hidream_image.py @@ -35,8 +35,8 @@ from vllm_omni.diffusion.models.hidream_image import HiDreamImageTransformer2DModel from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import download_weights_from_hf_specific logger = logging.get_logger(__name__) @@ -866,7 +866,7 @@ def interrupt(self): def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] = None, prompt_2: str | list[str] | None = None, prompt_3: str | list[str] | None = None, @@ -896,7 +896,7 @@ def forward( callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 128, **kwargs, - ): + ) -> DiffusionOutput: extracted_prompt, negative_prompt = self._extract_prompts(req.prompts) prompt = extracted_prompt or prompt diff --git a/vllm_omni/diffusion/models/hunyuan_image3/pipeline_hunyuan_image3.py b/vllm_omni/diffusion/models/hunyuan_image3/pipeline_hunyuan_image3.py index 483f345fe25..73e9ff6664d 100644 --- a/vllm_omni/diffusion/models/hunyuan_image3/pipeline_hunyuan_image3.py +++ b/vllm_omni/diffusion/models/hunyuan_image3/pipeline_hunyuan_image3.py @@ -29,6 +29,7 @@ DiffusionPipelineProfilerMixin, ) from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.model_executor.models.hunyuan_image3.siglip2 import Siglip2VisionTransformer @@ -300,32 +301,32 @@ def _build_cond_joint_image(raw_image: Any) -> dict[str, Any]: ) def pre_process_func(request: OmniDiffusionRequest): - for i, prompt in enumerate(request.prompts): - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - multi_modal_data = prompt.get("multi_modal_data") or {} - raw_images = multi_modal_data.get("image") - if raw_images is None: - raw_images = prompt.get("pil_image") - has_images = raw_images is not None and (not isinstance(raw_images, list) or len(raw_images) > 0) - if has_images: - image_list = raw_images if isinstance(raw_images, list) else [raw_images] - cond_image_infos = [_build_cond_joint_image(image) for image in image_list] - prompt["additional_information"]["batch_cond_image_info"] = cond_image_infos - - bridge_h = prompt.get("height") if isinstance(prompt, dict) else None - bridge_w = prompt.get("width") if isinstance(prompt, dict) else None - first_image_w, first_image_h = _to_pil_image(image_list[0]).size - if request.sampling_params.width is None: - request.sampling_params.width = int(bridge_w or first_image_w) - if request.sampling_params.height is None: - request.sampling_params.height = int(bridge_h or first_image_h) - - request.prompts[i] = prompt + prompt = request.prompt + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + multi_modal_data = prompt.get("multi_modal_data") or {} + raw_images = multi_modal_data.get("image") + if raw_images is None: + raw_images = prompt.get("pil_image") + has_images = raw_images is not None and (not isinstance(raw_images, list) or len(raw_images) > 0) + if has_images: + image_list = raw_images if isinstance(raw_images, list) else [raw_images] + cond_image_infos = [_build_cond_joint_image(image) for image in image_list] + prompt["additional_information"]["batch_cond_image_info"] = cond_image_infos + + bridge_h = prompt.get("height") if isinstance(prompt, dict) else None + bridge_w = prompt.get("width") if isinstance(prompt, dict) else None + first_image_w, first_image_h = _to_pil_image(image_list[0]).size + if request.sampling_params.width is None: + request.sampling_params.width = int(bridge_w or first_image_w) + if request.sampling_params.height is None: + request.sampling_params.height = int(bridge_h or first_image_h) + + request.prompt = prompt return request @@ -340,6 +341,7 @@ class HunyuanImage3Pipeline( DiffusionPipelineProfilerMixin, ): supports_step_execution: ClassVar[bool] = True + supports_request_batch = False support_image_input = True _dit_modules: ClassVar[list[str]] = ["model"] _encoder_modules: ClassVar[list[str]] = ["vision_model"] @@ -485,9 +487,9 @@ def pipeline(self): return self._pipeline def _validate_step_request(self, state: "DiffusionRequestState") -> None: - prompts = state.prompts or [] + prompt = state.prompt sampling = state.sampling - if len(prompts) != 1: + if prompt is None: raise ValueError("HunyuanImage3 step execution currently requires exactly one prompt per request.") if sampling.timesteps is not None or sampling.sigmas is not None: raise ValueError("HunyuanImage3 step execution does not support custom timesteps or sigmas yet.") @@ -586,7 +588,7 @@ def _extract_step_prompt_inputs( ) -> tuple[list[str], list[str | None], str | None, list[list[JointImageInfo]] | None, str]: sampling = state.sampling return self._extract_prompt_inputs( - state.prompts or [], + [state.prompt] if state.prompt is not None else [], getattr(sampling, "extra_args", {}) or {}, request_id=state.request_id, allow_cond_image=False, @@ -2241,7 +2243,7 @@ def post_decode( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] = "", image_size="auto", height: int = 1024, @@ -2299,11 +2301,13 @@ def forward( model_inputs.update(ar_kv_kwargs) outputs = self._generate(**model_inputs, **kwargs) + image = outputs[0] custom_output = {} if any(t is not None for t in cot_text_list): custom_output["ar_generated_text"] = cot_text_list[0] if len(cot_text_list) == 1 else cot_text_list + stage_durations = self.stage_durations if hasattr(self, "stage_durations") else None return DiffusionOutput( - output=outputs[0], + output=image, custom_output=custom_output, - stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None, + stage_durations=stage_durations, ) diff --git a/vllm_omni/diffusion/models/hunyuan_video/pipeline_hunyuan_video_1_5.py b/vllm_omni/diffusion/models/hunyuan_video/pipeline_hunyuan_video_1_5.py index aa5b9b83cd2..942e0b28435 100644 --- a/vllm_omni/diffusion/models/hunyuan_video/pipeline_hunyuan_video_1_5.py +++ b/vllm_omni/diffusion/models/hunyuan_video/pipeline_hunyuan_video_1_5.py @@ -31,8 +31,8 @@ from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin from vllm_omni.diffusion.models.t5_encoder import T5EncoderModel from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.platforms import current_omni_platform logger = logging.getLogger(__name__) @@ -394,7 +394,7 @@ def predict_noise(self, **kwargs: Any) -> torch.Tensor: def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, num_inference_steps: int = 50, guidance_scale: float = 6.0, height: int = 480, diff --git a/vllm_omni/diffusion/models/hunyuan_video/pipeline_hunyuan_video_1_5_i2v.py b/vllm_omni/diffusion/models/hunyuan_video/pipeline_hunyuan_video_1_5_i2v.py index 3a3cc46dd22..1281a9e39d6 100644 --- a/vllm_omni/diffusion/models/hunyuan_video/pipeline_hunyuan_video_1_5_i2v.py +++ b/vllm_omni/diffusion/models/hunyuan_video/pipeline_hunyuan_video_1_5_i2v.py @@ -46,6 +46,7 @@ from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.platforms import current_omni_platform logger = logging.getLogger(__name__) @@ -62,10 +63,10 @@ def get_hunyuan_video_15_i2v_pre_process_func(od_config: OmniDiffusionConfig): divisor = 16 # Must be divisible by VAE spatial compression def pre_process_func(req: OmniDiffusionRequest) -> OmniDiffusionRequest: - if not req.prompts: + if not req.prompt: return req - prompt_data = req.prompts[0] + prompt_data = req.prompt if isinstance(prompt_data, str): return req @@ -474,7 +475,7 @@ def predict_noise(self, **kwargs: Any) -> torch.Tensor: def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, num_inference_steps: int = 50, guidance_scale: float = 6.0, height: int = 480, diff --git a/vllm_omni/diffusion/models/internvla_a1/pipeline_internvla_a1.py b/vllm_omni/diffusion/models/internvla_a1/pipeline_internvla_a1.py index b761ae8c668..6c28e760c5a 100644 --- a/vllm_omni/diffusion/models/internvla_a1/pipeline_internvla_a1.py +++ b/vllm_omni/diffusion/models/internvla_a1/pipeline_internvla_a1.py @@ -16,7 +16,7 @@ DiffusionPipelineProfilerMixin, wrap_methods_by_paths, ) -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from .config import ( DEFAULT_QWEN3_VL_MODEL, @@ -228,7 +228,7 @@ def _predict_actions( ) @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: if len(req.prompts) > 1: logger.warning("InternVLAA1Pipeline only supports a single prompt/request; taking the first sample.") extra_args = getattr(req.sampling_params, "extra_args", {}) or {} diff --git a/vllm_omni/diffusion/models/longcat_image/pipeline_longcat_image.py b/vllm_omni/diffusion/models/longcat_image/pipeline_longcat_image.py index e8c23c1f093..2884adfcaaa 100644 --- a/vllm_omni/diffusion/models/longcat_image/pipeline_longcat_image.py +++ b/vllm_omni/diffusion/models/longcat_image/pipeline_longcat_image.py @@ -31,7 +31,7 @@ from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery from vllm_omni.diffusion.models.longcat_image.longcat_image_transformer import LongCatImageTransformer2DModel from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, ) @@ -203,6 +203,8 @@ def get_prompt_language(prompt): class LongCatImagePipeline(nn.Module, CFGParallelMixin, DiffusionPipelineProfilerMixin, SupportsComponentDiscovery): + supports_request_batch = False + _dit_modules: ClassVar[list[str]] = ["transformer"] _encoder_modules: ClassVar[list[str]] = ["text_encoder"] _vae_modules: ClassVar[list[str]] = ["vae"] @@ -495,7 +497,7 @@ def cfg_normalize_function(self, noise_pred, comb_pred, cfg_renorm_min=0.0): def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, height: int | None = None, @@ -693,9 +695,8 @@ def forward( image = self.vae.decode(latents, return_dict=False)[0] - return DiffusionOutput( - output=image, stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None - ) + stage_durations = self.stage_durations if hasattr(self, "stage_durations") else None + return DiffusionOutput(output=image, stage_durations=stage_durations) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: """Load weights using AutoWeightsLoader for vLLM integration.""" diff --git a/vllm_omni/diffusion/models/longcat_image/pipeline_longcat_image_edit.py b/vllm_omni/diffusion/models/longcat_image/pipeline_longcat_image_edit.py index 54de7e9f981..195aa89a403 100644 --- a/vllm_omni/diffusion/models/longcat_image/pipeline_longcat_image_edit.py +++ b/vllm_omni/diffusion/models/longcat_image/pipeline_longcat_image_edit.py @@ -37,6 +37,7 @@ from vllm_omni.diffusion.models.longcat_image.pipeline_longcat_image import calculate_shift from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, @@ -66,50 +67,48 @@ def pre_process_func( request: OmniDiffusionRequest, ): """Pre-process requests for LongCatImageEditPipeline.""" - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - if not raw_image: # None or empty list - raise ValueError("""Received no input image. This model requires one input image to run.""") - elif isinstance(raw_image, list): - if len(raw_image) > 1: - raise ValueError( - """Received multiple input images. Only a single image is supported by this model.""" - ) - else: - raw_image = raw_image[0] - - if isinstance(raw_image, str): - image = PIL.Image.open(raw_image) + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + if not raw_image: # None or empty list + raise ValueError("""Received no input image. This model requires one input image to run.""") + elif isinstance(raw_image, list): + if len(raw_image) > 1: + raise ValueError("""Received multiple input images. Only a single image is supported by this model.""") else: - image = cast(PIL.Image.Image | torch.Tensor | np.ndarray, raw_image) + raw_image = raw_image[0] - image_size = image.size - calculated_width, calculated_height = calculate_dimensions(1024 * 1024, image_size[0] * 1.0 / image_size[1]) - height = request.sampling_params.height or calculated_height - width = request.sampling_params.width or calculated_width - - # Store calculated dimensions in request - prompt["additional_information"]["calculated_height"] = calculated_height - prompt["additional_information"]["calculated_width"] = calculated_width - request.sampling_params.height = height - request.sampling_params.width = width - - # Preprocess image - if image is not None and not (isinstance(image, torch.Tensor) and image.size(1) == latent_channels): - image = image_processor.resize(image, calculated_height, calculated_width) - prompt_image = image_processor.resize(image, calculated_height // 2, calculated_width // 2) - image = image_processor.preprocess(image, calculated_height, calculated_width) - - # Store preprocessed image and prompt image in request - prompt["additional_information"]["preprocessed_image"] = image - prompt["additional_information"]["prompt_image"] = prompt_image - request.prompts[i] = prompt + if isinstance(raw_image, str): + image = PIL.Image.open(raw_image) + else: + image = cast(PIL.Image.Image | torch.Tensor | np.ndarray, raw_image) + + image_size = image.size + calculated_width, calculated_height = calculate_dimensions(1024 * 1024, image_size[0] * 1.0 / image_size[1]) + height = request.sampling_params.height or calculated_height + width = request.sampling_params.width or calculated_width + + # Store calculated dimensions in request + prompt["additional_information"]["calculated_height"] = calculated_height + prompt["additional_information"]["calculated_width"] = calculated_width + request.sampling_params.height = height + request.sampling_params.width = width + + # Preprocess image + if image is not None and not (isinstance(image, torch.Tensor) and image.size(1) == latent_channels): + image = image_processor.resize(image, calculated_height, calculated_width) + prompt_image = image_processor.resize(image, calculated_height // 2, calculated_width // 2) + image = image_processor.preprocess(image, calculated_height, calculated_width) + + # Store preprocessed image and prompt image in request + prompt["additional_information"]["preprocessed_image"] = image + prompt["additional_information"]["prompt_image"] = prompt_image + request.prompt = prompt return request return pre_process_func @@ -545,7 +544,7 @@ def check_inputs( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, image: PIL.Image.Image | torch.Tensor | None = None, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, @@ -560,7 +559,7 @@ def forward( output_type: str | None = "pil", return_dict: bool = True, joint_attention_kwargs: dict[str, Any] | None = None, - ): + ) -> DiffusionOutput: # TODO: In online mode, sometimes it receives [{"negative_prompt": None}, {...}], so cannot use .get("...", "") # TODO: May be some data formatting operations on the API side. Hack for now. if len(req.prompts) > 1: diff --git a/vllm_omni/diffusion/models/ltx2/pipeline_ltx2.py b/vllm_omni/diffusion/models/ltx2/pipeline_ltx2.py index 85927354fd0..daefa850d18 100644 --- a/vllm_omni/diffusion/models/ltx2/pipeline_ltx2.py +++ b/vllm_omni/diffusion/models/ltx2/pipeline_ltx2.py @@ -40,7 +40,7 @@ from vllm_omni.diffusion.models.dmd2 import DMD2PipelineMixin from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.lora.request import LoRARequest from .ltx2_transformer import LTX2VideoTransformer3DModel @@ -158,6 +158,8 @@ def step(self, noise_pred, t, latents, return_dict=False, generator=None): class LTX2Pipeline(nn.Module, CFGParallelMixin, ProgressBarMixin, SupportsComponentDiscovery): + supports_request_batch = False + _dit_modules: ClassVar[list[str]] = ["transformer"] _encoder_modules: ClassVar[list[str]] = ["text_encoder"] _vae_modules: ClassVar[list[str]] = ["vae", "audio_vae"] @@ -746,7 +748,7 @@ def _synchronize_cfg_parallel_step_output( @torch.no_grad() def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, height: int | None = None, @@ -1136,9 +1138,6 @@ def forward( generated_mel_spectrograms = self.audio_vae.decode(audio_latents, return_dict=False)[0] audio = self.vocoder(generated_mel_spectrograms) - if not return_dict: - return DiffusionOutput(output=(video, audio)) - return DiffusionOutput(output=(video, audio)) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: @@ -1150,6 +1149,7 @@ class LTX2TwoStagesPipeline(nn.Module, SupportsComponentDiscovery): """LTX2TwoStagesPipeline is for two stages image to video generation""" dummy_run_num_frames = 2 + supports_request_batch = False _dit_modules: ClassVar[list[str]] = ["pipe.transformer"] _encoder_modules: ClassVar[list[str]] = ["pipe.text_encoder"] @@ -1199,7 +1199,7 @@ def __init__( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, height: int | None = None, @@ -1225,8 +1225,8 @@ def forward( return_dict: bool = True, attention_kwargs: dict[str, Any] | None = None, max_sequence_length: int | None = None, - ): - video_latent, audio_latent = self.pipe( + ) -> DiffusionOutput: + stage1_output = self.pipe( req=req, prompt=prompt, negative_prompt=negative_prompt, @@ -1254,7 +1254,8 @@ def forward( return_dict=return_dict, attention_kwargs=attention_kwargs, max_sequence_length=max_sequence_length, - ).output + ) + video_latent, audio_latent = stage1_output.output upscaled_video_latent = self.upsample_pipe( latents=video_latent, @@ -1286,7 +1287,7 @@ def forward( stage_2_req.sampling_params = req.sampling_params.clone() stage_2_req.sampling_params.num_inference_steps = 3 - video, audio = self.pipe( + stage2_output = self.pipe( req=stage_2_req, latents=upscaled_video_latent, audio_latents=audio_latent, @@ -1298,7 +1299,8 @@ def forward( generator=generator, output_type="np", return_dict=False, - ).output + ) + video, audio = stage2_output.output return DiffusionOutput(output=(video, audio)) diff --git a/vllm_omni/diffusion/models/ltx2/pipeline_ltx2_3.py b/vllm_omni/diffusion/models/ltx2/pipeline_ltx2_3.py index f333db8ce39..9183065dc25 100644 --- a/vllm_omni/diffusion/models/ltx2/pipeline_ltx2_3.py +++ b/vllm_omni/diffusion/models/ltx2/pipeline_ltx2_3.py @@ -50,10 +50,9 @@ from vllm_omni.diffusion.models.progress_bar import ProgressBarMixin from vllm_omni.diffusion.offloader.module_collector import ModuleDiscovery from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch, split_diffusion_output_by_request from .pipeline_ltx2 import ( - _get_prompt_field, _VideoAudioScheduler, calculate_shift, create_transformer_from_config, @@ -62,6 +61,13 @@ logger = init_logger(__name__) + +def _get_audio_latents_from_sampling(sampling: Any) -> torch.Tensor | None: + if sampling.audio_latents is not None: + return sampling.audio_latents + return sampling.extra_args.get("audio_latents") + + # Try to import LTX2VocoderWithBWE (diffusers >= 0.38.0) try: from diffusers.pipelines.ltx2.vocoder import LTX2VocoderWithBWE @@ -131,6 +137,7 @@ class LTX23Pipeline( - Transformer: passes sigma for prompt_adaln """ + supports_request_batch = True # Audio is diffused jointly with video; warmup must size audio tokens. dummy_run_num_frames = 2 _dit_modules: ClassVar[list[str]] = ["transformer"] @@ -830,7 +837,7 @@ def _synchronize_cfg_parallel_step_output( @torch.no_grad() def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, height: int | None = None, @@ -856,76 +863,80 @@ def forward( return_dict: bool = True, attention_kwargs: dict[str, Any] | None = None, max_sequence_length: int | None = None, - ) -> DiffusionOutput: + ) -> list[DiffusionOutput]: # ---- Extract from request ---- + sampling_params_list = req.sampling_params_list + common_sampling_params = sampling_params_list[0] prompt = [p if isinstance(p, str) else (p.get("prompt") or "") for p in req.prompts] or prompt if all(isinstance(p, str) or p.get("negative_prompt") is None for p in req.prompts): negative_prompt = None elif req.prompts: negative_prompt = ["" if isinstance(p, str) else (p.get("negative_prompt") or "") for p in req.prompts] - height = req.sampling_params.height or height or 512 - width = req.sampling_params.width or width or 768 - num_frames = req.sampling_params.num_frames or num_frames or 121 - frame_rate = req.sampling_params.resolved_frame_rate or frame_rate or 24.0 - num_inference_steps = req.sampling_params.num_inference_steps or num_inference_steps or 40 + height = common_sampling_params.height or height or 512 + width = common_sampling_params.width or width or 768 + num_frames = common_sampling_params.num_frames or num_frames or 121 + frame_rate = common_sampling_params.resolved_frame_rate or frame_rate or 24.0 + num_inference_steps = common_sampling_params.num_inference_steps or num_inference_steps or 40 # Enforce minimum of 2 timesteps for flow matching scheduler if timesteps is None: num_inference_steps = max(int(num_inference_steps), 2) elif len(timesteps) < 2: raise ValueError("`timesteps` must contain at least 2 values for FlowMatchEulerDiscreteScheduler.") num_videos_per_prompt = ( - req.sampling_params.num_outputs_per_prompt - if req.sampling_params.num_outputs_per_prompt > 0 + common_sampling_params.num_outputs_per_prompt + if common_sampling_params.num_outputs_per_prompt > 0 else num_videos_per_prompt or 1 ) max_sequence_length = ( - req.sampling_params.max_sequence_length or max_sequence_length or self.tokenizer_max_length + common_sampling_params.max_sequence_length or max_sequence_length or self.tokenizer_max_length ) - if req.sampling_params.guidance_scale_provided: - guidance_scale = req.sampling_params.guidance_scale + if common_sampling_params.guidance_scale_provided: + guidance_scale = common_sampling_params.guidance_scale if generator is None: - generator = req.sampling_params.generator - if generator is None and req.sampling_params.seed is not None: - generator = torch.Generator(device=self.device).manual_seed(req.sampling_params.seed) - - latents = req.sampling_params.latents if req.sampling_params.latents is not None else latents - audio_latents = ( - req.sampling_params.audio_latents - if req.sampling_params.audio_latents is not None - else req.sampling_params.extra_args.get("audio_latents", audio_latents) + generator = req.collate_request_generators(num_videos_per_prompt, generator) + + latents = req.collate_request_tensors("latents", latents) + audio_latents = DiffusionRequestBatch.collate_tensors( + [_get_audio_latents_from_sampling(sampling) for sampling in sampling_params_list], + "audio_latents", + audio_latents, ) # Override with pre-computed embeddings if provided in request - req_prompt_embeds = [_get_prompt_field(p, "prompt_embeds") for p in req.prompts] - if any(p is not None for p in req_prompt_embeds): - prompt_embeds = torch.stack(req_prompt_embeds) - - req_negative_prompt_embeds = [_get_prompt_field(p, "negative_prompt_embeds") for p in req.prompts] - if any(p is not None for p in req_negative_prompt_embeds): - negative_prompt_embeds = torch.stack(req_negative_prompt_embeds) - - req_prompt_attention_masks = [ - _get_prompt_field(p, "prompt_attention_mask") or _get_prompt_field(p, "attention_mask") for p in req.prompts - ] - if any(m is not None for m in req_prompt_attention_masks): - prompt_attention_mask = torch.stack(req_prompt_attention_masks) - - req_negative_attention_masks = [ - _get_prompt_field(p, "negative_prompt_attention_mask") or _get_prompt_field(p, "negative_attention_mask") - for p in req.prompts - ] - if any(m is not None for m in req_negative_attention_masks): - negative_prompt_attention_mask = torch.stack(req_negative_attention_masks) + prompt_fields = DiffusionRequestBatch.collate_prompt_field_map( + req.prompts, + { + "prompt_embeds": prompt_embeds, + "negative_prompt_embeds": negative_prompt_embeds, + "prompt_attention_mask": prompt_attention_mask, + "negative_prompt_attention_mask": negative_prompt_attention_mask, + }, + field_aliases={ + "prompt_attention_mask": ("prompt_attention_mask", "attention_mask"), + "negative_prompt_attention_mask": ( + "negative_prompt_attention_mask", + "negative_attention_mask", + ), + }, + ) + prompt_embeds = prompt_fields["prompt_embeds"] + negative_prompt_embeds = prompt_fields["negative_prompt_embeds"] + prompt_attention_mask = prompt_fields["prompt_attention_mask"] + negative_prompt_attention_mask = prompt_fields["negative_prompt_attention_mask"] + if prompt_embeds is not None: + prompt = None + if negative_prompt_embeds is not None: + negative_prompt = None - if req.sampling_params.decode_timestep is not None: - decode_timestep = req.sampling_params.decode_timestep - if req.sampling_params.decode_noise_scale is not None: - decode_noise_scale = req.sampling_params.decode_noise_scale - if req.sampling_params.output_type is not None: - output_type = req.sampling_params.output_type + if common_sampling_params.decode_timestep is not None: + decode_timestep = common_sampling_params.decode_timestep + if common_sampling_params.decode_noise_scale is not None: + decode_noise_scale = common_sampling_params.decode_noise_scale + if common_sampling_params.output_type is not None: + output_type = common_sampling_params.output_type self.check_inputs( prompt=prompt, @@ -1273,9 +1284,13 @@ def forward( audio = self.vocoder(generated_mel_spectrograms) - return DiffusionOutput( - output=(video, audio), - stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None, + return split_diffusion_output_by_request( + DiffusionOutput( + output=(video, audio), + stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None, + ), + req, + num_outputs_per_prompt=num_videos_per_prompt, ) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: diff --git a/vllm_omni/diffusion/models/ltx2/pipeline_ltx2_image2video.py b/vllm_omni/diffusion/models/ltx2/pipeline_ltx2_image2video.py index e93f4e7447f..a21d9c785bd 100644 --- a/vllm_omni/diffusion/models/ltx2/pipeline_ltx2_image2video.py +++ b/vllm_omni/diffusion/models/ltx2/pipeline_ltx2_image2video.py @@ -27,7 +27,7 @@ from vllm_omni.diffusion.model_loader.diffusers_loader import DiffusersPipelineLoader from vllm_omni.diffusion.models.dmd2 import DMD2PipelineMixin from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.lora.request import LoRARequest from .pipeline_ltx2 import ( @@ -74,6 +74,7 @@ def step(self, noise_pred, t, latents, return_dict=False, generator=None): class LTX2ImageToVideoPipeline(LTX2Pipeline): + supports_request_batch = False support_image_input = True def __init__( @@ -285,7 +286,7 @@ def _step_video_latents_i2v( @torch.no_grad() def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, image: PIL.Image.Image | torch.Tensor | None = None, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, @@ -725,15 +726,13 @@ def forward( generated_mel_spectrograms = self.audio_vae.decode(audio_latents, return_dict=False)[0] audio = self.vocoder(generated_mel_spectrograms) - if not return_dict: - return DiffusionOutput(output=(video, audio)) - return DiffusionOutput(output=(video, audio)) class LTX2ImageToVideoTwoStagesPipeline(nn.Module, SupportsComponentDiscovery): """LTXImageToVideoTwoStagesPipeline is for two stages image to video generation""" + supports_request_batch = False support_image_input = True dummy_run_num_frames = 2 @@ -786,7 +785,7 @@ def __init__( @torch.no_grad() def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, image: PIL.Image.Image | torch.Tensor | None = None, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, @@ -814,8 +813,8 @@ def forward( return_dict: bool = True, attention_kwargs: dict[str, Any] | None = None, max_sequence_length: int | None = None, - ): - video_latent, audio_latent = self.pipe( + ) -> DiffusionOutput: + stage1_output = self.pipe( req=req, image=image, prompt=prompt, @@ -844,7 +843,8 @@ def forward( return_dict=return_dict, attention_kwargs=attention_kwargs, max_sequence_length=max_sequence_length, - ).output + ) + video_latent, audio_latent = stage1_output.output upscaled_video_latent = self.upsample_pipe( latents=video_latent, @@ -874,7 +874,7 @@ def forward( stage_2_req.sampling_params = req.sampling_params.clone() stage_2_req.sampling_params.num_inference_steps = 3 - video, audio = self.pipe( + stage2_output = self.pipe( req=stage_2_req, latents=upscaled_video_latent, audio_latents=audio_latent, @@ -886,7 +886,8 @@ def forward( generator=generator, output_type="np", return_dict=False, - ).output + ) + video, audio = stage2_output.output return DiffusionOutput(output=(video, audio)) diff --git a/vllm_omni/diffusion/models/magi_human/pipeline_magi_human.py b/vllm_omni/diffusion/models/magi_human/pipeline_magi_human.py index 2794ddddd4a..513d9333547 100644 --- a/vllm_omni/diffusion/models/magi_human/pipeline_magi_human.py +++ b/vllm_omni/diffusion/models/magi_human/pipeline_magi_human.py @@ -53,6 +53,7 @@ DiffusionPipelineProfilerMixin, ) from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from .magi_human_dit import ( DiTModel, @@ -2160,7 +2161,7 @@ def encode_prompt( @torch.inference_mode() def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | None = None, height: int = 256, width: int = 448, diff --git a/vllm_omni/diffusion/models/ming_flash_omni/pipeline_ming_imagegen.py b/vllm_omni/diffusion/models/ming_flash_omni/pipeline_ming_imagegen.py index 2c77351dd80..a55a7998758 100644 --- a/vllm_omni/diffusion/models/ming_flash_omni/pipeline_ming_imagegen.py +++ b/vllm_omni/diffusion/models/ming_flash_omni/pipeline_ming_imagegen.py @@ -49,7 +49,7 @@ MingZImageTransformer2DModel, ) from vllm_omni.diffusion.models.z_image.pipeline_z_image import ZImagePipeline -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, ) @@ -85,6 +85,10 @@ class _ZPipelineRequest: sampling_params: _ZPipelineSamplingParams prompts: list[dict[str, Any]] = field(default_factory=lambda: [{"prompt": "", "negative_prompt": ""}]) + @property + def num_reqs(self) -> int: + return 1 + class MingImagePipeline(ZImagePipeline): """Ming-flash-omni-2.0 text-to-image diffusion pipeline. @@ -95,6 +99,8 @@ class MingImagePipeline(ZImagePipeline): ships ``byt5/``) """ + supports_request_batch = False + def __init__( self, *, @@ -272,18 +278,18 @@ def _encode_reference_image(self, ref, height: int, width: int) -> torch.Tensor # ------------------------------------------------------------------ @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: """Run one text-to-image generation request. Args: - req: Diffusion request. The cross-stage thinker hidden states + req: Single-request batch. The cross-stage thinker hidden states must be present at ``req.prompts[0]["extra"]["thinker_hidden_states"]`` as a ``[N, H]`` (or ``[1, N, H]``) tensor, placed there by ``thinker2imagegen``. Returns: - DiffusionOutput with ``.output`` set to a ``[B, 3, H, W]`` + One DiffusionOutput with ``.output`` set to a ``[B, 3, H, W]`` image tensor in ``[-1, 1]``. The vllm-omni diffusion engine's output adapter converts this to PIL/base64 downstream. """ @@ -440,7 +446,7 @@ def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: None if ref_latent is None else tuple(ref_latent.shape), ) try: - output = super().forward( + outputs = super().forward( z_req, prompt=None, height=height, @@ -456,6 +462,7 @@ def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: finally: set_forward_context_ref_latent(None) + output = outputs[0] if hasattr(output, "output") and output.output is not None: raw = output.output elif hasattr(output, "images"): diff --git a/vllm_omni/diffusion/models/nextstep_1_1/pipeline_nextstep_1_1.py b/vllm_omni/diffusion/models/nextstep_1_1/pipeline_nextstep_1_1.py index a0d27c9d038..392e4286888 100644 --- a/vllm_omni/diffusion/models/nextstep_1_1/pipeline_nextstep_1_1.py +++ b/vllm_omni/diffusion/models/nextstep_1_1/pipeline_nextstep_1_1.py @@ -35,7 +35,7 @@ NextStepModel, ) from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, ) @@ -559,7 +559,7 @@ def decoding( @torch.no_grad() def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, height: int | None = None, width: int | None = None, @@ -712,7 +712,8 @@ def forward( sampled_images = sampled_images.detach().cpu().to(torch.float32) return DiffusionOutput( - output=sampled_images, stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None + output=sampled_images, + stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None, ) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: diff --git a/vllm_omni/diffusion/models/omnigen2/pipeline_omnigen2.py b/vllm_omni/diffusion/models/omnigen2/pipeline_omnigen2.py index 83ea7ed7363..2bdd513ef77 100644 --- a/vllm_omni/diffusion/models/omnigen2/pipeline_omnigen2.py +++ b/vllm_omni/diffusion/models/omnigen2/pipeline_omnigen2.py @@ -40,6 +40,7 @@ ) from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, @@ -258,44 +259,42 @@ def pre_process_func( request: OmniDiffusionRequest, ) -> OmniDiffusionRequest: """Pre-process requests for OmniGen2Pipeline.""" - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - if raw_image is not None: - if isinstance(raw_image, list): - images = [PIL.Image.open(img) if isinstance(img, str) else img for img in raw_image] - elif isinstance(raw_image, str): - images = [PIL.Image.open(raw_image)] - else: - images = [raw_image] - - first_raw = images[0] - if isinstance(first_raw, PIL.Image.Image): - new_h, new_w = image_processor.get_new_height_width( - first_raw, max_pixels=1024 * 1024, max_side_length=1024 - ) - if request.sampling_params.height is None: - request.sampling_params.height = new_h - if request.sampling_params.width is None: - request.sampling_params.width = new_w - - preprocessed_images = [] - for image in images: - if not ( - isinstance(image, torch.Tensor) and len(image.shape) > 1 and image.shape[1] == latent_channels - ): - image = image_processor.preprocess(image, max_pixels=1024 * 1024, max_side_length=1024) - preprocessed_images.append(image) - - prompt["additional_information"]["preprocessed_images"] = preprocessed_images - - request.prompts[i] = prompt + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + if raw_image is not None: + if isinstance(raw_image, list): + images = [PIL.Image.open(img) if isinstance(img, str) else img for img in raw_image] + elif isinstance(raw_image, str): + images = [PIL.Image.open(raw_image)] + else: + images = [raw_image] + + first_raw = images[0] + if isinstance(first_raw, PIL.Image.Image): + new_h, new_w = image_processor.get_new_height_width( + first_raw, max_pixels=1024 * 1024, max_side_length=1024 + ) + if request.sampling_params.height is None: + request.sampling_params.height = new_h + if request.sampling_params.width is None: + request.sampling_params.width = new_w + + preprocessed_images = [] + for image in images: + if not (isinstance(image, torch.Tensor) and len(image.shape) > 1 and image.shape[1] == latent_channels): + image = image_processor.preprocess(image, max_pixels=1024 * 1024, max_side_length=1024) + preprocessed_images.append(image) + + prompt["additional_information"]["preprocessed_images"] = preprocessed_images + + request.prompt = prompt return request return pre_process_func @@ -1005,7 +1004,7 @@ def cfg_range(self): @torch.no_grad() def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, prompt_embeds: torch.FloatTensor | None = None, diff --git a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py index 9671e275845..62ac027ca8b 100644 --- a/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py +++ b/vllm_omni/diffusion/models/omnivoice/pipeline_omnivoice.py @@ -26,7 +26,7 @@ from vllm_omni.diffusion.data import DiffusionOutput, OmniDiffusionConfig from vllm_omni.diffusion.distributed.utils import get_local_device from vllm_omni.diffusion.models.interface import SupportAudioOutput -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.models.omnivoice.duration import RuleDurationEstimator from vllm_omni.model_executor.models.omnivoice.omnivoice_decoder import OmniVoiceDecoder from vllm_omni.model_executor.models.omnivoice.omnivoice_generator import OmniVoiceGenerator @@ -206,7 +206,7 @@ def _encode_ref_audio(self, audio_signal: torch.Tensor, sr: int) -> torch.Tensor return tokens @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: """Generate speech audio from text, optionally with voice cloning. Accepts either a plain text prompt or a structured dict: @@ -248,7 +248,7 @@ def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: audio_field = audio_field[0] elif len(audio_field) > 1: return DiffusionOutput( - error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" + error=f"OmniVoice voice cloning supports a single reference audio; got {len(audio_field)}" # noqa: E501 ) else: audio_field = None diff --git a/vllm_omni/diffusion/models/ovis_image/pipeline_ovis_image.py b/vllm_omni/diffusion/models/ovis_image/pipeline_ovis_image.py index cfe0159c189..371af3f34d4 100644 --- a/vllm_omni/diffusion/models/ovis_image/pipeline_ovis_image.py +++ b/vllm_omni/diffusion/models/ovis_image/pipeline_ovis_image.py @@ -42,7 +42,7 @@ from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery from vllm_omni.diffusion.models.ovis_image.ovis_image_transformer import OvisImageTransformer2DModel from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import download_weights_from_hf_specific logger = init_logger(__name__) @@ -144,6 +144,8 @@ def retrieve_timesteps( class OvisImagePipeline(nn.Module, CFGParallelMixin, DiffusionPipelineProfilerMixin, SupportsComponentDiscovery): + supports_request_batch = False + _dit_modules: ClassVar[list[str]] = ["transformer"] _encoder_modules: ClassVar[list[str]] = ["text_encoder"] _vae_modules: ClassVar[list[str]] = ["vae"] @@ -551,7 +553,7 @@ def interrupt(self): def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, guidance_scale: float = 5.0, @@ -760,9 +762,8 @@ def forward( latents = (latents / self.vae.config.scaling_factor) + self.vae.config.shift_factor image = self.vae.decode(latents, return_dict=False)[0] - return DiffusionOutput( - output=image, stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None - ) + stage_durations = self.stage_durations if hasattr(self, "stage_durations") else None + return DiffusionOutput(output=image, stage_durations=stage_durations) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) diff --git a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image.py b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image.py index 498d185037a..974436b21ba 100644 --- a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image.py +++ b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image.py @@ -37,7 +37,6 @@ ) from vllm_omni.diffusion.models.qwen_image.rope_utils import txt_seq_lens_from_embeds from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.prompt_utils import ( validate_prompt_sequence_lengths, ) @@ -45,6 +44,7 @@ normalize_min_aligned_size, ) from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch, split_diffusion_output_by_request if TYPE_CHECKING: from vllm_omni.diffusion.worker.input_batch import InputBatch @@ -255,6 +255,7 @@ def apply_rotary_emb_qwen( class QwenImagePipeline( nn.Module, QwenImageCFGParallelMixin, DiffusionPipelineProfilerMixin, SupportsComponentDiscovery ): + supports_request_batch = True _dit_modules: ClassVar[list[str]] = ["transformer"] _encoder_modules: ClassVar[list[str]] = ["text_encoder"] _vae_modules: ClassVar[list[str]] = ["vae"] @@ -762,7 +763,7 @@ def prepare_encode( ) -> "DiffusionRequestState": """Populate *state* with encoded prompts, latents, timesteps, and CFG config.""" sampling = state.sampling - prompt, negative_prompt = self._extract_prompts(state.prompts or []) + prompt, negative_prompt = self._extract_prompts([state.prompt] if state.prompt is not None else []) ctx = self._prepare_generation_context( prompt=prompt, @@ -981,7 +982,7 @@ def post_decode( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, true_cfg_scale: float = 4.0, @@ -1001,25 +1002,45 @@ def forward( attention_kwargs: dict[str, Any] | None = None, callback_on_step_end_tensor_inputs: list[str] = ["latents"], max_sequence_length: int = 1024, - ) -> DiffusionOutput: + ) -> list[DiffusionOutput]: + sampling_params_list = req.sampling_params_list + common_sampling_params = sampling_params_list[0] extracted_prompt, negative_prompt = self._extract_prompts(req.prompts) prompt = extracted_prompt or prompt - height = req.sampling_params.height or self.default_sample_size * self.vae_scale_factor - width = req.sampling_params.width or self.default_sample_size * self.vae_scale_factor + height = common_sampling_params.height or self.default_sample_size * self.vae_scale_factor + width = common_sampling_params.width or self.default_sample_size * self.vae_scale_factor height, width = normalize_min_aligned_size(height, width, self.vae_scale_factor * 2) - num_inference_steps = req.sampling_params.num_inference_steps or num_inference_steps - sigmas = req.sampling_params.sigmas or sigmas - max_sequence_length = req.sampling_params.max_sequence_length or max_sequence_length - generator = req.sampling_params.generator or generator - true_cfg_scale = req.sampling_params.true_cfg_scale or true_cfg_scale - if req.sampling_params.guidance_scale_provided: - guidance_scale = req.sampling_params.guidance_scale + num_inference_steps = common_sampling_params.num_inference_steps or num_inference_steps + sigmas = common_sampling_params.sigmas or sigmas + max_sequence_length = common_sampling_params.max_sequence_length or max_sequence_length num_images_per_prompt = ( - req.sampling_params.num_outputs_per_prompt - if req.sampling_params.num_outputs_per_prompt > 0 + common_sampling_params.num_outputs_per_prompt + if common_sampling_params.num_outputs_per_prompt > 0 else num_images_per_prompt ) + generator = req.collate_request_generators(num_images_per_prompt, generator) + latents = req.collate_request_tensors("latents", latents) + prompt_fields = DiffusionRequestBatch.collate_prompt_field_map( + req.prompts, + { + "prompt_embeds": prompt_embeds, + "prompt_embeds_mask": prompt_embeds_mask, + "negative_prompt_embeds": negative_prompt_embeds, + "negative_prompt_embeds_mask": negative_prompt_embeds_mask, + }, + ) + prompt_embeds = prompt_fields["prompt_embeds"] + prompt_embeds_mask = prompt_fields["prompt_embeds_mask"] + negative_prompt_embeds = prompt_fields["negative_prompt_embeds"] + negative_prompt_embeds_mask = prompt_fields["negative_prompt_embeds_mask"] + if prompt_embeds is not None: + prompt = None + if negative_prompt_embeds is not None: + negative_prompt = None + true_cfg_scale = common_sampling_params.true_cfg_scale or true_cfg_scale + if common_sampling_params.guidance_scale_provided: + guidance_scale = common_sampling_params.guidance_scale ctx = self._prepare_generation_context( prompt=prompt, @@ -1064,7 +1085,12 @@ def forward( ) self._current_timestep = None - return self._decode_latents(latents, height, width, output_type) + result = self._decode_latents(latents, height, width, output_type) + return split_diffusion_output_by_request( + result, + req, + num_outputs_per_prompt=num_images_per_prompt, + ) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) diff --git a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_edit.py b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_edit.py index 482dbac7fa6..41475c65d26 100644 --- a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_edit.py +++ b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_edit.py @@ -46,6 +46,7 @@ normalize_min_aligned_size, ) from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, @@ -75,57 +76,55 @@ def pre_process_func( request: OmniDiffusionRequest, ): """Pre-process requests for QwenImageEditPipeline.""" - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - # Only handles single image - if not raw_image: # None or empty list - raise ValueError("""Received no input image. This model requires one input image to run.""") - elif isinstance(raw_image, list): - if len(raw_image) > 1: - raise ValueError( - """Received multiple input images. Only a single image is supported by this model.""" - ) - else: - raw_image = raw_image[0] - - if isinstance(raw_image, str): - image = PIL.Image.open(raw_image) + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + # Only handles single image + if not raw_image: # None or empty list + raise ValueError("""Received no input image. This model requires one input image to run.""") + elif isinstance(raw_image, list): + if len(raw_image) > 1: + raise ValueError("""Received multiple input images. Only a single image is supported by this model.""") else: - image = cast(PIL.Image.Image | torch.Tensor | np.ndarray, raw_image) - - image_size = image.size - calculated_width, calculated_height = calculate_dimensions(1024 * 1024, image_size[0] / image_size[1]) - height = request.sampling_params.height or calculated_height - width = request.sampling_params.width or calculated_width - - # Ensure dimensions are multiples of vae_scale_factor * 2 - height, width = normalize_min_aligned_size(height, width, vae_scale_factor * 2) - - # Store calculated dimensions in request - prompt["additional_information"]["calculated_height"] = calculated_height - prompt["additional_information"]["calculated_width"] = calculated_width - request.sampling_params.height = height - request.sampling_params.width = width - - # Preprocess image - if image is not None and not ( - isinstance(image, torch.Tensor) and len(image.shape) > 1 and image.shape[1] == latent_channels - ): - image = image_processor.resize(image, calculated_height, calculated_width) - prompt_image = image - image = image_processor.preprocess(image, calculated_height, calculated_width) - image = image.unsqueeze(2) + raw_image = raw_image[0] + + if isinstance(raw_image, str): + image = PIL.Image.open(raw_image) + else: + image = cast(PIL.Image.Image | torch.Tensor | np.ndarray, raw_image) - # Store preprocessed image and prompt image in request - prompt["additional_information"]["preprocessed_image"] = image - prompt["additional_information"]["prompt_image"] = prompt_image - request.prompts[i] = prompt + image_size = image.size + calculated_width, calculated_height = calculate_dimensions(1024 * 1024, image_size[0] / image_size[1]) + height = request.sampling_params.height or calculated_height + width = request.sampling_params.width or calculated_width + + # Ensure dimensions are multiples of vae_scale_factor * 2 + height, width = normalize_min_aligned_size(height, width, vae_scale_factor * 2) + + # Store calculated dimensions in request + prompt["additional_information"]["calculated_height"] = calculated_height + prompt["additional_information"]["calculated_width"] = calculated_width + request.sampling_params.height = height + request.sampling_params.width = width + + # Preprocess image + if image is not None and not ( + isinstance(image, torch.Tensor) and len(image.shape) > 1 and image.shape[1] == latent_channels + ): + image = image_processor.resize(image, calculated_height, calculated_width) + prompt_image = image + image = image_processor.preprocess(image, calculated_height, calculated_width) + image = image.unsqueeze(2) + + # Store preprocessed image and prompt image in request + prompt["additional_information"]["preprocessed_image"] = image + prompt["additional_information"]["prompt_image"] = prompt_image + request.prompt = prompt return request return pre_process_func @@ -676,7 +675,7 @@ def interrupt(self): def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, image: PIL.Image.Image | torch.Tensor | None = None, diff --git a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_edit_plus.py b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_edit_plus.py index ffae53ff552..17f4a8323d5 100644 --- a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_edit_plus.py +++ b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_edit_plus.py @@ -50,6 +50,7 @@ normalize_min_aligned_size, ) from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, @@ -88,74 +89,72 @@ def pre_process_func( request: OmniDiffusionRequest, ): """Pre-process requests for QwenImageEditPlusPipeline.""" - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - # Handle single image or list of images - if raw_image is None: - continue - - if not isinstance(raw_image, list): - raw_image = [raw_image] - if len(raw_image) > MAX_QWEN_IMAGE_EDIT_PLUS_INPUT_IMAGES: - raise ValueError( - f"Received {len(raw_image)} input images. " - f"At most {MAX_QWEN_IMAGE_EDIT_PLUS_INPUT_IMAGES} images are supported by this model." - ) - image = [ - PIL.Image.open(im) if isinstance(im, str) else cast(PIL.Image.Image | np.ndarray | torch.Tensor, im) - for im in raw_image - ] - - # Calculate dimensions based on first image - image_size = image[0].size - calculated_width, calculated_height = calculate_dimensions(VAE_IMAGE_SIZE, image_size[0] / image_size[1]) - height = request.sampling_params.height or calculated_height - width = request.sampling_params.width or calculated_width - - # Ensure dimensions are multiples of vae_scale_factor * 2 - height, width = normalize_min_aligned_size(height, width, vae_scale_factor * 2) - - # Store calculated dimensions in request - prompt["additional_information"]["calculated_height"] = calculated_height - prompt["additional_information"]["calculated_width"] = calculated_width - request.sampling_params.height = height - request.sampling_params.width = width - - # Preprocess images into condition_images (for prompt encoding) and vae_images (for VAE encoding) - condition_images = [] - vae_images = [] - condition_image_sizes = [] - vae_image_sizes = [] + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + # Handle single image or list of images + if raw_image is None: + return request + + if not isinstance(raw_image, list): + raw_image = [raw_image] + if len(raw_image) > MAX_QWEN_IMAGE_EDIT_PLUS_INPUT_IMAGES: + raise ValueError( + f"Received {len(raw_image)} input images. " + f"At most {MAX_QWEN_IMAGE_EDIT_PLUS_INPUT_IMAGES} images are supported by this model." + ) + image = [ + PIL.Image.open(im) if isinstance(im, str) else cast(PIL.Image.Image | np.ndarray | torch.Tensor, im) + for im in raw_image + ] - for img in image: - if isinstance(img, torch.Tensor) and len(img.shape) > 1 and img.shape[1] == latent_channels: - # Already a latent tensor - continue + # Calculate dimensions based on first image + image_size = image[0].size + calculated_width, calculated_height = calculate_dimensions(VAE_IMAGE_SIZE, image_size[0] / image_size[1]) + height = request.sampling_params.height or calculated_height + width = request.sampling_params.width or calculated_width + + # Ensure dimensions are multiples of vae_scale_factor * 2 + height, width = normalize_min_aligned_size(height, width, vae_scale_factor * 2) + + # Store calculated dimensions in request + prompt["additional_information"]["calculated_height"] = calculated_height + prompt["additional_information"]["calculated_width"] = calculated_width + request.sampling_params.height = height + request.sampling_params.width = width + + # Preprocess images into condition_images (for prompt encoding) and vae_images (for VAE encoding) + condition_images = [] + vae_images = [] + condition_image_sizes = [] + vae_image_sizes = [] + + for img in image: + if isinstance(img, torch.Tensor) and len(img.shape) > 1 and img.shape[1] == latent_channels: + # Already a latent tensor + continue - image_width, image_height = img.size - condition_width, condition_height = calculate_dimensions( - CONDITION_IMAGE_SIZE, image_width / image_height - ) - vae_width, vae_height = calculate_dimensions(VAE_IMAGE_SIZE, image_width / image_height) + image_width, image_height = img.size + condition_width, condition_height = calculate_dimensions(CONDITION_IMAGE_SIZE, image_width / image_height) + vae_width, vae_height = calculate_dimensions(VAE_IMAGE_SIZE, image_width / image_height) - condition_image_sizes.append((condition_width, condition_height)) - vae_image_sizes.append((vae_width, vae_height)) + condition_image_sizes.append((condition_width, condition_height)) + vae_image_sizes.append((vae_width, vae_height)) - condition_images.append(image_processor.resize(img, condition_height, condition_width)) - vae_images.append(image_processor.preprocess(img, vae_height, vae_width).unsqueeze(2)) + condition_images.append(image_processor.resize(img, condition_height, condition_width)) + vae_images.append(image_processor.preprocess(img, vae_height, vae_width).unsqueeze(2)) - # Store preprocessed images in request - prompt["additional_information"]["condition_images"] = condition_images - prompt["additional_information"]["vae_images"] = vae_images - prompt["additional_information"]["condition_image_sizes"] = condition_image_sizes - prompt["additional_information"]["vae_image_sizes"] = vae_image_sizes - request.prompts[i] = prompt + # Store preprocessed images in request + prompt["additional_information"]["condition_images"] = condition_images + prompt["additional_information"]["vae_images"] = vae_images + prompt["additional_information"]["condition_image_sizes"] = condition_image_sizes + prompt["additional_information"]["vae_image_sizes"] = vae_image_sizes + request.prompt = prompt return request return pre_process_func @@ -622,7 +621,7 @@ def interrupt(self): def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, image: PIL.Image.Image | list[PIL.Image.Image] | torch.Tensor | None = None, diff --git a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_layered.py b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_layered.py index 055edac3573..10363f42bdb 100644 --- a/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_layered.py +++ b/vllm_omni/diffusion/models/qwen_image/pipeline_qwen_image_layered.py @@ -45,6 +45,7 @@ normalize_min_aligned_size, ) from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, @@ -80,63 +81,61 @@ def pre_process_func( request: OmniDiffusionRequest, ): """Pre-process requests for QwenImageLayeredPipeline.""" - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - if not raw_image: # None or empty list - raise ValueError("""Received no input image. This model requires one input image to run.""") - elif isinstance(raw_image, list): - if len(raw_image) > 1: - raise ValueError( - """Received multiple input images. Only a single image is supported by this model.""" - ) - else: - raw_image = raw_image[0] - - if isinstance(raw_image, str): - image = PIL.Image.open(raw_image) + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + if not raw_image: # None or empty list + raise ValueError("""Received no input image. This model requires one input image to run.""") + elif isinstance(raw_image, list): + if len(raw_image) > 1: + raise ValueError("""Received multiple input images. Only a single image is supported by this model.""") else: - image = cast(PIL.Image.Image | torch.Tensor | np.ndarray, raw_image) + raw_image = raw_image[0] - if isinstance(image, PIL.Image.Image) and image.mode != "RGBA": - image = image.convert("RGBA") - - # 1. calculate dimensions - image_size = image.size - assert request.sampling_params.resolution in [640, 1024], ( - f"resolution must be either 640 or 1024, but got {request.sampling_params.resolution}" - ) - calculated_width, calculated_height = calculate_dimensions( - request.sampling_params.resolution * request.sampling_params.resolution, image_size[0] / image_size[1] - ) - height = calculated_height - width = calculated_width - - height, width = normalize_min_aligned_size(height, width, vae_scale_factor * 2) + if isinstance(raw_image, str): + image = PIL.Image.open(raw_image) + else: + image = cast(PIL.Image.Image | torch.Tensor | np.ndarray, raw_image) - # Store calculated dimensions in request - prompt["additional_information"]["calculated_height"] = calculated_height - prompt["additional_information"]["calculated_width"] = calculated_width - request.sampling_params.height = height - request.sampling_params.width = width + if isinstance(image, PIL.Image.Image) and image.mode != "RGBA": + image = image.convert("RGBA") - # 2. Preprocess image - if image is not None and not (isinstance(image, torch.Tensor) and image.size(1) == latent_channels): - image = image_processor.resize(image, calculated_height, calculated_width) - prompt_image = image - image = image_processor.preprocess(image, calculated_height, calculated_width) - image = image.unsqueeze(2) - # image = image.to(dtype=self.text_encoder.dtype) # do it later - - # Store preprocessed image and prompt image in request - prompt["additional_information"]["preprocessed_image"] = image - prompt["additional_information"]["prompt_image"] = prompt_image - request.prompts[i] = prompt + # 1. calculate dimensions + image_size = image.size + assert request.sampling_params.resolution in [640, 1024], ( + f"resolution must be either 640 or 1024, but got {request.sampling_params.resolution}" + ) + calculated_width, calculated_height = calculate_dimensions( + request.sampling_params.resolution * request.sampling_params.resolution, image_size[0] / image_size[1] + ) + height = calculated_height + width = calculated_width + + height, width = normalize_min_aligned_size(height, width, vae_scale_factor * 2) + + # Store calculated dimensions in request + prompt["additional_information"]["calculated_height"] = calculated_height + prompt["additional_information"]["calculated_width"] = calculated_width + request.sampling_params.height = height + request.sampling_params.width = width + + # 2. Preprocess image + if image is not None and not (isinstance(image, torch.Tensor) and image.size(1) == latent_channels): + image = image_processor.resize(image, calculated_height, calculated_width) + prompt_image = image + image = image_processor.preprocess(image, calculated_height, calculated_width) + image = image.unsqueeze(2) + # image = image.to(dtype=self.text_encoder.dtype) # do it later + + # Store preprocessed image and prompt image in request + prompt["additional_information"]["preprocessed_image"] = image + prompt["additional_information"]["prompt_image"] = prompt_image + request.prompt = prompt return request return pre_process_func @@ -647,7 +646,7 @@ def interrupt(self): def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, image: PIL.Image.Image | torch.Tensor | None = None, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, diff --git a/vllm_omni/diffusion/models/sd3/pipeline_sd3.py b/vllm_omni/diffusion/models/sd3/pipeline_sd3.py index 7dce56aac0a..e45614a42e6 100644 --- a/vllm_omni/diffusion/models/sd3/pipeline_sd3.py +++ b/vllm_omni/diffusion/models/sd3/pipeline_sd3.py @@ -26,7 +26,7 @@ SD3Transformer2DModel, ) from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch, split_diffusion_output_by_request from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, ) @@ -132,6 +132,8 @@ def retrieve_timesteps( class StableDiffusion3Pipeline(nn.Module, CFGParallelMixin, DiffusionPipelineProfilerMixin, SupportsComponentDiscovery): + supports_request_batch = True + _dit_modules: ClassVar[list[str]] = ["transformer"] _encoder_modules: ClassVar[list[str]] = ["text_encoder", "text_encoder_2", "text_encoder_3"] _vae_modules: ClassVar[list[str]] = ["vae"] @@ -434,6 +436,7 @@ def encode_prompt( prompt_2: str | list[str], prompt_3: str | list[str], prompt_embeds: torch.Tensor | None = None, + pooled_prompt_embeds: torch.Tensor | None = None, max_sequence_length: int = 256, num_images_per_prompt: int = 1, ): @@ -457,7 +460,6 @@ def encode_prompt( prompt = [prompt] if isinstance(prompt, str) else prompt - pooled_prompt_embeds = None if prompt_embeds is None: prompt_2 = prompt_2 or prompt prompt_2 = [prompt_2] if isinstance(prompt_2, str) else prompt_2 @@ -628,7 +630,7 @@ def diffuse( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] = "", prompt_2: str | list[str] = "", prompt_3: str | list[str] = "", @@ -647,25 +649,50 @@ def forward( pooled_prompt_embeds: torch.Tensor | None = None, negative_pooled_prompt_embeds: torch.Tensor | None = None, max_sequence_length: int = 256, - ) -> DiffusionOutput: + ) -> list[DiffusionOutput]: # TODO: In online mode, sometimes it receives [{"negative_prompt": None}, {...}], so cannot use .get("...", "") # TODO: May be some data formatting operations on the API side. Hack for now. + sampling_params_list = req.sampling_params_list + common_sampling_params = sampling_params_list[0] prompt = [p if isinstance(p, str) else (p.get("prompt") or "") for p in req.prompts] or prompt negative_prompt = [ "" if isinstance(p, str) else (p.get("negative_prompt") or "") for p in req.prompts ] or negative_prompt - height = req.sampling_params.height or self.default_sample_size * self.vae_scale_factor - width = req.sampling_params.width or self.default_sample_size * self.vae_scale_factor - sigmas = req.sampling_params.sigmas or sigmas - max_sequence_length = req.sampling_params.max_sequence_length or max_sequence_length - num_inference_steps = req.sampling_params.num_inference_steps or num_inference_steps - generator = req.sampling_params.generator or generator + height = common_sampling_params.height or self.default_sample_size * self.vae_scale_factor + width = common_sampling_params.width or self.default_sample_size * self.vae_scale_factor + sigmas = common_sampling_params.sigmas or sigmas + max_sequence_length = common_sampling_params.max_sequence_length or max_sequence_length + num_inference_steps = common_sampling_params.num_inference_steps or num_inference_steps num_images_per_prompt = ( - req.sampling_params.num_outputs_per_prompt - if req.sampling_params.num_outputs_per_prompt > 0 + common_sampling_params.num_outputs_per_prompt + if common_sampling_params.num_outputs_per_prompt > 0 else num_images_per_prompt ) + if generator is None: + generator = req.collate_request_generators(num_images_per_prompt, generator) + latents = req.collate_request_tensors("latents", latents) + prompt_fields = DiffusionRequestBatch.collate_prompt_field_map( + req.prompts, + { + "prompt_embeds": prompt_embeds, + "negative_prompt_embeds": negative_prompt_embeds, + "pooled_prompt_embeds": pooled_prompt_embeds, + "negative_pooled_prompt_embeds": negative_pooled_prompt_embeds, + }, + ) + prompt_embeds = prompt_fields["prompt_embeds"] + negative_prompt_embeds = prompt_fields["negative_prompt_embeds"] + pooled_prompt_embeds = prompt_fields["pooled_prompt_embeds"] + negative_pooled_prompt_embeds = prompt_fields["negative_pooled_prompt_embeds"] + if prompt_embeds is not None: + prompt = None + prompt_2 = None + prompt_3 = None + if negative_prompt_embeds is not None: + negative_prompt = None + negative_prompt_2 = None + negative_prompt_3 = None # 1. check inputs # 2. encode prompts # 3. prepare latents and timesteps @@ -686,7 +713,7 @@ def forward( max_sequence_length=max_sequence_length, ) - self._guidance_scale = req.sampling_params.guidance_scale + self._guidance_scale = common_sampling_params.guidance_scale self._current_timestep = None self._interrupt = False @@ -702,6 +729,7 @@ def forward( prompt_2=prompt_2, prompt_3=prompt_3, prompt_embeds=prompt_embeds, + pooled_prompt_embeds=pooled_prompt_embeds, max_sequence_length=max_sequence_length, ) @@ -712,6 +740,7 @@ def forward( prompt_2=negative_prompt_2, prompt_3=negative_prompt_3, prompt_embeds=negative_prompt_embeds, + pooled_prompt_embeds=negative_pooled_prompt_embeds, max_sequence_length=max_sequence_length, ) @@ -751,8 +780,13 @@ def forward( image = self.vae.decode(latents, return_dict=False)[0] - return DiffusionOutput( - output=image, stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None + return split_diffusion_output_by_request( + DiffusionOutput( + output=image, + stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None, + ), + req, + num_outputs_per_prompt=num_images_per_prompt, ) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: diff --git a/vllm_omni/diffusion/models/sdxl/pipeline_sdxl.py b/vllm_omni/diffusion/models/sdxl/pipeline_sdxl.py index d4a966cb04f..0bb8a48cd7c 100644 --- a/vllm_omni/diffusion/models/sdxl/pipeline_sdxl.py +++ b/vllm_omni/diffusion/models/sdxl/pipeline_sdxl.py @@ -22,7 +22,7 @@ from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery from vllm_omni.diffusion.models.sdxl.sdxl_unet import SDXLUNet2DConditionModel from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch logger = logging.getLogger(__name__) @@ -297,7 +297,7 @@ def diffuse( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] = "", negative_prompt: str | list[str] = "", height: int | None = None, diff --git a/vllm_omni/diffusion/models/sensenova_u1/pipeline_sensenova_u1.py b/vllm_omni/diffusion/models/sensenova_u1/pipeline_sensenova_u1.py index c1ba258b28c..080ed10e603 100644 --- a/vllm_omni/diffusion/models/sensenova_u1/pipeline_sensenova_u1.py +++ b/vllm_omni/diffusion/models/sensenova_u1/pipeline_sensenova_u1.py @@ -38,6 +38,7 @@ from vllm_omni.diffusion.models.interface import SupportsComponentDiscovery from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from .sensenova_u1_transformer import ( SenseNovaU1ForCausalLM, @@ -1193,7 +1194,7 @@ def _expand_and_prepare_kv(self, kv, token_hw, batch_size): prepare_flash_kv_cache(kv, current_len=token_hw, batch_size=batch_size) @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: p = self._parse_request(req) self.top_cfg.t_eps = p.t_eps diff --git a/vllm_omni/diffusion/models/soulx_singer/pipeline_soulx_singer_svc.py b/vllm_omni/diffusion/models/soulx_singer/pipeline_soulx_singer_svc.py index 91d4437c2a1..829256134bf 100644 --- a/vllm_omni/diffusion/models/soulx_singer/pipeline_soulx_singer_svc.py +++ b/vllm_omni/diffusion/models/soulx_singer/pipeline_soulx_singer_svc.py @@ -37,6 +37,7 @@ validate_soulx_extra_args, ) from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch logger = init_logger(__name__) @@ -58,7 +59,7 @@ def pre_process_func(request: OmniDiffusionRequest) -> OmniDiffusionRequest: # Inline build: when no warmup/no precomputed paths/no IPC payload, # build the preprocess payload directly from audio file paths. if not (is_warmup_request(request) or has_precomputed(extra_args, "svc")): - prompt = request.prompts[0] + prompt = request.prompt if not isinstance(prompt, str) and not get_soulx_preprocessed_payload(prompt): # type: ignore[arg-type] prompt_audio, target_audio = resolve_preprocess_audio(prompt, extra_args) # type: ignore[arg-type] if prompt_audio is not None and target_audio is not None: @@ -395,7 +396,7 @@ def infer_svc_batch( return generated_audio.unsqueeze(0), pitch_shift @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: return self._forward_batch_from_request( req, kind="svc", diff --git a/vllm_omni/diffusion/models/soulx_singer/pipeline_soulx_singer_svs.py b/vllm_omni/diffusion/models/soulx_singer/pipeline_soulx_singer_svs.py index ecf6f94a55b..9be70e651fb 100644 --- a/vllm_omni/diffusion/models/soulx_singer/pipeline_soulx_singer_svs.py +++ b/vllm_omni/diffusion/models/soulx_singer/pipeline_soulx_singer_svs.py @@ -34,6 +34,7 @@ validate_soulx_extra_args, ) from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch logger = init_logger(__name__) @@ -54,7 +55,7 @@ def pre_process_func(request: OmniDiffusionRequest) -> OmniDiffusionRequest: # Inline build: when no warmup/no precomputed paths/no IPC payload, # build the preprocess payload directly from audio file paths. if not (is_warmup_request(request) or has_precomputed(extra_args, "svs")): - prompt = request.prompts[0] + prompt = request.prompt if not isinstance(prompt, str) and not get_soulx_preprocessed_payload(prompt): # type: ignore[arg-type] prompt_audio, target_audio = resolve_preprocess_audio(prompt, extra_args) # type: ignore[arg-type] if prompt_audio is not None and target_audio is not None: @@ -338,7 +339,7 @@ def _ensure_processed_meta(self, meta: dict) -> dict: return self.metadata_processor.process(meta, None) @torch.inference_mode() - def forward(self, req: OmniDiffusionRequest) -> DiffusionOutput: + def forward(self, req: DiffusionRequestBatch) -> DiffusionOutput: return self._forward_batch_from_request( req, kind="svs", diff --git a/vllm_omni/diffusion/models/soulx_singer/preprocess/pre_process.py b/vllm_omni/diffusion/models/soulx_singer/preprocess/pre_process.py index 750175fc5ca..a8784fcba04 100644 --- a/vllm_omni/diffusion/models/soulx_singer/preprocess/pre_process.py +++ b/vllm_omni/diffusion/models/soulx_singer/preprocess/pre_process.py @@ -159,35 +159,38 @@ def attach_preprocess_for_diffusion_request( kind, dict(getattr(request.sampling_params, "extra_args", None) or {}), ) - for i, prompt in enumerate(request.prompts): - prompt = OmniTextPrompt(prompt=prompt) if isinstance(prompt, str) else prompt - - if is_warmup_request(request): - request.sampling_params.num_inference_steps = 1 - payload = build_warmup_payload( - kind, - metadata_processor=metadata_processor, - device=device, - sample_rate=sample_rate, - ) - elif get_soulx_preprocessed_payload(prompt): - request.prompts[i] = prompt - continue - elif has_precomputed(extra_args, kind): - payload = build_precomputed_payload( - kind, - extra_args, - metadata_processor=metadata_processor, - sample_rate=sample_rate, - device=device, - ) - else: - raise _preprocess_inputs_missing_error(kind) - - if payload.get("kind") != kind: - raise ValueError(f"Invalid {kind} preprocess payload kind: {payload.get('kind')}") - prompt.setdefault("additional_information", {})[SOULX_PREPROCESSED_KEY] = payload - request.prompts[i] = prompt + prompt = request.prompt + prompt = OmniTextPrompt(prompt=prompt) if isinstance(prompt, str) else prompt + + if is_warmup_request(request): + request.sampling_params.num_inference_steps = 1 + payload = build_warmup_payload( + kind, + metadata_processor=metadata_processor, + device=device, + sample_rate=sample_rate, + ) + elif get_soulx_preprocessed_payload(prompt): + request.prompt = prompt + if kind == "svs": + extra_args = normalize_svs_control_extra_args(extra_args) + request.sampling_params.extra_args = extra_args + return request + elif has_precomputed(extra_args, kind): + payload = build_precomputed_payload( + kind, + extra_args, + metadata_processor=metadata_processor, + sample_rate=sample_rate, + device=device, + ) + else: + raise _preprocess_inputs_missing_error(kind) + + if payload.get("kind") != kind: + raise ValueError(f"Invalid {kind} preprocess payload kind: {payload.get('kind')}") + prompt.setdefault("additional_information", {})[SOULX_PREPROCESSED_KEY] = payload + request.prompt = prompt if kind == "svs": extra_args = normalize_svs_control_extra_args(extra_args) diff --git a/vllm_omni/diffusion/models/stable_audio/pipeline_stable_audio.py b/vllm_omni/diffusion/models/stable_audio/pipeline_stable_audio.py index 3f49450e58c..23b61209aee 100644 --- a/vllm_omni/diffusion/models/stable_audio/pipeline_stable_audio.py +++ b/vllm_omni/diffusion/models/stable_audio/pipeline_stable_audio.py @@ -35,8 +35,8 @@ StableAudioSchedulerWrapper, ) from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.utils.tf_utils import get_transformer_config_kwargs +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch logger = init_logger(__name__) @@ -77,6 +77,8 @@ class StableAudioPipeline(nn.Module, SupportAudioOutput, SupportsComponentDiscov prefix: Weight prefix for loading (default: "") """ + supports_request_batch = False + # Picked up by ``supports_audio_output`` in the diffusion engine so the # default stage metadata reports ``final_output_type="audio"`` and the # ``multimodal_output`` payload includes the sample rate (mirrors the @@ -385,7 +387,7 @@ def prepare_latents( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, negative_prompt: str | list[str] | None = None, audio_end_in_s: float | None = None, @@ -605,9 +607,8 @@ def forward( # Trim to requested length audio = audio[:, :, waveform_start:waveform_end] - return DiffusionOutput( - output=audio, stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None - ) + stage_durations = self.stage_durations if hasattr(self, "stage_durations") else None + return DiffusionOutput(output=audio, stage_durations=stage_durations) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: """Load weights using AutoWeightsLoader for vLLM integration.""" diff --git a/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2.py b/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2.py index 1b4cee1232d..b2870549c01 100644 --- a/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2.py +++ b/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2.py @@ -36,6 +36,7 @@ from vllm_omni.diffusion.postprocess import interpolate_video_tensor from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.platforms import current_omni_platform @@ -210,48 +211,49 @@ def get_wan22_pre_process_func( import numpy as np def pre_process_func(request: OmniDiffusionRequest) -> OmniDiffusionRequest: - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - if raw_image is None: - continue - - if not isinstance(raw_image, (str, PIL.Image.Image)): - raise TypeError( - f"""Unsupported image format {raw_image.__class__}.""", - """Please correctly set `"multi_modal_data": {"image": , …}`""", - ) - image = PIL.Image.open(raw_image).convert("RGB") if isinstance(raw_image, str) else raw_image - - # Calculate dimensions based on aspect ratio if not provided - if request.sampling_params.height is None or request.sampling_params.width is None: - # Default max area for 720P - max_area = 720 * 1280 - aspect_ratio = image.height / image.width - - # Calculate dimensions maintaining aspect ratio - mod_value = 16 # Must be divisible by 16 - height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value - width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value - - if request.sampling_params.height is None: - request.sampling_params.height = height - if request.sampling_params.width is None: - request.sampling_params.width = width - - # Resize image to target dimensions - image = image.resize( - (request.sampling_params.width, request.sampling_params.height), # type: ignore # Above has ensured that width & height are not None - PIL.Image.Resampling.LANCZOS, + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + if raw_image is None: + request.prompt = prompt + return request + + if not isinstance(raw_image, (str, PIL.Image.Image)): + raise TypeError( + f"""Unsupported image format {raw_image.__class__}.""", + """Please correctly set `"multi_modal_data": {"image": , …}`""", ) - prompt["multi_modal_data"]["image"] = image # type: ignore # key existence already checked above + image = PIL.Image.open(raw_image).convert("RGB") if isinstance(raw_image, str) else raw_image + + # Calculate dimensions based on aspect ratio if not provided + if request.sampling_params.height is None or request.sampling_params.width is None: + # Default max area for 720P + max_area = 720 * 1280 + aspect_ratio = image.height / image.width + + # Calculate dimensions maintaining aspect ratio + mod_value = 16 # Must be divisible by 16 + height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value + width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value + + if request.sampling_params.height is None: + request.sampling_params.height = height + if request.sampling_params.width is None: + request.sampling_params.width = width + + # Resize image to target dimensions + image = image.resize( + (request.sampling_params.width, request.sampling_params.height), # type: ignore # Above has ensured that width & height are not None + PIL.Image.Resampling.LANCZOS, + ) + prompt["multi_modal_data"]["image"] = image # type: ignore # key existence already checked above - request.prompts[i] = prompt + request.prompt = prompt return request return pre_process_func @@ -540,7 +542,7 @@ def diffuse( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | None = None, negative_prompt: str | None = None, height: int = 480, diff --git a/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_i2v.py b/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_i2v.py index ffb21b245d5..fc2d79cec73 100644 --- a/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_i2v.py +++ b/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_i2v.py @@ -43,6 +43,7 @@ from vllm_omni.diffusion.postprocess import interpolate_video_tensor from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.platforms import current_omni_platform @@ -87,50 +88,50 @@ def get_wan22_i2v_pre_process_func( """Pre-process function for I2V: load and resize input image.""" def pre_process_func(request: OmniDiffusionRequest) -> OmniDiffusionRequest: - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - if raw_image is None: - raise ValueError( - """No image is provided. This model requires an image to run.""", - """Please correctly set `"multi_modal_data": {"image": , …}`""", - ) - if not isinstance(raw_image, (str, PIL.Image.Image)): - raise TypeError( - f"""Unsupported image format {raw_image.__class__}.""", - """Please correctly set `"multi_modal_data": {"image": , …}`""", - ) - image = PIL.Image.open(raw_image).convert("RGB") if isinstance(raw_image, str) else raw_image - - # Calculate dimensions based on aspect ratio if not provided - if request.sampling_params.height is None or request.sampling_params.width is None: - # Default max area for 480P - max_area = 480 * 832 - aspect_ratio = image.height / image.width - - # Calculate dimensions maintaining aspect ratio - mod_value = 16 # Must be divisible by 16 - height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value - width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value - - if request.sampling_params.height is None: - request.sampling_params.height = height - if request.sampling_params.width is None: - request.sampling_params.width = width - - # Resize image to target dimensions - image = image.resize( - (request.sampling_params.width, request.sampling_params.height), # type: ignore # Above has ensured that width & height are not None - PIL.Image.Resampling.LANCZOS, + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + if raw_image is None: + raise ValueError( + """No image is provided. This model requires an image to run.""", + """Please correctly set `"multi_modal_data": {"image": , …}`""", ) - prompt["multi_modal_data"]["image"] = image # type: ignore # key existence already checked above + if not isinstance(raw_image, (str, PIL.Image.Image)): + raise TypeError( + f"""Unsupported image format {raw_image.__class__}.""", + """Please correctly set `"multi_modal_data": {"image": , …}`""", + ) + image = PIL.Image.open(raw_image).convert("RGB") if isinstance(raw_image, str) else raw_image + + # Calculate dimensions based on aspect ratio if not provided + if request.sampling_params.height is None or request.sampling_params.width is None: + # Default max area for 480P + max_area = 480 * 832 + aspect_ratio = image.height / image.width + + # Calculate dimensions maintaining aspect ratio + mod_value = 16 # Must be divisible by 16 + height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value + width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value + + if request.sampling_params.height is None: + request.sampling_params.height = height + if request.sampling_params.width is None: + request.sampling_params.width = width + + # Resize image to target dimensions + image = image.resize( + (request.sampling_params.width, request.sampling_params.height), # type: ignore # Above has ensured that width & height are not None + PIL.Image.Resampling.LANCZOS, + ) + prompt["multi_modal_data"]["image"] = image # type: ignore # key existence already checked above - request.prompts[i] = prompt + request.prompt = prompt return request return pre_process_func @@ -430,7 +431,7 @@ def _create_transformer(self, config: dict) -> WanTransformer3DModel: def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | None = None, negative_prompt: str | None = None, image: PIL.Image.Image | torch.Tensor | None = None, diff --git a/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_s2v.py b/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_s2v.py index eb4ddc40cf0..47b8e02d158 100644 --- a/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_s2v.py +++ b/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_s2v.py @@ -42,6 +42,7 @@ ) from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.platforms import current_omni_platform @@ -282,61 +283,60 @@ def get_wan22_s2v_pre_process_func( """ def pre_process_func(request: OmniDiffusionRequest) -> OmniDiffusionRequest: - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - # -- Reference image -- - raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None - if raw_image is None: - raise ValueError( - "No reference image provided. S2V requires a reference image. " - 'Set `"multi_modal_data": {"image": , ...}`' - ) - if isinstance(raw_image, str): - image = PIL.Image.open(raw_image).convert("RGB") - elif isinstance(raw_image, PIL.Image.Image): - image = raw_image - else: - raise TypeError(f"Unsupported image type {type(raw_image)}") - - # -- Audio -- - raw_audio = multi_modal_data.get("audio", None) if multi_modal_data is not None else None - if raw_audio is None: - raise ValueError( - "No audio provided. S2V requires an audio file path. " - 'Set `"multi_modal_data": {"audio": "", ...}`' - ) - - # -- Compute target size -- - max_area = 720 * 1280 - if request.sampling_params.height is not None and request.sampling_params.width is not None: - height, width = request.sampling_params.height, request.sampling_params.width - else: - ref_h, ref_w = image.height, image.width - height, width = _get_size_less_than_area(ref_h, ref_w, target_area=max_area) - if request.sampling_params.height is None: - request.sampling_params.height = height - if request.sampling_params.width is None: - request.sampling_params.width = width - - # Resize + center-crop reference image to target size - resize_op = transforms.Resize(min(height, width)) - crop_op = transforms.CenterCrop((height, width)) - ref_pil = crop_op(resize_op(image)) - - prompt["multi_modal_data"]["image"] = ref_pil - prompt["additional_information"]["audio_path"] = raw_audio - prompt["additional_information"]["pose_video"] = ( - multi_modal_data.get("pose_video", None) if multi_modal_data is not None else None + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + # -- Reference image -- + raw_image = multi_modal_data.get("image", None) if multi_modal_data is not None else None + if raw_image is None: + raise ValueError( + "No reference image provided. S2V requires a reference image. " + 'Set `"multi_modal_data": {"image": , ...}`' ) - prompt["additional_information"]["init_first_frame"] = ( - multi_modal_data.get("init_first_frame", False) if multi_modal_data is not None else False + if isinstance(raw_image, str): + image = PIL.Image.open(raw_image).convert("RGB") + elif isinstance(raw_image, PIL.Image.Image): + image = raw_image + else: + raise TypeError(f"Unsupported image type {type(raw_image)}") + + # -- Audio -- + raw_audio = multi_modal_data.get("audio", None) if multi_modal_data is not None else None + if raw_audio is None: + raise ValueError( + 'No audio provided. S2V requires an audio file path. Set `"multi_modal_data": {"audio": "", ...}`' ) - request.prompts[i] = prompt + + # -- Compute target size -- + max_area = 720 * 1280 + if request.sampling_params.height is not None and request.sampling_params.width is not None: + height, width = request.sampling_params.height, request.sampling_params.width + else: + ref_h, ref_w = image.height, image.width + height, width = _get_size_less_than_area(ref_h, ref_w, target_area=max_area) + if request.sampling_params.height is None: + request.sampling_params.height = height + if request.sampling_params.width is None: + request.sampling_params.width = width + + # Resize + center-crop reference image to target size + resize_op = transforms.Resize(min(height, width)) + crop_op = transforms.CenterCrop((height, width)) + ref_pil = crop_op(resize_op(image)) + + prompt["multi_modal_data"]["image"] = ref_pil + prompt["additional_information"]["audio_path"] = raw_audio + prompt["additional_information"]["pose_video"] = ( + multi_modal_data.get("pose_video", None) if multi_modal_data is not None else None + ) + prompt["additional_information"]["init_first_frame"] = ( + multi_modal_data.get("init_first_frame", False) if multi_modal_data is not None else False + ) + request.prompt = prompt return request @@ -1054,7 +1054,7 @@ def diffuse( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | None = None, negative_prompt: str | None = None, image: PIL.Image.Image | None = None, diff --git a/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_vace.py b/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_vace.py index 7f3458d5e83..4173b86011b 100644 --- a/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_vace.py +++ b/vllm_omni/diffusion/models/wan2_2/pipeline_wan2_2_vace.py @@ -36,6 +36,7 @@ ) from vllm_omni.diffusion.models.wan2_2.wan2_2_vace_transformer import WanVACETransformer3DModel from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.inputs.data import OmniTextPrompt from vllm_omni.platforms import current_omni_platform @@ -96,68 +97,66 @@ def get_wan22_vace_pre_process_func(od_config: OmniDiffusionConfig): import numpy as np def pre_process_func(request: OmniDiffusionRequest) -> OmniDiffusionRequest: - for i, prompt in enumerate(request.prompts): - multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None - if isinstance(prompt, str): - prompt = OmniTextPrompt(prompt=prompt) - if "additional_information" not in prompt: - prompt["additional_information"] = {} - - if not multi_modal_data: - request.prompts[i] = prompt - continue - - # Handle reference images for R2V - # "image" is the standard key from online serving (SupportImageInput convention) - # "reference_images" is the offline API key for backwards compatibility - ref_images = multi_modal_data.get("image") or multi_modal_data.get("reference_images") - if ref_images is not None: - if isinstance(ref_images, str): - ref_images = [PIL.Image.open(ref_images).convert("RGB")] - elif isinstance(ref_images, PIL.Image.Image): - ref_images = [ref_images] - elif isinstance(ref_images, list): - ref_images = [ - PIL.Image.open(img).convert("RGB") if isinstance(img, str) else img for img in ref_images - ] - - # Calculate dimensions from first reference image if not provided - if request.sampling_params.height is None or request.sampling_params.width is None: - first_img = ref_images[0] - max_area = 480 * 832 # VACE default is 480p - aspect_ratio = first_img.height / first_img.width - mod_value = 16 - height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value - width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value - - if request.sampling_params.height is None: - request.sampling_params.height = height - if request.sampling_params.width is None: - request.sampling_params.width = width - - prompt["additional_information"]["reference_images"] = ref_images - - # Handle source video for V2V / MV2V - source_video = multi_modal_data.get("video") - if source_video is not None: - if isinstance(source_video, list) and len(source_video) > 0: - if isinstance(source_video[0], str): - source_video = [PIL.Image.open(f).convert("RGB") for f in source_video] - prompt["additional_information"]["source_video"] = source_video - - # Handle mask for MV2V / inpainting - mask = multi_modal_data.get("mask") - if mask is not None: - if isinstance(mask, list) and len(mask) > 0: - if isinstance(mask[0], str): - mask = [PIL.Image.open(m).convert("L") for m in mask] - elif isinstance(mask, str): - mask = [PIL.Image.open(mask).convert("L")] - elif isinstance(mask, PIL.Image.Image): - mask = [mask] - prompt["additional_information"]["mask"] = mask - - request.prompts[i] = prompt + prompt = request.prompt + multi_modal_data = prompt.get("multi_modal_data", {}) if not isinstance(prompt, str) else None + if isinstance(prompt, str): + prompt = OmniTextPrompt(prompt=prompt) + if "additional_information" not in prompt: + prompt["additional_information"] = {} + + if not multi_modal_data: + request.prompt = prompt + return request + + # Handle reference images for R2V + # "image" is the standard key from online serving (SupportImageInput convention) + # "reference_images" is the offline API key for backwards compatibility + ref_images = multi_modal_data.get("image") or multi_modal_data.get("reference_images") + if ref_images is not None: + if isinstance(ref_images, str): + ref_images = [PIL.Image.open(ref_images).convert("RGB")] + elif isinstance(ref_images, PIL.Image.Image): + ref_images = [ref_images] + elif isinstance(ref_images, list): + ref_images = [PIL.Image.open(img).convert("RGB") if isinstance(img, str) else img for img in ref_images] + + # Calculate dimensions from first reference image if not provided + if request.sampling_params.height is None or request.sampling_params.width is None: + first_img = ref_images[0] + max_area = 480 * 832 # VACE default is 480p + aspect_ratio = first_img.height / first_img.width + mod_value = 16 + height = round(np.sqrt(max_area * aspect_ratio)) // mod_value * mod_value + width = round(np.sqrt(max_area / aspect_ratio)) // mod_value * mod_value + + if request.sampling_params.height is None: + request.sampling_params.height = height + if request.sampling_params.width is None: + request.sampling_params.width = width + + prompt["additional_information"]["reference_images"] = ref_images + + # Handle source video for V2V / MV2V + source_video = multi_modal_data.get("video") + if source_video is not None: + if isinstance(source_video, list) and len(source_video) > 0: + if isinstance(source_video[0], str): + source_video = [PIL.Image.open(f).convert("RGB") for f in source_video] + prompt["additional_information"]["source_video"] = source_video + + # Handle mask for MV2V / inpainting + mask = multi_modal_data.get("mask") + if mask is not None: + if isinstance(mask, list) and len(mask) > 0: + if isinstance(mask[0], str): + mask = [PIL.Image.open(m).convert("L") for m in mask] + elif isinstance(mask, str): + mask = [PIL.Image.open(mask).convert("L")] + elif isinstance(mask, PIL.Image.Image): + mask = [mask] + prompt["additional_information"]["mask"] = mask + + request.prompt = prompt return request return pre_process_func @@ -469,7 +468,7 @@ def prepare_masks( def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | None = None, negative_prompt: str | None = None, height: int = 480, diff --git a/vllm_omni/diffusion/models/z_image/pipeline_z_image.py b/vllm_omni/diffusion/models/z_image/pipeline_z_image.py index d9f383249c9..104576f41d3 100644 --- a/vllm_omni/diffusion/models/z_image/pipeline_z_image.py +++ b/vllm_omni/diffusion/models/z_image/pipeline_z_image.py @@ -42,7 +42,7 @@ ZImageTransformer2DModel, ) from vllm_omni.diffusion.profiler.diffusion_pipeline_profiler import DiffusionPipelineProfilerMixin -from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.model_executor.model_loader.weight_utils import ( download_weights_from_hf_specific, ) @@ -162,6 +162,8 @@ def retrieve_timesteps( class ZImagePipeline(nn.Module, DiffusionPipelineProfilerMixin, SupportsComponentDiscovery): + supports_request_batch = False + _dit_modules: ClassVar[list[str]] = ["transformer"] _encoder_modules: ClassVar[list[str]] = ["text_encoder"] _vae_modules: ClassVar[list[str]] = ["vae"] @@ -410,7 +412,7 @@ def interrupt(self): def forward( self, - req: OmniDiffusionRequest, + req: DiffusionRequestBatch, prompt: str | list[str] | None = None, image: PipelineImageInput = None, strength: float = 0.6, @@ -828,9 +830,8 @@ def forward( image = self.vae.decode(latents, return_dict=False)[0] # image = self.image_processor.postprocess(image, output_type=output_type) - return DiffusionOutput( - output=image, stage_durations=self.stage_durations if hasattr(self, "stage_durations") else None - ) + stage_durations = self.stage_durations if hasattr(self, "stage_durations") else None + return DiffusionOutput(output=image, stage_durations=stage_durations) def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: loader = AutoWeightsLoader(self) diff --git a/vllm_omni/diffusion/output_formatter.py b/vllm_omni/diffusion/output_formatter.py index 26cfa992e06..0611c6b041a 100644 --- a/vllm_omni/diffusion/output_formatter.py +++ b/vllm_omni/diffusion/output_formatter.py @@ -10,6 +10,7 @@ from vllm_omni.diffusion.io_support import supports_audio_output from vllm_omni.diffusion.registry import DiffusionModelRegistry from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.inputs.data import OmniPromptType from vllm_omni.outputs import OmniRequestOutput @@ -79,12 +80,11 @@ def format_empty_diffusion_outputs( OmniRequestOutput.from_diffusion( request_id=request.request_id, images=[], - prompt=prompt, + prompt=request.prompt, metrics={}, latents=None, finished=finished, ) - for prompt in request.prompts ] @@ -120,25 +120,14 @@ def format_diffusion_outputs( model_cls = DiffusionModelRegistry._try_load_model_cls(od_config.model_class_name) audio_sample_rate = getattr(model_cls, "audio_sample_rate", None) - if len(request.prompts) == 1: - return _format_single_prompt_output( - request=request, - diffusion_output=diffusion_output, - outputs=outputs, - metrics=metrics, - postprocess_output=postprocess_output, - is_text_output=is_text_output, - is_audio_output=is_audio_output, - audio_sample_rate=audio_sample_rate, - finished=diffusion_output.finished, - ) - - return _format_multi_prompt_outputs( + return _format_single_prompt_output( request=request, + prompt=request.prompt, diffusion_output=diffusion_output, outputs=outputs, metrics=metrics, postprocess_output=postprocess_output, + is_text_output=is_text_output, is_audio_output=is_audio_output, audio_sample_rate=audio_sample_rate, finished=diffusion_output.finished, @@ -185,6 +174,7 @@ def _build_multimodal_output( def _format_single_prompt_output( *, request: OmniDiffusionRequest, + prompt: OmniPromptType, diffusion_output: DiffusionOutput, outputs: list[Any], metrics: dict[str, Any], @@ -194,7 +184,6 @@ def _format_single_prompt_output( audio_sample_rate: int | None, finished: bool = True, ) -> list[OmniRequestOutput]: - prompt = request.prompts[0] request_id = request.request_id mm_output = _build_multimodal_output(postprocess_output, audio_sample_rate) @@ -258,114 +247,3 @@ def _format_single_prompt_output( finished=finished, ), ] - - -def _format_multi_prompt_outputs( - *, - request: OmniDiffusionRequest, - diffusion_output: DiffusionOutput, - outputs: list[Any], - metrics: dict[str, Any], - postprocess_output: DiffusionPostprocessOutput, - is_audio_output: bool, - audio_sample_rate: int | None, - finished: bool = True, -) -> list[OmniRequestOutput]: - results = [] - output_idx = 0 - request_id = request.request_id - - for prompt in request.prompts: - num_outputs = request.sampling_params.num_outputs_per_prompt - start_idx = output_idx - end_idx = start_idx + num_outputs - request_outputs = outputs[start_idx:end_idx] if output_idx < len(outputs) else [] - output_idx = end_idx - - if is_audio_output and not _has_non_audio_postprocess_payload(postprocess_output): - if postprocess_output.audio_payload is not None: - request_audio_payload = _slice_batch_payload( - postprocess_output.audio_payload, - start_idx, - end_idx, - num_outputs, - ) - else: - request_audio_payload = request_outputs[0] if len(request_outputs) == 1 else request_outputs - results.append( - OmniRequestOutput.from_diffusion( - request_id=request_id, - images=[], - prompt=prompt, - metrics=metrics, - latents=diffusion_output.trajectory_latents, - trajectory_latents=diffusion_output.trajectory_latents, - trajectory_timesteps=diffusion_output.trajectory_timesteps, - trajectory_log_probs=diffusion_output.trajectory_log_probs, - trajectory_decoded=diffusion_output.trajectory_decoded, - multimodal_output=_format_audio_multimodal_output( - request_audio_payload, - audio_sample_rate, - ), - final_output_type="audio", - stage_durations=diffusion_output.stage_durations, - peak_memory_mb=diffusion_output.peak_memory_mb, - finished=finished, - ), - ) - continue - - mm_output: dict[str, Any] = {} - if postprocess_output.audio_payload is not None: - mm_output["audio"] = _slice_batch_payload( - postprocess_output.audio_payload, - start_idx, - end_idx, - num_outputs, - ) - if audio_sample_rate is not None: - mm_output["audio_sample_rate"] = audio_sample_rate - if postprocess_output.fps is not None: - mm_output["fps"] = postprocess_output.fps - if postprocess_output.action_payload is not None: - mm_output["actions"] = _slice_batch_payload( - postprocess_output.action_payload, - start_idx, - end_idx, - num_outputs, - ) - - results.append( - OmniRequestOutput.from_diffusion( - request_id=request_id, - images=request_outputs, - prompt=prompt, - metrics=metrics, - latents=diffusion_output.trajectory_latents, - trajectory_latents=diffusion_output.trajectory_latents, - trajectory_timesteps=diffusion_output.trajectory_timesteps, - trajectory_log_probs=diffusion_output.trajectory_log_probs, - trajectory_decoded=diffusion_output.trajectory_decoded, - custom_output=postprocess_output.custom_output, - multimodal_output=mm_output, - stage_durations=diffusion_output.stage_durations, - peak_memory_mb=diffusion_output.peak_memory_mb, - finished=finished, - ), - ) - - return results - - -def _slice_batch_payload(payload: Any, start_idx: int, end_idx: int, num_outputs: int) -> Any: - sliced = payload - if isinstance(payload, (list, tuple)): - sliced = payload[start_idx:end_idx] - if len(sliced) == 1: - sliced = sliced[0] - elif hasattr(payload, "shape") and getattr(payload, "shape", None) is not None: - if len(payload.shape) > 0 and payload.shape[0] >= end_idx: - sliced = payload[start_idx:end_idx] - if num_outputs == 1: - sliced = sliced[0] - return sliced diff --git a/vllm_omni/diffusion/request.py b/vllm_omni/diffusion/request.py index 79198c3ec44..42181419c6f 100644 --- a/vllm_omni/diffusion/request.py +++ b/vllm_omni/diffusion/request.py @@ -13,10 +13,11 @@ @dataclass class OmniDiffusionRequest: """ - Complete state passed through the pipeline execution. + Input payload for a single diffusion request. - This dataclass contains the prompts and sampling parameters for the diffusion pipeline + This dataclass contains the prompt and sampling parameters for the diffusion pipeline execution. It also contains a request_id for other components to trace this request and its outputs. + The runner wraps one or more requests into a DiffusionRequestBatch before pipeline execution. """ # TODO(will): double check that args are separate from server_args @@ -24,7 +25,7 @@ class OmniDiffusionRequest: # specific arguments. # data_type: DataType - prompts: list[OmniPromptType] # Actually supporting str-based prompts + prompt: OmniPromptType sampling_params: OmniDiffusionSamplingParams request_id: str kv_sender_info: dict | None = None @@ -50,8 +51,8 @@ def __post_init__(self): self.sampling_params.guidance_scale = 1.0 # Set do_classifier_free_guidance based on guidance scale and negative prompt - if self.sampling_params.guidance_scale > 1.0 and any( - (not isinstance(p, str) and p.get("negative_prompt")) for p in self.prompts + if self.sampling_params.guidance_scale > 1.0 and ( + not isinstance(self.prompt, str) and self.prompt.get("negative_prompt") ): self.sampling_params.do_classifier_free_guidance = True diff --git a/vllm_omni/diffusion/sched/base_scheduler.py b/vllm_omni/diffusion/sched/base_scheduler.py index fe7e212292e..9e4b6af5b6b 100644 --- a/vllm_omni/diffusion/sched/base_scheduler.py +++ b/vllm_omni/diffusion/sched/base_scheduler.py @@ -16,6 +16,7 @@ DiffusionRequestStatus, DiffusionSchedulerOutput, NewRequestData, + RequestBatchSamplingParamsKey, SamplingParamsKey, SchedulerInterface, ) @@ -24,22 +25,31 @@ # LoRA identity is derived from `sampling.lora_request`, not a same-named field # on sampling params, so it must be resolved separately from the bulk lookup. -_KEY_FIELD_NAMES = frozenset(f.name for f in fields(SamplingParamsKey)) - {"lora_int_id"} +_SAMPLING_PARAMS_KEY_FIELD_NAMES = frozenset(f.name for f in fields(SamplingParamsKey)) - {"lora_int_id"} +_REQUEST_BATCH_SAMPLING_PARAMS_KEY_FIELD_NAMES = frozenset(f.name for f in fields(RequestBatchSamplingParamsKey)) - { + "lora_int_id" +} -def get_sampling_params_key(request: OmniDiffusionRequest) -> SamplingParamsKey | None: +def get_sampling_params_key(request: OmniDiffusionRequest) -> SamplingParamsKey: """Build a batch-compatibility key from the request's sampling params.""" - if len(request.prompts) != 1: - return None - sampling = request.sampling_params lora_request = getattr(sampling, "lora_request", None) return SamplingParamsKey( lora_int_id=lora_request.lora_int_id if lora_request is not None else None, - **{name: getattr(sampling, name) for name in _KEY_FIELD_NAMES}, + **{name: getattr(sampling, name) for name in _SAMPLING_PARAMS_KEY_FIELD_NAMES}, ) +def get_request_batch_sampling_params_key(request: OmniDiffusionRequest) -> RequestBatchSamplingParamsKey: + """Build a request-batch compatibility key from the request's sampling params.""" + sampling = request.sampling_params + lora_request = getattr(sampling, "lora_request", None) + key_kwargs = {name: getattr(sampling, name) for name in _REQUEST_BATCH_SAMPLING_PARAMS_KEY_FIELD_NAMES} + key_kwargs["lora_int_id"] = lora_request.lora_int_id if lora_request is not None else None + return RequestBatchSamplingParamsKey(**key_kwargs) + + class _BaseScheduler(SchedulerInterface): """Shared queue/state bookkeeping for diffusion schedulers.""" @@ -49,7 +59,7 @@ def __init__(self) -> None: self._step_id: int = 0 self._waiting: deque[str] = deque() self._running: list[str] = [] - self._running_sampling_params_key: SamplingParamsKey | None = None + self._running_sampling_params_key: SamplingParamsKey | RequestBatchSamplingParamsKey | None = None self._finished_req_ids: set[str] = set() self.max_num_running_reqs: int = 1 self._prefetch_enabled: bool = False @@ -148,6 +158,12 @@ def schedule(self) -> DiffusionSchedulerOutput: def has_requests(self) -> bool: return bool(self._waiting or self._running) + def num_waiting_requests(self) -> int: + return len(self._waiting) + + def num_running_requests(self) -> int: + return len(self._running) + def get_request_state(self, request_id: str) -> DiffusionRequestState | None: return self._request_states.get(request_id) @@ -250,7 +266,7 @@ def _make_request_state(self, request_id: str, request: OmniDiffusionRequest) -> return DiffusionRequestState( request_id=request_id, req=request, - sampling_params_key=get_sampling_params_key(request), + sampling_params_key=self._build_sampling_params_key(request), ) def _can_schedule_waiting(self, state: DiffusionRequestState) -> bool: @@ -260,9 +276,14 @@ def _can_schedule_waiting(self, state: DiffusionRequestState) -> bool: current_key = self._current_sampling_params_key() return current_key is not None and current_key == state.sampling_params_key - def _current_sampling_params_key(self) -> SamplingParamsKey | None: + def _current_sampling_params_key(self) -> SamplingParamsKey | RequestBatchSamplingParamsKey | None: if self._running_sampling_params_key is not None or not self._running: return self._running_sampling_params_key state = self._request_states.get(self._running[0]) self._running_sampling_params_key = None if state is None else state.sampling_params_key return self._running_sampling_params_key + + def _build_sampling_params_key( + self, request: OmniDiffusionRequest + ) -> SamplingParamsKey | RequestBatchSamplingParamsKey | None: + return get_sampling_params_key(request) diff --git a/vllm_omni/diffusion/sched/interface.py b/vllm_omni/diffusion/sched/interface.py index 11dc66a2181..695641b5d94 100644 --- a/vllm_omni/diffusion/sched/interface.py +++ b/vllm_omni/diffusion/sched/interface.py @@ -35,7 +35,7 @@ def is_finished(status: DiffusionRequestStatus) -> bool: @dataclass(frozen=True, eq=True) class SamplingParamsKey: - """Batch-compatibility key derived from ``OmniDiffusionSamplingParams``. + """Denoise step level Batch-compatibility key derived from ``OmniDiffusionSamplingParams``. Only requests with the same key can be batched together. Fields not included here are treated as request-local and do not @@ -60,19 +60,71 @@ class SamplingParamsKey: true_cfg_scale: float | None = None cfg_normalize: bool = False + # Output count. Requests with different num_outputs_per_prompt produce + # differently shaped outputs and cannot share a batch. + num_outputs_per_prompt: int = 1 + # LoRA identity. Requests with different adapters or scales must run in # separate batches so the worker can activate exactly one adapter per step. lora_int_id: int | None = None lora_scale: float = 1.0 +@dataclass(frozen=True, eq=True) +class RequestBatchSamplingParamsKey: + """Request level Batch-compatibility key derived from ``OmniDiffusionSamplingParams``. + + Only request-batch-wide fields belong here. Request-local values such as + seeds, generators, latent tensors, timesteps, and pipeline-specific + ``extra_args`` are read per request from + ``DiffusionRequestBatch.sampling_params_list``. + """ + + # Spatial / temporal shape. + height: object = None + width: object = None + num_frames: int = 1 + resolution: object = 640 + fps: object = None + frame_rate: object = None + boundary_ratio: object = None + + # CFG / guidance. + do_classifier_free_guidance: bool = False + guidance_scale: float = 0.0 + guidance_scale_provided: bool = False + guidance_scale_2: object = None + guidance_rescale: float = 0.0 + true_cfg_scale: object = None + cfg_normalize: bool = False + strength: object = None + + # Scheduling / output shape. + num_inference_steps: object = None + sigmas: object = None + max_sequence_length: object = None + num_outputs_per_prompt: int = 1 + eta: float = 0.0 + decode_timestep: object = None + decode_noise_scale: object = None + output_type: object = None + + # Model-specific batch defaults used by request-mode pipelines. + layers: int = 4 + use_en_prompt: bool = False + + # LoRA identity. + lora_int_id: int | None = None + lora_scale: float = 1.0 + + @dataclass class DiffusionRequestState: """Scheduler-owned state for one queued OmniDiffusionRequest.""" request_id: str req: OmniDiffusionRequest - sampling_params_key: SamplingParamsKey | None = None + sampling_params_key: SamplingParamsKey | RequestBatchSamplingParamsKey | None = None status: DiffusionRequestStatus = DiffusionRequestStatus.WAITING error: str | None = None @@ -82,7 +134,12 @@ def is_finished(self) -> bool: @dataclass class NewRequestData: - """Full request payload for a newly scheduled diffusion request.""" + """Payload for a newly scheduled diffusion request. + + Carries the already-initialized request object so executors and workers do + not re-run ``OmniDiffusionRequest.__post_init__`` and mutate sentinel-based + fields like ``guidance_scale_provided``. + """ request_id: str req: OmniDiffusionRequest @@ -162,6 +219,14 @@ def get_request_state(self, request_id: str) -> DiffusionRequestState | None: def has_requests(self) -> bool: """Return whether the scheduler still owns runnable requests.""" + @abstractmethod + def num_waiting_requests(self) -> int: + """Return the number of requests waiting to be scheduled.""" + + @abstractmethod + def num_running_requests(self) -> int: + """Return the number of requests currently running.""" + @abstractmethod def pop_request_state(self, request_id: str) -> DiffusionRequestState | None: """Remove and return request state if present.""" diff --git a/vllm_omni/diffusion/sched/request_scheduler.py b/vllm_omni/diffusion/sched/request_scheduler.py index ba9cda61615..9a02eb8f3c0 100644 --- a/vllm_omni/diffusion/sched/request_scheduler.py +++ b/vllm_omni/diffusion/sched/request_scheduler.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING from vllm_omni.diffusion.request import OmniDiffusionRequest -from vllm_omni.diffusion.sched.base_scheduler import _BaseScheduler +from vllm_omni.diffusion.sched.base_scheduler import _BaseScheduler, get_request_batch_sampling_params_key from vllm_omni.diffusion.sched.interface import ( DiffusionRequestStatus, DiffusionSchedulerOutput, @@ -19,6 +19,9 @@ class RequestScheduler(_BaseScheduler): """Diffusion scheduler with vLLM-style waiting/running queues.""" + def _build_sampling_params_key(self, request: OmniDiffusionRequest): + return get_request_batch_sampling_params_key(request) + def add_request(self, request: OmniDiffusionRequest) -> str: return super().add_request(request) @@ -41,6 +44,9 @@ def update_from_output(self, sched_output: DiffusionSchedulerOutput, output: Run if result is None: terminal_statuses[request_id] = DiffusionRequestStatus.FINISHED_ERROR terminal_errors[request_id] = "No output result" + elif result.aborted: + terminal_statuses[request_id] = DiffusionRequestStatus.FINISHED_ABORTED + terminal_errors[request_id] = None elif result.error: terminal_statuses[request_id] = DiffusionRequestStatus.FINISHED_ERROR terminal_errors[request_id] = result.error diff --git a/vllm_omni/diffusion/sched/step_scheduler.py b/vllm_omni/diffusion/sched/step_scheduler.py index 0bc2d4a3cda..6cb6a9a28f1 100644 --- a/vllm_omni/diffusion/sched/step_scheduler.py +++ b/vllm_omni/diffusion/sched/step_scheduler.py @@ -88,6 +88,10 @@ def update_from_output(self, sched_output: DiffusionSchedulerOutput, output: Run continue req_result = req_output.result + if req_result is not None and req_result.aborted: + terminal_statuses[request_id] = DiffusionRequestStatus.FINISHED_ABORTED + terminal_errors[request_id] = None + continue output_error = req_result.error if req_result is not None else None if output_error is not None: terminal_statuses[request_id] = DiffusionRequestStatus.FINISHED_ERROR diff --git a/vllm_omni/diffusion/stage_diffusion_client.py b/vllm_omni/diffusion/stage_diffusion_client.py index 1e443b973d2..2b13ddc9d66 100644 --- a/vllm_omni/diffusion/stage_diffusion_client.py +++ b/vllm_omni/diffusion/stage_diffusion_client.py @@ -355,72 +355,6 @@ async def add_request_async( ) ) - # TODO(Long): Temporary solution to boost performance of diffusion stages. - # Remove this after scheduling algorithm is implemented - async def add_batch_request_async( - self, - request_id: str, - prompts: list[OmniPromptType], - sampling_params: OmniDiffusionSamplingParams, - kv_sender_info: dict[int, dict[str, Any]] | None = None, - ) -> None: - """Submit a list of prompts as a single batched engine call. - - All prompts are processed in one ``DiffusionEngine.step()`` call - and the combined result is placed on the output queue with a single - *request_id*. - """ - if self._engine_dead: - raise EngineDeadError() - logger.info( - "[StageDiffusionClient] stage-%s [rep-%s] add batch request: %s (%d prompts)", - self.stage_id, - self.replica_id, - request_id, - len(prompts), - ) - task = asyncio.create_task( - self._run_batch( - request_id, - prompts, - sampling_params, - kv_sender_info, - ), - name=f"diffusion-batch-{request_id}", - ) - self._tasks[request_id] = task - - async def _run_batch( - self, - request_id: str, - prompts: list[OmniPromptType], - sampling_params: OmniDiffusionSamplingParams, - kv_sender_info: dict[int, dict[str, Any]] | None = None, - ) -> None: - try: - self._request_socket.send( - self._encoder.encode( - { - "type": "add_batch_request", - "request_id": request_id, - "prompts": prompts, - "sampling_params": self._sampling_params_to_dict(sampling_params), - "kv_sender_info": kv_sender_info, - } - ) - ) - except Exception as e: - logger.exception( - "[StageDiffusionClient] stage-%s [rep-%s] batch req=%s failed: %s", - self.stage_id, - self.replica_id, - request_id, - e, - ) - await self._output_queue.put(OmniRequestOutput.from_error(request_id, str(e))) - finally: - self._tasks.pop(request_id, None) - def get_diffusion_output_nowait(self) -> OmniRequestOutput | None: self._drain_responses() try: diff --git a/vllm_omni/diffusion/stage_diffusion_proc.py b/vllm_omni/diffusion/stage_diffusion_proc.py index aa058c8d021..e74a34b597b 100644 --- a/vllm_omni/diffusion/stage_diffusion_proc.py +++ b/vllm_omni/diffusion/stage_diffusion_proc.py @@ -16,10 +16,8 @@ from typing import TYPE_CHECKING, Any import msgspec -import torch import zmq import zmq.asyncio -from PIL import Image from vllm.logger import init_logger from vllm.utils.network_utils import get_open_zmq_ipc_path, zmq_socket_ctx from vllm.utils.system_utils import get_mp_context @@ -161,7 +159,7 @@ async def _process_request( sampling_params = self._reconstruct_sampling_params(sampling_params_dict) request = OmniDiffusionRequest( - prompts=[prompt], + prompt=prompt, sampling_params=sampling_params, request_id=request_id, kv_sender_info=kv_sender_info, @@ -184,7 +182,7 @@ async def _process_streaming_request( sampling_params = self._reconstruct_sampling_params(sampling_params_dict) request = OmniDiffusionRequest( - prompts=[prompt], + prompt=prompt, sampling_params=sampling_params, request_id=request_id, kv_sender_info=kv_sender_info, @@ -196,85 +194,6 @@ async def _process_streaming_request( result.request_id = request_id yield result - async def _process_batch_request( - self, - request_id: str, - prompts: list[Any], - sampling_params_dict: dict, - kv_sender_info: dict[str, Any] | None = None, - ) -> OmniRequestOutput: - """Build a batched diffusion request and run DiffusionEngine.step(). - - All prompts are processed in a single step() call. The per-prompt - results are merged into one :class:`OmniRequestOutput` whose - ``images`` list contains every generated image, matching the - contract expected by the orchestrator and tests. - """ - if self._od_config.streaming_output: - raise NotImplementedError("Streaming output is not supported for batched requests") - - sampling_params = self._reconstruct_sampling_params(sampling_params_dict) - - request = OmniDiffusionRequest( - prompts=prompts, - sampling_params=sampling_params, - request_id=request_id, - kv_sender_info=kv_sender_info, - ) - - results = await self._engine.step(request) - - # Merge per-prompt results into a single combined output. - all_images: list = [] - merged_mm: dict[str, Any] = {} - merged_metrics: dict[str, Any] = {} - merged_durations: dict[str, float] = {} - merged_custom: dict[str, Any] = {} - peak_mem = 0.0 - latents = None - trajectory_latents: list[torch.Tensor] | None = None - trajectory_timesteps: list[torch.Tensor] | None = None - trajectory_log_probs: torch.Tensor | None = None - trajectory_decoded: list[Image.Image] | None = None - final_output_type = "image" - - for r in results: - all_images.extend(r.images) - merged_mm.update(r._multimodal_output) - merged_metrics.update(r.metrics) - merged_durations.update(r.stage_durations) - merged_custom.update(r._custom_output) - peak_mem = max(peak_mem, r.peak_memory_mb) - if latents is None and r.latents is not None: - latents = r.latents - if trajectory_latents is None: - trajectory_latents = r.trajectory_latents - if trajectory_timesteps is None: - trajectory_timesteps = r.trajectory_timesteps - if trajectory_log_probs is None: - trajectory_log_probs = r.trajectory_log_probs - if trajectory_decoded is None: - trajectory_decoded = r.trajectory_decoded - if r.final_output_type != "image": - final_output_type = r.final_output_type - - return OmniRequestOutput.from_diffusion( - request_id=request_id, - images=all_images, - prompt=prompts[0] if len(prompts) == 1 else None, - metrics=merged_metrics, - latents=latents, - trajectory_latents=trajectory_latents, - trajectory_timesteps=trajectory_timesteps, - trajectory_log_probs=trajectory_log_probs, - trajectory_decoded=trajectory_decoded, - custom_output=merged_custom or None, - multimodal_output=merged_mm or None, - final_output_type=final_output_type, - stage_durations=merged_durations, - peak_memory_mb=peak_mem, - ) - # ------------------------------------------------------------------ # Collective RPC dispatch # ------------------------------------------------------------------ @@ -500,61 +419,6 @@ async def _dispatch_request( ) tasks[request_id] = task - elif msg_type == "add_batch_request": - request_id = msg["request_id"] - - async def _dispatch_batch( - rid: str, - prompts: list, - sp_dict: dict, - kv_sender_info: dict[str, Any] | None = None, - ) -> None: - try: - result = await self._process_batch_request( - rid, - prompts, - sp_dict, - kv_sender_info=kv_sender_info, - ) - await response_socket.send(encoder.encode({"type": "result", "output": result})) - except DiffusionRequestAbortedError as e: - logger.info( - "request_id: %s aborted: %s", - rid, - str(e), - ) - except Exception as e: - logger.exception("Batch diffusion request %s failed: %s", rid, e) - status_code, error_type = client_error_metadata(e) - await response_socket.send( - encoder.encode( - { - "type": "error", - "request_id": rid, - "error": str(e), - "status_code": status_code, - "error_type": error_type, - } - ) - ) - # Same rationale as the single-request path: a - # closed executor turns every subsequent batch - # into a 500, so escalate now. - if self._is_executor_dead(): - self._signal_fatal_engine_failure(f"add_batch_request {rid}: {e!s}") - finally: - tasks.pop(rid, None) - - task = asyncio.create_task( - _dispatch_batch( - request_id, - msg["prompts"], - msg["sampling_params"], - msg.get("kv_sender_info"), - ) - ) - tasks[request_id] = task - elif msg_type == "abort": for rid in msg.get("request_ids", []): task = tasks.pop(rid, None) diff --git a/vllm_omni/diffusion/worker/diffusion_model_runner.py b/vllm_omni/diffusion/worker/diffusion_model_runner.py index 96a3013669f..84eb3204479 100644 --- a/vllm_omni/diffusion/worker/diffusion_model_runner.py +++ b/vllm_omni/diffusion/worker/diffusion_model_runner.py @@ -38,6 +38,7 @@ from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.sched.interface import DiffusionSchedulerOutput from vllm_omni.diffusion.worker.input_batch import InputBatch, scatter_latents +from vllm_omni.diffusion.worker.request_batch import DiffusionRequestBatch from vllm_omni.diffusion.worker.utils import BatchRunnerOutput, DiffusionRequestState, RunnerOutput from vllm_omni.distributed.omni_connectors.kv_transfer_manager import OmniKVTransferManager from vllm_omni.platforms import current_omni_platform @@ -46,6 +47,43 @@ logger = init_logger(__name__) +def _normalize_pipeline_outputs( + outputs: object, + *, + expected_count: int, + allow_single_output: bool, + pipeline_name: str, +) -> list[DiffusionOutput]: + if isinstance(outputs, DiffusionOutput): + if allow_single_output and expected_count == 1: + return [outputs] + raise RuntimeError( + f"{pipeline_name}.forward returned a single DiffusionOutput; " + "request-batch forward must return list[DiffusionOutput]." + ) + + if not isinstance(outputs, list): + raise RuntimeError( + f"{pipeline_name}.forward returned {type(outputs).__name__}; " + "expected DiffusionOutput or list[DiffusionOutput]." + ) + + if len(outputs) != expected_count: + raise RuntimeError( + f"{pipeline_name}.forward returned {len(outputs)} outputs for {expected_count} requests; " + "expected exactly one DiffusionOutput per request." + ) + + bad_index = next((idx for idx, output in enumerate(outputs) if not isinstance(output, DiffusionOutput)), None) + if bad_index is not None: + raise RuntimeError( + f"{pipeline_name}.forward returned list item {bad_index} with type " + f"{type(outputs[bad_index]).__name__}; expected DiffusionOutput." + ) + + return outputs + + class DiffusionModelRunner(OmniConnectorModelRunnerMixin): """ Model runner that handles model loading and execution for diffusion models. @@ -256,12 +294,13 @@ def get_prompt_embed_cache_stats(self) -> dict | None: return None return self.prompt_embed_cache.stats() - def _record_peak_memory(self, output: DiffusionOutput) -> None: - """Record peak GPU memory for the current forward pass into output. + def _sample_peak_memory_mb(self) -> float: + """Return peak GPU memory for the current forward pass in MB. - Must be called immediately after pipeline.forward(), with + Must be called immediately after the measured forward/step work, with reset_peak_memory_stats() called just before it, so the measurement - reflects this request only and not the global historical maximum. + reflects the current execution slice and not the global historical + maximum. Uses max_memory_reserved (CUDA memory pool high-water mark) rather than max_memory_allocated so that allocator fragmentation is also visible. @@ -270,7 +309,7 @@ def _record_peak_memory(self, output: DiffusionOutput) -> None: peak_reserved_bytes = current_omni_platform.max_memory_reserved() peak_allocated_bytes = current_omni_platform.max_memory_allocated() - output.peak_memory_mb = peak_reserved_bytes / (1024**2) + peak_memory_mb = peak_reserved_bytes / (1024**2) peak_reserved_gb = peak_reserved_bytes / (1024**3) peak_allocated_gb = peak_allocated_bytes / (1024**3) pool_overhead_gb = peak_reserved_gb - peak_allocated_gb @@ -282,120 +321,219 @@ def _record_peak_memory(self, output: DiffusionOutput) -> None: pool_overhead_gb, pool_overhead_gb / peak_reserved_gb * 100 if peak_reserved_gb > 0 else 0.0, ) + return peak_memory_mb - def execute_model(self, req: OmniDiffusionRequest, kv_prefetch_jobs: dict | None = None) -> DiffusionOutput: - """ - Execute a forward pass for the given requests. + def _prepare_request_for_forward( + self, + req: OmniDiffusionRequest, + *, + od_config: OmniDiffusionConfig, + kv_prefetch_jobs: dict | None = None, + use_prefetch: bool = False, + ) -> None: + # Receive AR KV. Single-request execution can use the prefetch path: + # consume prior-forward payload, sync-fallback on miss; request-batch + # execution keeps the synchronous per-request receive path. + kv_recv_t0 = time.perf_counter() + if use_prefetch and self._kv_prefetch_enabled: + self.kv_transfer_manager.consume_and_distribute_kv_cache( + req, + target_device=self.target_device, + ) + else: + self.kv_transfer_manager.receive_multi_kv_cache_distributed( + req, + cfg_kv_collect_func=getattr(od_config, "cfg_kv_collect_func", None), + target_device=self.target_device if use_prefetch else getattr(self.pipeline, "device", None), + ) + kv_recv_ms = (time.perf_counter() - kv_recv_t0) * 1000 + logger.debug("KV recv for %s %.1fms", req.request_id, kv_recv_ms) + + # Kick off the next request's prefetch (+ H2D) to overlap this forward. + if use_prefetch and self._kv_prefetch_enabled and kv_prefetch_jobs is not None: + self.kv_transfer_manager.start_prefetch(kv_prefetch_jobs, self.target_device) + + if req.sampling_params.generator is None and req.sampling_params.seed is not None: + if req.sampling_params.generator_device is not None: + gen_device = req.sampling_params.generator_device + elif self.device.type == "cpu": + gen_device = "cpu" + else: + gen_device = self.device + req.sampling_params.generator = torch.Generator(device=gen_device).manual_seed(req.sampling_params.seed) - Args: - req: A diffusion request containing a list of prompts to process. + def _refresh_cache_for_requests( + self, + reqs: list[OmniDiffusionRequest], + *, + od_config: OmniDiffusionConfig, + ) -> None: + first_req = reqs[0] + if ( + getattr(first_req, "skip_cache_refresh", False) + or self.cache_backend is None + or not self.cache_backend.is_enabled() + ): + return - Returns: - DiffusionOutput with generated results. + # Refresh cache context if needed. Batch admission groups requests by + # SamplingParamsKey, so the first request's num_inference_steps applies + # to the whole runner batch. + num_inference_steps = first_req.sampling_params.num_inference_steps + if num_inference_steps is None and od_config.cache_backend in ( + "tea_cache", + "step_cache", + ): + # When num_inference_steps is None, some pipelines defer to their + # own defaults. TeaCache refresh ignores this value; step_cache + # refresh is a no-op because per-chunk state resets in the denoise + # loop. Use the pipeline default when available to keep refresh + # behavior aligned with single-request execution. + num_inference_steps = getattr(self.pipeline, "num_inference_steps", 0) or 0 + + if num_inference_steps is not None: + self.cache_backend.refresh(self.pipeline, num_inference_steps) + else: + logger.warning( + "Failed to refresh the diffusion transformer cache; backend %s " + "currently requires num_inference_steps to be passed explicitly", + od_config.cache_backend, + ) - Note: - We use torch.no_grad() for HSDP because HSDP2's fully_shard requires access - to tensor version counters in pre_forward hooks, which inference tensors do - not track. For non-HSDP inference, we use torch.inference_mode() for better - performance. - """ - assert self.pipeline is not None, "Model not loaded. Call load_model() first." - if len(req.prompts) == 0: - raise ValueError("Cannot execute model with empty request list") + def _runner_output_from_outputs( + self, + reqs: list[OmniDiffusionRequest], + outputs: list[DiffusionOutput], + ) -> BatchRunnerOutput: + return BatchRunnerOutput.from_list( + [ + RunnerOutput( + request_id=reqs[i].request_id, + step_index=None, + finished=True, + result=outputs[i], + ) + for i in range(len(reqs)) + ] + ) - # Use no_grad() for HSDP compatibility, inference_mode() otherwise for better perf - use_hsdp = self.od_config.parallel_config.use_hsdp + def _execute_request_list( + self, + reqs: list[OmniDiffusionRequest], + *, + od_config: OmniDiffusionConfig, + allow_single_output: bool, + require_request_batch_support: bool, + kv_prefetch_jobs: dict | None = None, + record_name: str, + ) -> BatchRunnerOutput: + assert self.pipeline is not None, "Model not loaded. Call load_model() first." + if not reqs: + return BatchRunnerOutput.from_list([]) + for req in reqs: + if req.prompt is None: + raise ValueError("Cannot execute model with empty prompt") + if require_request_batch_support and not getattr(self.pipeline, "supports_request_batch", False): + raise RuntimeError(f"{type(self.pipeline).__name__} does not support request-batch forward.") + + # Use no_grad() for HSDP compatibility, inference_mode() otherwise for + # better perf. HSDP2's fully_shard pre-forward hooks need tensor version + # counters, which inference tensors do not track. + use_hsdp = od_config.parallel_config.use_hsdp grad_context = torch.no_grad() if use_hsdp else torch.inference_mode() with grad_context: - # Receive AR KV (fetch → distribute → apply inside the entry). prefetch on: - # consume prior-forward payload, sync-fallback on miss; else sync receive. - kv_recv_t0 = time.perf_counter() - if self._kv_prefetch_enabled: - self.kv_transfer_manager.consume_and_distribute_kv_cache( + for req in reqs: + self._prepare_request_for_forward( req, - target_device=self.target_device, + od_config=od_config, + kv_prefetch_jobs=kv_prefetch_jobs, + use_prefetch=allow_single_output, ) - else: - self.kv_transfer_manager.receive_multi_kv_cache_distributed( - req, - cfg_kv_collect_func=getattr(self.od_config, "cfg_kv_collect_func", None), - target_device=self.target_device, - ) - kv_recv_ms = (time.perf_counter() - kv_recv_t0) * 1000 - logger.debug("KV recv for %s %.1fms", req.request_id, kv_recv_ms) - - # Kick off the next request's prefetch (+ H2D) to overlap this forward. - if self._kv_prefetch_enabled and kv_prefetch_jobs is not None: - self.kv_transfer_manager.start_prefetch(kv_prefetch_jobs, self.target_device) - - if req.sampling_params.generator is None and req.sampling_params.seed is not None: - if req.sampling_params.generator_device is not None: - gen_device = req.sampling_params.generator_device - elif self.device.type == "cpu": - gen_device = "cpu" - else: - gen_device = self.device - req.sampling_params.generator = torch.Generator(device=gen_device).manual_seed(req.sampling_params.seed) - # Refresh cache context if needed - if ( - not getattr(req, "skip_cache_refresh", False) - and self.cache_backend is not None - and self.cache_backend.is_enabled() - ): - # FIXME (Alex): When num_inference_steps is None, we defer to - # pipelines for default, but don't refresh the cache; the right - # way to do this is to merge the sampling params first. - # - # For now, if num_inference_steps is not set, we pass 0 to allow - # TeaCache to refresh to align with the param signature. This is - # okay to force refresh TeaCache because the refresh does not use - # num_inference_steps at all (i.e., just resets state and clears - # stale residuals). - num_inference_steps = req.sampling_params.num_inference_steps - if num_inference_steps is None and self.od_config.cache_backend in ( - "tea_cache", - "step_cache", - ): - # TeaCache refresh ignores the value; step_cache refresh is a - # no-op (per-chunk state resets in the denoise loop). DreamZero often - # leaves sampling_params.num_inference_steps unset and uses the - # pipeline default instead. - num_inference_steps = getattr(self.pipeline, "num_inference_steps", 0) or 0 - - if num_inference_steps is not None: - self.cache_backend.refresh(self.pipeline, num_inference_steps) - else: - logger.warning( - "Failed to refresh the diffusion transformer cache; backend %s " - "currently requires num_inference_steps to be passed explicitly", - self.od_config.cache_backend, - ) + self._refresh_cache_for_requests(reqs, od_config=od_config) + batch = DiffusionRequestBatch(requests=reqs) is_primary = not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0 if is_primary: current_omni_platform.reset_peak_memory_stats() - with set_forward_context(vllm_config=self.vllm_config, omni_diffusion_config=self.od_config): - with record_function("pipeline_forward"): - output = self.pipeline.forward(req) + with set_forward_context(vllm_config=self.vllm_config, omni_diffusion_config=od_config): + with record_function(record_name): + raw_outputs = self.pipeline.forward(batch) + outputs = _normalize_pipeline_outputs( + raw_outputs, + expected_count=len(reqs), + allow_single_output=allow_single_output, + pipeline_name=type(self.pipeline).__name__, + ) - if is_primary: - self._record_peak_memory(output) + if is_primary and outputs: + batch_peak_memory_mb = self._sample_peak_memory_mb() + for output in outputs: + output.peak_memory_mb = max(output.peak_memory_mb, batch_peak_memory_mb) - # Log prompt-embed cache activity (hits/misses accumulate across requests). - if is_primary and self.prompt_embed_cache is not None: - logger.debug("prompt-embed cache: %s", self.prompt_embed_cache.stats()) + # Log prompt-embed cache activity; hits/misses accumulate across requests. + prompt_embed_cache = getattr(self, "prompt_embed_cache", None) + if is_primary and prompt_embed_cache is not None: + logger.debug("prompt-embed cache: %s", prompt_embed_cache.stats()) - # NOTE: if ( self.cache_backend is not None and self.cache_backend.is_enabled() - and self.od_config.cache_backend == "cache_dit" - and self.od_config.enable_cache_dit_summary + and od_config.cache_backend == "cache_dit" + and od_config.enable_cache_dit_summary ): cache_summary(self.pipeline, details=True) - return output + + return self._runner_output_from_outputs(reqs, outputs) + + def execute_model(self, req: OmniDiffusionRequest, kv_prefetch_jobs: dict | None = None) -> DiffusionOutput: + """ + Execute a forward pass for the given requests. + + Args: + req: A diffusion request containing a list of prompts to process. + + Returns: + DiffusionOutput with generated results. + + Note: + We use torch.no_grad() for HSDP because HSDP2's fully_shard requires access + to tensor version counters in pre_forward hooks, which inference tensors do + not track. For non-HSDP inference, we use torch.inference_mode() for better + performance. + """ + runner_output = self._execute_request_list( + [req], + od_config=self.od_config, + allow_single_output=True, + require_request_batch_support=False, + kv_prefetch_jobs=kv_prefetch_jobs, + record_name="pipeline_forward", + ) + output = runner_output.runner_outputs[0].result + assert output is not None + return output + + def execute_model_batch( + self, + scheduler_output: DiffusionSchedulerOutput, + od_config: OmniDiffusionConfig, + ) -> BatchRunnerOutput: + """Execute scheduled request-mode requests through the batch forward path. + + Builds a ``DiffusionRequestBatch`` from scheduled new requests, runs + per-request setup, and calls ``pipeline.forward(batch)``. The pipeline + must declare ``supports_request_batch = True``. + """ + reqs = [nr.req for nr in scheduler_output.scheduled_new_reqs] + return self._execute_request_list( + reqs, + od_config=od_config, + allow_single_output=False, + require_request_batch_support=True, + record_name="pipeline_forward_batch", + ) # ------------------------------------------------------------------ # Step-wise execution @@ -418,16 +556,16 @@ def _update_states( # process new requests for sched_new_req in scheduler_output.scheduled_new_reqs: request_id = sched_new_req.request_id - req = sched_new_req.req new_request_ids.append(request_id) if request_id in self.state_cache: raise ValueError(f"Received duplicate new-request payload for cached request {request_id}.") new_state = DiffusionRequestState( request_id=request_id, - sampling=copy.deepcopy(req.sampling_params), - prompts=req.prompts, + sampling=copy.deepcopy(sched_new_req.req.sampling_params), + prompt=sched_new_req.req.prompt, + kv_sender_info=sched_new_req.req.kv_sender_info, ) - state_req = copy.copy(req) + state_req = copy.copy(sched_new_req.req) state_req.sampling_params = new_state.sampling self.kv_transfer_manager.receive_multi_kv_cache_distributed( state_req, @@ -523,6 +661,9 @@ def execute_stepwise(self, scheduler_output: DiffusionSchedulerOutput) -> BatchR states, new_request_ids = self._update_states(scheduler_output) input_batch = self._prepare_batch_inputs(states, new_request_ids) attn_metadata = self._prepare_attn_metadata(input_batch) + is_primary = not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0 + if is_primary: + current_omni_platform.reset_peak_memory_stats() with set_forward_context( vllm_config=self.vllm_config, @@ -548,31 +689,50 @@ def execute_stepwise(self, scheduler_output: DiffusionSchedulerOutput) -> BatchR offset = 0 for req in states: row_num = req.latents.shape[0] - self.pipeline.step_scheduler( - req, noise_pred[offset : offset + row_num] if noise_pred is not None else None - ) - offset = offset + row_num - if self.od_config.streaming_output: - should_decode = req.chunk_denoise_completed - else: - should_decode = req.denoise_completed - - if should_decode: - result = self.pipeline.post_decode(req) - else: - result = None - # finished should be computed after post_decode() advanced chunk_index - finished = ( - req.request_denoise_completed if self.od_config.streaming_output else req.denoise_completed - ) - runner_output_list.append( - RunnerOutput( - request_id=req.request_id, - step_index=req.step_index, - finished=finished, - result=result, + try: + self.pipeline.step_scheduler( + req, noise_pred[offset : offset + row_num] if noise_pred is not None else None + ) + offset = offset + row_num + if self.od_config.streaming_output: + should_decode = req.chunk_denoise_completed + else: + should_decode = req.denoise_completed + + if should_decode: + result = self.pipeline.post_decode(req) + else: + result = None + # finished should be computed after post_decode() advanced chunk_index + finished = ( + req.request_denoise_completed + if self.od_config.streaming_output + else req.denoise_completed + ) + runner_output_list.append( + RunnerOutput( + request_id=req.request_id, + step_index=req.step_index, + finished=finished, + result=result, + ) + ) + except Exception as per_req_exc: + offset = offset + row_num + logger.error( + "Stepwise per-request error for %s: %s", + req.request_id, + per_req_exc, + exc_info=True, + ) + runner_output_list.append( + RunnerOutput( + request_id=req.request_id, + step_index=req.step_index, + finished=True, + result=DiffusionOutput(error=str(per_req_exc)), + ) ) - ) if noise_pred is not None and offset != noise_pred.shape[0]: raise ValueError( @@ -580,6 +740,22 @@ def execute_stepwise(self, scheduler_output: DiffusionSchedulerOutput) -> BatchR f"but batched noise_pred has {noise_pred.shape[0]} rows." ) + if is_primary: + batch_peak_memory_mb = self._sample_peak_memory_mb() + states_by_id = {state.request_id: state for state in states} + for state in states: + state.peak_memory_mb = max(state.peak_memory_mb, batch_peak_memory_mb) + for runner_output in runner_output_list: + if runner_output.result is None: + continue + state = states_by_id.get(runner_output.request_id) + if state is None: + continue + runner_output.result.peak_memory_mb = max( + runner_output.result.peak_memory_mb, + state.peak_memory_mb, + ) + self._update_states_after(states, input_batch, pipeline_interrupted) return BatchRunnerOutput.from_list(runner_output_list) diff --git a/vllm_omni/diffusion/worker/diffusion_worker.py b/vllm_omni/diffusion/worker/diffusion_worker.py index 6b7e8863eb5..6518464c71e 100644 --- a/vllm_omni/diffusion/worker/diffusion_worker.py +++ b/vllm_omni/diffusion/worker/diffusion_worker.py @@ -48,7 +48,7 @@ from vllm_omni.diffusion.request import OmniDiffusionRequest from vllm_omni.diffusion.sched.interface import DiffusionSchedulerOutput from vllm_omni.diffusion.worker.diffusion_model_runner import DiffusionModelRunner -from vllm_omni.diffusion.worker.utils import BaseRunnerOutput +from vllm_omni.diffusion.worker.utils import BaseRunnerOutput, BatchRunnerOutput from vllm_omni.engine.stage_init_utils import set_death_signal from vllm_omni.lora.request import LoRARequest from vllm_omni.platforms import current_omni_platform @@ -392,6 +392,28 @@ def execute_model( profiler.step() return output + def execute_model_batch( + self, scheduler_output: DiffusionSchedulerOutput, od_config: OmniDiffusionConfig + ) -> BatchRunnerOutput: + """Batch forward: LoRA activate once, delegate to model runner.""" + assert self.model_runner is not None, "Model runner not initialized" + # LoRA: same adapter/scale within batch guaranteed by SamplingParamsKey + if self.lora_manager is not None and scheduler_output.scheduled_new_reqs: + sp = scheduler_output.scheduled_new_reqs[0].req.sampling_params + try: + self.lora_manager.set_active_adapter(sp.lora_request, sp.lora_scale) + except Exception as exc: + if sp.lora_request is not None: + raise + logger.warning("LoRA activation skipped: %s", exc) + profiler = self._get_profiler() + ctx = profiler.annotate_context_manager("diffusion_forward_batch") if profiler else nullcontext() + with ctx: + output = self.model_runner.execute_model_batch(scheduler_output, od_config) + if profiler: + profiler.step() + return output + def execute_stepwise(self, scheduler_output: DiffusionSchedulerOutput) -> BaseRunnerOutput: """Execute one diffusion step by delegating to the model runner.""" assert self.model_runner is not None, "Model runner not initialized" @@ -887,7 +909,7 @@ def worker_busy_loop(self) -> None: continue else: - # Handle generation request + # Handle direct generation requests. try: output = self.worker.execute_model(msg, self.od_config) except Exception as e: diff --git a/vllm_omni/diffusion/worker/input_batch.py b/vllm_omni/diffusion/worker/input_batch.py index 37554677445..511f2dd6963 100644 --- a/vllm_omni/diffusion/worker/input_batch.py +++ b/vllm_omni/diffusion/worker/input_batch.py @@ -754,4 +754,7 @@ def scatter_latents( ) -DiffusionInputBatch = InputBatch +# Alias: InputBatch is the step/tensor-level batch. +# DiffusionRequestBatch (in request_batch.py) is the request-level batch. +StepInputBatch = InputBatch +DiffusionInputBatch = StepInputBatch diff --git a/vllm_omni/diffusion/worker/request_batch.py b/vllm_omni/diffusion/worker/request_batch.py new file mode 100644 index 00000000000..1f7c2acaabd --- /dev/null +++ b/vllm_omni/diffusion/worker/request_batch.py @@ -0,0 +1,266 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project +"""Request-level batch abstraction for diffusion runner.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Any + +import torch + +from vllm_omni.diffusion.data import DiffusionOutput +from vllm_omni.diffusion.request import OmniDiffusionRequest +from vllm_omni.inputs.data import OmniDiffusionSamplingParams, OmniPromptType + + +def _slice_request_output(value: Any, start: int, stop: int) -> Any: + if isinstance(value, tuple): + return tuple(_slice_request_output(item, start, stop) for item in value) + if isinstance(value, list): + return value[start:stop] + if isinstance(value, torch.Tensor): + return value[start:stop] + return value + + +def split_diffusion_output_by_request( + result: DiffusionOutput, + req: DiffusionRequestBatch, + *, + num_outputs_per_prompt: int, +) -> list[DiffusionOutput]: + """Split a batched DiffusionOutput into one output per request.""" + if num_outputs_per_prompt <= 0: + raise ValueError(f"num_outputs_per_prompt must be positive, got {num_outputs_per_prompt}.") + + return [ + DiffusionOutput( + output=_slice_request_output( + result.output, + idx * num_outputs_per_prompt, + (idx + 1) * num_outputs_per_prompt, + ), + error=result.error, + finished=result.finished, + stage_durations=result.stage_durations, + peak_memory_mb=result.peak_memory_mb, + chunk_index=result.chunk_index, + total_chunks=result.total_chunks, + ) + for idx in range(req.num_reqs) + ] + + +@dataclass +class DiffusionRequestBatch: + """Request-level batch wrapping original diffusion requests. + + Each :class:`~vllm_omni.diffusion.request.OmniDiffusionRequest` represents + one logical diffusion request with one prompt. The scheduler and runner use + this wrapper to present a compatible request batch to pipeline + ``forward()`` methods without reintroducing list-shaped request payloads. + + This is distinct from ``InputBatch`` (aliased as ``StepInputBatch``), + which manages step/tensor-level data for stepwise execution. + + Args: + requests: Independent diffusion requests scheduled together for + request-mode execution. + + Attributes: + requests: Original request objects in scheduler order. + num_reqs: Number of requests in the batch. + request_ids: Request IDs in the same order as ``requests``. + prompts: Prompt list assembled from each request's single ``prompt``. + sampling_params_list: Per-request sampling parameters in scheduler + order. Request-batch pipelines read request-local values here. + sampling_params: Sampling parameters for single-request legacy paths. + request_id: First request ID, kept as a compatibility convenience for + code paths that handle a single-request batch. + kv_sender_info: KV-transfer metadata from the first request. + """ + + requests: list[OmniDiffusionRequest] + + @property + def num_reqs(self) -> int: + return len(self.requests) + + @property + def request_ids(self) -> list[str]: + return [req.request_id for req in self.requests] + + @property + def prompts(self) -> list[OmniPromptType]: + return [req.prompt for req in self.requests] + + @property + def sampling_params_list(self) -> list[OmniDiffusionSamplingParams]: + return [req.sampling_params for req in self.requests] + + @property + def sampling_params(self) -> OmniDiffusionSamplingParams: + # Legacy pipelines do not accept RequestBatch, so they are invoked with + # one request at a time. In that path, this batch is expected to contain + # a single request, and we expose its sampling params for compatibility. + assert len(self.requests) == 1, "RequestBatch with multiple requests does not have a single sampling_params" + return self.requests[0].sampling_params + + @property + def request_id(self) -> str: + return self.requests[0].request_id + + @property + def kv_sender_info(self) -> dict | None: + return self.requests[0].kv_sender_info + + def is_dummy_run(self) -> bool: + return self.requests[0].is_dummy_run_request_id(self.request_id) + + def get(self, request_id: str) -> OmniDiffusionRequest | None: + for req in self.requests: + if req.request_id == request_id: + return req + return None + + def collate_request_generators( + self, + num_outputs_per_prompt: int, + default_generator: torch.Generator | list[torch.Generator] | None, + ) -> torch.Generator | list[torch.Generator] | None: + return self.collate_sampling_param_generators( + self.sampling_params_list, + num_outputs_per_prompt, + default_generator, + ) + + def collate_request_tensors( + self, + attr: str, + default_tensor: torch.Tensor | None, + ) -> torch.Tensor | None: + return self.collate_tensors( + [getattr(sampling, attr) for sampling in self.sampling_params_list], + attr, + default_tensor, + ) + + @staticmethod + def collate_tensors( + tensors: list[torch.Tensor | None], + name: str, + default_tensor: torch.Tensor | None, + ) -> torch.Tensor | None: + validated_tensors = DiffusionRequestBatch._validate_tensor_sequence(tensors, name) + if validated_tensors is None: + return default_tensor + return torch.cat(validated_tensors, dim=0) + + @staticmethod + def collate_prompt_tensors( + tensors: list[torch.Tensor | None], + name: str, + default_tensor: torch.Tensor | None, + ) -> torch.Tensor | None: + validated_tensors = DiffusionRequestBatch._validate_tensor_sequence(tensors, name) + if validated_tensors is None: + return default_tensor + return torch.stack(validated_tensors, dim=0) + + @staticmethod + def get_prompt_field(prompt: OmniPromptType, name: str) -> Any: + if isinstance(prompt, str): + return None + value = prompt.get(name) + if value is None: + additional = prompt.get("additional_information") + if isinstance(additional, dict): + value = additional.get(name) + if isinstance(value, list): + return value[0] if value else None + return value + + @staticmethod + def collate_prompt_fields( + prompts: list[OmniPromptType], + name: str, + default_tensor: torch.Tensor | None, + ) -> torch.Tensor | None: + return DiffusionRequestBatch.collate_prompt_tensors( + [DiffusionRequestBatch.get_prompt_field(prompt, name) for prompt in prompts], + name, + default_tensor, + ) + + @staticmethod + def get_prompt_field_with_aliases(prompt: OmniPromptType, names: Sequence[str]) -> Any: + for name in names: + value = DiffusionRequestBatch.get_prompt_field(prompt, name) + if value is not None: + return value + return None + + @staticmethod + def collate_prompt_field_map( + prompts: list[OmniPromptType], + field_defaults: Mapping[str, torch.Tensor | None], + field_aliases: Mapping[str, Sequence[str]] | None = None, + ) -> dict[str, torch.Tensor | None]: + collated_fields: dict[str, torch.Tensor | None] = {} + for name, default_tensor in field_defaults.items(): + aliases = field_aliases.get(name, (name,)) if field_aliases is not None else (name,) + collated_fields[name] = DiffusionRequestBatch.collate_prompt_tensors( + [DiffusionRequestBatch.get_prompt_field_with_aliases(prompt, aliases) for prompt in prompts], + name, + default_tensor, + ) + return collated_fields + + @staticmethod + def _validate_tensor_sequence( + tensors: list[torch.Tensor | None], + name: str, + ) -> list[torch.Tensor] | None: + if not any(tensor is not None for tensor in tensors): + return None + if not all(isinstance(tensor, torch.Tensor) for tensor in tensors): + raise ValueError(f"Cannot batch requests with a mix of provided and missing {name}.") + + first = tensors[0] + assert isinstance(first, torch.Tensor) + for tensor in tensors[1:]: + assert isinstance(tensor, torch.Tensor) + if tensor.shape != first.shape or tensor.dtype != first.dtype or tensor.device != first.device: + raise ValueError( + f"Batched request {name} must have matching shape, dtype, and device; " + f"got {tensor.shape}/{tensor.dtype}/{tensor.device} and " + f"{first.shape}/{first.dtype}/{first.device}." + ) + return tensors + + @staticmethod + def collate_sampling_param_generators( + sampling_params_list: list[Any], + num_outputs_per_prompt: int, + default_generator: torch.Generator | list[torch.Generator] | None, + ) -> torch.Generator | list[torch.Generator] | None: + request_generators = [sampling.generator for sampling in sampling_params_list] + if not any(generator is not None for generator in request_generators): + return default_generator + if not all(generator is not None for generator in request_generators): + raise ValueError("Cannot batch requests with a mix of provided and missing generators.") + + generators: list[torch.Generator] = [] + for generator in request_generators: + if isinstance(generator, list): + if len(generator) != num_outputs_per_prompt: + raise ValueError( + "Per-request generator lists must match num_outputs_per_prompt, " + f"got {len(generator)} and {num_outputs_per_prompt}." + ) + generators.extend(generator) + else: + generators.extend([generator] * num_outputs_per_prompt) + return generators diff --git a/vllm_omni/diffusion/worker/utils.py b/vllm_omni/diffusion/worker/utils.py index 02cb0b14ac8..4fd374d34b5 100644 --- a/vllm_omni/diffusion/worker/utils.py +++ b/vllm_omni/diffusion/worker/utils.py @@ -40,7 +40,8 @@ class DiffusionRequestState: # ── Identity / request-level inputs ── request_id: str sampling: OmniDiffusionSamplingParams - prompts: list[OmniPromptType] | None = None + prompt: OmniPromptType | None = None + kv_sender_info: dict | None = None # ── Encoded prompts (set once by prepare_encode) ── prompt_embeds: torch.Tensor | None = None @@ -78,6 +79,9 @@ class DiffusionRequestState: # For example: Wan condition tensors / masks, or Bagel KV contexts. extra: dict[str, Any] = field(default_factory=dict) + # Peak device memory observed while this request is active in step mode. + peak_memory_mb: float = 0.0 + # ── Properties ── @property diff --git a/vllm_omni/engine/arg_utils.py b/vllm_omni/engine/arg_utils.py index acfc0a83f63..725424854a9 100644 --- a/vllm_omni/engine/arg_utils.py +++ b/vllm_omni/engine/arg_utils.py @@ -172,6 +172,8 @@ class OmniEngineArgs(EngineArgs): # in __post_init__ based on worker_type (ar/generation), so None is safe here. enable_sleep_mode: bool = False omni: bool = False + # Diffusion request-mode batch admission (forwarded to OmniDiffusionConfig). + request_batch_max_wait_ms: float = 0.0 @classmethod def _add_omni_specific_args(cls, parser: argparse.ArgumentParser) -> argparse.ArgumentParser: diff --git a/vllm_omni/engine/async_omni_engine.py b/vllm_omni/engine/async_omni_engine.py index 5b8d2430feb..f7305f5ec61 100644 --- a/vllm_omni/engine/async_omni_engine.py +++ b/vllm_omni/engine/async_omni_engine.py @@ -1000,6 +1000,7 @@ def _create_default_diffusion_stage_cfg(kwargs: dict[str, Any]) -> list: "model_config": kwargs.get("model_config", None), "additional_config": kwargs.get("additional_config", None), "step_execution": kwargs.get("step_execution", False), + "request_batch_max_wait_ms": kwargs.get("request_batch_max_wait_ms", 0.0), "vae_use_slicing": kwargs.get("vae_use_slicing", False), "vae_use_tiling": kwargs.get("vae_use_tiling", False), "cache_backend": cache_backend, diff --git a/vllm_omni/engine/orchestrator.py b/vllm_omni/engine/orchestrator.py index 41207102b1c..e95293f47a7 100644 --- a/vllm_omni/engine/orchestrator.py +++ b/vllm_omni/engine/orchestrator.py @@ -1216,6 +1216,31 @@ async def _forward_to_next_stage( src_stage_id, next_logical, ) + if diffusion_prompt is None: + error_output = OmniRequestOutput.from_error( + req_id, + f"Stage-{src_stage_id} produced no valid inputs for diffusion stage-{next_logical}", + ) + logger.warning( + "[Orchestrator] req=%s stage=%d produced empty diffusion inputs for stage=%d; " + "routing terminal error output", + req_id, + src_stage_id, + next_logical, + ) + await self.output_async_queue.put( + OutputMessage( + request_id=req_id, + stage_id=next_logical, + engine_outputs=error_output, + metrics=None, + finished=True, + ) + ) + await self._cleanup_request_ids( + [req_id, *self._cfg_tracker.cleanup_parent(req_id)], + ) + return if isinstance(diffusion_prompt, list): if not diffusion_prompt: error_output = OmniRequestOutput.from_error( diff --git a/vllm_omni/engine/stage_client.py b/vllm_omni/engine/stage_client.py index 9236fbc1802..276775f6674 100644 --- a/vllm_omni/engine/stage_client.py +++ b/vllm_omni/engine/stage_client.py @@ -100,12 +100,4 @@ async def add_request_async( kv_sender_info: dict[int, dict[str, Any]] | None = None, ) -> None: ... - async def add_batch_request_async( - self, - request_id: str, - prompts: list[OmniPromptType], - sampling_params: OmniDiffusionSamplingParams, - kv_sender_info: dict[int, dict[str, Any]] | None = None, - ) -> None: ... - def get_diffusion_output_nowait(self) -> OmniRequestOutput | None: ... diff --git a/vllm_omni/engine/stage_pool.py b/vllm_omni/engine/stage_pool.py index 4e1f3d5a55b..1c116b776ed 100644 --- a/vllm_omni/engine/stage_pool.py +++ b/vllm_omni/engine/stage_pool.py @@ -918,15 +918,17 @@ async def submit_initial( params = OmniDiffusionSamplingParams() submit_kwargs = dict(submit_kwargs or {}) if self.stage_type == "diffusion": + if isinstance(request, list): + raise ValueError( + "Diffusion list-prompt batch requests are no longer supported. " + "Submit multiple independent requests to use scheduler batching." + ) replica_id = await self._pick_or_select( request_id, affinity_request_id=affinity_request_id, ) client = self._diffusion_client(replica_id) - if isinstance(request, list): - await client.add_batch_request_async(request_id, request, params, **submit_kwargs) - else: - await client.add_request_async(request_id, request, params, **submit_kwargs) + await client.add_request_async(request_id, request, params, **submit_kwargs) return replica_id replica_id = await self._pick_or_select( @@ -987,6 +989,11 @@ async def submit_update( raise RuntimeError(f"stage {self.stage_id} replica {replica_id} is not attached") if self.stage_type == "diffusion": + if isinstance(request, list): + raise ValueError( + "Diffusion list-prompt batch requests are no longer supported. " + "Submit multiple independent requests to use scheduler batching." + ) await self._diffusion_client(replica_id).add_request_async(request_id, request, params) else: # Refresh the shared output-processor state before yielding to the diff --git a/vllm_omni/entrypoints/async_omni.py b/vllm_omni/entrypoints/async_omni.py index 6a1f505e49b..ab2d8836b3a 100644 --- a/vllm_omni/entrypoints/async_omni.py +++ b/vllm_omni/entrypoints/async_omni.py @@ -270,15 +270,16 @@ async def generate( through all stages in the pipeline and yields outputs as they become available. - **Batch mode (diffusion only):** - When *prompt* is a ``list``, all prompts are dispatched in a single - ``DiffusionEngine.step()`` call at the diffusion stage. The combined - result is yielded as one ``OmniRequestOutput`` with all generated - images. Only a single *request_id* is used for the whole batch. + **Diffusion batching:** + Diffusion stages accept only a single prompt per request. Passing a + ``list`` of prompts to a diffusion stage will raise ``ValueError``. + To batch multiple diffusion prompts, submit each as an independent + request; the scheduler will automatically co-batch compatible requests. Args: - prompt: A single prompt **or** a list of prompts. A list - triggers batch mode when the diffusion stage is reached. + prompt: A single prompt **or** a list of prompts. For diffusion + stages, only a single prompt is accepted; a list will be + rejected with an error. request_id: Unique identifier for this request. If one is not provided, a random one will be generated. sampling_params_list: List of SamplingParams, one per stage. @@ -288,11 +289,10 @@ async def generate( Yields: OmniRequestOutput objects as they are produced by each stage. - In batch mode the diffusion stage yields one output containing - all generated images. Raises: - ValueError: If sampling_params_list has incorrect length. + ValueError: If sampling_params_list has incorrect length, or + if a list prompt is submitted to a diffusion stage. """ # Append a random UUID suffix to the request_id to ensure it is unique # and non-empty, similar to vLLM's input processor. The suffix is used @@ -306,6 +306,15 @@ async def generate( logger.debug(f"[AsyncOmni] generate() called for request {external_request_id}") + # Reject diffusion list-prompt early with a clear API error. + if isinstance(prompt, list) and any( + getattr(client, "stage_type", "") == "diffusion" for client in getattr(self.engine, "stage_clients", []) + ): + raise ValueError( + "Diffusion stages accept only a single prompt per request. " + "Submit multiple independent requests to use scheduler batching." + ) + input_stream_task: asyncio.Task | None = None try: # Start final output dispatcher on the first call to generate() diff --git a/vllm_omni/entrypoints/cli/serve.py b/vllm_omni/entrypoints/cli/serve.py index d4c0b462d42..5f94632460e 100644 --- a/vllm_omni/entrypoints/cli/serve.py +++ b/vllm_omni/entrypoints/cli/serve.py @@ -556,6 +556,14 @@ def subparser_init(self, subparsers: argparse._SubParsersAction) -> TrackingArgu action="store_true", help="Enable per-step diffusion execution so running requests can be aborted between denoise steps.", ) + omni_config_group.add_argument( + "--request-batch-max-wait-ms", + type=float, + default=0.0, + help="Request-mode batch admission: max milliseconds to wait for compatible " + "requests to accumulate before scheduling a fused forward wave. " + "0 disables admission (default).", + ) # VAE memory optimization parameters omni_config_group.add_argument( diff --git a/vllm_omni/entrypoints/openpi/serving.py b/vllm_omni/entrypoints/openpi/serving.py index ec5c88408ea..4f1d5808144 100644 --- a/vllm_omni/entrypoints/openpi/serving.py +++ b/vllm_omni/entrypoints/openpi/serving.py @@ -128,7 +128,7 @@ async def infer(self, obs: dict, *, session_id: str, reset: bool) -> ActionOutpu # exposes an async iterator, so consume it to completion and use the # final output, matching other non-streaming OpenAI serving paths. async for output in self.engine_client.generate( - prompt=request.prompts[0], + prompt=request.prompt, request_id=request.request_id, sampling_params_list=[request.sampling_params], ): @@ -159,7 +159,7 @@ def _build_request(self, obs: dict, *, session_id: str, reset: bool) -> Any: prompt = obs.get("prompt", "") sampling_params = OmniDiffusionSamplingParams(extra_args=extra_args) return OmniDiffusionRequest( - prompts=[prompt], + prompt=prompt, sampling_params=sampling_params, request_id=self._next_request_id(session_id), ) diff --git a/vllm_omni/model_executor/stage_input_processors/glm_image.py b/vllm_omni/model_executor/stage_input_processors/glm_image.py index bf644abd0e3..d31d8c44169 100644 --- a/vllm_omni/model_executor/stage_input_processors/glm_image.py +++ b/vllm_omni/model_executor/stage_input_processors/glm_image.py @@ -217,192 +217,192 @@ def ar2diffusion( prompt: OmniTokensPrompt | TextPrompt | list | None = None, requires_multimodal_data: bool = False, streaming_context: Any | None = None, -) -> list[dict[str, Any]]: +) -> dict[str, Any] | None: """Process AR stage outputs to create Diffusion stage inputs. - This processor accepts the stage-pool transition interface: - ``ar2diffusion(source_outputs, prompt, requires_multimodal_data)``. + GLM-Image only produces one downstream diffusion request per AR request. + ``source_outputs`` may still include CFG companion outputs, but only the + first AR output is used to build the diffusion payload. """ del streaming_context _t_total = time.perf_counter() - ar_outputs = source_outputs - diffusion_inputs = [] - - # Normalize prompt to list - if not isinstance(prompt, list): - prompt = [prompt] if prompt is not None else [{}] - - for i, ar_output in enumerate(ar_outputs): - _t_req = time.perf_counter() - output = ar_output.outputs[0] - generated_token_ids = output.cumulative_token_ids - - # Get original prompt info - original_prompt = prompt[i] if i < len(prompt) else {} - if isinstance(original_prompt, dict): - pass - elif hasattr(original_prompt, "_asdict"): - original_prompt = original_prompt._asdict() - elif hasattr(original_prompt, "__dict__"): - original_prompt = vars(original_prompt) - else: - original_prompt = {} - - mm_processor_kwargs = original_prompt.get("mm_processor_kwargs") - - def _coerce_dim(v: Any, default: int) -> int: - try: - iv = int(v) - return iv if iv > 0 else default - except (TypeError, ValueError): - return default - - # Prefer GLM-Image target size from mm_processor_kwargs (set by serving layer), - # then fall back to top-level fields for backward compatibility. - height = _coerce_dim( - mm_processor_kwargs.get("target_h") if isinstance(mm_processor_kwargs, dict) else None, - _coerce_dim(original_prompt.get("height"), 1024), - ) - width = _coerce_dim( - mm_processor_kwargs.get("target_w") if isinstance(mm_processor_kwargs, dict) else None, - _coerce_dim(original_prompt.get("width"), 1024), - ) - text_prompt = original_prompt.get("prompt", "") - - # Detect i2i mode. - # Prefer normalized prompt multi_modal_data source-image presence, with - # multimodal output as secondary signal. - _t_mode = time.perf_counter() - is_i2i = False + if not source_outputs: + return None - prompt_modalities = original_prompt.get("modalities") - if isinstance(prompt_modalities, list) and "img2img" in prompt_modalities: - is_i2i = True + ar_output = source_outputs[0] + _t_req = time.perf_counter() + output = ar_output.outputs[0] + generated_token_ids = output.cumulative_token_ids - prompt_mm_data = original_prompt.get("multi_modal_data") - if _has_source_image(prompt_mm_data): - is_i2i = True + if isinstance(prompt, list): + original_prompt = prompt[0] if prompt else {} + elif prompt is not None: + original_prompt = prompt + else: + original_prompt = {} + + if isinstance(original_prompt, dict): + pass + elif hasattr(original_prompt, "_asdict"): + original_prompt = original_prompt._asdict() + elif hasattr(original_prompt, "__dict__"): + original_prompt = vars(original_prompt) + else: + original_prompt = {} - if hasattr(ar_output, "multimodal_output") and ar_output.multimodal_output: - mm_output = ar_output.multimodal_output - if isinstance(mm_output, Mapping) and mm_output.get("ids", {}).get("prior_image") is not None: - is_i2i = True - _dt_mode = (time.perf_counter() - _t_mode) * 1000 + mm_processor_kwargs = original_prompt.get("mm_processor_kwargs") - # Parse and upsample prior tokens - _t_parse = time.perf_counter() + def _coerce_dim(v: Any, default: int) -> int: try: - prior_token_ids, pixel_h, pixel_w = _parse_generated_tokens( - generated_token_ids, - height, - width, - is_i2i=is_i2i, - ) - except ValueError as e: - logger.warning( - "[ar2diffusion] Request %s: skip due to token parse failure: %s " - "(target=%sx%s, mode=%s, raw_tokens=%s, tail=%s)", - i, - e, - height, - width, - "i2i" if is_i2i else "t2i", - len(generated_token_ids), - generated_token_ids[-8:] if len(generated_token_ids) >= 8 else generated_token_ids, - ) - continue - _dt_parse = (time.perf_counter() - _t_parse) * 1000 - - # Get prior_token_image_ids from AR model output (for i2i mode) - # This contains VQ-VAE tokens from input image, used for KV cache conditioning - # NOTE: multimodal_output is attached to ar_output (RequestOutput), NOT output (CompletionOutput) - _t_prior_img = time.perf_counter() - prior_token_image_ids = None - - # Check ar_output (RequestOutput) for multimodal_output - this is the correct location - if hasattr(ar_output, "multimodal_output") and ar_output.multimodal_output: - mm_output = ar_output.multimodal_output + iv = int(v) + return iv if iv > 0 else default + except (TypeError, ValueError): + return default + + # Prefer GLM-Image target size from mm_processor_kwargs (set by serving layer), + # then fall back to top-level fields for backward compatibility. + height = _coerce_dim( + mm_processor_kwargs.get("target_h") if isinstance(mm_processor_kwargs, dict) else None, + _coerce_dim(original_prompt.get("height"), 1024), + ) + width = _coerce_dim( + mm_processor_kwargs.get("target_w") if isinstance(mm_processor_kwargs, dict) else None, + _coerce_dim(original_prompt.get("width"), 1024), + ) + text_prompt = original_prompt.get("prompt", "") + + # Detect i2i mode. + # Prefer normalized prompt multi_modal_data source-image presence, with + # multimodal output as secondary signal. + _t_mode = time.perf_counter() + is_i2i = False + + prompt_modalities = original_prompt.get("modalities") + if isinstance(prompt_modalities, list) and "img2img" in prompt_modalities: + is_i2i = True + + prompt_mm_data = original_prompt.get("multi_modal_data") + if _has_source_image(prompt_mm_data): + is_i2i = True + + if hasattr(ar_output, "multimodal_output") and ar_output.multimodal_output: + mm_output = ar_output.multimodal_output + if isinstance(mm_output, Mapping) and mm_output.get("ids", {}).get("prior_image") is not None: + is_i2i = True + _dt_mode = (time.perf_counter() - _t_mode) * 1000 + + # Parse and upsample prior tokens + _t_parse = time.perf_counter() + try: + prior_token_ids, pixel_h, pixel_w = _parse_generated_tokens( + generated_token_ids, + height, + width, + is_i2i=is_i2i, + ) + except ValueError as e: + logger.warning( + "[ar2diffusion] Request %s: skip due to token parse failure: %s " + "(target=%sx%s, mode=%s, raw_tokens=%s, tail=%s)", + 0, + e, + height, + width, + "i2i" if is_i2i else "t2i", + len(generated_token_ids), + generated_token_ids[-8:] if len(generated_token_ids) >= 8 else generated_token_ids, + ) + return None + _dt_parse = (time.perf_counter() - _t_parse) * 1000 + + # Get prior_token_image_ids from AR model output (for i2i mode) + # This contains VQ-VAE tokens from input image, used for KV cache conditioning + # NOTE: multimodal_output is attached to ar_output (RequestOutput), NOT output (CompletionOutput) + _t_prior_img = time.perf_counter() + prior_token_image_ids = None + + # Check ar_output (RequestOutput) for multimodal_output - this is the correct location + if hasattr(ar_output, "multimodal_output") and ar_output.multimodal_output: + mm_output = ar_output.multimodal_output + if isinstance(mm_output, Mapping): + raw_prior_image_ids = mm_output.get("ids", {}).get("prior_image") + if raw_prior_image_ids is not None: + # Handle different formats: + # 1. Single tensor -> wrap in list + # 2. List of tensors -> use as-is + # 3. List of Python lists (from serialization) -> convert to tensors + if isinstance(raw_prior_image_ids, torch.Tensor): + prior_token_image_ids = [raw_prior_image_ids] + elif isinstance(raw_prior_image_ids, list): + # Check if elements are tensors or Python lists + if raw_prior_image_ids and isinstance(raw_prior_image_ids[0], torch.Tensor): + prior_token_image_ids = raw_prior_image_ids + elif raw_prior_image_ids and isinstance(raw_prior_image_ids[0], list): + # Convert Python lists back to tensors + prior_token_image_ids = [torch.tensor(ids, dtype=torch.long) for ids in raw_prior_image_ids] + else: + logger.warning( + f"[ar2diffusion] Request 0: unexpected prior_token_image_ids format: " + f"{type(raw_prior_image_ids[0]) if raw_prior_image_ids else 'empty'}" + ) + else: + # Fallback: also check output (CompletionOutput) in case of different vLLM versions + if hasattr(output, "multimodal_output") and output.multimodal_output: + mm_output = output.multimodal_output + logger.debug("[ar2diffusion] Request 0: found multimodal_output on CompletionOutput (fallback)") if isinstance(mm_output, Mapping): raw_prior_image_ids = mm_output.get("ids", {}).get("prior_image") if raw_prior_image_ids is not None: - # Handle different formats: - # 1. Single tensor -> wrap in list - # 2. List of tensors -> use as-is - # 3. List of Python lists (from serialization) -> convert to tensors if isinstance(raw_prior_image_ids, torch.Tensor): prior_token_image_ids = [raw_prior_image_ids] elif isinstance(raw_prior_image_ids, list): - # Check if elements are tensors or Python lists - if raw_prior_image_ids and isinstance(raw_prior_image_ids[0], torch.Tensor): - prior_token_image_ids = raw_prior_image_ids - elif raw_prior_image_ids and isinstance(raw_prior_image_ids[0], list): - # Convert Python lists back to tensors - prior_token_image_ids = [torch.tensor(ids, dtype=torch.long) for ids in raw_prior_image_ids] - else: - logger.warning( - f"[ar2diffusion] Request {i}: unexpected prior_token_image_ids format: " - f"{type(raw_prior_image_ids[0]) if raw_prior_image_ids else 'empty'}" - ) - else: - # Fallback: also check output (CompletionOutput) in case of different vLLM versions - if hasattr(output, "multimodal_output") and output.multimodal_output: - mm_output = output.multimodal_output - logger.debug(f"[ar2diffusion] Request {i}: found multimodal_output on CompletionOutput (fallback)") - if isinstance(mm_output, Mapping): - raw_prior_image_ids = mm_output.get("ids", {}).get("prior_image") - if raw_prior_image_ids is not None: - if isinstance(raw_prior_image_ids, torch.Tensor): - prior_token_image_ids = [raw_prior_image_ids] - elif isinstance(raw_prior_image_ids, list): - prior_token_image_ids = raw_prior_image_ids - _dt_prior_img = (time.perf_counter() - _t_prior_img) * 1000 - - diffusion_input = { - "prompt": text_prompt, - "height": pixel_h, - "width": pixel_w, - "extra": { - "prior_token_ids": prior_token_ids, - "prior_token_image_ids": prior_token_image_ids, - }, - } - - if requires_multimodal_data: - mm_data = original_prompt.get("multi_modal_data") - if mm_data: - pil_image = _first_source_image(mm_data) - diffusion_input["pil_image"] = pil_image - - for key in ["seed", "num_inference_steps", "guidance_scale", "negative_prompt"]: - if key in original_prompt: - diffusion_input[key] = original_prompt[key] - - _dt_req = (time.perf_counter() - _t_req) * 1000 - logger.info( - "[ar2diffusion] req=%d mode=%s target=%dx%d " - "raw_tokens=%d prior_tokens=%d prior_image_ids=%s " - "timing: mode_detect=%.3fms parse+upsample=%.3fms " - "prior_image_ids_extract=%.3fms req_total=%.3fms", - i, - "i2i" if is_i2i else "t2i", - pixel_h, - pixel_w, - len(generated_token_ids), - len(prior_token_ids), - "yes" if prior_token_image_ids is not None else "no", - _dt_mode, - _dt_parse, - _dt_prior_img, - _dt_req, - ) - diffusion_inputs.append(diffusion_input) + prior_token_image_ids = raw_prior_image_ids + _dt_prior_img = (time.perf_counter() - _t_prior_img) * 1000 + + diffusion_input = { + "prompt": text_prompt, + "height": pixel_h, + "width": pixel_w, + "extra": { + "prior_token_ids": prior_token_ids, + "prior_token_image_ids": prior_token_image_ids, + }, + } + + if requires_multimodal_data: + mm_data = original_prompt.get("multi_modal_data") + if mm_data: + pil_image = _first_source_image(mm_data) + diffusion_input["pil_image"] = pil_image + + for key in ["seed", "num_inference_steps", "guidance_scale", "negative_prompt"]: + if key in original_prompt: + diffusion_input[key] = original_prompt[key] + + _dt_req = (time.perf_counter() - _t_req) * 1000 + logger.info( + "[ar2diffusion] req=%d mode=%s target=%dx%d " + "raw_tokens=%d prior_tokens=%d prior_image_ids=%s " + "timing: mode_detect=%.3fms parse+upsample=%.3fms " + "prior_image_ids_extract=%.3fms req_total=%.3fms", + 0, + "i2i" if is_i2i else "t2i", + pixel_h, + pixel_w, + len(generated_token_ids), + len(prior_token_ids), + "yes" if prior_token_image_ids is not None else "no", + _dt_mode, + _dt_parse, + _dt_prior_img, + _dt_req, + ) _dt_total = (time.perf_counter() - _t_total) * 1000 logger.info( - "[ar2diffusion] batch done: %d reqs, total=%.3fms", - len(diffusion_inputs), + "[ar2diffusion] request done: 1 req, total=%.3fms", _dt_total, ) - return diffusion_inputs + return diffusion_input