[diffusion] logging: log available mem when each stage starts in debug level (#18998)
This commit is contained in:
@@ -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())
|
||||||
|
|||||||
+3
-3
@@ -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
|
||||||
},
|
},
|
||||||
|
|||||||
Reference in New Issue
Block a user