diff --git a/python/sglang/multimodal_gen/benchmarks/bench_serving.py b/python/sglang/multimodal_gen/benchmarks/bench_serving.py new file mode 100644 index 000000000..5044581eb --- /dev/null +++ b/python/sglang/multimodal_gen/benchmarks/bench_serving.py @@ -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))