From 5e7f91d45151a8c9e8a23cc08d0445d342d845e4 Mon Sep 17 00:00:00 2001 From: Mick Date: Fri, 21 Nov 2025 18:47:16 +0800 Subject: [PATCH] [diffusion] profile: support performance metric dumping and comparison (#13630) --- .../multimodal_gen/benchmarks/compare_perf.py | 216 +++++ .../multimodal_gen/configs/sample/base.py | 1 + python/sglang/multimodal_gen/envs.py | 10 +- .../runtime/entrypoints/cli/generate.py | 54 +- .../entrypoints/diffusion_generator.py | 6 +- .../runtime/managers/gpu_worker.py | 43 +- .../runtime/managers/scheduler.py | 2 +- .../executors/pipeline_executor.py | 24 +- .../runtime/pipelines_core/schedule_batch.py | 64 +- .../runtime/pipelines_core/stages/base.py | 54 +- .../runtime/pipelines_core/stages/decoding.py | 1 + .../pipelines_core/stages/denoising.py | 166 ++-- .../multimodal_gen/runtime/utils/common.py | 6 +- .../runtime/utils/perf_logger.py | 265 ++++++ .../runtime/utils/performance_logger.py | 204 ---- .../test/server/perf_baselines.json | 220 ++--- .../test/server/test_server_performance.py | 140 +-- .../test/server/test_server_utils.py | 57 +- .../test/server/testcase_configs.py | 56 +- .../sglang/multimodal_gen/test/test_utils.py | 874 +++++++++--------- 20 files changed, 1315 insertions(+), 1148 deletions(-) create mode 100644 python/sglang/multimodal_gen/benchmarks/compare_perf.py create mode 100644 python/sglang/multimodal_gen/runtime/utils/perf_logger.py delete mode 100644 python/sglang/multimodal_gen/runtime/utils/performance_logger.py diff --git a/python/sglang/multimodal_gen/benchmarks/compare_perf.py b/python/sglang/multimodal_gen/benchmarks/compare_perf.py new file mode 100644 index 000000000..2dfb087c7 --- /dev/null +++ b/python/sglang/multimodal_gen/benchmarks/compare_perf.py @@ -0,0 +1,216 @@ +import argparse +import json +import re +from datetime import datetime +from typing import Any, Dict, List, Tuple + + +def calculate_diff(base: float, new: float) -> Tuple[float, float]: + """Returns (diff, diff_percent).""" + diff = new - base + if base == 0: + percent = 0.0 + else: + percent = (diff / base) * 100 + return diff, percent + + +def calculate_upper_bound(baseline: float, rel_tol: float, min_abs_tol: float) -> float: + """Calculates the upper bound for performance regression check.""" + rel_limit = baseline * (1 + rel_tol) + abs_limit = baseline + min_abs_tol + return max(rel_limit, abs_limit) + + +def calculate_lower_bound(baseline: float, rel_tol: float, min_abs_tol: float) -> float: + """Calculates the lower bound for performance improvement check.""" + rel_lower = baseline * (1 - rel_tol) + abs_lower = baseline - min_abs_tol + return min(rel_lower, abs_lower) + + +def get_perf_status_emoji( + baseline: float, + new: float, + rel_tol: float = 0.1, + min_abs_tol: float = 120.0, +) -> str: + """ + Determines the status emoji based on performance difference. + + Logic: + Upper bound (Slower): max(baseline * (1 + rel_tol), baseline + min_abs_tol) + Lower bound (Faster): min(baseline * (1 - rel_tol), baseline - min_abs_tol) + """ + upper_bound = calculate_upper_bound(baseline, rel_tol, min_abs_tol) + lower_bound = calculate_lower_bound(baseline, rel_tol, min_abs_tol) + + if new > upper_bound: + return "🔴" + elif new < lower_bound: + return "🟢" + else: + return "⚪️" + + +def consolidate_steps( + steps_list: List[Dict[str, Any]], +) -> Tuple[Dict[str, float], List[str], Dict[str, int]]: + """ + Aggregates specific repeating steps (like denoising_step_*) into groups. + Returns: + - aggregated_durations: {name: duration_ms} + - ordered_names: list of names in execution order + - counts: {name: count_of_steps_aggregated} + """ + durations = {} + counts = {} + ordered_names = [] + seen_names = set() + + # Regex for steps to group + # Group "denoising_step_0", "denoising_step_1" -> "Denoising Loop" + denoise_pattern = re.compile(r"^denoising_step_(\d+)$") + denoising_group_name = "Denoising Loop" + + for step in steps_list: + name = step.get("name", "unknown") + dur = step.get("duration_ms", 0.0) + + match = denoise_pattern.match(name) + if match: + key = denoising_group_name + if key not in durations: + durations[key] = 0.0 + counts[key] = 0 + if key not in seen_names: + ordered_names.append(key) + seen_names.add(key) + durations[key] += dur + counts[key] += 1 + else: + # Standard stage (preserve order) + if name not in durations: + durations[name] = 0.0 + counts[name] = 0 + if name not in seen_names: + ordered_names.append(name) + seen_names.add(name) + durations[name] += dur + counts[name] += 1 + + return durations, ordered_names, counts + + +def _load_benchmark_file(file_path: str) -> Dict[str, Any]: + """Loads a benchmark JSON file.""" + with open(file_path, "r", encoding="utf-8") as f: + return json.load(f) + + +def compare_benchmarks( + baseline_path: str, new_path: str, output_format: str = "markdown" +): + """ + Compares two benchmark JSON files and prints a report. + """ + try: + base_data = _load_benchmark_file(baseline_path) + new_data = _load_benchmark_file(new_path) + except Exception as e: + print(f"Error loading benchmark files: {e}") + return + + base_e2e = base_data.get("total_duration_ms", 0) + new_e2e = new_data.get("total_duration_ms", 0) + + diff_ms, diff_pct = calculate_diff(base_e2e, new_e2e) + + if diff_pct < -2.0: + status = "✅" + elif diff_pct > 2.0: + status = "❌" + else: + status = "" + + # --- Stage Breakdown --- + base_durations, base_order, base_counts = consolidate_steps( + base_data.get("steps", []) + ) + new_durations, new_order, new_counts = consolidate_steps(new_data.get("steps", [])) + + # Merge orders: Start with New order (execution order), append any missing from Base + combined_order = list(new_order) + for name in base_order: + if name not in combined_order: + combined_order.append(name) + + stage_rows = [] + for stage in combined_order: + b_val = base_durations.get(stage, 0.0) + n_val = new_durations.get(stage, 0.0) + b_count = base_counts.get(stage, 1) + n_count = new_counts.get(stage, 1) + + s_diff, s_pct = calculate_diff(b_val, n_val) + + # Format count string if aggregated + count_str = "" + if stage == "Denoising Loop": + count_str = ( + f" ({n_count} steps)" + if n_count == b_count + else f" ({b_count}->{n_count} steps)" + ) + + # filter noise: show if diff is > 0.5ms OR if it's a major stage (like Denoising Loop) + # always show Denoising Loop or stages with significant duration/diff + stage_rows.append((stage + count_str, b_val, n_val, s_diff, s_pct)) + + if output_format == "markdown": + print("### Performance Comparison Report\n") + + # Summary Table + print("#### 1. High-level Summary") + print("| Metric | Baseline | New | Diff | Status |") + print("| :--- | :--- | :--- | :--- | :--- |") + print( + f"| **E2E Latency** | {base_e2e:.2f} ms | {new_e2e:.2f} ms | **{diff_ms:+.2f} ms ({diff_pct:+.1f}%)** | {status} |" + ) + print( + f"| **Throughput** | {1000 / base_e2e if base_e2e else 0:.2f} req/s | {1000 / new_e2e if new_e2e else 0:.2f} req/s | - | - |" + ) + print("\n") + + # Detailed Breakdown + print("#### 2. Stage Breakdown") + print( + "| Stage Name | Baseline (ms) | New (ms) | Diff (ms) | Diff (%) | Status |" + ) + print("| :--- | :--- | :--- | :--- | :--- | :--- |") + for name, b, n, d, p in stage_rows: + name_str = name + status_emoji = get_perf_status_emoji(b, n) + print( + f"| {name_str} | {b:.2f} | {n:.2f} | {d:+.2f} | {p:+.1f}% | {status_emoji} |" + ) + print("\n") + + # Metadata + print("
") + print("Metadata\n") + print(f"- Baseline Commit: `{base_data.get('commit_hash', 'N/A')}`") + print(f"- New Commit: `{new_data.get('commit_hash', 'N/A')}`") + print(f"- Timestamp: {datetime.now().isoformat()}") + print("
") + + +if __name__ == "__main__": + parser = argparse.ArgumentParser( + description="Compare two sglang-diffusion performance JSON files." + ) + parser.add_argument("baseline", help="Path to the baseline JSON file") + parser.add_argument("new", help="Path to the new JSON file") + args = parser.parse_args() + + compare_benchmarks(args.baseline, args.new) diff --git a/python/sglang/multimodal_gen/configs/sample/base.py b/python/sglang/multimodal_gen/configs/sample/base.py index 11de5bc87..d15f2e427 100644 --- a/python/sglang/multimodal_gen/configs/sample/base.py +++ b/python/sglang/multimodal_gen/configs/sample/base.py @@ -129,6 +129,7 @@ class SamplingParams: # Debugging debug: bool = False + perf_dump_path: str | None = None # Misc save_output: bool = True diff --git a/python/sglang/multimodal_gen/envs.py b/python/sglang/multimodal_gen/envs.py index c07786090..56418e72d 100644 --- a/python/sglang/multimodal_gen/envs.py +++ b/python/sglang/multimodal_gen/envs.py @@ -12,6 +12,8 @@ import diffusers import torch from packaging import version +from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var + logger = logging.getLogger(__name__) if TYPE_CHECKING: @@ -277,13 +279,13 @@ environment_variables: dict[str, Callable[[], Any]] = { # If set, sgl_diffusion will run in development mode, which will enable # some additional endpoints for developing and debugging, # e.g. `/reset_prefix_cache` - "SGLANG_DIFFUSION_SERVER_DEV_MODE": lambda: bool( - int(os.getenv("SGLANG_DIFFUSION_SERVER_DEV_MODE", "0")) + "SGLANG_DIFFUSION_SERVER_DEV_MODE": lambda: get_bool_env_var( + "SGLANG_DIFFUSION_SERVER_DEV_MODE" ), # If set, sgl_diffusion will enable stage logging, which will print the time # taken for each stage - "SGLANG_DIFFUSION_STAGE_LOGGING": lambda: bool( - int(os.getenv("SGLANG_DIFFUSION_STAGE_LOGGING", "0")) + "SGLANG_DIFFUSION_STAGE_LOGGING": lambda: get_bool_env_var( + "SGLANG_DIFFUSION_STAGE_LOGGING" ), } diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py b/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py index 948b86375..b557ae2a8 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/cli/generate.py @@ -8,6 +8,7 @@ import dataclasses import os from typing import cast +import sglang.multimodal_gen.envs as envs from sglang.multimodal_gen import DiffGenerator from sglang.multimodal_gen.configs.sample.base import ( SamplingParams, @@ -19,6 +20,10 @@ from sglang.multimodal_gen.runtime.entrypoints.cli.utils import ( ) from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.perf_logger import ( + PerformanceLogger, + RequestTimings, +) from sglang.multimodal_gen.utils import FlexibleArgumentParser logger = init_logger(__name__) @@ -33,6 +38,13 @@ def add_multimodal_gen_generate_args(parser: argparse.ArgumentParser): required=False, help="Read CLI options from a config JSON or YAML file. If provided, --model-path and --prompt are optional.", ) + parser.add_argument( + "--perf-dump-path", + type=str, + default=None, + required=False, + help="Path to dump the performance metrics (JSON) for the run.", + ) parser = ServerArgs.add_cli_args(parser) parser = SamplingParams.add_cli_args(parser) @@ -46,11 +58,47 @@ def add_multimodal_gen_generate_args(parser: argparse.ArgumentParser): return parser +def maybe_dump_performance( + args: argparse.Namespace, server_args, sampling_params, results +): + """dump performance if necessary""" + if not (args.perf_dump_path and results): + return + + if isinstance(results, list): + result = results[0] if results else {} + else: + result = results + + timings_dict = result.get("timings") + if not (args.perf_dump_path and timings_dict): + return + + timings = RequestTimings(request_id=timings_dict.get("request_id")) + timings.stages = timings_dict.get("stages", {}) + timings.total_duration_ms = timings_dict.get("total_duration_ms", 0) + + PerformanceLogger.dump_benchmark_report( + file_path=args.perf_dump_path, + timings=timings, + meta={ + "prompt": sampling_params.prompt, + "model": server_args.model_path, + }, + tag="cli_generate", + ) + + def generate_cmd(args: argparse.Namespace): """The entry point for the generate command.""" # FIXME(mick): do not hard code args.request_id = generate_request_id() + # Auto-enable stage logging if dump path is provided + if args.perf_dump_path: + os.environ["SGLANG_DIFFUSION_STAGE_LOGGING"] = "True" + envs.SGLANG_DIFFUSION_STAGE_LOGGING = True + server_args = ServerArgs.from_cli_args(args) sampling_params = SamplingParams.from_cli_args(args) sampling_params.request_id = generate_request_id() @@ -58,7 +106,11 @@ def generate_cmd(args: argparse.Namespace): model_path=server_args.model_path, server_args=server_args ) - generator.generate(prompt=sampling_params.prompt, sampling_params=sampling_params) + results = generator.generate( + prompt=sampling_params.prompt, sampling_params=sampling_params + ) + + maybe_dump_performance(args, server_args, sampling_params, results) class GenerateSubcommand(CLISubcommand): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py index c1a95d930..0d47d0839 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/diffusion_generator.py @@ -334,7 +334,11 @@ class DiffGenerator: "prompts": req.prompt, "size": (req.height, req.width, req.num_frames), "generation_time": gen_time, - "logging_info": output_batch.logging_info, + "timings": ( + output_batch.timings.to_dict() + if output_batch.timings + else {} + ), "trajectory": output_batch.trajectory_latents, "trajectory_timesteps": output_batch.trajectory_timesteps, "trajectory_decoded": output_batch.trajectory_decoded, diff --git a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py index 8ab64af0f..aeca02b0e 100644 --- a/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py +++ b/python/sglang/multimodal_gen/runtime/managers/gpu_worker.py @@ -3,6 +3,7 @@ # SPDX-License-Identifier: Apache-2.0 import multiprocessing as mp import os +import time from typing import List import torch @@ -25,6 +26,10 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import ( init_logger, suppress_other_loggers, ) +from sglang.multimodal_gen.runtime.utils.perf_logger import ( + PerformanceLogger, + RequestTimings, +) logger = init_logger(__name__) @@ -86,27 +91,35 @@ class GPUWorker: f"Worker {self.rank}: Initialized device, model, and distributed environment." ) - def execute_forward(self, batch: List[Req], server_args: ServerArgs) -> OutputBatch: + def execute_forward(self, batch: List[Req]) -> OutputBatch: """ Execute a forward pass. """ assert self.pipeline is not None # TODO: dealing with first req for now req = batch[0] - output_batch = self.pipeline.forward(req, server_args) - if req.perf_logger: - logging_info = getattr(output_batch, "logging_info", None) or getattr( - req, "logging_info", None - ) - if logging_info: - try: - req.perf_logger.log_stage_metrics(logging_info) - except Exception: - logger.exception( - "Failed to log stage metrics for request %s", req.request_id - ) - req.perf_logger.log_total_duration("total_inference_time") - return output_batch + output_batch = None + try: + start_time = time.monotonic() + timings = RequestTimings(request_id=req.request_id) + req.timings = timings + + output_batch = self.pipeline.forward(req, self.server_args) + duration_ms = (time.monotonic() - start_time) * 1000 + + if output_batch.timings: + output_batch.timings.total_duration_ms = duration_ms + PerformanceLogger.log_request_summary(timings=output_batch.timings) + except Exception as e: + if output_batch is None: + from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import ( + OutputBatch, + ) + + output_batch = OutputBatch() + output_batch.error = f"Error executing request {req.request_id}: {e}" + finally: + return output_batch def set_lora_adapter( self, lora_nickname: str, lora_path: str | None = None diff --git a/python/sglang/multimodal_gen/runtime/managers/scheduler.py b/python/sglang/multimodal_gen/runtime/managers/scheduler.py index e7789a4b5..8b2e33f58 100644 --- a/python/sglang/multimodal_gen/runtime/managers/scheduler.py +++ b/python/sglang/multimodal_gen/runtime/managers/scheduler.py @@ -133,7 +133,7 @@ class Scheduler: # 2: execute, make sure a reply is always sent try: - output_batch = self.worker.execute_forward(reqs, self.server_args) + output_batch = self.worker.execute_forward(reqs) except Exception as e: logger.error( f"Error executing forward in scheduler event loop: {e}", diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py index 3a3ec051a..917af3203 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/executors/pipeline_executor.py @@ -4,7 +4,7 @@ """ Base class for all pipeline executors. """ -import time + from abc import ABC, abstractmethod from typing import List @@ -12,31 +12,19 @@ from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages import PipelineStage from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler logger = init_logger(__name__) -class Timer: +class Timer(StageProfiler): """ - A very simple timer that doesn't for cuda-stream to be synced + A wrapper around StageProfiler to maintain backward compatibility. + It forces simple logging behavior (log start/end) regardless of env vars. """ def __init__(self, name="Stage"): - self.name = name - self.start = None - self.end = None - self.elapsed = None - - def __enter__(self): - self.start = time.perf_counter() - logger.info(f"[{self.name}] started...") - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - self.end = time.perf_counter() - self.elapsed = self.end - self.start - logger.info(f"[{self.name}] finished in {self.elapsed:.4f} seconds") - return False + super().__init__(stage_name=name, timings=None, simple_log=True, logger=logger) class PipelineExecutor(ABC): diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py index ef60d8f5e..566e3e244 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/schedule_batch.py @@ -9,60 +9,26 @@ This module defines the dataclasses used to pass state between pipeline componen in a functional manner, reducing the need for explicit parameter passing. """ +from __future__ import annotations + import pprint from dataclasses import asdict, dataclass, field -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Optional import PIL.Image import torch from sglang.multimodal_gen.configs.sample.base import DataType -from sglang.multimodal_gen.runtime.server_args import ServerArgs -from sglang.multimodal_gen.runtime.utils.performance_logger import PerformanceLogger - -if TYPE_CHECKING: - from torchcodec.decoders import VideoDecoder - -import time -from collections import OrderedDict - from sglang.multimodal_gen.configs.sample.teacache import ( TeaCacheParams, WanTeaCacheParams, ) +from sglang.multimodal_gen.runtime.server_args import ServerArgs +if TYPE_CHECKING: + from torchcodec.decoders import VideoDecoder -class PipelineLoggingInfo: - """Simple approach using OrderedDict to track stage metrics.""" - - def __init__(self): - # OrderedDict preserves insertion order and allows easy access - self.stages: OrderedDict[str, dict[str, Any]] = OrderedDict() - - def add_stage_execution_time(self, stage_name: str, execution_time: float): - """Add execution time for a stage.""" - if stage_name not in self.stages: - self.stages[stage_name] = {} - self.stages[stage_name]["execution_time"] = execution_time - self.stages[stage_name]["timestamp"] = time.time() - - def add_stage_metric(self, stage_name: str, metric_name: str, value: Any): - """Add any metric for a stage.""" - if stage_name not in self.stages: - self.stages[stage_name] = {} - self.stages[stage_name][metric_name] = value - - def get_stage_info(self, stage_name: str) -> dict[str, Any]: - """Get all info for a specific stage.""" - return self.stages.get(stage_name, {}) - - def get_execution_order(self) -> list[str]: - """Get stages in execution order.""" - return list(self.stages.keys()) - - def get_total_execution_time(self) -> float: - """Get total pipeline execution time.""" - return sum(stage.get("execution_time", 0) for stage in self.stages.values()) + from sglang.multimodal_gen.runtime.utils.perf_logger import RequestTimings @dataclass @@ -191,10 +157,9 @@ class Req: # VSA parameters VSA_sparsity: float = 0.0 - perf_logger: PerformanceLogger | None = None # stage logging - logging_info: PipelineLoggingInfo = field(default_factory=PipelineLoggingInfo) + timings: Optional["RequestTimings"] = None # profile profile: bool = False @@ -202,6 +167,8 @@ class Req: # debugging debug: bool = False + # dummy for now + perf_dump_path: str | None = None # results output: torch.Tensor | None = None @@ -230,9 +197,6 @@ class Req: if self.guidance_scale_2 is None: self.guidance_scale_2 = self.guidance_scale - if self.perf_logger is None: - self.perf_logger = PerformanceLogger(self.request_id) - def set_width_and_height(self, server_args: ServerArgs): if self.height is None or self.width is None: width, height = server_args.pipeline_config.adjust_size( @@ -248,10 +212,6 @@ class Req: return pprint.pformat(asdict(self), indent=2, width=120) -@dataclass -class ForwardBatch: ... - - @dataclass class OutputBatch: """ @@ -264,8 +224,8 @@ class OutputBatch: trajectory_decoded: list[torch.Tensor] | None = None error: str | None = None - # Logging info - logging_info: PipelineLoggingInfo = field(default_factory=PipelineLoggingInfo) + # logged timings info, directly from Req.timings + timings: Optional["RequestTimings"] = None @dataclass diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py index 5f66ddaa9..6fef33510 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/base.py @@ -8,20 +8,18 @@ This module defines the abstract base classes for pipeline stages that can be composed to create complete diffusion pipelines. """ -import time -import traceback from abc import ABC, abstractmethod from enum import Enum, auto import torch -import sglang.multimodal_gen.envs as envs from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( VerificationResult, ) from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler logger = init_logger(__name__) @@ -186,54 +184,8 @@ class PipelineStage(ABC): logger.error("Input verification failed for %s: %s", stage_name, str(e)) raise - # Execute the actual stage logic - logging_info = getattr(batch, "logging_info", None) - - if envs.SGLANG_DIFFUSION_STAGE_LOGGING: - logger.info("[%s] Starting execution", stage_name) - start_time = time.perf_counter() - - try: - result = self.forward(batch, server_args) - execution_time = time.perf_counter() - start_time - logger.info( - "[%s] Execution completed in %s ms", - stage_name, - execution_time * 1000, - ) - if logging_info is not None: - try: - logging_info.add_stage_execution_time( - stage_name, execution_time - ) - except Exception: - logger.warning( - "[%s] Failed to record stage timing on batch.logging_info", - stage_name, - exc_info=True, - ) - perf_logger = getattr(batch, "perf_logger", None) - if perf_logger is not None: - try: - perf_logger.log_stage_metric(stage_name, execution_time * 1000) - except Exception: - logger.warning( - "[%s] Failed to log stage metric to performance logger", - stage_name, - exc_info=True, - ) - except Exception as e: - execution_time = time.perf_counter() - start_time - logger.error( - "[%s] Error during execution after %s ms: %s", - stage_name, - execution_time * 1000, - e, - ) - logger.error("[%s] Traceback: %s", stage_name, traceback.format_exc()) - raise - else: - # Direct execution (current behavior) + # Execute the actual stage logic with unified profiling + with StageProfiler(stage_name, logger=logger, timings=batch.timings): result = self.forward(batch, server_args) if enable_verification: diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py index d54e64f6a..5fa5bc16b 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/decoding.py @@ -216,6 +216,7 @@ class DecodingStage(PipelineStage): trajectory_timesteps=batch.trajectory_timesteps, trajectory_latents=batch.trajectory_latents, trajectory_decoded=trajectory_decoded, + timings=batch.timings, ) # Offload models if needed diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index bb0c6c3a5..4c5f32540 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -58,6 +58,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( from sglang.multimodal_gen.runtime.platforms.interface import AttentionBackendEnum from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler from sglang.multimodal_gen.utils import dict_to_3d_list, masks_like try: @@ -804,95 +805,94 @@ class DenoisingStage(PipelineStage): ): with self.progress_bar(total=num_inference_steps) as progress_bar: for i, t_host in enumerate(timesteps_cpu): - if batch.perf_logger: - batch.perf_logger.record_step_start() # Skip if interrupted if hasattr(self, "interrupt") and self.interrupt: continue - t_int = int(t_host.item()) - t_device = timesteps[i] - current_model, current_guidance_scale = ( - self._select_and_manage_model( - t_int=t_int, - boundary_timestep=boundary_timestep, - server_args=server_args, - batch=batch, - ) - ) - - # Expand latents for I2V - latent_model_input = latents.to(target_dtype) - if batch.image_latent is not None: - assert ( - not server_args.pipeline_config.task_type - == ModelTaskType.TI2V - ), "image latents should not be provided for TI2V task" - latent_model_input = torch.cat( - [latent_model_input, batch.image_latent], dim=1 - ).to(target_dtype) - - timestep = self.expand_timestep_before_forward( - batch, - server_args, - t_device, - target_dtype, - seq_len, - reserved_frames_mask, - ) - - latent_model_input = self.scheduler.scale_model_input( - latent_model_input, t_device - ) - - # Predict noise residual - attn_metadata = self._build_attn_metadata(i, batch, server_args) - noise_pred = self._predict_noise_with_cfg( - current_model, - latent_model_input, - timestep, - batch, - i, - attn_metadata, - target_dtype, - current_guidance_scale, - image_kwargs, - pos_cond_kwargs, - neg_cond_kwargs, - server_args, - guidance=guidance, - latents=latents, - ) - - if batch.perf_logger: - batch.perf_logger.record_step_end("denoising_step_guided", i) - # Compute the previous noisy sample - latents = self.scheduler.step( - model_output=noise_pred, - timestep=t_device, - sample=latents, - **extra_step_kwargs, - return_dict=False, - )[0] - - latents = self.post_forward_for_ti2v_task( - batch, server_args, reserved_frames_mask, latents, z - ) - - # save trajectory latents if needed - if batch.return_trajectory_latents: - trajectory_timesteps.append(t_host) - trajectory_latents.append(latents) - - # Update progress bar - if i == num_timesteps - 1 or ( - (i + 1) > num_warmup_steps - and (i + 1) % self.scheduler.order == 0 - and progress_bar is not None + with StageProfiler( + f"denoising_step_{i}", logger=logger, timings=batch.timings ): - progress_bar.update() + t_int = int(t_host.item()) + t_device = timesteps[i] + current_model, current_guidance_scale = ( + self._select_and_manage_model( + t_int=t_int, + boundary_timestep=boundary_timestep, + server_args=server_args, + batch=batch, + ) + ) - self.step_profile() + # Expand latents for I2V + latent_model_input = latents.to(target_dtype) + if batch.image_latent is not None: + assert ( + not server_args.pipeline_config.task_type + == ModelTaskType.TI2V + ), "image latents should not be provided for TI2V task" + latent_model_input = torch.cat( + [latent_model_input, batch.image_latent], dim=1 + ).to(target_dtype) + + timestep = self.expand_timestep_before_forward( + batch, + server_args, + t_device, + target_dtype, + seq_len, + reserved_frames_mask, + ) + + latent_model_input = self.scheduler.scale_model_input( + latent_model_input, t_device + ) + + # Predict noise residual + attn_metadata = self._build_attn_metadata(i, batch, server_args) + noise_pred = self._predict_noise_with_cfg( + current_model, + latent_model_input, + timestep, + batch, + i, + attn_metadata, + target_dtype, + current_guidance_scale, + image_kwargs, + pos_cond_kwargs, + neg_cond_kwargs, + server_args, + guidance=guidance, + latents=latents, + ) + + # Compute the previous noisy sample + latents = self.scheduler.step( + model_output=noise_pred, + timestep=t_device, + sample=latents, + **extra_step_kwargs, + return_dict=False, + )[0] + + latents = self.post_forward_for_ti2v_task( + batch, server_args, reserved_frames_mask, latents, z + ) + + # save trajectory latents if needed + if batch.return_trajectory_latents: + trajectory_timesteps.append(t_host) + trajectory_latents.append(latents) + + # Update progress bar + if i == num_timesteps - 1 or ( + (i + 1) > num_warmup_steps + and (i + 1) % self.scheduler.order == 0 + and progress_bar is not None + ): + progress_bar.update() + + self.step_profile() self.stop_profile(batch) diff --git a/python/sglang/multimodal_gen/runtime/utils/common.py b/python/sglang/multimodal_gen/runtime/utils/common.py index f3dc0d71a..6907756e2 100644 --- a/python/sglang/multimodal_gen/runtime/utils/common.py +++ b/python/sglang/multimodal_gen/runtime/utils/common.py @@ -2,6 +2,7 @@ import importlib import ipaddress +import logging import os import platform import signal @@ -14,9 +15,8 @@ import psutil import torch import zmq -from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger - -logger = init_logger(__name__) +# use the native logger to avoid circular import +logger = logging.getLogger(__name__) def kill_process_tree(parent_pid, include_parent: bool = True, skip_pid: int = None): diff --git a/python/sglang/multimodal_gen/runtime/utils/perf_logger.py b/python/sglang/multimodal_gen/runtime/utils/perf_logger.py new file mode 100644 index 000000000..f840bdf21 --- /dev/null +++ b/python/sglang/multimodal_gen/runtime/utils/perf_logger.py @@ -0,0 +1,265 @@ +# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +import dataclasses +import json +import logging +import os +import subprocess +import sys +import time +import traceback +from datetime import datetime +from functools import lru_cache +from pathlib import Path +from typing import Any, Dict, Optional + +from dateutil.tz import UTC + +import sglang +import sglang.multimodal_gen.envs as envs + + +@dataclasses.dataclass +class RequestTimings: + """A lightweight data class to store performance timings for a single request.""" + + def __init__(self, request_id: str): + self.request_id = request_id + self.stages: Dict[str, float] = {} + self.steps: list[float] = [] + self.total_duration_ms: float = 0.0 + + def record_stage(self, stage_name: str, duration_s: float): + """Records the duration of a pipeline stage""" + self.stages[stage_name] = duration_s * 1000 # Store as milliseconds + + def record_steps(self, index: int, duration_s: float): + """Records the duration of a denoising step""" + assert index == len(self.steps) + self.steps.append(duration_s * 1000) + + def to_dict(self) -> Dict[str, Any]: + """Serializes the timing data to a dictionary.""" + return { + "request_id": self.request_id, + "stages": self.stages, + "steps": self.steps, + "total_duration_ms": self.total_duration_ms, + } + + +def get_diffusion_perf_log_dir() -> str: + """ + Determines the directory for performance logs. + """ + log_dir = os.environ.get("SGLANG_PERF_LOG_DIR") + if log_dir: + return os.path.abspath(log_dir) + if log_dir is None: + sglang_path = Path(sglang.__file__).resolve() + target_path = (sglang_path.parent / "../../.cache/logs").resolve() + return str(target_path) + return "" + + +@lru_cache(maxsize=1) +def get_git_commit_hash() -> str: + try: + commit_hash = os.environ.get("SGLANG_GIT_COMMIT") + if not commit_hash: + commit_hash = ( + subprocess.check_output( + ["git", "rev-parse", "HEAD"], stderr=subprocess.DEVNULL + ) + .strip() + .decode("utf-8") + ) + _CACHED_COMMIT_HASH = commit_hash + return commit_hash + except (subprocess.CalledProcessError, FileNotFoundError): + _CACHED_COMMIT_HASH = "N/A" + return "N/A" + + +@dataclasses.dataclass +class RequestPerfRecord: + request_id: str + + timestamp: str + commit_hash: str + tag: str + + stages: list[dict] + steps: list[float] + total_duration_ms: float + + def __init__( + self, + request_id, + commit_hash, + tag, + stages, + steps, + total_duration_ms, + timestamp=None, + ): + self.request_id = request_id + if timestamp is not None: + self.timestamp = timestamp + else: + self.timestamp = datetime.now(UTC).isoformat() + + self.commit_hash = commit_hash + self.tag = tag + self.stages = stages + self.steps = steps + self.total_duration_ms = total_duration_ms + + +class StageProfiler: + """ + A unified context manager, records timing information (usually of a single Stage or a step) into a provided RequestTimings object (usually from a Req). + """ + + def __init__( + self, + stage_name: str, + logger: logging.Logger, + timings: Optional["RequestTimings"], + simple_log: bool = False, + ): + self.stage_name = stage_name + self.timings = timings + self.logger = logger + self.simple_log = simple_log + self.logger = logging.getLogger(__name__) + self.start_time = 0.0 + + # Check env var at runtime to ensure we pick up changes (e.g. from CLI args) + self.metrics_enabled = envs.SGLANG_DIFFUSION_STAGE_LOGGING + + def __enter__(self): + if self.simple_log: + self.logger.info(f"[{self.stage_name}] started...") + + if (self.metrics_enabled and self.timings) or self.simple_log: + self.start_time = time.perf_counter() + + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + if not ((self.metrics_enabled and self.timings) or self.simple_log): + return False + + execution_time_s = time.perf_counter() - self.start_time + + if exc_type: + self.logger.error( + "[%s] Error during execution after %.4f ms: %s", + self.stage_name, + execution_time_s * 1000, + exc_val, + ) + if self.metrics_enabled: + self.logger.error( + "[%s] Traceback: %s", + self.stage_name, + "".join(traceback.format_tb(exc_tb)), + ) + return False + + if self.simple_log: + self.logger.info( + f"[{self.stage_name}] finished in {execution_time_s:.4f} seconds" + ) + + if self.metrics_enabled and self.timings: + if "denoising_step_" in self.stage_name: + index = int(self.stage_name[len("denoising_step_") :]) + self.timings.record_steps(index, execution_time_s) + else: + self.timings.record_stage(self.stage_name, execution_time_s) + + return False + + +class PerformanceLogger: + """ + A global utility class for logging performance metrics for all request, categorized by request-id. + + Serves both as a runtime logger (stream to file) and a dump utility. + + Notice that ""RequestTimings"" stores the performance metrics of a single request + """ + + @classmethod + def dump_benchmark_report( + cls, + file_path: str, + timings: "RequestTimings", + meta: Optional[Dict[str, Any]] = None, + tag: str = "benchmark_dump", + ): + """ + Static method to dump a standardized benchmark report to a file. + Eliminates duplicate logic in CLI/Client code. + """ + formatted_steps = [ + {"name": name, "duration_ms": duration_ms} + for name, duration_ms in timings.stages.items() + ] + + report = { + "timestamp": datetime.now(UTC).isoformat(), + "request_id": timings.request_id, + "commit_hash": get_git_commit_hash(), + "tag": tag, + "total_duration_ms": timings.total_duration_ms, + "steps": formatted_steps, + "meta": meta or {}, + } + + try: + abs_path = os.path.abspath(file_path) + os.makedirs(os.path.dirname(abs_path), exist_ok=True) + with open(abs_path, "w", encoding="utf-8") as f: + json.dump(report, f, indent=2) + print(f"[Performance] Metrics dumped to: {abs_path}") + except IOError as e: + print(f"[Performance] Failed to dump metrics to {abs_path}: {e}") + logging.getLogger(__name__).error(f"Dump failed: {e}") + + @classmethod + def log_request_summary( + cls, + timings: "RequestTimings", + tag: str = "total_inference_time", + ): + """logs the stage metrics and total duration for a completed request + to the performance_log file. + """ + formatted_stages = [ + {"name": name, "execution_time_ms": duration_ms} + for name, duration_ms in timings.stages.items() + ] + + record = RequestPerfRecord( + timings.request_id, + commit_hash=get_git_commit_hash(), + tag="pipeline_stage_metrics", + stages=formatted_stages, + steps=timings.steps, + total_duration_ms=timings.total_duration_ms, + ) + + try: + log_dir = get_diffusion_perf_log_dir() + if not os.path.exists(log_dir): + os.makedirs(log_dir, exist_ok=True) + + log_file = os.path.join(log_dir, "performance.log") + + with open(log_file, "a", encoding="utf-8") as f: + f.write(json.dumps(dataclasses.asdict(record)) + "\n") + + except (OSError, PermissionError) as e: + print(f"WARNING: Failed to log performance record: {e}", file=sys.stderr) diff --git a/python/sglang/multimodal_gen/runtime/utils/performance_logger.py b/python/sglang/multimodal_gen/runtime/utils/performance_logger.py deleted file mode 100644 index c5c1d99bf..000000000 --- a/python/sglang/multimodal_gen/runtime/utils/performance_logger.py +++ /dev/null @@ -1,204 +0,0 @@ -# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo - -import json -import logging -import os -import subprocess -import time -from datetime import datetime -from pathlib import Path -from typing import Any - -from dateutil.tz import UTC - -import sglang - - -def get_diffusion_perf_log_dir() -> str: - """ - Determines the directory for performance logs, centralizing the logic. - - Resolution order: - 1. SGLANG_PERF_LOG_DIR environment variable, if set and not empty. - 2. Default to ~/.cache/sglang/logs if the environment variable is not set. - 3. Returns an empty string if SGLANG_PERF_LOG_DIR is set to an empty string, - which effectively disables file logging. - """ - log_dir = os.environ.get("SGLANG_PERF_LOG_DIR") - if log_dir: - return os.path.abspath(log_dir) - if log_dir is None: - # Not set, use default - sglang_path = Path(sglang.__file__).resolve() - # .gitignore - target_path = (sglang_path.parent / "../../.cache/logs").resolve() - return str(target_path) - # Is set, but is an empty string - return "" - - -LOG_DIR = get_diffusion_perf_log_dir() - -# Configure a specific logger for performance metrics -perf_logger = logging.getLogger("performance") -perf_logger.setLevel(logging.INFO) -perf_logger.propagate = False # Prevent perf logs from going to the main logger - -_perf_logger_initialized = False - - -class OnDemandFileHandler(logging.Handler): - """ - A logging handler that opens the file for each log record, writes, and closes it. - This is less performant than FileHandler but avoids long-lived file handles, - which can be problematic on certain filesystems like NFS. - """ - - def __init__(self, filename: str, mode: str = "a", encoding: str | None = None): - super().__init__() - self.baseFilename = os.path.abspath(filename) - self.mode = mode - self.encoding = encoding - self.terminator = "\n" - - def emit(self, record: logging.LogRecord): - """Emit a record.""" - try: - msg = self.format(record) - with open( - self.baseFilename, self.mode, encoding=self.encoding, errors="replace" - ) as f: - f.write(msg + self.terminator) - except Exception: - self.handleError(record) - - -def _initialize_perf_logger(): - """Initialize the performance logger with a file handler.""" - global _perf_logger_initialized - if _perf_logger_initialized or not LOG_DIR: - return - - try: - # Ensure the logs directory exists - if not os.path.exists(LOG_DIR): - os.makedirs(LOG_DIR) - - # Set up a file handler for the performance logger - handler = OnDemandFileHandler(os.path.join(LOG_DIR, "performance.log")) - handler.setFormatter(logging.Formatter("%(message)s")) - perf_logger.addHandler(handler) - except (OSError, PermissionError) as e: - perf_logger.warning(f"Failed to initialize performance logger: {e}") - # Disable file logging if initialization fails - globals()["LOG_DIR"] = "" - finally: - _perf_logger_initialized = True - - -def get_git_commit_hash() -> str: - """Get the current git commit hash.""" - try: - commit_hash = ( - subprocess.check_output(["git", "rev-parse", "HEAD"]) - .strip() - .decode("utf-8") - ) - return commit_hash - except (subprocess.CalledProcessError, FileNotFoundError): - return "N/A" - - -class PerformanceLogger: - """ - A utility class for logging performance metrics. - """ - - def __init__(self, request_id: str): - self.request_id = request_id - self.start_time = time.monotonic() - self.step_timings = [] - self.commit_hash = get_git_commit_hash() - - def record_step_start(self): - """Records the start time of a step.""" - self.step_start_time = time.monotonic() - - def record_step_end(self, step_name: str, step_index: int | None = None): - """Records the end time of a step and calculates the duration.""" - duration = time.monotonic() - self.step_start_time - self.step_timings.append( - {"name": step_name, "index": step_index, "duration_ms": duration * 1000} - ) - - def log_total_duration(self, tag: str): - """Logs the total duration of the operation and all recorded steps.""" - _initialize_perf_logger() - total_duration = time.monotonic() - self.start_time - log_entry = { - "timestamp": datetime.now(UTC).isoformat(), - "request_id": self.request_id, - "commit_hash": self.commit_hash, - "tag": tag, - "total_duration_ms": total_duration * 1000, - "steps": self.step_timings, - } - perf_logger.info(json.dumps(log_entry)) - - def log_stage_metric(self, stage_name: str, duration_ms: float): - """Logs a single pipeline stage timing entry.""" - _initialize_perf_logger() - log_entry = { - "timestamp": datetime.now(UTC).isoformat(), - "request_id": self.request_id, - "commit_hash": self.commit_hash, - "tag": "pipeline_stage_metric", - "stage": stage_name, - "duration_ms": duration_ms, - } - perf_logger.info(json.dumps(log_entry)) - - def log_stage_metrics(self, stages: Any): - """ - Persist per-stage execution stats to performance.log. - - Args: - stages: Either a PipelineLoggingInfo instance or any object exposing - a mapping of stage metadata via a `stages` attribute/dict. - """ - _initialize_perf_logger() - if stages is None: - return - - if hasattr(stages, "stages"): - stage_items = getattr(stages, "stages", {}).items() - elif isinstance(stages, dict): - stage_items = stages.items() - else: - return - - formatted_stages: list[dict[str, Any]] = [] - for name, info in stage_items: - if not info: - continue - entry = {"name": name} - execution_time = info.get("execution_time") - if execution_time is not None: - entry["execution_time_ms"] = execution_time * 1000 - for key, value in info.items(): - if key == "execution_time": - continue - entry[key] = value - formatted_stages.append(entry) - - if not formatted_stages: - return - - log_entry = { - "timestamp": datetime.now(UTC).isoformat(), - "request_id": self.request_id, - "commit_hash": self.commit_hash, - "tag": "pipeline_stage_metrics", - "stages": formatted_stages, - } - perf_logger.info(json.dumps(log_entry)) diff --git a/python/sglang/multimodal_gen/test/server/perf_baselines.json b/python/sglang/multimodal_gen/test/server/perf_baselines.json index 07a1e745d..aab4f6bfc 100644 --- a/python/sglang/multimodal_gen/test/server/perf_baselines.json +++ b/python/sglang/multimodal_gen/test/server/perf_baselines.json @@ -101,65 +101,65 @@ "InputValidationStage": 0.03, "TextEncodingStage": 81.49, "ConditioningStage": 0.01, - "TimestepPreparationStage": 2.32, + "TimestepPreparationStage": 2.43, "LatentPreparationStage": 6.29, "DenoisingStage": 8381.3, "DecodingStage": 653.03 }, "denoise_step_ms": { - "0": 57.18, - "1": 58.71, - "2": 166.91, - "3": 167.39, - "4": 166.7, - "5": 172.02, - "6": 179.04, - "7": 168.91, - "8": 165.77, - "9": 166.79, - "10": 167.45, - "11": 171.15, - "12": 171.31, - "13": 169.56, - "14": 169.67, - "15": 166.97, - "16": 169.15, - "17": 169.68, - "18": 170.1, - "19": 169.59, - "20": 168.52, - "21": 167.19, - "22": 169.36, - "23": 172.21, - "24": 171.8, - "25": 169.29, - "26": 169.67, - "27": 169.19, - "28": 169.46, - "29": 171.16, - "30": 170.98, - "31": 169.38, - "32": 168.55, - "33": 169.64, - "34": 168.16, - "35": 170.85, - "36": 170.21, - "37": 168.42, - "38": 168.17, - "39": 168.25, - "40": 167.47, - "41": 169.53, - "42": 171.65, - "43": 169.1, - "44": 172.15, - "45": 171.81, - "46": 171.26, - "47": 167.78, - "48": 168.44, - "49": 168.31 + "0": 165.27, + "1": 58.88, + "2": 166.85, + "3": 166.51, + "4": 166.77, + "5": 167.55, + "6": 172.4, + "7": 167.77, + "8": 167.51, + "9": 167.22, + "10": 168.19, + "11": 167.74, + "12": 168.48, + "13": 168.08, + "14": 168.16, + "15": 167.15, + "16": 167.05, + "17": 169.27, + "18": 167.96, + "19": 167.74, + "20": 168.21, + "21": 167.07, + "22": 167.35, + "23": 167.06, + "24": 169.28, + "25": 169.41, + "26": 168.92, + "27": 167.59, + "28": 167.57, + "29": 170.42, + "30": 166.24, + "31": 168.33, + "32": 168.56, + "33": 168.62, + "34": 167.28, + "35": 167.12, + "36": 168.21, + "37": 168.78, + "38": 168.89, + "39": 167.74, + "40": 168.57, + "41": 167.89, + "42": 168.03, + "43": 167.61, + "44": 167.75, + "45": 168.03, + "46": 168.81, + "47": 168.29, + "48": 168.64, + "49": 168.78 }, "expected_e2e_ms": 9275.51, - "expected_avg_denoise_ms": 165.04, + "expected_avg_denoise_ms": 165.83, "expected_median_denoise_ms": 169.33 }, "qwen_image_edit_ti2i": { @@ -486,70 +486,70 @@ }, "wan2_2_ti2v_5b": { "stages_ms": { - "InputValidationStage": 79.9, - "TextEncodingStage": 2241.88, - "ConditioningStage": 0.01, - "TimestepPreparationStage": 1.63, + "InputValidationStage": 94.18, + "TextEncodingStage": 3413.73, + "ConditioningStage": 0.02, + "TimestepPreparationStage": 2.48, "LatentPreparationStage": 22.52, - "DenoisingStage": 130094.11, + "DenoisingStage": 30000, "DecodingStage": 12793.58, "per_frame_generation": null }, "denoise_step_ms": { - "0": 3092.88, - "1": 2526.43, - "2": 2544.76, - "3": 2543.33, - "4": 2545.33, - "5": 2542.44, - "6": 2540.33, - "7": 2542.23, - "8": 2544.87, - "9": 2547.86, - "10": 2548.18, - "11": 2551.13, - "12": 2547.02, - "13": 2551.31, - "14": 2551.97, - "15": 2549.61, - "16": 2551.75, - "17": 2552.97, - "18": 2551.2, - "19": 2555.07, - "20": 2552.72, - "21": 2551.24, - "22": 2554.63, - "23": 2555.52, - "24": 2555.06, - "25": 2550.04, - "26": 2554.88, - "27": 2553.69, - "28": 2550.75, - "29": 2555.17, - "30": 2556.75, - "31": 2554.22, - "32": 2552.74, - "33": 2554.31, - "34": 2554.98, - "35": 2553.65, - "36": 2552.21, - "37": 2554.85, - "38": 2555.96, - "39": 2553.78, - "40": 2553.5, - "41": 2550.98, - "42": 2555.66, - "43": 2551.91, - "44": 2551.23, - "45": 2555.91, - "46": 2556.11, - "47": 2548.55, - "48": 2552.78, - "49": 2553.49 + "0": 1021.97, + "1": 407.7, + "2": 410.44, + "3": 411.69, + "4": 411.45, + "5": 410.84, + "6": 411.26, + "7": 412.68, + "8": 412.24, + "9": 410.15, + "10": 413.4, + "11": 410.18, + "12": 411.94, + "13": 411.54, + "14": 409.91, + "15": 412.89, + "16": 412.08, + "17": 411.64, + "18": 411.58, + "19": 410.54, + "20": 411.42, + "21": 412.88, + "22": 412.22, + "23": 412.97, + "24": 412.13, + "25": 413.21, + "26": 413.07, + "27": 410.89, + "28": 411.56, + "29": 414.19, + "30": 412.68, + "31": 411.13, + "32": 412.41, + "33": 412.79, + "34": 411.02, + "35": 410.2, + "36": 410.27, + "37": 411.63, + "38": 410.8, + "39": 411.52, + "40": 411.6, + "41": 411.75, + "42": 410.56, + "43": 411.65, + "44": 411.82, + "45": 410.83, + "46": 410.61, + "47": 411.5, + "48": 410.61, + "49": 411.6 }, - "expected_e2e_ms": 145253.72, - "expected_avg_denoise_ms": 2561.76, - "expected_median_denoise_ms": 2552.46 + "expected_e2e_ms": 32954.73, + "expected_avg_denoise_ms": 423.75, + "expected_median_denoise_ms": 411.59 } } } diff --git a/python/sglang/multimodal_gen/test/server/test_server_performance.py b/python/sglang/multimodal_gen/test/server/test_server_performance.py index 7f44ef0f3..585e010a3 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_performance.py +++ b/python/sglang/multimodal_gen/test/server/test_server_performance.py @@ -16,6 +16,7 @@ import pytest from openai import OpenAI from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.perf_logger import RequestPerfRecord from sglang.multimodal_gen.test.server.conftest import _GLOBAL_PERF_RESULTS from sglang.multimodal_gen.test.server.test_server_utils import ( VALIDATOR_REGISTRY, @@ -34,11 +35,10 @@ from sglang.multimodal_gen.test.server.testcase_configs import ( ) from sglang.multimodal_gen.test.test_utils import ( get_dynamic_server_port, - read_perf_records, + read_perf_logs, validate_image, validate_openai_video, - wait_for_perf_record, - wait_for_stage_metrics, + wait_for_req_perf_record, ) logger = init_logger(__name__) @@ -140,42 +140,32 @@ Consider updating perf_baselines.json with the snippets below: base_url=f"http://localhost:{ctx.port}/v1", ) - def _run_and_collect( + def run_and_collect( self, ctx: ServerContext, - case: DiffusionTestCase, - generate_fn: Callable[[], None], - ) -> tuple[dict, dict]: + generate_fn: Callable[[], str], + ) -> RequestPerfRecord: """Run generation and collect performance records.""" log_path = ctx.perf_log_path - prev_len = len(read_perf_records(log_path)) + prev_len = len(read_perf_logs(log_path)) log_wait_timeout = 1200 - generate_fn() + rid = generate_fn() - perf_record, _ = wait_for_perf_record( - "total_inference_time", + req_perf_record, _ = wait_for_req_perf_record( + rid, prev_len, log_path, timeout=log_wait_timeout, ) - stage_metrics = {} - if perf_record: - stage_metrics, _ = wait_for_stage_metrics( - perf_record.get("request_id", ""), - prev_len, - log_path, - timeout=log_wait_timeout, - ) + return req_perf_record - return perf_record, stage_metrics - - def _generate_for_case( + def get_generate_fn( self, ctx: ServerContext, case: DiffusionTestCase, - ) -> Callable[[], None]: + ) -> Callable[[], str]: """Return appropriate generation function for the case.""" client = self._client(ctx) @@ -186,7 +176,7 @@ Consider updating perf_baselines.json with the snippets below: prompt: str | None = None, seconds: int | None = None, input_reference: Any | None = None, - ) -> bytes: + ) -> str: """ Create a video job via /v1/videos, poll until completion, then download the binary content and validate it. @@ -231,7 +221,7 @@ Consider updating perf_baselines.json with the snippets below: f"{case.id}: video job {video_id} timed out during baseline generation. " "Attempting to collect performance data anyway." ) - return b"" + return video_id pytest.fail(f"{case.id}: video job {video_id} did not complete in time") @@ -239,7 +229,7 @@ Consider updating perf_baselines.json with the snippets below: resp = client.videos.download_content(video_id=video_id) # type: ignore[attr-defined] content = resp.read() validate_openai_video(content) - return content + return video_id # for all tests, seconds = case.seconds or fallback 4 seconds video_seconds = case.seconds or 4 @@ -248,20 +238,23 @@ Consider updating perf_baselines.json with the snippets below: # IMAGE MODE # ------------------------- - def generate_image(): + def generate_image() -> str: """T2I: Text to Image generation.""" if not case.prompt: pytest.skip(f"{case.id}: no text prompt configured") - result = client.images.generate( + + response = client.images.with_raw_response.generate( model=case.model_path, prompt=case.prompt, n=1, size=case.output_size, response_format="b64_json", ) + result = response.parse() validate_image(result.data[0].b64_json) + return str(result.created) - def generate_image_edit(): + def generate_image_edit() -> str: """TI2I: Text + Image ? Image edit.""" if not case.edit_prompt or not case.image_path: pytest.skip(f"{case.id}: no edit config") @@ -275,7 +268,7 @@ Consider updating perf_baselines.json with the snippets below: pytest.skip(f"{case.id}: file missing: {image_path}") with image_path.open("rb") as fh: - result = client.images.edit( + response = client.images.with_raw_response.edit( model=case.model_path, image=fh, prompt=case.edit_prompt, @@ -283,25 +276,30 @@ Consider updating perf_baselines.json with the snippets below: size=case.output_size, response_format="b64_json", ) + rid = response.headers.get("x-request-id", "") + print(f"{response=}") + + result = response.parse() validate_image(result.data[0].b64_json) + return rid # ------------------------- # VIDEO MODE # ------------------------- - def generate_video(): + def generate_video() -> str: """T2V: Text ? Video.""" if not case.prompt: pytest.skip(f"{case.id}: no text prompt configured") - _create_and_download_video( + return _create_and_download_video( model=case.model_path, prompt=case.prompt, size=case.output_size, seconds=video_seconds, ) - def generate_image_to_video(): + def generate_image_to_video() -> str: """I2V: Image ? Video (optional prompt).""" if not case.image_path: pytest.skip(f"{case.id}: no input image configured") @@ -315,7 +313,7 @@ Consider updating perf_baselines.json with the snippets below: pytest.skip(f"{case.id}: file missing: {image_path}") with image_path.open("rb") as fh: - _create_and_download_video( + return _create_and_download_video( model=case.model_path, prompt=case.edit_prompt, size=case.output_size, @@ -323,7 +321,7 @@ Consider updating perf_baselines.json with the snippets below: input_reference=fh, ) - def generate_text_image_to_video(): + def generate_text_image_to_video() -> str: """TI2V: Text + Image ? Video.""" if not case.edit_prompt or not case.image_path: pytest.skip(f"{case.id}: no edit config") @@ -337,7 +335,7 @@ Consider updating perf_baselines.json with the snippets below: pytest.skip(f"{case.id}: file missing: {image_path}") with image_path.open("rb") as fh: - _create_and_download_video( + return _create_and_download_video( model=case.model_path, prompt=case.edit_prompt, size=case.output_size, @@ -362,8 +360,7 @@ Consider updating perf_baselines.json with the snippets below: def _validate_and_record( self, case: DiffusionTestCase, - perf_record: dict, - stage_metrics: dict, + perf_record: RequestPerfRecord, ) -> None: """Validate metrics and record results.""" is_baseline_generation_mode = os.environ.get("SGLANG_GEN_BASELINE", "0") == "1" @@ -395,7 +392,7 @@ Consider updating perf_baselines.json with the snippets below: step_fractions=BASELINE_CONFIG.step_fractions, ) - summary = validator.collect_metrics(perf_record, stage_metrics) + summary = validator.collect_metrics(perf_record) if is_baseline_generation_mode or missing_scenario: self._dump_baseline_for_testcase(case, summary) @@ -406,30 +403,12 @@ Consider updating perf_baselines.json with the snippets below: self._check_for_improvement(case, summary, scenario) try: - validator.validate(perf_record, stage_metrics, case.num_frames) + validator.validate(perf_record, case.num_frames) except AssertionError as e: logger.error(f"Performance validation failed for {case.id}:\n{e}") self._dump_baseline_for_testcase(case, summary) raise - if case.modality == "video" and summary.frames_per_second: - logger.info( - "[Perf] %s: E2E %.2f ms; Avg %.2f ms; FPS %.2f; Frames %d", - case.id, - summary.e2e_ms, - summary.avg_denoise_ms, - summary.frames_per_second, - summary.total_frames or 0, - ) - else: - logger.info( - "[Perf] %s: E2E %.2f ms; Avg %.2f ms; Median %.2f ms", - case.id, - summary.e2e_ms, - summary.avg_denoise_ms, - summary.median_denoise_ms, - ) - result = { "test_name": case.id, "modality": case.modality, @@ -452,40 +431,6 @@ Consider updating perf_baselines.json with the snippets below: self.__class__._perf_results.append(result) - logger.info("[BASELINE] %s expected_e2e_ms = %.2f", case.id, summary.e2e_ms) - logger.info( - "[BASELINE] %s expected_avg_denoise_ms = %.2f", - case.id, - summary.avg_denoise_ms, - ) - logger.info( - "[BASELINE] %s expected_median_denoise_ms = %.2f", - case.id, - summary.median_denoise_ms, - ) - logger.info("[BASELINE] %s stages_ms = %r", case.id, summary.stage_metrics) - logger.info( - "[BASELINE] %s denoise_step_ms = %r", case.id, summary.sampled_steps - ) - - # Only log video-specific metrics when they exist - if summary.frames_per_second is not None: - logger.info( - "[BASELINE] %s frames_per_second = %.2f", - case.id, - summary.frames_per_second, - ) - if summary.total_frames is not None: - logger.info( - "[BASELINE] %s total_frames = %d", case.id, summary.total_frames - ) - if summary.avg_frame_time_ms is not None: - logger.info( - "[BASELINE] %s avg_frame_time_ms = %.2f", - case.id, - summary.avg_frame_time_ms, - ) - def _check_for_improvement( self, case: DiffusionTestCase, @@ -514,7 +459,6 @@ Consider updating perf_baselines.json with the snippets below: ) ): is_improved = True - # Combine metrics, always taking the better (lower) value new_stages = { stage: min( @@ -528,7 +472,8 @@ Consider updating perf_baselines.json with the snippets below: safe_get_metric(summary.all_denoise_steps, step), safe_get_metric(scenario.denoise_step_ms, step), ) - for step in set(summary.all_denoise_steps) | set(scenario.denoise_step_ms) + for step in set(summary.all_denoise_steps.keys()) + | set(scenario.denoise_step_ms) } # Check for stage-level improvements @@ -614,10 +559,9 @@ the "scenarios" section of perf_baselines.json: - test_diffusion_perf[qwen_image_edit] - etc. """ - generate_fn = self._generate_for_case(diffusion_server, case) - perf_record, stage_metrics = self._run_and_collect( + generate_fn = self.get_generate_fn(diffusion_server, case) + perf_record = self.run_and_collect( diffusion_server, - case, generate_fn, ) - self._validate_and_record(case, perf_record, stage_metrics) + self._validate_and_record(case, perf_record) diff --git a/python/sglang/multimodal_gen/test/server/test_server_utils.py b/python/sglang/multimodal_gen/test/server/test_server_utils.py index d9aee5f9e..c5d33dc3c 100644 --- a/python/sglang/multimodal_gen/test/server/test_server_utils.py +++ b/python/sglang/multimodal_gen/test/server/test_server_utils.py @@ -5,7 +5,6 @@ Server management and performance validation for diffusion tests. from __future__ import annotations import os -import statistics import subprocess import sys import tempfile @@ -18,18 +17,16 @@ from urllib.request import urlopen from openai import OpenAI +from sglang.multimodal_gen.benchmarks.compare_perf import calculate_upper_bound from sglang.multimodal_gen.runtime.utils.common import kill_process_tree from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.perf_logger import RequestPerfRecord from sglang.multimodal_gen.test.server.testcase_configs import ( PerformanceSummary, ScenarioConfig, ToleranceConfig, ) -from sglang.multimodal_gen.test.test_utils import ( - prepare_perf_log, - sample_step_indices, - validate_image, -) +from sglang.multimodal_gen.test.test_utils import prepare_perf_log, validate_image logger = init_logger(__name__) @@ -319,9 +316,7 @@ class PerformanceValidator: Uses the larger of relative tolerance or absolute tolerance to prevent flaky failures on very fast operations. """ - rel_limit = expected * (1 + tolerance) - abs_limit = expected + min_abs_tolerance_ms - upper_bound = max(rel_limit, abs_limit) + upper_bound = calculate_upper_bound(expected, tolerance, min_abs_tolerance_ms) assert actual <= upper_bound, ( f"Validation failed for '{name}'.\n" f" Actual: {actual:.4f}ms\n" @@ -331,10 +326,10 @@ class PerformanceValidator: ) def validate( - self, perf_record: dict, stage_metrics: dict, *args, **kwargs + self, perf_record: RequestPerfRecord, *args, **kwargs ) -> PerformanceSummary: """Validate all performance metrics and return summary.""" - summary = self.collect_metrics(perf_record, stage_metrics) + summary = self.collect_metrics(perf_record) if self.is_baseline_generation_mode: return summary @@ -347,40 +342,9 @@ class PerformanceValidator: def collect_metrics( self, - perf_record: dict, - stage_metrics: dict, + perf_record: RequestPerfRecord, ) -> PerformanceSummary: - """Collect all performance metrics into a summary without validation.""" - e2e_ms = float(perf_record.get("total_duration_ms", 0.0)) - steps = [ - s - for s in perf_record.get("steps", []) or [] - if s.get("name") == "denoising_step_guided" and "duration_ms" in s - ] - - avg_denoise = 0.0 - median_denoise = 0.0 - if steps: - durations = [float(s["duration_ms"]) for s in steps] - avg_denoise = sum(durations) / len(durations) - median_denoise = statistics.median(durations) - - per_step = { - int(s["index"]): float(s["duration_ms"]) - for s in steps - if s.get("index") is not None - } - sample_indices = sample_step_indices(per_step, self.step_fractions) - sampled_steps = {idx: per_step[idx] for idx in sample_indices} - - return PerformanceSummary( - e2e_ms=e2e_ms, - avg_denoise_ms=avg_denoise, - median_denoise_ms=median_denoise, - stage_metrics=stage_metrics, - sampled_steps=sampled_steps, - all_denoise_steps=per_step, - ) + return PerformanceSummary.from_req_perf_record(perf_record, self.step_fractions) def _validate_e2e(self, summary: PerformanceSummary) -> None: """Validate end-to-end performance.""" @@ -455,12 +419,11 @@ class VideoPerformanceValidator(PerformanceValidator): def validate( self, - perf_record: dict, - stage_metrics: dict, + perf_record: RequestPerfRecord, num_frames: int | None = None, ) -> PerformanceSummary: """Validate video metrics including frame generation rates.""" - summary = super().validate(perf_record, stage_metrics) + summary = super().validate(perf_record) if num_frames and summary.e2e_ms > 0: summary.total_frames = num_frames diff --git a/python/sglang/multimodal_gen/test/server/testcase_configs.py b/python/sglang/multimodal_gen/test/server/testcase_configs.py index 3cc5b37ab..fa8f281fd 100644 --- a/python/sglang/multimodal_gen/test/server/testcase_configs.py +++ b/python/sglang/multimodal_gen/test/server/testcase_configs.py @@ -20,10 +20,13 @@ from __future__ import annotations import json import os +import statistics from dataclasses import dataclass from pathlib import Path from typing import Sequence +from sglang.multimodal_gen.runtime.utils.perf_logger import RequestPerfRecord + @dataclass class ToleranceConfig: @@ -137,20 +140,71 @@ class DiffusionTestCase: ) +def sample_step_indices( + step_map: dict[int, float], fractions: Sequence[float] +) -> list[int]: + if not step_map: + return [] + max_idx = max(step_map.keys()) + indices = set() + for fraction in fractions: + idx = min(max_idx, max(0, int(round(fraction * max_idx)))) + if idx in step_map: + indices.add(idx) + return sorted(indices) + + @dataclass class PerformanceSummary: - """Summary of performance metrics.""" + """Summary of performance of a request, built from RequestPerfRecord""" e2e_ms: float avg_denoise_ms: float median_denoise_ms: float + # { "stage_1": time_1, "stage_2": time_2 } stage_metrics: dict[str, float] + step_metrics: list[float] sampled_steps: dict[int, float] all_denoise_steps: dict[int, float] frames_per_second: float | None = None total_frames: int | None = None avg_frame_time_ms: float | None = None + @staticmethod + def from_req_perf_record( + record: RequestPerfRecord, step_fractions: Sequence[float] + ): + """Collect all performance metrics into a summary without validation.""" + e2e_ms = record.total_duration_ms + + step_durations = record.steps + avg_denoise = 0.0 + median_denoise = 0.0 + if step_durations: + avg_denoise = sum(step_durations) / len(step_durations) + median_denoise = statistics.median(step_durations) + + per_step = {index: s for index, s in enumerate(step_durations)} + sample_indices = sample_step_indices(per_step, step_fractions) + sampled_steps = {idx: per_step[idx] for idx in sample_indices} + + # convert from list to dict + stage_metrics = {} + for item in record.stages: + if isinstance(item, dict) and "name" in item: + val = item.get("execution_time_ms", 0.0) + stage_metrics[item["name"]] = val + + return PerformanceSummary( + e2e_ms=e2e_ms, + avg_denoise_ms=avg_denoise, + median_denoise_ms=median_denoise, + stage_metrics=stage_metrics, + step_metrics=step_durations, + sampled_steps=sampled_steps, + all_denoise_steps=per_step, + ) + # Common paths IMAGE_INPUT_FILE = Path(__file__).resolve().parents[1] / "test_files" / "girl.jpg" diff --git a/python/sglang/multimodal_gen/test/test_utils.py b/python/sglang/multimodal_gen/test/test_utils.py index 6d81203ad..457406892 100644 --- a/python/sglang/multimodal_gen/test/test_utils.py +++ b/python/sglang/multimodal_gen/test/test_utils.py @@ -1,459 +1,415 @@ -# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo -import base64 -import dataclasses -import json -import os -import shlex -import socket -import subprocess -import sys -import time -import unittest -from pathlib import Path -from typing import Optional, Sequence - -from PIL import Image - -from sglang.multimodal_gen.configs.sample.base import DataType -from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var -from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger -from sglang.multimodal_gen.runtime.utils.performance_logger import ( - get_diffusion_perf_log_dir, -) - -logger = init_logger(__name__) - - -def run_command(command) -> Optional[float]: - """Runs a command and returns the execution time and status.""" - print(f"Running command: {shlex.join(command)}") - - duration = None - with subprocess.Popen( - command, - stdout=subprocess.PIPE, - stderr=subprocess.STDOUT, - text=True, - encoding="utf-8", - ) as process: - for line in process.stdout: - sys.stdout.write(line) - if "Pixel data generated" in line: - words = line.split(" ") - duration = float(words[-2]) - - if process.returncode == 0: - return duration - else: - print(f"Command failed with exit code {process.returncode}") - return None - - -def probe_port(host="127.0.0.1", port=30010, timeout=2.0) -> bool: - with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: - s.settimeout(timeout) - try: - s.connect((host, port)) - return True - except OSError: - return False - - -def is_in_ci() -> bool: - return get_bool_env_var("SGLANG_IS_IN_CI") - - -def get_dynamic_server_port() -> int: - cuda_devices = os.environ.get("CUDA_VISIBLE_DEVICES", "0") - if not cuda_devices: - cuda_devices = "0" - try: - first_device_id = int(cuda_devices.split(",")[0].strip()[0]) - except (ValueError, IndexError): - first_device_id = 0 - - if is_in_ci(): - base_port = 10000 + first_device_id * 2000 - else: - base_port = 20000 + first_device_id * 1000 - - return base_port + 1000 - - -def is_mp4(data): - idx = data.find(b"ftyp") - return 0 <= idx <= 32 - - -def is_jpeg(data: bytes) -> bool: - # JPEG files start with: FF D8 FF - return data.startswith(b"\xff\xd8\xff") - - -def is_png(data): - # PNG files start with: 89 50 4E 47 0D 0A 1A 0A - return data.startswith(b"\x89PNG\r\n\x1a\n") - - -def wait_for_port(host="127.0.0.1", port=30010, deadline=300.0, interval=0.5): - end = time.time() + deadline - last_err = None - while time.time() < end: - if probe_port(host, port, timeout=interval): - return True - time.sleep(interval) - raise TimeoutError(f"Port {host}:{port} not ready. Last error: {last_err}") - - -def check_image_size(ut, image, width, height): - # check image size - ut.assertEqual(image.size, (width, height)) - - -def get_perf_log_dir() -> Path: - """Gets the performance log directory from the centralized sglang utility.""" - log_dir_str = get_diffusion_perf_log_dir() - if not log_dir_str: - raise RuntimeError( - "Performance logging is disabled (SGLANG_PERF_LOG_DIR is empty), " - "but a test tried to access the log directory." - ) - return Path(log_dir_str) - - -def _ensure_log_path(log_dir: Path) -> Path: - log_dir.mkdir(parents=True, exist_ok=True) - return log_dir / "performance.log" - - -def clear_perf_log(log_dir: Path) -> Path: - """Delete the perf log file so tests can watch for fresh entries.""" - log_path = _ensure_log_path(log_dir) - if log_path.exists(): - log_path.unlink() - logger.info("[server-test] Monitoring perf log at %s", log_path.as_posix()) - return log_path - - -def prepare_perf_log() -> tuple[Path, Path]: - """Convenience helper to resolve and clear the perf log in one call.""" - log_dir = get_perf_log_dir() - log_path = clear_perf_log(log_dir) - return log_dir, log_path - - -def read_perf_records(log_path: Path) -> list[dict]: - if not log_path.exists(): - return [] - records: list[dict] = [] - with log_path.open("r", encoding="utf-8") as fh: - for line in fh: - line = line.strip() - if not line: - continue - try: - records.append(json.loads(line)) - except json.JSONDecodeError: - continue - return records - - -def wait_for_perf_record( - tag: str, - prev_len: int, - log_path: Path, - timeout: float = 120.0, -) -> tuple[dict, int]: - deadline = time.time() + timeout - while time.time() < deadline: - records = read_perf_records(log_path) - if len(records) > prev_len: - for rec in records[prev_len:]: - if rec.get("tag") == tag: - return rec, len(records) - time.sleep(0.5) - - if os.environ.get("SGLANG_GEN_BASELINE", "0") == "1": - records = read_perf_records(log_path) - return {}, len(records) - - raise AssertionError( - f"Timeout waiting for perf log entry '{tag}' (start_len={prev_len})" - ) - - -def wait_for_stage_metrics( - request_id: str, - prev_len: int, - log_path: Path, - timeout: float = 300.0, -) -> tuple[dict[str, float], int]: - deadline = time.time() + timeout - metrics: dict[str, float] = {} - while time.time() < deadline: - records = read_perf_records(log_path) - for rec in records[prev_len:]: - # Check if the request is completed - if ( - rec.get("tag") == "total_inference_time" - and rec.get("request_id") == request_id - ): - return metrics, len(records) - - if ( - rec.get("tag") == "pipeline_stage_metric" - and rec.get("request_id") == request_id - ): - stage = rec.get("stage") - duration = rec.get("duration_ms") - if stage is not None and duration is not None: - metrics[str(stage)] = float(duration) - time.sleep(0.5) - - if os.environ.get("SGLANG_GEN_BASELINE", "0") == "1": - records = read_perf_records(log_path) - return {}, len(records) - raise AssertionError(f"Timeout waiting for stage metrics for request {request_id} ") - - -def sample_step_indices( - step_map: dict[int, float], fractions: Sequence[float] -) -> list[int]: - if not step_map: - return [] - max_idx = max(step_map.keys()) - indices = set() - for fraction in fractions: - idx = min(max_idx, max(0, int(round(fraction * max_idx)))) - if idx in step_map: - indices.add(idx) - return sorted(indices) - - -def validate_image(b64_json: str) -> None: - """Decode and validate that image is PNG or JPEG.""" - image_bytes = base64.b64decode(b64_json) - assert is_png(image_bytes) or is_jpeg(image_bytes), "Image must be PNG or JPEG" - - -def validate_video(b64_json: str) -> None: - """Decode and validate that video is a valid format.""" - video_bytes = base64.b64decode(b64_json) - is_mp4 = ( - video_bytes[:4] == b"\x00\x00\x00\x18" or video_bytes[:4] == b"\x00\x00\x00\x1c" - ) - is_webm = video_bytes[:4] == b"\x1a\x45\xdf\xa3" - assert is_mp4 or is_webm, "Video must be MP4 or WebM" - - -def validate_openai_video(video_bytes: bytes) -> None: - """Validate that video is MP4 or WebM by magic bytes.""" - is_mp4 = ( - video_bytes.startswith(b"\x00\x00\x00\x18") - or video_bytes.startswith(b"\x00\x00\x00\x1c") - or video_bytes[4:8] == b"ftyp" - ) - is_webm = video_bytes.startswith(b"\x1a\x45\xdf\xa3") - assert is_mp4 or is_webm, "Video must be MP4 or WebM" - - -@dataclasses.dataclass -class TestResult: - name: str - key: str - duration: Optional[float] - succeed: bool - - @property - def duration_str(self): - return f"{self.duration:.4f}" if self.duration else "NA" - - -class TestCLIBase(unittest.TestCase): - model_path: str = None - extra_args = [] - data_type: DataType = None - # tested on h100 - thresholds = {} - - width: int = 720 - height: int = 720 - output_path: str = "test_outputs" - - base_command = [ - "sglang", - "generate", - "--text-encoder-cpu-offload", - "--pin-cpu-memory", - "--prompt", - "A curious raccoon", - "--save-output", - "--log-level=debug", - f"--width={width}", - f"--height={height}", - f"--output-path={output_path}", - ] - - results = [] - - @classmethod - def setUpClass(cls): - cls.results = [] - - def _run_command(self, name: str, model_path: str, test_key: str = "", args=[]): - command = ( - self.base_command - + [f"--model-path={model_path}"] - + shlex.split(args or "") - + ["--output-file-name", f"{name}"] - + self.extra_args - ) - duration = run_command(command) - status = "Success" if duration else "Failed" - succeed = duration is not None - - duration = float(duration) if succeed else None - self.results.append(TestResult(name, test_key, duration, succeed)) - - return name, duration, status - - -class TestGenerateBase(TestCLIBase): - model_path: str = None - extra_args = [] - data_type: DataType = None - # tested on h100 - thresholds = {} - - width: int = 720 - height: int = 720 - output_path: str = "test_outputs" - image_path: str | None = None - prompt: str | None = "A curious raccoon" - - base_command = [ - "sglang", - "generate", - # "--text-encoder-cpu-offload", - # "--pin-cpu-memory", - f"--prompt", - f"{prompt}", - "--save-output", - "--log-level=debug", - f"--width={width}", - f"--height={height}", - f"--output-path={output_path}", - ] - - results: list[TestResult] = [] - - @classmethod - def setUpClass(cls): - cls.results = [] - - @classmethod - def tearDownClass(cls): - # Print markdown table - print("\n## Test Results\n") - print("| Test Case | Duration | Status |") - print("|--------------------------------|----------|---------|") - test_keys = ["test_single_gpu", "test_cfg_parallel", "test_usp", "test_mixed"] - test_key_to_order = { - test_key: order for order, test_key in enumerate(test_keys) - } - - ordered_results: list[TestResult] = [None] * len(test_keys) - for result in cls.results: - order = test_key_to_order[result.key] - ordered_results[order] = result - - for result in ordered_results: - if not result: - continue - status = ( - "Succeed" - if ( - result.succeed - and float(result.duration) <= float(cls.thresholds[result.key]) - ) - else "Failed" - ) - print(f"| {result.name:<30} | {result.duration_str:<8} | {status:<7} |") - print() - durations = [result.duration_str for result in cls.results] - print(" | ".join([""] + durations + [""])) - - def _run_test(self, name: str, args, model_path: str, test_key: str): - time_threshold = self.thresholds[test_key] - name, duration, status = self._run_command( - name, args=args, model_path=model_path, test_key=test_key - ) - self.verify(status, name, duration, time_threshold) - - def verify(self, status, name, duration, time_threshold): - print("-" * 80) - print("\n" * 3) - - # test task status - self.assertEqual(status, "Success", f"{name} command failed") - self.assertIsNotNone(duration, f"Could not parse duration for {name}") - self.assertLessEqual( - duration, - time_threshold, - f"{name} failed with {duration:.4f}s > {time_threshold}s", - ) - - # test output file - path = os.path.join( - self.output_path, f"{name}.{self.data_type.get_default_extension()}" - ) - self.assertTrue(os.path.exists(path), f"Output file not exist for {path}") - if self.data_type == DataType.IMAGE: - with Image.open(path) as image: - check_image_size(self, image, self.width, self.height) - logger.info(f"{name} passed in {duration:.4f}s (threshold: {time_threshold}s)") - - def model_name(self): - return self.model_path.split("/")[-1] - - def test_single_gpu(self): - """single gpu""" - self._run_test( - name=f"{self.model_name()}_single_gpu", - args=None, - model_path=self.model_path, - test_key="test_single_gpu", - ) - - def test_cfg_parallel(self): - """cfg parallel""" - if self.data_type == DataType.IMAGE: - return - self._run_test( - name=f"{self.model_name()}_cfg_parallel", - args="--num-gpus 2 --enable-cfg-parallel", - model_path=self.model_path, - test_key="test_cfg_parallel", - ) - - def test_usp(self): - """usp""" - if self.data_type == DataType.IMAGE: - return - self._run_test( - name=f"{self.model_name()}_usp", - args="--num-gpus 4 --ulysses-degree=2 --ring-degree=2", - model_path=self.model_path, - test_key="test_usp", - ) - - def test_mixed(self): - """mixed""" - if self.data_type == DataType.IMAGE: - return - self._run_test( - name=f"{self.model_name()}_mixed", - args="--num-gpus 4 --ulysses-degree=2 --ring-degree=1 --enable-cfg-parallel", - model_path=self.model_path, - test_key="test_mixed", - ) +# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo +import base64 +import dataclasses +import json +import os +import shlex +import socket +import subprocess +import sys +import time +import unittest +from pathlib import Path +from typing import Optional + +from PIL import Image + +from sglang.multimodal_gen.configs.sample.base import DataType +from sglang.multimodal_gen.runtime.utils.common import get_bool_env_var +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.runtime.utils.perf_logger import ( + RequestPerfRecord, + get_diffusion_perf_log_dir, +) + +logger = init_logger(__name__) + + +def run_command(command) -> Optional[float]: + """Runs a command and returns the execution time and status.""" + print(f"Running command: {shlex.join(command)}") + + duration = None + with subprocess.Popen( + command, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + encoding="utf-8", + ) as process: + for line in process.stdout: + sys.stdout.write(line) + if "Pixel data generated" in line: + words = line.split(" ") + duration = float(words[-2]) + + if process.returncode == 0: + return duration + else: + print(f"Command failed with exit code {process.returncode}") + return None + + +def probe_port(host="127.0.0.1", port=30010, timeout=2.0) -> bool: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.settimeout(timeout) + try: + s.connect((host, port)) + return True + except OSError: + return False + + +def is_in_ci() -> bool: + return get_bool_env_var("SGLANG_IS_IN_CI") + + +def get_dynamic_server_port() -> int: + cuda_devices = os.environ.get("CUDA_VISIBLE_DEVICES", "0") + if not cuda_devices: + cuda_devices = "0" + try: + first_device_id = int(cuda_devices.split(",")[0].strip()[0]) + except (ValueError, IndexError): + first_device_id = 0 + + if is_in_ci(): + base_port = 10000 + first_device_id * 2000 + else: + base_port = 20000 + first_device_id * 1000 + + return base_port + 1000 + + +def is_mp4(data): + idx = data.find(b"ftyp") + return 0 <= idx <= 32 + + +def is_jpeg(data: bytes) -> bool: + # JPEG files start with: FF D8 FF + return data.startswith(b"\xff\xd8\xff") + + +def is_png(data): + # PNG files start with: 89 50 4E 47 0D 0A 1A 0A + return data.startswith(b"\x89PNG\r\n\x1a\n") + + +def wait_for_port(host="127.0.0.1", port=30010, deadline=300.0, interval=0.5): + end = time.time() + deadline + last_err = None + while time.time() < end: + if probe_port(host, port, timeout=interval): + return True + time.sleep(interval) + raise TimeoutError(f"Port {host}:{port} not ready. Last error: {last_err}") + + +def check_image_size(ut, image, width, height): + # check image size + ut.assertEqual(image.size, (width, height)) + + +def get_perf_log_dir() -> Path: + """Gets the performance log directory from the centralized sglang utility.""" + log_dir_str = get_diffusion_perf_log_dir() + if not log_dir_str: + raise RuntimeError( + "Performance logging is disabled (SGLANG_PERF_LOG_DIR is empty), " + "but a test tried to access the log directory." + ) + return Path(log_dir_str) + + +def _ensure_log_path(log_dir: Path) -> Path: + log_dir.mkdir(parents=True, exist_ok=True) + return log_dir / "performance.log" + + +def clear_perf_log(log_dir: Path) -> Path: + """Delete the perf log file so tests can watch for fresh entries.""" + log_path = _ensure_log_path(log_dir) + if log_path.exists(): + log_path.unlink() + logger.info("[server-test] Monitoring perf log at %s", log_path.as_posix()) + return log_path + + +def prepare_perf_log() -> tuple[Path, Path]: + """Convenience helper to resolve and clear the perf log in one call.""" + log_dir = get_perf_log_dir() + log_path = clear_perf_log(log_dir) + return log_dir, log_path + + +def read_perf_logs(log_path: Path) -> list[RequestPerfRecord]: + if not log_path.exists(): + return [] + records: list[RequestPerfRecord] = [] + with log_path.open("r", encoding="utf-8") as fh: + for line in fh: + line = line.strip() + if not line: + continue + try: + record_dict = json.loads(line) + records.append(RequestPerfRecord(**record_dict)) + except json.JSONDecodeError: + continue + return records + + +def wait_for_req_perf_record( + request_id: str, + prev_len: int, + log_path: Path, + timeout: float = 300.0, +) -> tuple[RequestPerfRecord | None, int]: + """ + the stage metrics of this request should be in the performance_log file with {request-id} + """ + logger.info(f"Waiting for req perf record with request id: {request_id}") + deadline = time.time() + timeout + while time.time() < deadline: + records = read_perf_logs(log_path) + if len(records) == prev_len + 1: + # FIXME: unable to get rid from openai apis, this is a hack. we should compare rid + # potential error when there are multiple servers + return records[-1], len(records) + + time.sleep(0.5) + + if os.environ.get("SGLANG_GEN_BASELINE", "0") == "1": + records = read_perf_logs(log_path) + return None, len(records) + raise AssertionError(f"Timeout waiting for stage metrics for request {request_id} ") + + +def validate_image(b64_json: str) -> None: + """Decode and validate that image is PNG or JPEG.""" + image_bytes = base64.b64decode(b64_json) + assert is_png(image_bytes) or is_jpeg(image_bytes), "Image must be PNG or JPEG" + + +def validate_video(b64_json: str) -> None: + """Decode and validate that video is a valid format.""" + video_bytes = base64.b64decode(b64_json) + is_mp4 = ( + video_bytes[:4] == b"\x00\x00\x00\x18" or video_bytes[:4] == b"\x00\x00\x00\x1c" + ) + is_webm = video_bytes[:4] == b"\x1a\x45\xdf\xa3" + assert is_mp4 or is_webm, "Video must be MP4 or WebM" + + +def validate_openai_video(video_bytes: bytes) -> None: + """Validate that video is MP4 or WebM by magic bytes.""" + is_mp4 = ( + video_bytes.startswith(b"\x00\x00\x00\x18") + or video_bytes.startswith(b"\x00\x00\x00\x1c") + or video_bytes[4:8] == b"ftyp" + ) + is_webm = video_bytes.startswith(b"\x1a\x45\xdf\xa3") + assert is_mp4 or is_webm, "Video must be MP4 or WebM" + + +@dataclasses.dataclass +class TestResult: + name: str + key: str + duration: Optional[float] + succeed: bool + + @property + def duration_str(self): + return f"{self.duration:.4f}" if self.duration else "NA" + + +class TestCLIBase(unittest.TestCase): + model_path: str = None + extra_args = [] + data_type: DataType = None + # tested on h100 + thresholds = {} + + width: int = 720 + height: int = 720 + output_path: str = "test_outputs" + + base_command = [ + "sglang", + "generate", + "--text-encoder-cpu-offload", + "--pin-cpu-memory", + "--prompt", + "A curious raccoon", + "--save-output", + "--log-level=debug", + f"--width={width}", + f"--height={height}", + f"--output-path={output_path}", + ] + + results = [] + + @classmethod + def setUpClass(cls): + cls.results = [] + + def _run_command(self, name: str, model_path: str, test_key: str = "", args=[]): + command = ( + self.base_command + + [f"--model-path={model_path}"] + + shlex.split(args or "") + + ["--output-file-name", f"{name}"] + + self.extra_args + ) + duration = run_command(command) + status = "Success" if duration else "Failed" + succeed = duration is not None + + duration = float(duration) if succeed else None + self.results.append(TestResult(name, test_key, duration, succeed)) + + return name, duration, status + + +class TestGenerateBase(TestCLIBase): + model_path: str = None + extra_args = [] + data_type: DataType = None + # tested on h100 + thresholds = {} + + width: int = 720 + height: int = 720 + output_path: str = "test_outputs" + image_path: str | None = None + prompt: str | None = "A curious raccoon" + + base_command = [ + "sglang", + "generate", + # "--text-encoder-cpu-offload", + # "--pin-cpu-memory", + f"--prompt", + f"{prompt}", + "--save-output", + "--log-level=debug", + f"--width={width}", + f"--height={height}", + f"--output-path={output_path}", + ] + + results: list[TestResult] = [] + + @classmethod + def setUpClass(cls): + cls.results = [] + + @classmethod + def tearDownClass(cls): + # Print markdown table + print("\n## Test Results\n") + print("| Test Case | Duration | Status |") + print("|--------------------------------|----------|---------|") + test_keys = ["test_single_gpu", "test_cfg_parallel", "test_usp", "test_mixed"] + test_key_to_order = { + test_key: order for order, test_key in enumerate(test_keys) + } + + ordered_results: list[TestResult] = [None] * len(test_keys) + for result in cls.results: + order = test_key_to_order[result.key] + ordered_results[order] = result + + for result in ordered_results: + if not result: + continue + status = ( + "Succeed" + if ( + result.succeed + and float(result.duration) <= float(cls.thresholds[result.key]) + ) + else "Failed" + ) + print(f"| {result.name:<30} | {result.duration_str:<8} | {status:<7} |") + print() + durations = [result.duration_str for result in cls.results] + print(" | ".join([""] + durations + [""])) + + def _run_test(self, name: str, args, model_path: str, test_key: str): + time_threshold = self.thresholds[test_key] + name, duration, status = self._run_command( + name, args=args, model_path=model_path, test_key=test_key + ) + self.verify(status, name, duration, time_threshold) + + def verify(self, status, name, duration, time_threshold): + print("-" * 80) + print("\n" * 3) + + # test task status + self.assertEqual(status, "Success", f"{name} command failed") + self.assertIsNotNone(duration, f"Could not parse duration for {name}") + self.assertLessEqual( + duration, + time_threshold, + f"{name} failed with {duration:.4f}s > {time_threshold}s", + ) + + # test output file + path = os.path.join( + self.output_path, f"{name}.{self.data_type.get_default_extension()}" + ) + self.assertTrue(os.path.exists(path), f"Output file not exist for {path}") + if self.data_type == DataType.IMAGE: + with Image.open(path) as image: + check_image_size(self, image, self.width, self.height) + logger.info(f"{name} passed in {duration:.4f}s (threshold: {time_threshold}s)") + + def model_name(self): + return self.model_path.split("/")[-1] + + def test_single_gpu(self): + """single gpu""" + self._run_test( + name=f"{self.model_name()}_single_gpu", + args=None, + model_path=self.model_path, + test_key="test_single_gpu", + ) + + def test_cfg_parallel(self): + """cfg parallel""" + if self.data_type == DataType.IMAGE: + return + self._run_test( + name=f"{self.model_name()}_cfg_parallel", + args="--num-gpus 2 --enable-cfg-parallel", + model_path=self.model_path, + test_key="test_cfg_parallel", + ) + + def test_usp(self): + """usp""" + if self.data_type == DataType.IMAGE: + return + self._run_test( + name=f"{self.model_name()}_usp", + args="--num-gpus 4 --ulysses-degree=2 --ring-degree=2", + model_path=self.model_path, + test_key="test_usp", + ) + + def test_mixed(self): + """mixed""" + if self.data_type == DataType.IMAGE: + return + self._run_test( + name=f"{self.model_name()}_mixed", + args="--num-gpus 4 --ulysses-degree=2 --ring-degree=1 --enable-cfg-parallel", + model_path=self.model_path, + test_key="test_mixed", + )