[diffusion] fix: add profiling support and fix VBench dataset handling in bench_offline_throughput (#27704)
This commit is contained in:
@@ -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,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user