[diffusion] logging: log available mem when each stage starts in debug level (#18998)

This commit is contained in:
Mick
2026-02-20 19:57:06 +08:00
committed by GitHub
parent 0d20cf5a66
commit 38a69652e6
19 changed files with 109 additions and 74 deletions
@@ -18,9 +18,11 @@ from sglang.multimodal_gen.runtime.entrypoints.cli.cli_types import CLISubcomman
from sglang.multimodal_gen.runtime.entrypoints.cli.utils import ( from sglang.multimodal_gen.runtime.entrypoints.cli.utils import (
RaiseNotImplementedAction, RaiseNotImplementedAction,
) )
from sglang.multimodal_gen.runtime.entrypoints.utils import GenerationResult
from sglang.multimodal_gen.runtime.server_args import ServerArgs 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.logging_utils import init_logger
from sglang.multimodal_gen.runtime.utils.perf_logger import ( from sglang.multimodal_gen.runtime.utils.perf_logger import (
MemorySnapshot,
PerformanceLogger, PerformanceLogger,
RequestMetrics, RequestMetrics,
) )
@@ -58,7 +60,12 @@ def add_multimodal_gen_generate_args(parser: argparse.ArgumentParser):
return parser return parser
def maybe_dump_performance(args: argparse.Namespace, server_args, prompt: str, results): def maybe_dump_performance(
args: argparse.Namespace,
server_args,
prompt: str,
results: GenerationResult | list[GenerationResult] | None,
):
"""dump performance if necessary""" """dump performance if necessary"""
if not (args.perf_dump_path and results): if not (args.perf_dump_path and results):
return return
@@ -68,20 +75,29 @@ def maybe_dump_performance(args: argparse.Namespace, server_args, prompt: str, r
else: else:
result = results result = results
timings_dict = getattr(result, "timings", None) or ( metrics_dict = result.metrics
result.get("timings") if isinstance(result, dict) else None if not (args.perf_dump_path and metrics_dict):
)
if not (args.perf_dump_path and timings_dict):
return return
timings = RequestMetrics(request_id=timings_dict.get("request_id")) metrics = RequestMetrics(request_id=metrics_dict.get("request_id"))
timings.stages = timings_dict.get("stages", {}) metrics.stages = metrics_dict.get("stages", {})
timings.steps = timings_dict.get("steps", []) metrics.steps = metrics_dict.get("steps", [])
timings.total_duration_ms = timings_dict.get("total_duration_ms", 0) metrics.total_duration_ms = metrics_dict.get("total_duration_ms", 0)
# restore memory snapshots from serialized dict
memory_snapshots_dict = metrics_dict.get("memory_snapshots", {})
for checkpoint_name, snapshot_dict in memory_snapshots_dict.items():
snapshot = MemorySnapshot(
allocated_mb=snapshot_dict.get("allocated_mb", 0.0),
reserved_mb=snapshot_dict.get("reserved_mb", 0.0),
peak_allocated_mb=snapshot_dict.get("peak_allocated_mb", 0.0),
peak_reserved_mb=snapshot_dict.get("peak_reserved_mb", 0.0),
)
metrics.memory_snapshots[checkpoint_name] = snapshot
PerformanceLogger.dump_benchmark_report( PerformanceLogger.dump_benchmark_report(
file_path=args.perf_dump_path, file_path=args.perf_dump_path,
timings=timings, metrics=metrics,
meta={ meta={
"prompt": prompt, "prompt": prompt,
"model": server_args.model_path, "model": server_args.model_path,
@@ -206,9 +206,9 @@ class DiffGenerator:
size=(req.height, req.width, req.num_frames), size=(req.height, req.width, req.num_frames),
generation_time=timer.duration, generation_time=timer.duration,
peak_memory_mb=output_batch.peak_memory_mb, peak_memory_mb=output_batch.peak_memory_mb,
timings=( metrics=(
output_batch.timings.to_dict() output_batch.metrics.to_dict()
if output_batch.timings if output_batch.metrics
else {} else {}
), ),
trajectory_latents=output_batch.trajectory_latents, trajectory_latents=output_batch.trajectory_latents,
@@ -297,7 +297,7 @@ class DiffGenerator:
if not results: if not results:
return return
if self.server_args.warmup: if self.server_args.warmup:
total_duration_ms = results[0].timings.get("total_duration_ms", 0) total_duration_ms = results[0].metrics.get("total_duration_ms", 0)
logger.info( logger.info(
f"Warmed-up request processed in {GREEN}%.2f{RESET} seconds (with warmup excluded)", f"Warmed-up request processed in {GREEN}%.2f{RESET} seconds (with warmup excluded)",
total_duration_ms / 1000.0, total_duration_ms / 1000.0,
@@ -304,8 +304,8 @@ def add_common_data_to_response(
if result.peak_memory_mb and result.peak_memory_mb > 0: if result.peak_memory_mb and result.peak_memory_mb > 0:
response["peak_memory_mb"] = result.peak_memory_mb response["peak_memory_mb"] = result.peak_memory_mb
if result.timings and result.timings.total_duration_s > 0: if result.metrics and result.metrics.total_duration_s > 0:
response["inference_time_s"] = result.timings.total_duration_s response["inference_time_s"] = result.metrics.total_duration_s
response["id"] = request_id response["id"] = request_id
@@ -105,7 +105,7 @@ class GenerationResult:
size: tuple | None = None # (height, width, num_frames) size: tuple | None = None # (height, width, num_frames)
generation_time: float = 0.0 generation_time: float = 0.0
peak_memory_mb: float = 0.0 peak_memory_mb: float = 0.0
timings: dict = field(default_factory=dict) metrics: dict = field(default_factory=dict)
trajectory_latents: Any = None trajectory_latents: Any = None
trajectory_timesteps: Any = None trajectory_timesteps: Any = None
trajectory_decoded: Any = None trajectory_decoded: Any = None
@@ -19,7 +19,6 @@ from sglang.multimodal_gen.runtime.distributed import get_local_torch_device
from sglang.multimodal_gen.runtime.loader.utils import ( from sglang.multimodal_gen.runtime.loader.utils import (
_normalize_component_type, _normalize_component_type,
component_name_to_loader_cls, component_name_to_loader_cls,
get_memory_usage_of_component,
) )
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
@@ -125,11 +124,9 @@ class ComponentLoader(ABC):
if isinstance(component, nn.Module): if isinstance(component, nn.Module):
component = component.eval() component = component.eval()
current_gpu_mem = current_platform.get_available_gpu_memory() current_gpu_mem = current_platform.get_available_gpu_memory()
consumed = get_memory_usage_of_component(component) consumed = gpu_mem_before_loading - current_gpu_mem
if consumed is None or consumed == 0.0:
consumed = gpu_mem_before_loading - current_gpu_mem
logger.info( logger.info(
f"Loaded %s: %s ({source} version). model size: %.2f GB, avail mem: %.2f GB", f"Loaded %s: %s ({source} version). consumed: %.2f GB, avail mem: %.2f GB",
component_name, component_name,
component.__class__.__name__, component.__class__.__name__,
consumed, consumed,
@@ -158,8 +158,8 @@ class GPUWorker:
def do_mem_analysis(self, output_batch: OutputBatch): def do_mem_analysis(self, output_batch: OutputBatch):
final_snapshot = capture_memory_snapshot() final_snapshot = capture_memory_snapshot()
if output_batch.timings: if output_batch.metrics:
output_batch.timings.record_memory_snapshot("mem_analysis", final_snapshot) output_batch.metrics.record_memory_snapshot("mem_analysis", final_snapshot)
# for details on max_memory_reserved: https://docs.pytorch.org/docs/stable/generated/torch.cuda.memory.max_memory_reserved.html # for details on max_memory_reserved: https://docs.pytorch.org/docs/stable/generated/torch.cuda.memory.max_memory_reserved.html
peak_reserved_bytes = torch.get_device_module().max_memory_reserved() peak_reserved_bytes = torch.get_device_module().max_memory_reserved()
@@ -219,9 +219,9 @@ class GPUWorker:
start_time = time.monotonic() start_time = time.monotonic()
# capture memory baseline before forward # capture memory baseline before forward
if self.rank == 0 and req.timings: if self.rank == 0 and req.metrics:
baseline_snapshot = capture_memory_snapshot() baseline_snapshot = capture_memory_snapshot()
req.timings.record_memory_snapshot("before_forward", baseline_snapshot) req.metrics.record_memory_snapshot("before_forward", baseline_snapshot)
req.log(server_args=self.server_args) req.log(server_args=self.server_args)
result = self.pipeline.forward(req, self.server_args) result = self.pipeline.forward(req, self.server_args)
@@ -231,7 +231,7 @@ class GPUWorker:
output=result.output, output=result.output,
audio=getattr(result, "audio", None), audio=getattr(result, "audio", None),
audio_sample_rate=getattr(result, "audio_sample_rate", None), audio_sample_rate=getattr(result, "audio_sample_rate", None),
timings=result.timings, metrics=result.metrics,
trajectory_timesteps=getattr(result, "trajectory_timesteps", None), trajectory_timesteps=getattr(result, "trajectory_timesteps", None),
trajectory_latents=getattr(result, "trajectory_latents", None), trajectory_latents=getattr(result, "trajectory_latents", None),
noise_pred=getattr(result, "noise_pred", None), noise_pred=getattr(result, "noise_pred", None),
@@ -241,9 +241,9 @@ class GPUWorker:
output_batch = result output_batch = result
# capture memory after forward (peak) # capture memory after forward (peak)
if self.rank == 0 and output_batch.timings: if self.rank == 0 and output_batch.metrics:
peak_snapshot = capture_memory_snapshot() peak_snapshot = capture_memory_snapshot()
output_batch.timings.record_memory_snapshot( output_batch.metrics.record_memory_snapshot(
"after_forward", peak_snapshot "after_forward", peak_snapshot
) )
@@ -251,7 +251,7 @@ class GPUWorker:
self.do_mem_analysis(output_batch) self.do_mem_analysis(output_batch)
duration_ms = (time.monotonic() - start_time) * 1000 duration_ms = (time.monotonic() - start_time) * 1000
output_batch.timings.total_duration_ms = duration_ms output_batch.metrics.total_duration_ms = duration_ms
# Save output to file and return file path only if requested. Avoid the serialization # Save output to file and return file path only if requested. Avoid the serialization
# and deserialization overhead between scheduler_client and gpu_worker. # and deserialization overhead between scheduler_client and gpu_worker.
@@ -273,11 +273,13 @@ class GPUWorker:
if req.perf_dump_path is not None or envs.SGLANG_DIFFUSION_STAGE_LOGGING: if req.perf_dump_path is not None or envs.SGLANG_DIFFUSION_STAGE_LOGGING:
# Avoid logging warmup perf records that share the same request_id. # Avoid logging warmup perf records that share the same request_id.
if not req.is_warmup: if not req.is_warmup:
PerformanceLogger.log_request_summary(timings=output_batch.timings) PerformanceLogger.log_request_summary(metrics=output_batch.metrics)
except Exception as e: except Exception as e:
logger.error( logger.error(
f"Error executing request {req.request_id}: {e}", exc_info=True f"Error executing request {req.request_id}: {e}", exc_info=True
) )
if isinstance(e, _oom_exceptions()):
logger.warning(OOM_MSG)
if output_batch is None: if output_batch is None:
output_batch = OutputBatch() output_batch = OutputBatch()
output_batch.error = f"Error executing request {req.request_id}: {e}" output_batch.error = f"Error executing request {req.request_id}: {e}"
@@ -427,8 +429,8 @@ OOM detected. Possible solutions:
- If the OOM occurs during loading: - If the OOM occurs during loading:
1. Enable CPU offload for memory-intensive components, or use `--dit-layerwise-offload` for DiT 1. Enable CPU offload for memory-intensive components, or use `--dit-layerwise-offload` for DiT
- If the OOM occurs during runtime: - If the OOM occurs during runtime:
1. Reduce the number of output tokens by lowering resolution or decreasing `--num-frames` 1. Enable SP and/or TP (in a multi-GPU setup)
2. Enable SP and/or TP 2. Reduce the number of output tokens by lowering resolution or decreasing `--num-frames`
3. Opt for a sparse-attention backend 3. Opt for a sparse-attention backend
4. Enable FSDP by `--use-fsdp-inference` (in a multi-GPU setup) 4. Enable FSDP by `--use-fsdp-inference` (in a multi-GPU setup)
5. Enable quantization (e.g. nunchaku) 5. Enable quantization (e.g. nunchaku)
@@ -436,6 +438,14 @@ OOM detected. Possible solutions:
""" """
def _oom_exceptions():
# torch.OutOfMemoryError exists only in some PyTorch builds
types = [torch.cuda.OutOfMemoryError]
if hasattr(torch, "OutOfMemoryError"):
types.append(torch.OutOfMemoryError)
return tuple(types)
def run_scheduler_process( def run_scheduler_process(
local_rank: int, local_rank: int,
rank: int, rank: int,
@@ -483,7 +493,7 @@ def run_scheduler_process(
} }
) )
scheduler.event_loop() scheduler.event_loop()
except torch.OutOfMemoryError as _e: except _oom_exceptions() as _e:
logger.warning(OOM_MSG) logger.warning(OOM_MSG)
raise raise
finally: finally:
@@ -395,12 +395,12 @@ class Scheduler:
if self._warmup_total > 0: if self._warmup_total > 0:
logger.info( logger.info(
f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processed in {GREEN}%.2f{RESET} seconds", f"Warmup req ({self._warmup_processed}/{self._warmup_total}) processed in {GREEN}%.2f{RESET} seconds",
output_batch.timings.total_duration_s, output_batch.metrics.total_duration_s,
) )
else: else:
logger.info( logger.info(
f"Warmup req processed in {GREEN}%.2f{RESET} seconds", f"Warmup req processed in {GREEN}%.2f{RESET} seconds",
output_batch.timings.total_duration_s, output_batch.metrics.total_duration_s,
) )
else: else:
if self._warmup_total > 0: if self._warmup_total > 0:
@@ -22,6 +22,7 @@ from sglang.multimodal_gen.runtime.pipelines_core.executors.pipeline_executor im
) )
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req
from sglang.multimodal_gen.runtime.pipelines_core.stages import PipelineStage from sglang.multimodal_gen.runtime.pipelines_core.stages import PipelineStage
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import ( from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model, maybe_download_model,
@@ -311,7 +312,11 @@ class ComposedPipelineBase(ABC):
f"Required module: {module_name} was not found in loaded modules: {list(loaded_components.keys())}" f"Required module: {module_name} was not found in loaded modules: {list(loaded_components.keys())}"
) )
logger.debug("Memory usage of loaded modules: %s", self.memory_usages) logger.debug(
"Memory usage of loaded modules (GiB): %s. Available memory: %s",
self.memory_usages,
round(current_platform.get_available_gpu_memory(), 2),
)
return loaded_components return loaded_components
@@ -31,7 +31,7 @@ class Timer(StageProfiler):
def __init__(self, name="Stage"): def __init__(self, name="Stage"):
super().__init__( super().__init__(
stage_name=name, timings=None, log_stage_start_end=True, logger=logger stage_name=name, logger=logger, metrics=None, log_stage_start_end=True
) )
@@ -152,7 +152,7 @@ class Req:
VSA_sparsity: float = 0.0 VSA_sparsity: float = 0.0
# stage logging # stage logging
timings: Optional["RequestMetrics"] = None metrics: Optional["RequestMetrics"] = None
# results # results
output: torch.Tensor | None = None output: torch.Tensor | None = None
@@ -267,7 +267,7 @@ class Req:
if self.guidance_scale_2 is None: if self.guidance_scale_2 is None:
self.guidance_scale_2 = self.guidance_scale self.guidance_scale_2 = self.guidance_scale
self.timings = RequestMetrics(request_id=self.request_id) self.metrics = RequestMetrics(request_id=self.request_id)
if self.is_warmup: if self.is_warmup:
self.set_as_warmup() self.set_as_warmup()
@@ -330,7 +330,7 @@ class OutputBatch:
output_file_paths: list[str] | None = None output_file_paths: list[str] | None = None
# logged metrics info, directly from Req.timings # logged metrics info, directly from Req.timings
timings: Optional["RequestMetrics"] = None metrics: Optional["RequestMetrics"] = None
# For ComfyUI integration: noise prediction from denoising stage # For ComfyUI integration: noise prediction from denoising stage
noise_pred: torch.Tensor | None = None noise_pred: torch.Tensor | None = None
@@ -195,10 +195,10 @@ class PipelineStage(ABC):
with StageProfiler( with StageProfiler(
stage_name, stage_name,
logger=logger, logger=logger,
timings=batch.timings, metrics=batch.metrics,
perf_dump_path_provided=batch.perf_dump_path is not None,
log_stage_start_end=not batch.is_warmup log_stage_start_end=not batch.is_warmup
and not (self.server_args and self.server_args.comfyui_mode), and not (self.server_args and self.server_args.comfyui_mode),
perf_dump_path_provided=batch.perf_dump_path is not None,
): ):
result = self.forward(batch, server_args) result = self.forward(batch, server_args)
@@ -232,7 +232,7 @@ class DecodingStage(PipelineStage):
trajectory_timesteps=batch.trajectory_timesteps, trajectory_timesteps=batch.trajectory_timesteps,
trajectory_latents=batch.trajectory_latents, trajectory_latents=batch.trajectory_latents,
trajectory_decoded=trajectory_decoded, trajectory_decoded=trajectory_decoded,
timings=batch.timings, metrics=batch.metrics,
) )
self.offload_model() self.offload_model()
@@ -69,7 +69,7 @@ class LTX2AVDecodingStage(DecodingStage):
trajectory_timesteps=batch.trajectory_timesteps, trajectory_timesteps=batch.trajectory_timesteps,
trajectory_latents=batch.trajectory_latents, trajectory_latents=batch.trajectory_latents,
trajectory_decoded=None, trajectory_decoded=None,
timings=batch.timings, metrics=batch.metrics,
) )
# 2. Decode Audio # 2. Decode Audio
@@ -1017,7 +1017,7 @@ class DenoisingStage(PipelineStage):
with StageProfiler( with StageProfiler(
f"denoising_step_{i}", f"denoising_step_{i}",
logger=logger, logger=logger,
timings=batch.timings, metrics=batch.metrics,
perf_dump_path_provided=batch.perf_dump_path is not None, perf_dump_path_provided=batch.perf_dump_path is not None,
): ):
t_int = int(t_host.item()) t_int = int(t_host.item())
@@ -346,7 +346,7 @@ class LTX2AVDenoisingStage(DenoisingStage):
with StageProfiler( with StageProfiler(
f"denoising_step_{i}", f"denoising_step_{i}",
logger=logger, logger=logger,
timings=batch.timings, metrics=batch.metrics,
perf_dump_path_provided=batch.perf_dump_path is not None, perf_dump_path_provided=batch.perf_dump_path is not None,
): ):
t_int = int(t_host.item()) t_int = int(t_host.item())
@@ -102,7 +102,7 @@ class DmdDenoisingStage(DenoisingStage):
with StageProfiler( with StageProfiler(
f"denoising_step_{i}", f"denoising_step_{i}",
logger=logger, logger=logger,
timings=batch.timings, metrics=batch.metrics,
perf_dump_path_provided=batch.perf_dump_path is not None, perf_dump_path_provided=batch.perf_dump_path is not None,
): ):
t_int = int(t.item()) t_int = int(t.item())
@@ -399,7 +399,7 @@ class MOVADenoisingStage(PipelineStage):
getattr(batch, "extra_step_kwargs", None) or {}, getattr(batch, "extra_step_kwargs", None) or {},
) )
timings = getattr(batch, "timings", None) metrics = getattr(batch, "metrics", None)
perf_dump_path_provided = getattr(batch, "perf_dump_path", None) is not None perf_dump_path_provided = getattr(batch, "perf_dump_path", None) is not None
with self.progress_bar(total=total_steps) as progress_bar: with self.progress_bar(total=total_steps) as progress_bar:
@@ -407,7 +407,7 @@ class MOVADenoisingStage(PipelineStage):
with StageProfiler( with StageProfiler(
f"denoising_step_{idx_step}", f"denoising_step_{idx_step}",
logger=logger, logger=logger,
timings=timings, metrics=metrics,
perf_dump_path_provided=perf_dump_path_provided, perf_dump_path_provided=perf_dump_path_provided,
): ):
pair_t = paired_timesteps[idx_step] pair_t = paired_timesteps[idx_step]
@@ -908,6 +908,6 @@ class MOVADecodingStage(PipelineStage):
output=video, output=video,
audio=audio, audio=audio,
audio_sample_rate=getattr(self.audio_vae, "sample_rate", None), audio_sample_rate=getattr(self.audio_vae, "sample_rate", None),
timings=batch.timings, metrics=batch.metrics,
) )
return output_batch return output_batch
@@ -1,6 +1,7 @@
# Copied and adapted from: https://github.com/hao-ai-lab/FastVideo # Copied and adapted from: https://github.com/hao-ai-lab/FastVideo
import dataclasses import dataclasses
import json import json
import logging
import os import os
import subprocess import subprocess
import sys import sys
@@ -15,7 +16,10 @@ from dateutil.tz import UTC
import sglang import sglang
import sglang.multimodal_gen.envs as envs import sglang.multimodal_gen.envs as envs
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.multimodal_gen.runtime.utils.logging_utils import ( from sglang.multimodal_gen.runtime.utils.logging_utils import (
CYAN,
RESET,
_SGLDiffusionLogger, _SGLDiffusionLogger,
get_is_main_process, get_is_main_process,
init_logger, init_logger,
@@ -184,13 +188,13 @@ class StageProfiler:
self, self,
stage_name: str, stage_name: str,
logger: _SGLDiffusionLogger, logger: _SGLDiffusionLogger,
timings: Optional["RequestMetrics"], metrics: Optional["RequestMetrics"],
log_stage_start_end: bool = False, log_stage_start_end: bool = False,
perf_dump_path_provided: bool = False, perf_dump_path_provided: bool = False,
capture_memory: bool = False, capture_memory: bool = False,
): ):
self.stage_name = stage_name self.stage_name = stage_name
self.timings = timings self.metrics = metrics
self.logger = logger self.logger = logger
self.start_time = 0.0 self.start_time = 0.0
self.log_timing = perf_dump_path_provided or envs.SGLANG_DIFFUSION_STAGE_LOGGING self.log_timing = perf_dump_path_provided or envs.SGLANG_DIFFUSION_STAGE_LOGGING
@@ -199,9 +203,12 @@ class StageProfiler:
def __enter__(self): def __enter__(self):
if self.log_stage_start_end: if self.log_stage_start_end:
self.logger.info(f"[{self.stage_name}] started...") msg = f"[{self.stage_name}] started..."
if self.logger.isEnabledFor(logging.DEBUG):
msg += f" ({round(current_platform.get_available_gpu_memory(), 2)} GB left)"
self.logger.info(msg)
if (self.log_timing and self.timings) or self.log_stage_start_end: if (self.log_timing and self.metrics) or self.log_stage_start_end:
if ( if (
os.environ.get("SGLANG_DIFFUSION_SYNC_STAGE_PROFILING", "0") == "1" os.environ.get("SGLANG_DIFFUSION_SYNC_STAGE_PROFILING", "0") == "1"
and self.stage_name.startswith("denoising_step_") and self.stage_name.startswith("denoising_step_")
@@ -213,7 +220,7 @@ class StageProfiler:
return self return self
def __exit__(self, exc_type, exc_val, exc_tb): def __exit__(self, exc_type, exc_val, exc_tb):
if not ((self.log_timing and self.timings) or self.log_stage_start_end): if not ((self.log_timing and self.metrics) or self.log_stage_start_end):
return False return False
if ( if (
@@ -239,17 +246,17 @@ class StageProfiler:
f"[{self.stage_name}] finished in {execution_time_s:.4f} seconds", f"[{self.stage_name}] finished in {execution_time_s:.4f} seconds",
) )
if self.log_timing and self.timings: if self.log_timing and self.metrics:
if "denoising_step_" in self.stage_name: if "denoising_step_" in self.stage_name:
index = int(self.stage_name[len("denoising_step_") :]) index = int(self.stage_name[len("denoising_step_") :])
self.timings.record_steps(index, execution_time_s) self.metrics.record_steps(index, execution_time_s)
else: else:
self.timings.record_stage(self.stage_name, execution_time_s) self.metrics.record_stage(self.stage_name, execution_time_s)
# capture memory snapshot after stage if requested # capture memory snapshot after stage if requested
if self.capture_memory and torch.cuda.is_available(): if self.capture_memory and torch.cuda.is_available():
snapshot = capture_memory_snapshot() snapshot = capture_memory_snapshot()
self.timings.record_memory_snapshot( self.metrics.record_memory_snapshot(
f"after_{self.stage_name}", snapshot f"after_{self.stage_name}", snapshot
) )
@@ -269,7 +276,7 @@ class PerformanceLogger:
def dump_benchmark_report( def dump_benchmark_report(
cls, cls,
file_path: str, file_path: str,
timings: "RequestMetrics", metrics: "RequestMetrics",
meta: Optional[Dict[str, Any]] = None, meta: Optional[Dict[str, Any]] = None,
tag: str = "benchmark_dump", tag: str = "benchmark_dump",
): ):
@@ -279,25 +286,25 @@ class PerformanceLogger:
""" """
formatted_steps = [ formatted_steps = [
{"name": name, "duration_ms": duration_ms} {"name": name, "duration_ms": duration_ms}
for name, duration_ms in timings.stages.items() for name, duration_ms in metrics.stages.items()
] ]
denoise_steps_ms = [ denoise_steps_ms = [
{"step": idx, "duration_ms": duration_ms} {"step": idx, "duration_ms": duration_ms}
for idx, duration_ms in enumerate(timings.steps) for idx, duration_ms in enumerate(metrics.steps)
] ]
memory_checkpoints = { memory_checkpoints = {
name: snapshot.to_dict() name: snapshot.to_dict()
for name, snapshot in timings.memory_snapshots.items() for name, snapshot in metrics.memory_snapshots.items()
} }
report = { report = {
"timestamp": datetime.now(UTC).isoformat(), "timestamp": datetime.now(UTC).isoformat(),
"request_id": timings.request_id, "request_id": metrics.request_id,
"commit_hash": get_git_commit_hash(), "commit_hash": get_git_commit_hash(),
"tag": tag, "tag": tag,
"total_duration_ms": timings.total_duration_ms, "total_duration_ms": metrics.total_duration_ms,
"steps": formatted_steps, "steps": formatted_steps,
"denoise_steps_ms": denoise_steps_ms, "denoise_steps_ms": denoise_steps_ms,
"memory_checkpoints": memory_checkpoints, "memory_checkpoints": memory_checkpoints,
@@ -309,14 +316,14 @@ class PerformanceLogger:
os.makedirs(os.path.dirname(abs_path), exist_ok=True) os.makedirs(os.path.dirname(abs_path), exist_ok=True)
with open(abs_path, "w", encoding="utf-8") as f: with open(abs_path, "w", encoding="utf-8") as f:
json.dump(report, f, indent=2) json.dump(report, f, indent=2)
logger.info(f"Metrics dumped to: {abs_path}") logger.info(f"Metrics dumped to: {CYAN}{abs_path}{RESET}")
except IOError as e: except IOError as e:
logger.error(f"Failed to dump metrics to {abs_path}: {e}") logger.error(f"Failed to dump metrics to {abs_path}: {e}")
@classmethod @classmethod
def log_request_summary( def log_request_summary(
cls, cls,
timings: "RequestMetrics", metrics: "RequestMetrics",
tag: str = "total_inference_time", tag: str = "total_inference_time",
): ):
"""logs the stage metrics and total duration for a completed request """logs the stage metrics and total duration for a completed request
@@ -326,21 +333,21 @@ class PerformanceLogger:
""" """
formatted_stages = [ formatted_stages = [
{"name": name, "execution_time_ms": duration_ms} {"name": name, "execution_time_ms": duration_ms}
for name, duration_ms in timings.stages.items() for name, duration_ms in metrics.stages.items()
] ]
memory_checkpoints = { memory_checkpoints = {
name: snapshot.to_dict() name: snapshot.to_dict()
for name, snapshot in timings.memory_snapshots.items() for name, snapshot in metrics.memory_snapshots.items()
} }
record = RequestPerfRecord( record = RequestPerfRecord(
timings.request_id, metrics.request_id,
commit_hash=get_git_commit_hash(), commit_hash=get_git_commit_hash(),
tag="pipeline_stage_metrics", tag="pipeline_stage_metrics",
stages=formatted_stages, stages=formatted_stages,
steps=timings.steps, steps=metrics.steps,
total_duration_ms=timings.total_duration_ms, total_duration_ms=metrics.total_duration_ms,
memory_snapshots=memory_checkpoints, memory_snapshots=memory_checkpoints,
) )
@@ -542,7 +542,7 @@
"7": 93.97, "7": 93.97,
"8": 94.32 "8": 94.32
}, },
"expected_e2e_ms": 1192.92, "expected_e2e_ms": 1292.92,
"expected_avg_denoise_ms": 83.75, "expected_avg_denoise_ms": 83.75,
"expected_median_denoise_ms": 93.58 "expected_median_denoise_ms": 93.58
}, },