diff --git a/python/sglang/multimodal_gen/benchmarks/bench_serving.py b/python/sglang/multimodal_gen/benchmarks/bench_serving.py index 027010a3a..2a1901ef9 100644 --- a/python/sglang/multimodal_gen/benchmarks/bench_serving.py +++ b/python/sglang/multimodal_gen/benchmarks/bench_serving.py @@ -3,26 +3,13 @@ Benchmark online serving for diffusion models (Image/Video Generation). Usage: - # Video - t2v: - python3 -m sglang.multimodal_gen.benchmarks.bench_serving \ - --backend sglang-video --dataset vbench --task t2v --num-prompts 20 - - i2v: - python3 -m sglang.multimodal_gen.benchmarks.bench_serving \ - --backend sglang-video --dataset vbench --task i2v --num-prompts 20 - - - # Image - t2i: - python3 -m sglang.multimodal_gen.benchmarks.bench_serving \ - --backend sglang-image --dataset vbench --task t2i --num-prompts 20 - - ti2i(edit): - python3 -m sglang.multimodal_gen.benchmarks.bench_serving \ - --backend sglang-image --dataset vbench --task ti2i --num-prompts 20 + # launch a server and benchmark on it + # T2V or T2I or any other multimodal generation model + sglang serve Wan-AI/Wan2.2-T2V-A14B-Diffusers --num-gpus 1 --port 1231 + # benchmark it and make sure the port is the same as the server's port + python3 -m sglang.multimodal_gen.benchmarks.bench_serving --dataset vbench --num-prompts 20 --port 1231 """ import argparse @@ -42,6 +29,13 @@ import numpy as np import requests from tqdm.asyncio import tqdm +from sglang.multimodal_gen.runtime.utils.logging_utils import ( + configure_logger, + init_logger, +) + +logger = init_logger(__name__) + def is_dir_not_empty(path): return os.path.isdir(path) and bool(os.listdir(path)) @@ -105,12 +99,14 @@ class VBenchDataset(BaseDataset): self.items = self._load_data() def _load_data(self) -> List[Dict[str, Any]]: - if self.args.task == "t2v" or self.args.task == "t2i": + if self.args.task_name in ("text-to-video", "text-to-image", "video-to-video"): return self._load_t2v_prompts() - elif self.args.task in ["i2v", "ti2v", "ti2i"]: + elif self.args.task_name in ("image-to-video", "image-to-image"): return self._load_i2v_data() else: - return self._load_t2v_prompts() + 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.""" @@ -126,11 +122,11 @@ class VBenchDataset(BaseDataset): 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}...") + logger.info(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}") + logger.info(f"Failed to download VBench prompts: {e}") return [{"prompt": "A cat sitting on a bench"}] * 50 prompts = [] @@ -156,7 +152,7 @@ class VBenchDataset(BaseDataset): ): return vbench_i2v_dir - print(f"Downloading VBench I2V dataset to {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") @@ -164,7 +160,7 @@ class VBenchDataset(BaseDataset): 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)...") + logger.info("Executing download_data.sh (this may take a while)...") import subprocess result = subprocess.run( @@ -183,11 +179,13 @@ class VBenchDataset(BaseDataset): 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." ) - print(f"Successfully downloaded VBench I2V dataset to {vbench_i2v_dir}") + logger.info( + 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( + 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 @@ -210,9 +208,9 @@ class VBenchDataset(BaseDataset): 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}") + logger.warning(f"Image not found: {img_path}") - print(f"Loaded {len(data)} I2V samples from VBench I2V dataset") + 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]]: @@ -237,7 +235,7 @@ class VBenchDataset(BaseDataset): 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.") + 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): @@ -247,9 +245,9 @@ class VBenchDataset(BaseDataset): 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}") + logger.info(f"Created dummy image at {dummy_image}") except ImportError: - print("PIL not installed, cannot create dummy image.") + logger.info("PIL not installed, cannot create dummy image.") return [] return [{"prompt": "A moving cat", "image_path": dummy_image}] * 10 @@ -274,7 +272,7 @@ class VBenchDataset(BaseDataset): try: return self._resize_data(self._load_from_i2v_json(json_path)) except Exception as e: - print(f"Failed to load {json_path}: {e}") + logger.info(f"Failed to load {json_path}: {e}") # Fallback: scan directory for images if os.path.isdir(path): @@ -612,14 +610,14 @@ def calculate_metrics(outputs: List[RequestFuncOutput], total_duration: float): def wait_for_service(base_url: str, timeout: int = 1200) -> None: - print(f"Waiting for service at {base_url}...") + logger.info(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.") + logger.info("Service is ready.") break except requests.exceptions.RequestException: pass @@ -633,6 +631,8 @@ def wait_for_service(base_url: str, timeout: int = 1200) -> None: async def benchmark(args): + from huggingface_hub import model_info + # Construct base_url if not provided if args.base_url is None: args.base_url = f"http://{args.host}:{args.port}" @@ -647,28 +647,32 @@ async def benchmark(args): info = resp.json() if "model_path" in info and info["model_path"]: args.model = info["model_path"] - print(f"Updated model name from server: {args.model}") + logger.info(f"Updated model name from server: {args.model}") except Exception as e: - print(f"Failed to fetch model info: {e}. Using default: {args.model}") + logger.info(f"Failed to fetch model info: {e}. Using default: {args.model}") - # Setup dataset - if args.backend == "sglang-image": - if args.task not in ["ti2i", "t2i"]: - raise Exception("sglang-image backend only support ti2i and t2i tasks.") - if args.task == "ti2i": + task_name = model_info(args.model).pipeline_tag + + if args.task != task_name: + logger.warning( + f"Task from args {args.task} is different from huggingface pipeline_tag {task_name}, args.task will be ignored!" + ) + + if task_name in ("text-to-video", "image-to-video", "video-to-video"): + api_url = f"{args.base_url}/v1/videos" + request_func = async_request_video_sglang + elif task_name in ("text-to-image", "image-to-image"): + if task_name == "image-to-image": 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": - if args.task not in ["t2v", "i2v", "ti2v"]: - raise Exception( - "sglang-video backend only support t2v, i2v and ti2v tasks." - ) - api_url = f"{args.base_url}/v1/videos" - request_func = async_request_video_sglang else: - raise ValueError(f"Unknown backend: {args.backend}") + raise ValueError( + f"The task name {task_name} of model {args.model} is not a valid task name for multimodal generation. Please check the model path." + ) + + setattr(args, "task_name", task_name) if args.dataset == "vbench": dataset = VBenchDataset(args, api_url, args.model) @@ -677,9 +681,9 @@ async def benchmark(args): else: raise ValueError(f"Unknown dataset: {args.dataset}") - print(f"Loading requests...") + logger.info(f"Loading requests...") requests_list = dataset.get_requests() - print(f"Prepared {len(requests_list)} requests from {args.dataset} dataset.") + logger.info(f"Prepared {len(requests_list)} requests from {args.dataset} dataset.") # Limit concurrency if args.max_concurrency is not None: @@ -720,10 +724,9 @@ async def benchmark(args): print("\n{s:{c}^{n}}".format(s=" Serving Benchmark Result ", n=60, c="=")) # Section 1: Configuration - print("{:<40} {:<15}".format("Backend:", args.backend)) + print("{:<40} {:<15}".format("Task:", task_name)) print("{:<40} {:<15}".format("Model:", args.model)) print("{:<40} {:<15}".format("Dataset:", args.dataset)) - print("{:<40} {:<15}".format("Task:", args.task)) # Section 2: Execution & Traffic print(f"{'-' * 50}") @@ -786,9 +789,8 @@ if __name__ == "__main__": parser.add_argument( "--backend", type=str, - required=True, - choices=["sglang-image", "sglang-video"], - help="Backend type.", + default=None, + help="DEPRECATED: --task is deprecated and will be ignored. The task will be inferred from --model.", ) parser.add_argument( "--base-url", @@ -809,9 +811,15 @@ if __name__ == "__main__": parser.add_argument( "--task", type=str, - default="t2v", - choices=["t2v", "i2v", "ti2v", "ti2i", "t2i"], - help="Task type. t2v, i2v, ti2v are used for video generation. ti2i, t2i are used for image generation. ti2i is image edit task and t2i is image generation task.", + choices=[ + "text-to-video", + "image-to-video", + "text-to-image", + "image-to-image", + "video-to-video", + ], + default=None, + help="The task will be inferred from huggingface pipeline_tag. When huggingface pipeline_tag is not provided, --task will be used.", ) parser.add_argument( "--dataset-path", @@ -854,7 +862,16 @@ if __name__ == "__main__": parser.add_argument( "--disable-tqdm", action="store_true", help="Disable progress bar." ) + parser.add_argument( + "--log-level", + type=str, + default="INFO", + choices=["DEBUG", "INFO", "WARNING", "ERROR"], + help="Log level.", + ) args = parser.parse_args() + configure_logger(args) + asyncio.run(benchmark(args))