项目文件夹

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

540 行
20 KiB
Python

# 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))