Files
wehub-resource-sync eec33d25b2
pre-commit / pre-commit (push) Failing after 1s
Build Wheel / build (3.11) (push) Failing after 1s
Build Wheel / build (3.12) (push) Failing after 0s
chore: import upstream snapshot with attribution
2026-07-13 12:29:08 +08:00

540 lines
20 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""
End-to-end test for diffusion batching via AsyncOmni.
This test fires multiple concurrent ``AsyncOmni.generate()`` calls for a
diffusion model and validates that every caller receives its correct
individual result. When the underlying diffusion stage is configured with
``batch_size > 1`` (via stage config or ``StageDiffusionClient``), the
requests will be batched internally.
Even without explicit batching config this test is useful for verifying
that concurrent async requests are handled correctly.
Usage (standalone):
python tests/diffusion/batching/test_diffusion_batching.py \
--model <model_name_or_path> \
--num-prompts 8
Or via pytest:
pytest tests/diffusion/batching/test_diffusion_batching.py -s
"""
from __future__ import annotations
import argparse
import asyncio
import sys
import time
import uuid
from pathlib import Path
import pytest
import torch
from tests.helpers.mark import hardware_test
from tests.helpers.runtime import OmniRunner
from vllm_omni.entrypoints.async_omni import AsyncOmni
from vllm_omni.inputs.data import OmniDiffusionSamplingParams
from vllm_omni.outputs import OmniRequestOutput
from vllm_omni.platforms import current_omni_platform
# ruff: noqa: E402
REPO_ROOT = Path(__file__).resolve().parents[3]
if str(REPO_ROOT) not in sys.path:
sys.path.insert(0, str(REPO_ROOT))
# ------------------------------------------------------------------
models = ["tiny-random/Qwen-Image"]
# ------------------------------------------------------------------
# Prompt fixtures
# ------------------------------------------------------------------
WARMUP_PROMPTS: list[dict[str, str]] = [
{"prompt": "a sunflower in a glass vase", "negative_prompt": "blurry"},
{"prompt": "a rocket launching into space", "negative_prompt": "low detail"},
{"prompt": "a small cottage in the snowy mountains", "negative_prompt": "foggy"},
{"prompt": "a colorful parrot sitting on a tree branch", "negative_prompt": "low contrast"},
]
TEST_PROMPTS: list[dict[str, str]] = [
{"prompt": "a cup of coffee on a table", "negative_prompt": "low resolution"},
{"prompt": "a toy dinosaur on a sandy beach", "negative_prompt": "cinematic, realistic"},
{"prompt": "a futuristic city skyline at sunset", "negative_prompt": "blurry, foggy"},
{"prompt": "a bowl of fresh strawberries", "negative_prompt": "low detail"},
{"prompt": "a medieval knight standing in the rain", "negative_prompt": "modern clothing"},
{"prompt": "a cat wearing sunglasses lounging in a garden", "negative_prompt": "dark lighting"},
{"prompt": "a spaceship flying above a volcano", "negative_prompt": "low contrast"},
{"prompt": "a watercolor painting of a mountain lake", "negative_prompt": "photo, realistic"},
]
# ------------------------------------------------------------------
# Helpers
# ------------------------------------------------------------------
def _default_sampling_params(**overrides) -> OmniDiffusionSamplingParams:
defaults = dict(
num_inference_steps=2,
width=256,
height=256,
guidance_scale=0.0,
)
defaults.update(overrides)
return OmniDiffusionSamplingParams(**defaults)
def _default_sync_sampling_params(**overrides) -> OmniDiffusionSamplingParams:
"""Create sampling params for the synchronous Omni.generate() API."""
defaults = dict(
num_inference_steps=2,
width=256,
height=256,
guidance_scale=0.0,
generator=torch.Generator(current_omni_platform.device_type).manual_seed(42),
)
defaults.update(overrides)
return OmniDiffusionSamplingParams(**defaults)
async def _collect_generate(omni: AsyncOmni, prompt, request_id, sampling_params_list) -> OmniRequestOutput:
"""Consume the AsyncOmni.generate() async generator and return the last output."""
last_output: OmniRequestOutput | None = None
async for output in omni.generate(
prompt=prompt,
request_id=request_id,
sampling_params_list=sampling_params_list,
):
last_output = output
if last_output is None:
raise RuntimeError(f"No output received for request {request_id}")
return last_output
def _extract_images(output: OmniRequestOutput) -> list:
"""Extract images from an OmniRequestOutput, handling both direct
and nested request_output structures."""
if output.images:
return output.images
# When the output comes from the orchestrator pipeline, images may be
# nested inside request_output.
inner = getattr(output, "request_output", None)
if inner is not None and hasattr(inner, "images") and inner.images:
return inner.images
return []
# ------------------------------------------------------------------
# Warm-up (async)
# ------------------------------------------------------------------
async def warmup(omni: AsyncOmni, prompts: list[dict[str, str]]) -> None:
"""Warm-up: send prompts in parallel to pre-load the model."""
print(f"🔥 Warming up with {len(prompts)} prompts ...")
sp = _default_sampling_params(num_inference_steps=2)
start = time.perf_counter()
tasks = [
_collect_generate(
omni,
prompt=p,
request_id=f"warmup-{i}-{uuid.uuid4().hex[:8]}",
sampling_params_list=[sp],
)
for i, p in enumerate(prompts)
]
await asyncio.gather(*tasks)
elapsed = time.perf_counter() - start
print(f" Warm-up done in {elapsed:.2f}s\n")
# ------------------------------------------------------------------
# Single (sequential) benchmark
# ------------------------------------------------------------------
async def run_single(omni: AsyncOmni, prompts: list[dict[str, str]]) -> float:
"""Run prompts one-by-one sequentially."""
print(f"🧩 Running SINGLE (sequential) mode {len(prompts)} prompts ...")
sp = _default_sampling_params()
total_start = time.perf_counter()
for i, prompt in enumerate(prompts):
start = time.perf_counter()
result = await _collect_generate(
omni,
prompt=prompt,
request_id=f"single-{i}-{uuid.uuid4().hex[:8]}",
sampling_params_list=[sp],
)
elapsed = time.perf_counter() - start
images = _extract_images(result)
print(f" prompt {i}: {elapsed:.2f}s ({len(images)} images)")
total = time.perf_counter() - total_start
print(f" ✅ Total single-mode: {total:.2f}s\n")
return total
# ------------------------------------------------------------------
# Batch (parallel) benchmark — concurrent individual requests
# ------------------------------------------------------------------
async def run_batch(
omni: AsyncOmni,
prompts: list[dict[str, str]],
label: str = "batch",
) -> float:
"""Send all prompts concurrently via asyncio.gather (one request per prompt)."""
print(f"⚙️ Running {label.upper()} mode {len(prompts)} prompts concurrently ...")
sp = _default_sampling_params()
start = time.perf_counter()
tasks = [
_collect_generate(
omni,
prompt=p,
request_id=f"{label}-{i}-{uuid.uuid4().hex[:8]}",
sampling_params_list=[sp],
)
for i, p in enumerate(prompts)
]
results = await asyncio.gather(*tasks)
elapsed = time.perf_counter() - start
for i, result in enumerate(results):
images = _extract_images(result)
print(f" prompt {i}: {len(images)} images, request_id={result.request_id}")
print(f" ✅ Total {label} mode: {elapsed:.2f}s\n")
return elapsed
# ------------------------------------------------------------------
# Async validation helpers
# ------------------------------------------------------------------
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."""
print(f"🔍 Validating concurrent correctness with {len(prompts)} prompts ...")
sp = _default_sampling_params()
request_ids = [f"validate-{i}-{uuid.uuid4().hex[:8]}" for i in range(len(prompts))]
tasks = [
_collect_generate(omni, prompt=p, request_id=rid, sampling_params_list=[sp])
for p, rid in zip(prompts, request_ids)
]
results = await asyncio.gather(*tasks)
assert len(results) == len(prompts), f"Expected {len(prompts)} results, got {len(results)}"
returned_ids = [r.request_id for r in results]
for rid in request_ids:
assert rid in returned_ids, f"Missing request_id {rid} in results"
print(" ✅ All request_ids matched, results count correct.\n")
# ------------------------------------------------------------------
# Single vs Parallel comparison (CLI only)
# ------------------------------------------------------------------
async def compare_single_vs_parallel(
model: str,
prompts: list[dict[str, str]],
batch_size: int = 1,
) -> None:
"""Run the same prompts sequentially then in parallel and print a comparison."""
omni = AsyncOmni(model=model, diffusion_batch_size=batch_size)
try:
await warmup(omni, WARMUP_PROMPTS)
single_time = await run_single(omni, prompts)
parallel_time = await run_batch(omni, prompts, label="parallel")
finally:
omni.shutdown()
speedup_parallel = single_time / parallel_time if parallel_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("=" * 60)
# ------------------------------------------------------------------
# CLI main entrypoint
# ------------------------------------------------------------------
async def main(model: str, num_prompts: int, mode: str, batch_size: int = 1) -> None:
prompts = (TEST_PROMPTS * ((num_prompts // len(TEST_PROMPTS)) + 1))[:num_prompts]
if mode == "compare":
await compare_single_vs_parallel(model, prompts, batch_size=batch_size)
return
omni = AsyncOmni(model=model, diffusion_batch_size=batch_size)
try:
await warmup(omni, WARMUP_PROMPTS)
if mode == "validate":
await validate_concurrent(omni, prompts)
elif mode == "batch":
await run_batch(omni, prompts, label="measurement")
elif mode == "single":
await run_single(omni, prompts)
else:
raise ValueError(f"Unknown mode: {mode}")
finally:
omni.shutdown()
# ==================================================================
# pytest test cases
# ==================================================================
@pytest.mark.core_model
@pytest.mark.diffusion
@hardware_test(res={"cuda": "L4", "rocm": "MI325", "xpu": "B60"})
@pytest.mark.parametrize("model_name", models)
def test_diffusion_batching_sync_sequential(model_name: str):
"""Test that synchronous Omni can generate images for multiple prompts
submitted sequentially (one at a time) and each returns a valid image."""
try:
with OmniRunner(model_name) as runner:
m = runner.omni
sp = _default_sync_sampling_params()
prompts = TEST_PROMPTS[:4]
for i, prompt in enumerate(prompts):
outputs = m.generate(prompt, sp)
first_output = outputs[0]
assert first_output.final_output_type == "image", (
f"Expected 'image', got '{first_output.final_output_type}'"
)
# Images are surfaced both at top-level and inside request_output
images = _extract_images(first_output)
assert len(images) >= 1, f"Expected at least 1 image for prompt {i}, got {len(images)}"
assert images[0].width == 256
assert images[0].height == 256
print(f" prompt {i}: OK ({len(images)} images)")
except Exception as e:
print(f"Test failed with error: {e}")
raise
@pytest.mark.core_model
@pytest.mark.diffusion
@hardware_test(res={"cuda": "L4", "rocm": "MI325", "xpu": "B60"})
@pytest.mark.parametrize("model_name", models)
def test_diffusion_batching_sync_multi_prompt(model_name: str):
"""Test that synchronous Omni correctly handles a list of multiple
prompts submitted at once and returns one result per prompt.
Note: Omni.generate() iterates the list and submits each prompt
individually with its own request_id. This tests concurrent request
handling at the diffusion stage, not the explicit list-batch path
(which is only available via AsyncOmni).
"""
try:
with OmniRunner(model_name) as runner:
m = runner.omni
sp = _default_sync_sampling_params()
prompts = TEST_PROMPTS[:4]
outputs = m.generate(prompts, sp)
assert len(outputs) == len(prompts), f"Expected {len(prompts)} outputs, got {len(outputs)}"
for i, output in enumerate(outputs):
assert output.final_output_type == "image", (
f"Output {i} final_output_type expected 'image', got '{output.final_output_type}'"
)
images = _extract_images(output)
assert images and len(images) >= 1, f"Expected at least 1 image for prompt {i}"
assert images[0].width == 256
assert images[0].height == 256
print(f" prompt {i}: OK ({len(images)} images, request_id={output.request_id})")
# Verify all request_ids are distinct
request_ids = [o.request_id for o in outputs]
assert len(set(request_ids)) == len(request_ids), f"Duplicate request_ids found: {request_ids}"
except Exception as e:
print(f"Test failed with error: {e}")
raise
@pytest.mark.core_model
@pytest.mark.diffusion
@hardware_test(res={"cuda": "L4", "rocm": "MI325", "xpu": "B60"})
@pytest.mark.parametrize("model_name", models)
def test_diffusion_batching_async_concurrent(model_name: str):
"""Test that AsyncOmni correctly handles multiple concurrent requests
fired via asyncio.gather. Each request_id must appear in the results."""
async def _inner():
omni = AsyncOmni(model=model_name, diffusion_batch_size=1)
try:
prompts = TEST_PROMPTS[:4]
sp = _default_sampling_params()
request_ids = [f"async-concurrent-{i}-{uuid.uuid4().hex[:8]}" for i in range(len(prompts))]
tasks = [
_collect_generate(omni, prompt=p, request_id=rid, sampling_params_list=[sp])
for p, rid in zip(prompts, request_ids)
]
results = await asyncio.gather(*tasks)
assert len(results) == len(prompts), f"Expected {len(prompts)} results, got {len(results)}"
returned_ids = [r.request_id for r in results]
for rid in request_ids:
assert rid in returned_ids, f"Missing request_id {rid} in results"
for i, result in enumerate(results):
images = _extract_images(result)
assert len(images) >= 1, f"No images for prompt {i}"
assert images[0].width == 256
assert images[0].height == 256
print(f" prompt {i}: OK ({len(images)} images, request_id={result.request_id})")
finally:
omni.shutdown()
asyncio.run(_inner())
@pytest.mark.core_model
@pytest.mark.diffusion
@hardware_test(res={"cuda": "L4", "rocm": "MI325", "xpu": "B60"})
@pytest.mark.parametrize("model_name", models)
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():
omni = AsyncOmni(model=model_name, diffusion_batch_size=4)
try:
prompts = TEST_PROMPTS[:4]
sp = _default_sampling_params()
request_id = f"explicit-batch-{uuid.uuid4().hex[:8]}"
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()
asyncio.run(_inner())
@pytest.mark.core_model
@pytest.mark.diffusion
@hardware_test(res={"cuda": "L4", "rocm": "MI325", "xpu": "B60"})
@pytest.mark.parametrize("model_name", models)
def test_diffusion_batching_num_outputs(model_name: str):
"""Test that the diffusion model respects num_outputs_per_prompt and
generates the correct number of images per request."""
try:
with OmniRunner(model_name) as runner:
m = runner.omni
num_outputs = 2
sp = _default_sync_sampling_params(num_outputs_per_prompt=num_outputs)
outputs = m.generate(
"a photo of a cat sitting on a laptop keyboard",
sp,
)
first_output = outputs[0]
assert first_output.final_output_type == "image"
images = _extract_images(first_output)
assert images is not None and len(images) == num_outputs, (
f"Expected {num_outputs} images, got {len(images) if images else 0}"
)
for img in images:
assert img.width == 256
assert img.height == 256
except Exception as e:
print(f"Test failed with error: {e}")
raise
@pytest.mark.core_model
@pytest.mark.diffusion
@hardware_test(res={"cuda": "L4", "rocm": "MI325", "xpu": "B60"})
@pytest.mark.parametrize("model_name", models)
def test_diffusion_batching_distinct_results(model_name: str):
"""Test that different prompts produce distinct images when batched,
ensuring the batching logic does not mix up results across requests."""
try:
with OmniRunner(model_name) as runner:
m = runner.omni
sp = _default_sync_sampling_params()
prompts = [
{"prompt": "a bright red apple on a white table", "negative_prompt": "blurry"},
{"prompt": "a blue ocean with white waves crashing", "negative_prompt": "blurry"},
]
outputs = m.generate(prompts, sp)
assert len(outputs) == len(prompts), f"Expected {len(prompts)} outputs, got {len(outputs)}"
# Verify each output has a unique request_id
request_ids = [o.request_id for o in outputs]
assert len(set(request_ids)) == len(request_ids), f"Duplicate request_ids: {request_ids}"
# Verify each output has images
for i, output in enumerate(outputs):
images = _extract_images(output)
assert images and len(images) >= 1, f"No images for prompt {i}"
assert images[0].width == 256
assert images[0].height == 256
except Exception as e:
print(f"Test failed with error: {e}")
raise
# ------------------------------------------------------------------
# CLI
# ------------------------------------------------------------------
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="E2E diffusion concurrent benchmark / validation")
parser.add_argument("--model", type=str, required=True, help="Model name or path")
parser.add_argument("--num-prompts", type=int, default=8, help="Number of prompts to run")
parser.add_argument("--batch-size", type=int, default=1, help="Diffusion batch size (1 = no batching)")
parser.add_argument(
"--mode",
choices=["batch", "single", "compare", "validate"],
default="compare",
help=(
"Run mode: 'batch' (parallel gather), 'single' (sequential), "
"'compare' (single vs parallel), 'validate' (concurrent correctness)"
),
)
args = parser.parse_args()
asyncio.run(main(args.model, args.num_prompts, args.mode, batch_size=args.batch_size))