[diffusion] fix: add profiling support and fix VBench dataset handling in bench_offline_throughput (#27704)

This commit is contained in:
Chandrakant Khandelwal
2026-07-03 15:08:38 +08:00
committed by GitHub
parent 76f7f7c006
commit fe60764f54
2 changed files with 83 additions and 6 deletions
@@ -80,6 +80,7 @@ class BenchArgs:
dataset_path: str = "" dataset_path: str = ""
task_name: str = "unknown" task_name: str = "unknown"
num_prompts: int = 10 num_prompts: int = 10
num_outputs_per_prompt: int = 1
batch_size: int = 1 batch_size: int = 1
random_request_config: str = None random_request_config: str = None
random_request_seed: int = 42 random_request_seed: int = 42
@@ -89,6 +90,11 @@ class BenchArgs:
output_file: str = "" output_file: str = ""
disable_tqdm: bool = False disable_tqdm: bool = False
# Profiling
profile: bool = False
num_profiled_timesteps: int = 5
profile_all_stages: bool = False
@staticmethod @staticmethod
def add_cli_args(parser: argparse.ArgumentParser): def add_cli_args(parser: argparse.ArgumentParser):
"""Add benchmark-specific CLI arguments.""" """Add benchmark-specific CLI arguments."""
@@ -146,6 +152,12 @@ class BenchArgs:
default=10, default=10,
help="Total number of prompts to benchmark", help="Total number of prompts to benchmark",
) )
parser.add_argument(
"--num-outputs-per-prompt",
type=int,
default=1,
help="Number of generated outputs requested per prompt",
)
parser.add_argument( parser.add_argument(
"--batch-size", "--batch-size",
type=int, type=int,
@@ -185,6 +197,28 @@ class BenchArgs:
action="store_true", action="store_true",
help="Disable progress bar", help="Disable progress bar",
) )
parser.add_argument(
"--profile",
action="store_true",
help=(
"Enable PyTorch profiler for diffusion generation. "
"Set SGLANG_DIFFUSION_TORCH_PROFILER_DIR to control trace output directory."
),
)
parser.add_argument(
"--num-profiled-timesteps",
type=int,
default=5,
help=(
"Number of denoising timesteps to profile after warmup. "
"Use -1 to profile all denoising timesteps."
),
)
parser.add_argument(
"--profile-all-stages",
action="store_true",
help="Profile all diffusion pipeline stages instead of only denoising steps.",
)
@classmethod @classmethod
def from_cli_args(cls, args: argparse.Namespace): def from_cli_args(cls, args: argparse.Namespace):
@@ -216,7 +250,7 @@ def generate_batch(
output = BatchOutput() output = BatchOutput()
start_time = time.perf_counter() start_time = time.perf_counter()
torch.cuda.reset_peak_memory_stats() torch.get_device_module().reset_peak_memory_stats()
for prompt, params in zip(prompts, user_sampling_params): for prompt, params in zip(prompts, user_sampling_params):
try: try:
@@ -237,7 +271,9 @@ def generate_batch(
output.latency = time.perf_counter() - start_time output.latency = time.perf_counter() - start_time
output.latency_per_sample = output.latency / len(prompts) if prompts else 0.0 output.latency_per_sample = output.latency / len(prompts) if prompts else 0.0
output.success = output.num_samples > 0 output.success = output.num_samples > 0
output.peak_memory_mb = torch.cuda.max_memory_allocated() / (1024 * 1024) output.peak_memory_mb = torch.get_device_module().max_memory_allocated() / (
1024 * 1024
)
logger.debug( logger.debug(
f"Batch generated: {output.num_samples}/{len(prompts)} samples in {output.latency:.2f}s" f"Batch generated: {output.num_samples}/{len(prompts)} samples in {output.latency:.2f}s"
@@ -309,9 +345,14 @@ def throughput_test(
"--random-request-config can only be used with --dataset random" "--random-request-config can only be used with --dataset random"
) )
if bench_args.num_outputs_per_prompt != 1:
raise ValueError(
"bench_offline_throughput currently supports only --num-outputs-per-prompt 1"
)
logger.info(f"Loading {bench_args.dataset} dataset...") logger.info(f"Loading {bench_args.dataset} dataset...")
if bench_args.dataset == "vbench": if bench_args.dataset == "vbench":
bench_args.task_name = engine.server_args.pipeline_config.task_type bench_args.task_name = str(engine.server_args.pipeline_config.task_type)
dataset = VBenchDataset(bench_args) dataset = VBenchDataset(bench_args)
elif bench_args.dataset == "random": elif bench_args.dataset == "random":
dataset = RandomDataset(bench_args) dataset = RandomDataset(bench_args)
@@ -324,7 +365,11 @@ def throughput_test(
"height": bench_args.height, "height": bench_args.height,
"width": bench_args.width, "width": bench_args.width,
"num_frames": bench_args.num_frames, "num_frames": bench_args.num_frames,
"num_outputs_per_prompt": bench_args.num_outputs_per_prompt,
"seed": bench_args.seed, "seed": bench_args.seed,
"profile": bench_args.profile,
"num_profiled_timesteps": bench_args.num_profiled_timesteps,
"profile_all_stages": bench_args.profile_all_stages,
} }
if bench_args.disable_safety_checker: if bench_args.disable_safety_checker:
_sampling_params["safety_checker"] = None _sampling_params["safety_checker"] = None
@@ -345,7 +390,9 @@ def throughput_test(
logger.info("Running warmup batch...") logger.info("Running warmup batch...")
warmup_count = min(bench_args.batch_size, total_count) warmup_count = min(bench_args.batch_size, total_count)
warmup_prompts = all_prompts[:warmup_count] warmup_prompts = all_prompts[:warmup_count]
warmup_sampling_params = all_sampling_params[:warmup_count] warmup_sampling_params = [
{**p, "profile": False} for p in all_sampling_params[:warmup_count]
]
generate_batch(engine, bench_args, warmup_prompts, warmup_sampling_params) generate_batch(engine, bench_args, warmup_prompts, warmup_sampling_params)
logger.info(f"Running benchmark with {bench_args.num_prompts} prompts...") logger.info(f"Running benchmark with {bench_args.num_prompts} prompts...")
@@ -83,10 +83,39 @@ class VBenchDataset(BaseDataset):
self.cache_dir = os.path.join(os.path.expanduser("~"), ".cache", "sglang") self.cache_dir = os.path.join(os.path.expanduser("~"), ".cache", "sglang")
self.items = self._load_data() self.items = self._load_data()
@staticmethod
def _normalize_task_name(task_name: Any) -> Any:
"""Normalize enum-style task values to legacy benchmark task-name strings."""
enum_to_task_name = {
"T2V": "text-to-video",
"I2V": "image-to-video",
"TI2V": "image-to-video",
"T2I": "text-to-image",
"I2I": "image-to-image",
"TI2I": "image-to-image",
}
# Handle Enum-like objects, e.g., ModelTaskType.T2I
enum_name = getattr(task_name, "name", None)
if isinstance(enum_name, str):
return enum_to_task_name.get(enum_name, task_name)
# Handle direct string inputs or enum string repr
if isinstance(task_name, str):
if task_name in enum_to_task_name:
return enum_to_task_name[task_name]
if "." in task_name:
suffix = task_name.split(".")[-1]
return enum_to_task_name.get(suffix, task_name)
return task_name
def _load_data(self) -> List[Dict[str, Any]]: def _load_data(self) -> List[Dict[str, Any]]:
if self.args.task_name in ("text-to-video", "text-to-image", "video-to-video"): task_name = self._normalize_task_name(self.args.task_name)
if task_name in ("text-to-video", "text-to-image", "video-to-video"):
return self._load_t2v_prompts() return self._load_t2v_prompts()
elif self.args.task_name in ("image-to-video", "image-to-image"): elif task_name in ("image-to-video", "image-to-image"):
return self._load_i2v_data() return self._load_i2v_data()
else: else:
raise ValueError( raise ValueError(
@@ -281,6 +310,7 @@ class VBenchDataset(BaseDataset):
height=self.args.height, height=self.args.height,
num_frames=self.args.num_frames, num_frames=self.args.num_frames,
fps=self.args.fps, fps=self.args.fps,
num_inference_steps=self.args.num_inference_steps,
image_paths=[item["image_path"]] if "image_path" in item else None, image_paths=[item["image_path"]] if "image_path" in item else None,
) )