diff --git a/python/sglang/multimodal_gen/benchmarks/bench_offline_throughput.py b/python/sglang/multimodal_gen/benchmarks/bench_offline_throughput.py new file mode 100644 index 000000000..5f0a5996c --- /dev/null +++ b/python/sglang/multimodal_gen/benchmarks/bench_offline_throughput.py @@ -0,0 +1,443 @@ +""" +Benchmark offline throughput for multimodal generation models (Image/Video Generation). + +This script benchmarks generation throughput without running a server, using low-level APIs. +It provides detailed metrics on throughput, latency, and resource utilization. + +# Usage Examples + +## Text-to-Video with VBench dataset +python -m sglang.multimodal_gen.benchmarks.bench_offline_throughput \\ + --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \\ + --dataset vbench \\ + --num-prompts 20 \\ + --batch-size 1 \\ + --width 512 --height 512 --num-frames 16 + +## Random dataset for stress testing +python -m sglang.multimodal_gen.benchmarks.bench_offline_throughput \\ + --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \\ + --dataset random \\ + --num-prompts 100 \\ + --batch-size 1 \\ + --num-inference-steps 20 \\ + --output-file results.json +""" + +import argparse +import dataclasses +import json +import time +from dataclasses import dataclass +from typing import Any, Dict, List, Tuple + +import torch +from tqdm import tqdm + +from sglang.multimodal_gen.benchmarks.datasets import RandomDataset, VBenchDataset +from sglang.multimodal_gen.runtime.entrypoints.diffusion_generator import DiffGenerator +from sglang.multimodal_gen.runtime.server_args import ServerArgs, set_global_server_args +from sglang.multimodal_gen.runtime.utils.logging_utils import ( + configure_logger, + init_logger, +) +from sglang.multimodal_gen.test.test_utils import print_divider, print_value_formatted + +logger = init_logger(__name__) + + +@dataclass +class BatchOutput: + """Container for batch generation results.""" + + latency: float = 0.0 + latency_per_sample: float = 0.0 + num_samples: int = 0 + total_frames: int = 0 + peak_memory_mb: float = 0.0 + success: bool = False + error: str = "" + + +@dataclass +class BenchArgs: + """Benchmark configuration for multimodal generation.""" + + # Diffusion Model Configuration + num_inference_steps: int = 20 + guidance_scale: float = 7.5 + seed: int = 42 + disable_safety_checker: bool = False + + # Output Configuration + width: int = 32 + height: int = 32 + num_frames: int = 1 + fps: int = 24 + + # Dataset & Benchmark + dataset: str = "random" + dataset_path: str = "" + task_name: str = "unknown" + num_prompts: int = 10 + batch_size: int = 1 + + # Benchmark Execution + skip_warmup: bool = False + output_file: str = "" + disable_tqdm: bool = False + + @staticmethod + def add_cli_args(parser: argparse.ArgumentParser): + """Add benchmark-specific CLI arguments.""" + # Diffusion Model Configuration + parser.add_argument( + "--num-inference-steps", + type=int, + default=20, + help="Number of denoising steps", + ) + parser.add_argument( + "--guidance-scale", + type=float, + default=7.5, + help="Classifier-free guidance scale", + ) + parser.add_argument("--seed", type=int, default=42, help="Random seed") + parser.add_argument( + "--disable-safety-checker", + action="store_true", + help="Disable NSFW detection", + ) + + # Output Configuration + parser.add_argument("--width", type=int, default=32, help="Image/video width") + parser.add_argument("--height", type=int, default=32, help="Image/video height") + parser.add_argument( + "--num-frames", type=int, default=1, help="Number of frames for video" + ) + parser.add_argument("--fps", type=int, default=24, help="FPS for video") + + # Dataset & Benchmark + parser.add_argument( + "--dataset", + type=str, + default="random", + choices=["vbench", "random"], + help="Dataset to use", + ) + parser.add_argument( + "--dataset-path", + type=str, + default="", + help="Path to dataset (prompts file or image directory)", + ) + parser.add_argument( + "--task-name", + type=str, + default="unknown", + help="Task name for benchmark identification", + ) + parser.add_argument( + "--num-prompts", + type=int, + default=10, + help="Total number of prompts to benchmark", + ) + parser.add_argument( + "--batch-size", + type=int, + default=1, + help="Batch size per generation call (currently only bs=1 is supported)", + ) + + # Benchmark Execution + parser.add_argument( + "--skip-warmup", action="store_true", help="Skip warmup batch" + ) + parser.add_argument( + "--output-file", + type=str, + default="", + help="Output JSON file for results (append mode)", + ) + parser.add_argument( + "--disable-tqdm", + action="store_true", + help="Disable progress bar", + ) + + @classmethod + def from_cli_args(cls, args: argparse.Namespace): + """Create BenchArgs from parsed CLI arguments.""" + attrs = [attr.name for attr in dataclasses.fields(cls)] + return cls(**{attr: getattr(args, attr) for attr in attrs}) + + +def initialize_engine(server_args: ServerArgs) -> DiffGenerator: + """Initialize diffusion pipeline engine.""" + logger.info("Initializing engine...") + engine = DiffGenerator.from_server_args(server_args, local_mode=True) + logger.info("Engine initialized successfully") + return engine + + +def generate_batch( + engine: DiffGenerator, + bench_args: BenchArgs, + prompts: List[str], + user_sampling_params: Dict[str, Any], +) -> BatchOutput: + """Generate batch of images/videos synchronously.""" + output = BatchOutput() + start_time = time.perf_counter() + + torch.cuda.reset_peak_memory_stats() + + for prompt in prompts: + try: + sampling_params_kwargs = dict(user_sampling_params) + sampling_params_kwargs["prompt"] = prompt + result = engine.generate(sampling_params_kwargs=sampling_params_kwargs) + + if result is not None: + if isinstance(result, list): + output.total_frames += len(result) + else: + output.total_frames += 1 + output.num_samples += 1 + except Exception as e: + logger.error(f"Generation failed for prompt '{prompt[:50]}...': {e}") + output.error = str(e) + + output.latency = time.perf_counter() - start_time + output.latency_per_sample = output.latency / len(prompts) if prompts else 0.0 + output.success = output.num_samples > 0 + output.peak_memory_mb = torch.cuda.max_memory_allocated() / (1024 * 1024) + + logger.debug( + f"Batch generated: {output.num_samples}/{len(prompts)} samples in {output.latency:.2f}s" + ) + + return output + + +def calculate_metrics( + outputs: List[BatchOutput], + total_duration: float, + resolution: Tuple[int, int, int], + num_requests: int, +) -> Dict[str, Any]: + """Calculate generation-specific throughput metrics.""" + successful = [o for o in outputs if o.success] + num_success = sum(o.num_samples for o in successful) + total_frames = sum(o.total_frames for o in successful) + peak_memory = max((o.peak_memory_mb for o in outputs), default=0) + + width, height, frames = resolution + pixels_per_sample = width * height * frames + total_pixels = num_success * pixels_per_sample + + metrics = { + "num_requests": num_requests, + "successful_requests": num_success, + "failed_requests": num_requests - num_success, + "total_duration_seconds": total_duration, + "total_frames_generated": total_frames, + "total_pixels_generated": total_pixels, + "images_per_second": num_success / total_duration if total_duration > 0 else 0, + "frames_per_second": total_frames / total_duration if total_duration > 0 else 0, + "megapixels_per_second": ( + total_pixels / (total_duration * 1e6) if total_duration > 0 else 0 + ), + "requests_per_second": ( + num_success / total_duration if total_duration > 0 else 0 + ), + "latency_per_request_seconds": ( + total_duration / num_success if num_success > 0 else 0 + ), + "peak_memory_mb": peak_memory, + } + + return metrics + + +def throughput_test( + server_args: ServerArgs, + bench_args: BenchArgs, +) -> Dict[str, Any]: + """Main throughput benchmark function.""" + configure_logger(server_args=server_args) + logger.info("Starting offline throughput benchmark...") + + engine = initialize_engine(server_args) + + logger.info(f"Loading {bench_args.dataset} dataset...") + if bench_args.dataset == "vbench": + bench_args.task_name = engine.server_args.pipeline_config.task_type + dataset = VBenchDataset(bench_args) + elif bench_args.dataset == "random": + dataset = RandomDataset(bench_args) + else: + raise ValueError(f"Unknown dataset: {bench_args.dataset}") + + sampling_params = { + "guidance_scale": bench_args.guidance_scale, + "num_inference_steps": bench_args.num_inference_steps, + "height": bench_args.height, + "width": bench_args.width, + "num_frames": bench_args.num_frames, + "seed": bench_args.seed, + } + if bench_args.disable_safety_checker: + sampling_params["safety_checker"] = None + + if not bench_args.skip_warmup: + logger.info("Running warmup batch...") + warmup_count = min(bench_args.batch_size, len(dataset)) + warmup_prompts = [dataset[i].prompt for i in range(warmup_count)] + generate_batch(engine, bench_args, warmup_prompts, sampling_params) + + logger.info(f"Running benchmark with {bench_args.num_prompts} prompts...") + outputs: List[BatchOutput] = [] + total_count = min(bench_args.num_prompts, len(dataset)) + all_prompts = [dataset[i].prompt for i in range(total_count)] + + start_time = time.perf_counter() + + num_batches = (total_count + bench_args.batch_size - 1) // bench_args.batch_size + pbar = tqdm( + total=num_batches, + disable=bench_args.disable_tqdm, + desc="Benchmark", + ) + + for batch_start in range(0, total_count, bench_args.batch_size): + batch_end = min(batch_start + bench_args.batch_size, total_count) + batch_prompts = all_prompts[batch_start:batch_end] + + batch_output = generate_batch( + engine, bench_args, batch_prompts, sampling_params + ) + outputs.append(batch_output) + + pbar.update(1) + + pbar.close() + total_duration = time.perf_counter() - start_time + + resolution = (bench_args.width, bench_args.height, bench_args.num_frames) + metrics = calculate_metrics( + outputs, + total_duration, + resolution=resolution, + num_requests=total_count, + ) + + display_results( + metrics, + bench_args, + model_path=server_args.model_path, + ) + + if bench_args.output_file: + save_results(metrics, bench_args, server_args) + + return metrics + + +def display_results( + metrics: Dict[str, Any], + bench_args: BenchArgs, + model_path: str, +): + """Display benchmark results in console.""" + print( + "\n{s:{c}^{n}}".format(s=" Offline Throughput Benchmark Result ", n=110, c="=") + ) + print_value_formatted("Model:", model_path) + print_value_formatted("Dataset:", bench_args.dataset) + print_value_formatted( + "Resolution:", + f"{bench_args.width}x{bench_args.height}x{bench_args.num_frames}", + ) + print_value_formatted("Num Inference Steps:", bench_args.num_inference_steps) + print_divider(75) + print_value_formatted("Total Requests:", metrics["num_requests"]) + print_value_formatted("Successful Requests:", metrics["successful_requests"]) + print_value_formatted("Failed Requests:", metrics["failed_requests"]) + print_value_formatted( + "Total Duration (seconds):", metrics["total_duration_seconds"] + ) + print_divider(75) + print_value_formatted("Frames Generated:", metrics["total_frames_generated"]) + print_value_formatted( + "Megapixels Generated:", metrics["total_pixels_generated"] / 1e6 + ) + print_divider(75) + print_value_formatted( + "Frame Throughput (frames/sec):", metrics["frames_per_second"] + ) + print_value_formatted("MP Throughput (MP/sec):", metrics["megapixels_per_second"]) + print_value_formatted("Requests Per Second:", metrics["requests_per_second"]) + print_value_formatted( + "Latency Per Request (sec):", metrics["latency_per_request_seconds"] + ) + print_value_formatted("Peak Memory (MB):", metrics["peak_memory_mb"]) + print_divider(110, "=") + + +def save_results( + metrics: Dict[str, Any], + bench_args: BenchArgs, + server_args: ServerArgs, +): + """Save benchmark results to JSON file.""" + result = { + "metadata": { + "timestamp": time.strftime("%Y-%m-%dT%H:%M:%S"), + "model_path": server_args.model_path, + "task_type": bench_args.task_name, + "backend": "engine", + }, + "configuration": { + "num_inference_steps": bench_args.num_inference_steps, + "guidance_scale": bench_args.guidance_scale, + "seed": bench_args.seed, + "batch_size": bench_args.batch_size, + "num_prompts": bench_args.num_prompts, + "resolution": f"{bench_args.width}x{bench_args.height}x{bench_args.num_frames}", + "dataset": bench_args.dataset, + }, + "results": metrics, + } + + with open(bench_args.output_file, "a") as f: + f.write(json.dumps(result) + "\n") + + logger.info(f"Results saved to {bench_args.output_file}") + + +def main(): + """Main entry point.""" + parser = argparse.ArgumentParser( + description="Offline throughput benchmark for multimodal generation models" + ) + + ServerArgs.add_cli_args(parser) + BenchArgs.add_cli_args(parser) + + args = parser.parse_args() + + server_args = ServerArgs.from_cli_args(args) + bench_args = BenchArgs.from_cli_args(args) + + set_global_server_args(server_args) + + result = throughput_test(server_args, bench_args) + + return result + + +if __name__ == "__main__": + main() diff --git a/python/sglang/multimodal_gen/benchmarks/bench_serving.py b/python/sglang/multimodal_gen/benchmarks/bench_serving.py index 4ed132148..1ec9721fa 100644 --- a/python/sglang/multimodal_gen/benchmarks/bench_serving.py +++ b/python/sglang/multimodal_gen/benchmarks/bench_serving.py @@ -14,14 +14,9 @@ Usage: import argparse import asyncio -import glob import json import os -import re import time -import uuid -from abc import ABC, abstractmethod -from dataclasses import dataclass, field from typing import Any, Dict, List, Optional import aiohttp @@ -29,318 +24,21 @@ import numpy as np import requests from tqdm.asyncio import tqdm +from sglang.multimodal_gen.benchmarks.datasets import ( + RandomDataset, + RequestFuncInput, + RequestFuncOutput, + VBenchDataset, +) from sglang.multimodal_gen.runtime.utils.logging_utils import ( configure_logger, init_logger, ) +from sglang.multimodal_gen.test.test_utils import print_divider, print_value_formatted logger = init_logger(__name__) -def is_dir_not_empty(path): - return os.path.isdir(path) and bool(os.listdir(path)) - - -@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) - peak_memory_mb: float = 0.0 - - -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_name in ("text-to-video", "text-to-image", "video-to-video"): - return self._load_t2v_prompts() - elif self.args.task_name in ("image-to-video", "image-to-image"): - return self._load_i2v_data() - else: - raise ValueError( - f"Illegal task name is found in VBenchDataset {args.task_name}" - ) - - 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): - logger.info(f"Downloading VBench T2V prompts to {path}...") - try: - self._download_file(self.T2V_PROMPT_URL, path) - except Exception as e: - logger.info(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") - crop_dir = os.path.join(vbench_i2v_dir, "data", "crop") - origin_dir = os.path.join(vbench_i2v_dir, "data", "origin") - - if ( - os.path.exists(info_json_path) - and is_dir_not_empty(crop_dir) - and is_dir_not_empty(origin_dir) - ): - return vbench_i2v_dir - - logger.info(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) - - logger.info("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}") - missing_packages = re.findall(r"(\S+): command not found", result.stderr) - if missing_packages: - missing_packages = list(set(missing_packages)) - package_list = ", ".join(f"'{cmd}'" for cmd in missing_packages) - raise RuntimeError( - f"Download script failed because the following commands are not installed: {package_list}.\n" - "Please install them (e.g., on Ubuntu: `sudo apt install ...`) and try again." - ) - logger.info( - f"Successfully downloaded VBench I2V dataset to {vbench_i2v_dir}" - ) - except Exception as e: - logger.info(f"Failed to download VBench I2V dataset: {e}") - logger.info("Please manually download following instructions at:") - logger.info( - "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: - logger.warning(f"Image not found: {img_path}") - - logger.info(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.""" - logger.info("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) - logger.info(f"Created dummy image at {dummy_image}") - except ImportError: - logger.info("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: - logger.info(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, @@ -501,7 +199,7 @@ async def async_request_video_sglang( else: # Use JSON - payload = { + payload: Dict[str, Any] = { "model": input.model, "prompt": input.prompt, } @@ -746,57 +444,41 @@ async def benchmark(args): print("\n{s:{c}^{n}}".format(s=" Serving Benchmark Result ", n=60, c="=")) # Section 1: Configuration - print("{:<40} {:<15}".format("Task:", task_name)) - print("{:<40} {:<15}".format("Model:", args.model)) - print("{:<40} {:<15}".format("Dataset:", args.dataset)) + print_value_formatted("Task:", task_name) + print_value_formatted("Model:", args.model) + print_value_formatted("Dataset:", args.dataset) # Section 2: Execution & Traffic - print(f"{'-' * 50}") - print("{:<40} {:<15.2f}".format("Benchmark duration (s):", metrics["duration"])) - print("{:<40} {:<15}".format("Request rate:", str(args.request_rate))) - print( - "{:<40} {:<15}".format( - "Max request concurrency:", - str(args.max_concurrency) if args.max_concurrency else "not set", - ) + print_divider(50) + print_value_formatted("Benchmark duration (s):", metrics["duration"]) + print_value_formatted("Request rate:", str(args.request_rate)) + print_value_formatted( + "Max request concurrency:", + str(args.max_concurrency) if args.max_concurrency else "not set", ) - print( - "{:<40} {}/{:<15}".format( - "Successful requests:", metrics["completed_requests"], len(requests_list) - ) + print_value_formatted( + "Successful requests:", + f"{metrics['completed_requests']}/{len(requests_list)}", ) # Section 3: Performance Metrics - print(f"{'-' * 50}") + print_divider(50) - print( - "{:<40} {:<15.2f}".format( - "Request throughput (req/s):", metrics["throughput_qps"] - ) - ) - print("{:<40} {:<15.4f}".format("Latency Mean (s):", metrics["latency_mean"])) - print("{:<40} {:<15.4f}".format("Latency Median (s):", metrics["latency_median"])) - print("{:<40} {:<15.4f}".format("Latency P99 (s):", metrics["latency_p99"])) + print_value_formatted("Request throughput (req/s):", metrics["throughput_qps"]) + + print_value_formatted("Latency Mean (s):", metrics["latency_mean"]) + print_value_formatted("Latency Median (s):", metrics["latency_median"]) + print_value_formatted("Latency P99 (s):", metrics["latency_p99"]) if metrics["peak_memory_mb_max"] > 0: - print(f"{'-' * 50}") - print( - "{:<40} {:<15.2f}".format( - "Peak Memory Max (MB):", metrics["peak_memory_mb_max"] - ) - ) - print( - "{:<40} {:<15.2f}".format( - "Peak Memory Mean (MB):", metrics["peak_memory_mb_mean"] - ) - ) - print( - "{:<40} {:<15.2f}".format( - "Peak Memory Median (MB):", metrics["peak_memory_mb_median"] - ) + print_divider(50) + print_value_formatted("Peak Memory Max (MB):", metrics["peak_memory_mb_max"]) + print_value_formatted("Peak Memory Mean (MB):", metrics["peak_memory_mb_mean"]) + print_value_formatted( + "Peak Memory Median (MB):", metrics["peak_memory_mb_median"] ) - print("=" * 60) + print_divider(60) if args.output_file: with open(args.output_file, "w") as f: diff --git a/python/sglang/multimodal_gen/benchmarks/datasets.py b/python/sglang/multimodal_gen/benchmarks/datasets.py new file mode 100644 index 000000000..d3eb1d10b --- /dev/null +++ b/python/sglang/multimodal_gen/benchmarks/datasets.py @@ -0,0 +1,298 @@ +import glob +import json +import os +import re +import subprocess +import uuid +from abc import ABC, abstractmethod +from dataclasses import dataclass, field +from typing import Any, Dict, List, Optional + +import requests +from PIL import Image + +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger + +logger = init_logger(__name__) + + +@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) + peak_memory_mb: float = 0.0 + + +def is_dir_not_empty(path: str) -> bool: + return os.path.isdir(path) and bool(os.listdir(path)) + + +class BaseDataset(ABC): + def __init__(self, args, api_url: str = "", model: str = ""): + self.args = args + self.api_url = api_url + self.model = model + self.items: List[Dict[str, Any]] = [] + + @abstractmethod + def __len__(self) -> int: + pass + + @abstractmethod + def __getitem__(self, idx: int) -> RequestFuncInput: + pass + + def get_requests(self) -> List[RequestFuncInput]: + return [self[i] for i in range(len(self))] + + +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_name in ("text-to-video", "text-to-image", "video-to-video"): + return self._load_t2v_prompts() + elif self.args.task_name in ("image-to-video", "image-to-image"): + return self._load_i2v_data() + else: + raise ValueError( + f"Illegal task name is found in VBenchDataset {self.args.task_name}" + ) + + 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): + logger.info(f"Downloading VBench T2V prompts to {path}...") + try: + self._download_file(self.T2V_PROMPT_URL, path) + except Exception as e: + logger.info(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) -> Optional[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") + crop_dir = os.path.join(vbench_i2v_dir, "data", "crop") + origin_dir = os.path.join(vbench_i2v_dir, "data", "origin") + + if ( + os.path.exists(info_json_path) + and is_dir_not_empty(crop_dir) + and is_dir_not_empty(origin_dir) + ): + return vbench_i2v_dir + + logger.info(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) + + logger.info("Executing download_data.sh (this may take a while)...") + + 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}") + missing_packages = re.findall(r"(\S+): command not found", result.stderr) + if missing_packages: + missing_packages = list(set(missing_packages)) + package_list = ", ".join(f"'{cmd}'" for cmd in missing_packages) + raise RuntimeError( + f"Download script failed because the following commands are not installed: {package_list}.\n" + "Please install them (e.g., on Ubuntu: `sudo apt install ...`) and try again." + ) + logger.info( + f"Successfully downloaded VBench I2V dataset to {vbench_i2v_dir}" + ) + except Exception as e: + logger.info(f"Failed to download VBench I2V dataset: {e}") + logger.info("Please manually download following instructions at:") + logger.info( + "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: + logger.warning(f"Image not found: {img_path}") + + logger.info(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()))) + + 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.""" + logger.info("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): + os.makedirs(self.cache_dir, exist_ok=True) + img = Image.new("RGB", (100, 100), color="red") + img.save(dummy_image) + logger.info(f"Created dummy image at {dummy_image}") + + 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 + if not path: + path = self._auto_download_i2v_dataset() + if not path: + return self._resize_data(self._create_dummy_data()) + + 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: + logger.info(f"Failed to load {json_path}: {e}") + + if os.path.isdir(path): + data = self._scan_directory_for_images(path) + if data: + return self._resize_data(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] + 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=[item["image_path"]] if "image_path" in item else None, + ) + + +class RandomDataset(BaseDataset): + def __init__(self, args, api_url: str = "", model: str = ""): + super().__init__(args, api_url, 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, + ) diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index b4efafc65..073b36972 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -22,6 +22,28 @@ from sglang.multimodal_gen.runtime.utils.perf_logger import ( logger = init_logger(__name__) +def print_value_formatted(description: str, value: int | float | str): + """Helper function to print a metric value formatted.""" + if isinstance(value, int): + if value >= 1e6: + value_str = f"{value / 1e6:<30.2f}M" + elif value >= 1e3: + value_str = f"{value / 1e3:<30.2f}K" + else: + value_str = f"{value:<30}" + elif isinstance(value, float): + value_str = f"{value:<30.2f}" + else: + value_str = f"{value:<30}" + + print(f"{description:<45} {value_str}") + + +def print_divider(length: int, char: str = "-"): + """Helper function to print a divider line.""" + print(char * length) + + def is_image_url(image_path: str | Path | None) -> bool: """Check if image_path is a URL.""" if image_path is None: