[diffusion] chore: fix stage profiler for multi-stage denoising (#21955)

This commit is contained in:
Mick
2026-04-03 01:16:38 +08:00
committed by GitHub
parent df94cdcebb
commit 2278a321ca
5 changed files with 15 additions and 8 deletions
@@ -1025,6 +1025,7 @@ class DenoisingStage(PipelineStage):
logger=logger, logger=logger,
metrics=batch.metrics, 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,
record_as_step=True,
): ):
t_int = int(t_host.item()) t_int = int(t_host.item())
t_device = timesteps[i] t_device = timesteps[i]
@@ -407,6 +407,7 @@ class LTX2AVDenoisingStage(DenoisingStage):
logger=logger, logger=logger,
metrics=batch.metrics, 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,
record_as_step=True,
): ):
t_int = int(t_host.item()) t_int = int(t_host.item())
t_device = timesteps[i] t_device = timesteps[i]
@@ -104,6 +104,7 @@ class DmdDenoisingStage(DenoisingStage):
logger=logger, logger=logger,
metrics=batch.metrics, 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,
record_as_step=True,
): ):
t_int = int(t.item()) t_int = int(t.item())
if self.transformer_2 is not None: if self.transformer_2 is not None:
@@ -430,6 +430,7 @@ class MOVADenoisingStage(PipelineStage):
logger=logger, logger=logger,
metrics=metrics, metrics=metrics,
perf_dump_path_provided=perf_dump_path_provided, perf_dump_path_provided=perf_dump_path_provided,
record_as_step=True,
): ):
pair_t = paired_timesteps[idx_step] pair_t = paired_timesteps[idx_step]
if getattr(pair_t, "shape", None) == (2,): if getattr(pair_t, "shape", None) == (2,):
@@ -64,9 +64,8 @@ class RequestMetrics:
"""Records the duration of a pipeline stage""" """Records the duration of a pipeline stage"""
self.stages[stage_name] = duration_s * 1000 # Store as milliseconds self.stages[stage_name] = duration_s * 1000 # Store as milliseconds
def record_steps(self, index: int, duration_s: float): def record_step(self, duration_s: float):
"""Records the duration of a denoising step""" """Records the duration of a denoising step in execution order."""
assert index == len(self.steps)
self.steps.append(duration_s * 1000) self.steps.append(duration_s * 1000)
def record_memory_snapshot(self, checkpoint_name: str, snapshot: MemorySnapshot): def record_memory_snapshot(self, checkpoint_name: str, snapshot: MemorySnapshot):
@@ -192,6 +191,7 @@ class StageProfiler:
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,
record_as_step: bool = False,
): ):
self.stage_name = stage_name self.stage_name = stage_name
self.metrics = metrics self.metrics = metrics
@@ -200,6 +200,10 @@ class StageProfiler:
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
self.log_stage_start_end = log_stage_start_end self.log_stage_start_end = log_stage_start_end
self.capture_memory = capture_memory self.capture_memory = capture_memory
self.record_as_step = record_as_step
def _should_record_as_step(self) -> bool:
return self.record_as_step or self.stage_name.startswith("denoising_step_")
def __enter__(self): def __enter__(self):
if self.log_stage_start_end: if self.log_stage_start_end:
@@ -211,7 +215,7 @@ class StageProfiler:
if (self.log_timing and self.metrics) 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._should_record_as_step()
and torch.get_device_module().is_available() and torch.get_device_module().is_available()
): ):
torch.get_device_module().synchronize() torch.get_device_module().synchronize()
@@ -225,7 +229,7 @@ class StageProfiler:
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._should_record_as_step()
and torch.get_device_module().is_available() and torch.get_device_module().is_available()
): ):
torch.get_device_module().synchronize() torch.get_device_module().synchronize()
@@ -247,9 +251,8 @@ class StageProfiler:
) )
if self.log_timing and self.metrics: if self.log_timing and self.metrics:
if "denoising_step_" in self.stage_name: if self._should_record_as_step():
index = int(self.stage_name[len("denoising_step_") :]) self.metrics.record_step(execution_time_s)
self.metrics.record_steps(index, execution_time_s)
else: else:
self.metrics.record_stage(self.stage_name, execution_time_s) self.metrics.record_stage(self.stage_name, execution_time_s)