[diffusion] profiling: add bench_serving.py and VBench (#15410)
This commit is contained in:
@@ -0,0 +1,744 @@
|
|||||||
|
"""
|
||||||
|
Benchmark online serving for diffusion models (Image/Video Generation).
|
||||||
|
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
|
||||||
|
t2v:
|
||||||
|
python3 -m sglang.multimodal_gen.benchmarks.bench_serving \
|
||||||
|
--backend sglang-image --dataset vbench --task t2v --num-prompts 20
|
||||||
|
|
||||||
|
i2v:
|
||||||
|
python3 -m sglang.multimodal_gen.benchmarks.bench_serving \
|
||||||
|
--backend sglang-image --dataset vbench --task i2v --num-prompts 20
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
import asyncio
|
||||||
|
import glob
|
||||||
|
import json
|
||||||
|
import os
|
||||||
|
import time
|
||||||
|
import uuid
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from dataclasses import dataclass, field
|
||||||
|
from typing import Any, Dict, List, Optional
|
||||||
|
|
||||||
|
import aiohttp
|
||||||
|
import numpy as np
|
||||||
|
import requests
|
||||||
|
from tqdm.asyncio import tqdm
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class RequestFuncInput:
|
||||||
|
prompt: str
|
||||||
|
api_url: str
|
||||||
|
model: str
|
||||||
|
width: Optional[int] = None
|
||||||
|
height: Optional[int] = None
|
||||||
|
num_frames: Optional[int] = None
|
||||||
|
fps: Optional[int] = None
|
||||||
|
extra_body: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
image_paths: Optional[List[str]] = None
|
||||||
|
request_id: str = field(default_factory=lambda: str(uuid.uuid4()))
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class RequestFuncOutput:
|
||||||
|
success: bool = False
|
||||||
|
latency: float = 0.0
|
||||||
|
error: str = ""
|
||||||
|
start_time: float = 0.0
|
||||||
|
response_body: Dict[str, Any] = field(default_factory=dict)
|
||||||
|
|
||||||
|
|
||||||
|
class BaseDataset(ABC):
|
||||||
|
def __init__(self, args, api_url: str, model: str):
|
||||||
|
self.args = args
|
||||||
|
self.api_url = api_url
|
||||||
|
self.model = model
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def __len__(self) -> int:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def __getitem__(self, idx: int) -> RequestFuncInput:
|
||||||
|
pass
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_requests(self) -> List[RequestFuncInput]:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
class VBenchDataset(BaseDataset):
|
||||||
|
"""
|
||||||
|
Dataset loader for VBench prompts.
|
||||||
|
Supports t2v, i2v.
|
||||||
|
"""
|
||||||
|
|
||||||
|
T2V_PROMPT_URL = "https://raw.githubusercontent.com/Vchitect/VBench/master/prompts/prompts_per_dimension/subject_consistency.txt"
|
||||||
|
I2V_DOWNLOAD_SCRIPT_URL = "https://raw.githubusercontent.com/Vchitect/VBench/master/vbench2_beta_i2v/download_data.sh"
|
||||||
|
|
||||||
|
def __init__(self, args, api_url: str, model: str):
|
||||||
|
super().__init__(args, api_url, model)
|
||||||
|
self.cache_dir = os.path.join(os.path.expanduser("~"), ".cache", "sglang")
|
||||||
|
self.items = self._load_data()
|
||||||
|
|
||||||
|
def _load_data(self) -> List[Dict[str, Any]]:
|
||||||
|
if self.args.task == "t2v":
|
||||||
|
return self._load_t2v_prompts()
|
||||||
|
elif self.args.task in ["i2v", "ti2v", "ti2i"]:
|
||||||
|
return self._load_i2v_data()
|
||||||
|
else:
|
||||||
|
return self._load_t2v_prompts()
|
||||||
|
|
||||||
|
def _download_file(self, url: str, dest_path: str) -> None:
|
||||||
|
"""Download a file from URL to destination path."""
|
||||||
|
os.makedirs(os.path.dirname(dest_path), exist_ok=True)
|
||||||
|
resp = requests.get(url)
|
||||||
|
resp.raise_for_status()
|
||||||
|
with open(dest_path, "w") as f:
|
||||||
|
f.write(resp.text)
|
||||||
|
|
||||||
|
def _load_t2v_prompts(self) -> List[Dict[str, Any]]:
|
||||||
|
path = self.args.dataset_path
|
||||||
|
|
||||||
|
if not path:
|
||||||
|
path = os.path.join(self.cache_dir, "vbench_subject_consistency.txt")
|
||||||
|
if not os.path.exists(path):
|
||||||
|
print(f"Downloading VBench T2V prompts to {path}...")
|
||||||
|
try:
|
||||||
|
self._download_file(self.T2V_PROMPT_URL, path)
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Failed to download VBench prompts: {e}")
|
||||||
|
return [{"prompt": "A cat sitting on a bench"}] * 50
|
||||||
|
|
||||||
|
prompts = []
|
||||||
|
with open(path, "r") as f:
|
||||||
|
for line in f:
|
||||||
|
line = line.strip()
|
||||||
|
if line:
|
||||||
|
prompts.append({"prompt": line})
|
||||||
|
|
||||||
|
return self._resize_data(prompts)
|
||||||
|
|
||||||
|
def _auto_download_i2v_dataset(self) -> str:
|
||||||
|
"""Auto-download VBench I2V dataset and return the dataset directory."""
|
||||||
|
vbench_i2v_dir = os.path.join(self.cache_dir, "vbench_i2v", "vbench2_beta_i2v")
|
||||||
|
info_json_path = os.path.join(vbench_i2v_dir, "data", "i2v-bench-info.json")
|
||||||
|
|
||||||
|
if os.path.exists(info_json_path):
|
||||||
|
return vbench_i2v_dir
|
||||||
|
|
||||||
|
print(f"Downloading VBench I2V dataset to {vbench_i2v_dir}...")
|
||||||
|
try:
|
||||||
|
cache_root = os.path.join(self.cache_dir, "vbench_i2v")
|
||||||
|
script_path = os.path.join(cache_root, "download_data.sh")
|
||||||
|
|
||||||
|
self._download_file(self.I2V_DOWNLOAD_SCRIPT_URL, script_path)
|
||||||
|
os.chmod(script_path, 0o755)
|
||||||
|
|
||||||
|
print("Executing download_data.sh (this may take a while)...")
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
result = subprocess.run(
|
||||||
|
["bash", script_path],
|
||||||
|
cwd=cache_root,
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
if result.returncode != 0:
|
||||||
|
raise RuntimeError(f"Download script failed: {result.stderr}")
|
||||||
|
|
||||||
|
print(f"Successfully downloaded VBench I2V dataset to {vbench_i2v_dir}")
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Failed to download VBench I2V dataset: {e}")
|
||||||
|
print("Please manually download following instructions at:")
|
||||||
|
print(
|
||||||
|
"https://github.com/Vchitect/VBench/tree/master/vbench2_beta_i2v#22-download"
|
||||||
|
)
|
||||||
|
return None
|
||||||
|
|
||||||
|
return vbench_i2v_dir if os.path.exists(info_json_path) else None
|
||||||
|
|
||||||
|
def _load_from_i2v_json(self, json_path: str) -> List[Dict[str, Any]]:
|
||||||
|
"""Load I2V data from i2v-bench-info.json format."""
|
||||||
|
with open(json_path, "r") as f:
|
||||||
|
items = json.load(f)
|
||||||
|
|
||||||
|
base_dir = os.path.dirname(
|
||||||
|
os.path.dirname(json_path)
|
||||||
|
) # Go up to vbench2_beta_i2v
|
||||||
|
origin_dir = os.path.join(base_dir, "data", "origin")
|
||||||
|
|
||||||
|
data = []
|
||||||
|
for item in items:
|
||||||
|
img_path = os.path.join(origin_dir, item.get("file_name", ""))
|
||||||
|
if os.path.exists(img_path):
|
||||||
|
data.append({"prompt": item.get("caption", ""), "image_path": img_path})
|
||||||
|
else:
|
||||||
|
print(f"Warning: Image not found: {img_path}")
|
||||||
|
|
||||||
|
print(f"Loaded {len(data)} I2V samples from VBench I2V dataset")
|
||||||
|
return data
|
||||||
|
|
||||||
|
def _scan_directory_for_images(self, path: str) -> List[Dict[str, Any]]:
|
||||||
|
"""Scan directory for image files."""
|
||||||
|
exts = ["*.jpg", "*.jpeg", "*.png", "*.webp"]
|
||||||
|
files = []
|
||||||
|
|
||||||
|
for ext in exts:
|
||||||
|
files.extend(glob.glob(os.path.join(path, ext)))
|
||||||
|
files.extend(glob.glob(os.path.join(path, ext.upper())))
|
||||||
|
|
||||||
|
# Also check in data/origin subdirectory
|
||||||
|
origin_dir = os.path.join(path, "data", "origin")
|
||||||
|
if os.path.exists(origin_dir):
|
||||||
|
files.extend(glob.glob(os.path.join(origin_dir, ext)))
|
||||||
|
files.extend(glob.glob(os.path.join(origin_dir, ext.upper())))
|
||||||
|
|
||||||
|
return [
|
||||||
|
{"prompt": os.path.splitext(os.path.basename(f))[0], "image_path": f}
|
||||||
|
for f in files
|
||||||
|
]
|
||||||
|
|
||||||
|
def _create_dummy_data(self) -> List[Dict[str, Any]]:
|
||||||
|
"""Create dummy data with a placeholder image in cache directory."""
|
||||||
|
print("No I2V data found. Using dummy placeholders.")
|
||||||
|
|
||||||
|
dummy_image = os.path.join(self.cache_dir, "dummy_image.jpg")
|
||||||
|
if not os.path.exists(dummy_image):
|
||||||
|
try:
|
||||||
|
from PIL import Image
|
||||||
|
|
||||||
|
os.makedirs(self.cache_dir, exist_ok=True)
|
||||||
|
img = Image.new("RGB", (100, 100), color="red")
|
||||||
|
img.save(dummy_image)
|
||||||
|
print(f"Created dummy image at {dummy_image}")
|
||||||
|
except ImportError:
|
||||||
|
print("PIL not installed, cannot create dummy image.")
|
||||||
|
return []
|
||||||
|
|
||||||
|
return [{"prompt": "A moving cat", "image_path": dummy_image}] * 10
|
||||||
|
|
||||||
|
def _load_i2v_data(self) -> List[Dict[str, Any]]:
|
||||||
|
"""Load I2V data from VBench I2V dataset or user-provided path."""
|
||||||
|
path = self.args.dataset_path
|
||||||
|
|
||||||
|
# Auto-download if no path provided
|
||||||
|
if not path:
|
||||||
|
path = self._auto_download_i2v_dataset()
|
||||||
|
if not path:
|
||||||
|
return self._resize_data(self._create_dummy_data())
|
||||||
|
|
||||||
|
# Try to load from i2v-bench-info.json
|
||||||
|
info_json_candidates = [
|
||||||
|
os.path.join(path, "data", "i2v-bench-info.json"),
|
||||||
|
path if path.endswith(".json") else None,
|
||||||
|
]
|
||||||
|
|
||||||
|
for json_path in info_json_candidates:
|
||||||
|
if json_path and os.path.exists(json_path):
|
||||||
|
try:
|
||||||
|
return self._resize_data(self._load_from_i2v_json(json_path))
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Failed to load {json_path}: {e}")
|
||||||
|
|
||||||
|
# Fallback: scan directory for images
|
||||||
|
if os.path.isdir(path):
|
||||||
|
data = self._scan_directory_for_images(path)
|
||||||
|
if data:
|
||||||
|
return self._resize_data(data)
|
||||||
|
|
||||||
|
# Last resort: dummy data
|
||||||
|
return self._resize_data(self._create_dummy_data())
|
||||||
|
|
||||||
|
def _resize_data(self, data: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
|
||||||
|
"""Resize data to match num_prompts."""
|
||||||
|
if not self.args.num_prompts:
|
||||||
|
return data
|
||||||
|
|
||||||
|
if len(data) < self.args.num_prompts:
|
||||||
|
factor = (self.args.num_prompts // len(data)) + 1
|
||||||
|
data = data * factor
|
||||||
|
|
||||||
|
return data[: self.args.num_prompts]
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return len(self.items)
|
||||||
|
|
||||||
|
def __getitem__(self, idx: int) -> RequestFuncInput:
|
||||||
|
item = self.items[idx]
|
||||||
|
image_paths = [item["image_path"]] if "image_path" in item else None
|
||||||
|
|
||||||
|
return RequestFuncInput(
|
||||||
|
prompt=item.get("prompt", ""),
|
||||||
|
api_url=self.api_url,
|
||||||
|
model=self.model,
|
||||||
|
width=self.args.width,
|
||||||
|
height=self.args.height,
|
||||||
|
num_frames=self.args.num_frames,
|
||||||
|
fps=self.args.fps,
|
||||||
|
image_paths=image_paths,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_requests(self) -> List[RequestFuncInput]:
|
||||||
|
return [self[i] for i in range(len(self))]
|
||||||
|
|
||||||
|
|
||||||
|
class RandomDataset(BaseDataset):
|
||||||
|
def __init__(self, args, api_url: str, model: str):
|
||||||
|
self.args = args
|
||||||
|
self.api_url = api_url
|
||||||
|
self.model = model
|
||||||
|
self.num_prompts = args.num_prompts or 100
|
||||||
|
|
||||||
|
def __len__(self) -> int:
|
||||||
|
return self.num_prompts
|
||||||
|
|
||||||
|
def __getitem__(self, idx: int) -> RequestFuncInput:
|
||||||
|
return RequestFuncInput(
|
||||||
|
prompt=f"Random prompt {idx} for benchmarking diffusion models",
|
||||||
|
api_url=self.api_url,
|
||||||
|
model=self.model,
|
||||||
|
width=self.args.width,
|
||||||
|
height=self.args.height,
|
||||||
|
num_frames=self.args.num_frames,
|
||||||
|
fps=self.args.fps,
|
||||||
|
)
|
||||||
|
|
||||||
|
def get_requests(self) -> List[RequestFuncInput]:
|
||||||
|
return [self[i] for i in range(len(self))]
|
||||||
|
|
||||||
|
|
||||||
|
async def async_request_image_sglang(
|
||||||
|
input: RequestFuncInput,
|
||||||
|
session: aiohttp.ClientSession,
|
||||||
|
pbar: Optional[tqdm] = None,
|
||||||
|
) -> RequestFuncOutput:
|
||||||
|
output = RequestFuncOutput()
|
||||||
|
output.start_time = time.perf_counter()
|
||||||
|
|
||||||
|
# Check if we need to use multipart (for image edits with input images)
|
||||||
|
if input.image_paths and len(input.image_paths) > 0:
|
||||||
|
# Use multipart/form-data for image edits
|
||||||
|
data = aiohttp.FormData()
|
||||||
|
data.add_field("model", input.model)
|
||||||
|
data.add_field("prompt", input.prompt)
|
||||||
|
data.add_field("response_format", "b64_json")
|
||||||
|
|
||||||
|
if input.width and input.height:
|
||||||
|
data.add_field("size", f"{input.width}x{input.height}")
|
||||||
|
|
||||||
|
# Merge extra parameters
|
||||||
|
for key, value in input.extra_body.items():
|
||||||
|
data.add_field(key, str(value))
|
||||||
|
|
||||||
|
# Add image file(s)
|
||||||
|
for idx, img_path in enumerate(input.image_paths):
|
||||||
|
if os.path.exists(img_path):
|
||||||
|
data.add_field(
|
||||||
|
"image",
|
||||||
|
open(img_path, "rb"),
|
||||||
|
filename=os.path.basename(img_path),
|
||||||
|
content_type="application/octet-stream",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
output.error = f"Image file not found: {img_path}"
|
||||||
|
output.success = False
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with session.post(input.api_url, data=data) as response:
|
||||||
|
if response.status == 200:
|
||||||
|
resp_json = await response.json()
|
||||||
|
output.response_body = resp_json
|
||||||
|
output.success = True
|
||||||
|
else:
|
||||||
|
output.error = f"HTTP {response.status}: {await response.text()}"
|
||||||
|
output.success = False
|
||||||
|
except Exception as e:
|
||||||
|
output.error = str(e)
|
||||||
|
output.success = False
|
||||||
|
else:
|
||||||
|
# Use JSON for text-to-image generation
|
||||||
|
payload = {
|
||||||
|
"model": input.model,
|
||||||
|
"prompt": input.prompt,
|
||||||
|
"n": 1,
|
||||||
|
"response_format": "b64_json",
|
||||||
|
}
|
||||||
|
|
||||||
|
if input.width and input.height:
|
||||||
|
payload["size"] = f"{input.width}x{input.height}"
|
||||||
|
|
||||||
|
# Merge extra parameters
|
||||||
|
payload.update(input.extra_body)
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with session.post(input.api_url, json=payload) as response:
|
||||||
|
if response.status == 200:
|
||||||
|
resp_json = await response.json()
|
||||||
|
output.response_body = resp_json
|
||||||
|
output.success = True
|
||||||
|
else:
|
||||||
|
output.error = f"HTTP {response.status}: {await response.text()}"
|
||||||
|
output.success = False
|
||||||
|
except Exception as e:
|
||||||
|
output.error = str(e)
|
||||||
|
output.success = False
|
||||||
|
|
||||||
|
output.latency = time.perf_counter() - output.start_time
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
async def async_request_video_sglang(
|
||||||
|
input: RequestFuncInput,
|
||||||
|
session: aiohttp.ClientSession,
|
||||||
|
pbar: Optional[tqdm] = None,
|
||||||
|
) -> RequestFuncOutput:
|
||||||
|
output = RequestFuncOutput()
|
||||||
|
output.start_time = time.perf_counter()
|
||||||
|
|
||||||
|
# 1. Submit Job
|
||||||
|
job_id = None
|
||||||
|
|
||||||
|
# Check if we need to upload images (Multipart) or just send JSON
|
||||||
|
if input.image_paths and len(input.image_paths) > 0:
|
||||||
|
# Use multipart/form-data
|
||||||
|
data = aiohttp.FormData()
|
||||||
|
data.add_field("model", input.model)
|
||||||
|
data.add_field("prompt", input.prompt)
|
||||||
|
|
||||||
|
if input.width and input.height:
|
||||||
|
data.add_field("size", f"{input.width}x{input.height}")
|
||||||
|
|
||||||
|
# Add extra body fields to form data if possible, or assume simple key-values
|
||||||
|
# Note: Nested dicts in extra_body might need JSON serialization if API expects it stringified
|
||||||
|
if input.extra_body:
|
||||||
|
data.add_field("extra_body", json.dumps(input.extra_body))
|
||||||
|
|
||||||
|
# Explicitly add fps/num_frames if they are not in extra_body (bench_serving logic overrides)
|
||||||
|
if input.num_frames:
|
||||||
|
data.add_field("num_frames", str(input.num_frames))
|
||||||
|
if input.fps:
|
||||||
|
data.add_field("fps", str(input.fps))
|
||||||
|
|
||||||
|
# Add image file
|
||||||
|
# Currently only support single image upload as 'input_reference' per API spec
|
||||||
|
img_path = input.image_paths[0]
|
||||||
|
if os.path.exists(img_path):
|
||||||
|
data.add_field(
|
||||||
|
"input_reference",
|
||||||
|
open(img_path, "rb"),
|
||||||
|
filename=os.path.basename(img_path),
|
||||||
|
content_type="application/octet-stream",
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
output.error = f"Image file not found: {img_path}"
|
||||||
|
output.success = False
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with session.post(input.api_url, data=data) as response:
|
||||||
|
if response.status == 200:
|
||||||
|
resp_json = await response.json()
|
||||||
|
job_id = resp_json.get("id")
|
||||||
|
else:
|
||||||
|
output.error = (
|
||||||
|
f"Submit failed HTTP {response.status}: {await response.text()}"
|
||||||
|
)
|
||||||
|
output.success = False
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
except Exception as e:
|
||||||
|
output.error = f"Submit exception: {str(e)}"
|
||||||
|
output.success = False
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
|
||||||
|
else:
|
||||||
|
# Use JSON
|
||||||
|
payload = {
|
||||||
|
"model": input.model,
|
||||||
|
"prompt": input.prompt,
|
||||||
|
}
|
||||||
|
if input.width and input.height:
|
||||||
|
payload["size"] = f"{input.width}x{input.height}"
|
||||||
|
if input.num_frames:
|
||||||
|
payload["num_frames"] = input.num_frames
|
||||||
|
if input.fps:
|
||||||
|
payload["fps"] = input.fps
|
||||||
|
|
||||||
|
payload.update(input.extra_body)
|
||||||
|
|
||||||
|
try:
|
||||||
|
async with session.post(input.api_url, json=payload) as response:
|
||||||
|
if response.status == 200:
|
||||||
|
resp_json = await response.json()
|
||||||
|
job_id = resp_json.get("id")
|
||||||
|
else:
|
||||||
|
output.error = (
|
||||||
|
f"Submit failed HTTP {response.status}: {await response.text()}"
|
||||||
|
)
|
||||||
|
output.success = False
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
except Exception as e:
|
||||||
|
output.error = f"Submit exception: {str(e)}"
|
||||||
|
output.success = False
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
|
||||||
|
if not job_id:
|
||||||
|
output.error = "No job_id returned"
|
||||||
|
output.success = False
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
|
||||||
|
# 2. Poll for completion
|
||||||
|
# Assuming the API returns a 'status' field.
|
||||||
|
# We construct the check URL. Assuming api_url is like .../v1/videos
|
||||||
|
# The check url should be .../v1/videos/{id}
|
||||||
|
check_url = f"{input.api_url}/{job_id}"
|
||||||
|
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
async with session.get(check_url) as response:
|
||||||
|
if response.status == 200:
|
||||||
|
status_data = await response.json()
|
||||||
|
status = status_data.get("status")
|
||||||
|
if status == "completed":
|
||||||
|
output.success = True
|
||||||
|
output.response_body = status_data
|
||||||
|
break
|
||||||
|
elif status == "failed":
|
||||||
|
output.success = False
|
||||||
|
output.error = f"Job failed: {status_data.get('error')}"
|
||||||
|
break
|
||||||
|
else:
|
||||||
|
# queued or processing
|
||||||
|
await asyncio.sleep(1.0)
|
||||||
|
else:
|
||||||
|
output.success = False
|
||||||
|
output.error = (
|
||||||
|
f"Poll failed HTTP {response.status}: {await response.text()}"
|
||||||
|
)
|
||||||
|
break
|
||||||
|
except Exception as e:
|
||||||
|
output.success = False
|
||||||
|
output.error = f"Poll exception: {str(e)}"
|
||||||
|
break
|
||||||
|
|
||||||
|
output.latency = time.perf_counter() - output.start_time
|
||||||
|
if pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def calculate_metrics(outputs: List[RequestFuncOutput], total_duration: float):
|
||||||
|
success_outputs = [o for o in outputs if o.success]
|
||||||
|
error_outputs = [o for o in outputs if not o.success]
|
||||||
|
|
||||||
|
num_success = len(success_outputs)
|
||||||
|
latencies = [o.latency for o in success_outputs]
|
||||||
|
|
||||||
|
metrics = {
|
||||||
|
"duration": total_duration,
|
||||||
|
"completed_requests": num_success,
|
||||||
|
"failed_requests": len(error_outputs),
|
||||||
|
"throughput_qps": num_success / total_duration if total_duration > 0 else 0,
|
||||||
|
"latency_mean": np.mean(latencies) if latencies else 0,
|
||||||
|
"latency_median": np.median(latencies) if latencies else 0,
|
||||||
|
"latency_p99": np.percentile(latencies, 99) if latencies else 0,
|
||||||
|
"latency_p50": np.percentile(latencies, 50) if latencies else 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
return metrics
|
||||||
|
|
||||||
|
|
||||||
|
def wait_for_service(base_url: str, timeout: int = 120) -> None:
|
||||||
|
print(f"Waiting for service at {base_url}...")
|
||||||
|
start_time = time.time()
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
# Try /health endpoint first
|
||||||
|
resp = requests.get(f"{base_url}/health", timeout=1)
|
||||||
|
if resp.status_code == 200:
|
||||||
|
print("Service is ready.")
|
||||||
|
break
|
||||||
|
except requests.exceptions.RequestException:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if time.time() - start_time > timeout:
|
||||||
|
raise TimeoutError(
|
||||||
|
f"Service at {base_url} did not start within {timeout} seconds."
|
||||||
|
)
|
||||||
|
|
||||||
|
time.sleep(1)
|
||||||
|
|
||||||
|
|
||||||
|
async def benchmark(args):
|
||||||
|
# Construct base_url if not provided
|
||||||
|
if args.base_url is None:
|
||||||
|
args.base_url = f"http://{args.host}:{args.port}"
|
||||||
|
|
||||||
|
# Wait for service
|
||||||
|
wait_for_service(args.base_url)
|
||||||
|
|
||||||
|
# Setup dataset
|
||||||
|
if args.backend == "sglang-image":
|
||||||
|
if args.task == "i2v":
|
||||||
|
api_url = f"{args.base_url}/v1/images/edits"
|
||||||
|
else:
|
||||||
|
api_url = f"{args.base_url}/v1/images/generations"
|
||||||
|
request_func = async_request_image_sglang
|
||||||
|
elif args.backend == "sglang-video":
|
||||||
|
api_url = f"{args.base_url}/v1/videos"
|
||||||
|
request_func = async_request_video_sglang
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown backend: {args.backend}")
|
||||||
|
|
||||||
|
if args.dataset == "vbench":
|
||||||
|
dataset = VBenchDataset(args, api_url, args.model)
|
||||||
|
elif args.dataset == "random":
|
||||||
|
dataset = RandomDataset(args, api_url, args.model)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown dataset: {args.dataset}")
|
||||||
|
|
||||||
|
requests_list = dataset.get_requests()
|
||||||
|
print(f"Prepared {len(requests_list)} requests from {args.dataset} dataset.")
|
||||||
|
|
||||||
|
# Limit concurrency
|
||||||
|
semaphore = asyncio.Semaphore(args.max_concurrency)
|
||||||
|
|
||||||
|
async def limited_request_func(req, session, pbar):
|
||||||
|
async with semaphore:
|
||||||
|
return await request_func(req, session, pbar)
|
||||||
|
|
||||||
|
# Run benchmark
|
||||||
|
pbar = tqdm(total=len(requests_list), disable=args.disable_tqdm)
|
||||||
|
|
||||||
|
async with aiohttp.ClientSession() as session:
|
||||||
|
start_time = time.perf_counter()
|
||||||
|
tasks = []
|
||||||
|
for req in requests_list:
|
||||||
|
if args.request_rate != float("inf"):
|
||||||
|
# Simple rate limiting
|
||||||
|
interval = 1.0 / args.request_rate
|
||||||
|
await asyncio.sleep(interval)
|
||||||
|
|
||||||
|
task = asyncio.create_task(limited_request_func(req, session, pbar))
|
||||||
|
tasks.append(task)
|
||||||
|
|
||||||
|
outputs = await asyncio.gather(*tasks)
|
||||||
|
total_duration = time.perf_counter() - start_time
|
||||||
|
|
||||||
|
pbar.close()
|
||||||
|
|
||||||
|
# Calculate metrics
|
||||||
|
metrics = calculate_metrics(outputs, total_duration)
|
||||||
|
|
||||||
|
print("\n" + "=" * 40)
|
||||||
|
print("Benchmark Results")
|
||||||
|
print("=" * 40)
|
||||||
|
print(f"Backend: {args.backend}")
|
||||||
|
print(f"Model: {args.model}")
|
||||||
|
print(f"Dataset: {args.dataset}")
|
||||||
|
print(f"Total Duration: {metrics['duration']:.2f} s")
|
||||||
|
print(f"Throughput: {metrics['throughput_qps']:.2f} req/s")
|
||||||
|
print(f"Success Rate: {metrics['completed_requests']}/{len(requests_list)}")
|
||||||
|
print(f"Latency Mean: {metrics['latency_mean']:.4f} s")
|
||||||
|
print(f"Latency Median: {metrics['latency_median']:.4f} s")
|
||||||
|
print(f"Latency P99: {metrics['latency_p99']:.4f} s")
|
||||||
|
print("=" * 40)
|
||||||
|
|
||||||
|
if args.output_file:
|
||||||
|
with open(args.output_file, "w") as f:
|
||||||
|
json.dump(metrics, f, indent=2)
|
||||||
|
print(f"Metrics saved to {args.output_file}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Benchmark serving for diffusion models."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--backend",
|
||||||
|
type=str,
|
||||||
|
required=True,
|
||||||
|
choices=["sglang-image", "sglang-video"],
|
||||||
|
help="Backend type.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--base-url",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Base URL of the server (e.g., http://localhost:30000). Overrides host/port.",
|
||||||
|
)
|
||||||
|
parser.add_argument("--host", type=str, default="localhost", help="Server host.")
|
||||||
|
parser.add_argument("--port", type=int, default=30000, help="Server port.")
|
||||||
|
parser.add_argument("--model", type=str, default="default", help="Model name.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--dataset",
|
||||||
|
type=str,
|
||||||
|
default="vbench",
|
||||||
|
choices=["vbench", "random"],
|
||||||
|
help="Dataset to use.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--task",
|
||||||
|
type=str,
|
||||||
|
default="t2v",
|
||||||
|
choices=["t2v", "i2v", "ti2v", "ti2i"],
|
||||||
|
help="Task type.",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--dataset-path",
|
||||||
|
type=str,
|
||||||
|
default=None,
|
||||||
|
help="Path to local dataset file (optional).",
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--num-prompts", type=int, default=10, help="Number of prompts to benchmark."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--max-concurrency", type=int, default=10, help="Maximum concurrent requests."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--request-rate", type=float, default=float("inf"), help="Request rate (req/s)."
|
||||||
|
)
|
||||||
|
parser.add_argument("--width", type=int, default=None, help="Image/Video width.")
|
||||||
|
parser.add_argument("--height", type=int, default=None, help="Image/Video height.")
|
||||||
|
parser.add_argument(
|
||||||
|
"--num-frames", type=int, default=None, help="Number of frames (for video)."
|
||||||
|
)
|
||||||
|
parser.add_argument("--fps", type=int, default=None, help="FPS (for video).")
|
||||||
|
parser.add_argument(
|
||||||
|
"--output-file", type=str, default=None, help="Output JSON file for metrics."
|
||||||
|
)
|
||||||
|
parser.add_argument(
|
||||||
|
"--disable-tqdm", action="store_true", help="Disable progress bar."
|
||||||
|
)
|
||||||
|
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
asyncio.run(benchmark(args))
|
||||||
Reference in New Issue
Block a user