[diffusion] profiling: add bench_serving.py and VBench (#15410)
This commit is contained in:
744
python/sglang/multimodal_gen/benchmarks/bench_serving.py
Normal file
744
python/sglang/multimodal_gen/benchmarks/bench_serving.py
Normal file
@@ -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