Compare commits

...
Author SHA1 Message Date
SolitaryThinker 3cc2420e84 make the region system actually usable: short names, summaries, nvtx, dead-code removal
- FASTVIDEO_TORCH_PROFILE_REGIONS accepts short names (model_loading ==
  profiler_region_model_loading)
- warn once on unregistered region names instead of silently profiling nothing
- write key_averages summary (txt + json) per rank next to the trace at stop
- nvtx range around regions so the same names show up in nsys timelines
- region marker record_function now opens after collection is enabled so it
  actually appears in the trace
- with_stack off by default (~1.5x runtime overhead; the declared default and
  the env lambda disagreed)
- module-level profiler_region() helper; h3 denoising uses it
- delete dead code: pipeline.profile()/self.profiler, controller.profiler
  property, get/set_global_profiler, default_enabled
- CPU contract test: subprocess five-window matrix proving regions gate
  collection and atexit exports the trace
2026-08-09 17:55:09 -07:00
SolitaryThinker ccaafd327f fix the profiler region system: it recorded nothing with default env
Three bugs, found while wiring the inference-denoising region:

- the profiler was built with schedule(wait=2, warmup=1, active=2) but
  nothing ever calls profiler.step(), so it never left the wait phase and
  recorded nothing. Delete the schedule and the three *_STEPS env knobs
  (they could only brick the feature).
- _collection_enabled initialized False while torch's profiler actually
  collects from start(), so the initial toggle-off short-circuited and
  everything before the first region exit was captured. Initialize True so
  start() really disables collection.
- traces only export at stop(), which no inference path calls — regions on
  inference stages could never produce a trace. Register an atexit stop
  (made idempotent).

Also: with_stack default off (measured 1.46x overhead and +44% trace size
vs 1.27x without), fix the module docstring's non-resolving region example,
and un-swap the train/train_one_step descriptions.

Verified end to end on CPU in a child process: pre/post-region ops excluded,
in-region ops captured, trace exports at interpreter exit without an
explicit stop().
2026-08-09 17:54:08 -07:00
5 changed files with 232 additions and 90 deletions
+5 -14
View File
@@ -30,11 +30,8 @@ if TYPE_CHECKING:
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = True
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_FLOPS: bool = False
FASTVIDEO_TORCH_PROFILER_WAIT_STEPS: int = 2
FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS: int = 1
FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS: int = 2
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
FASTVIDEO_TRACE_ACTIVATIONS: bool = False
FASTVIDEO_TRACE_LAYERS: str = ""
@@ -242,11 +239,11 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY", "0") != "0"),
# Enable torch profiler to profile stack if set
# FASTVIDEO_TORCH_PROFILER_WITH_STACK=1. If not set, torch profiler WILL
# profile stack by default.
# Enable torch profiler stack capture with
# FASTVIDEO_TORCH_PROFILER_WITH_STACK=1. Off by default: stack capture
# costs ~1.5x runtime overhead and ~1.4x trace size.
"FASTVIDEO_TORCH_PROFILER_WITH_STACK":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_STACK", "1") != "0"),
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_STACK", "0") != "0"),
# Enable torch profiler to profile flops if set
# FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1. If not set, torch profiler will
@@ -255,16 +252,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_FLOPS", "0") != "0"),
# Wait steps per profiling cycle (torch.profiler.schedule wait parameter)
# Defaults to 2 if not set.
"FASTVIDEO_TORCH_PROFILER_WAIT_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_WAIT_STEPS", "2")),
# Warmup steps per profiling cycle (torch.profiler.schedule warmup parameter)
# Defaults to 1 if not set.
"FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS", "1")),
# Active steps per profiling cycle (torch.profiler.schedule active parameter)
# Defaults to 2 if not set.
"FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS", "2")),
"FASTVIDEO_TORCH_PROFILE_REGIONS":
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
@@ -3,7 +3,6 @@
from __future__ import annotations
import contextlib
from typing import Any
import torch
@@ -12,7 +11,7 @@ from fastvideo.attention.selector import component_attention_backend, get_attn_b
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.profiler import get_global_controller
from fastvideo.profiler import profiler_region
from fastvideo.hooks.activation_trace import trace_step
from fastvideo.pipelines.basic.minimax_h3.packing import (
MINIMAX_H3_KEYFRAME_NOISE_AUG,
@@ -147,13 +146,10 @@ class MiniMaxH3DenoisingStage(PipelineStage):
vsa_dense_layers = tuple(batch.extra.get("vsa_dense_layers", ()))
vsa_dense_first_n = int(batch.extra.get("vsa_dense_first_n_steps", 0))
controller = get_global_controller()
denoise_region = (controller.region("profiler_region_inference_denoising")
if controller is not None else contextlib.nullcontext())
try:
with denoise_region:
for index, (video_timestep,
audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps, strict=True)):
with profiler_region("inference_denoising"):
for index, (video_timestep, audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps,
strict=True)):
unique_timesteps, timestep_indices = row_timestep_plan[index]
attn_metadata = None
if vsa_metadata_builder is not None:
@@ -94,7 +94,6 @@ class ComposedPipelineBase(ABC):
# FASTVIDEO_TORCH_PROFILER_DIR=/path/to/save/trace
trace_dir = envs.FASTVIDEO_TORCH_PROFILER_DIR
self.profiler_controller = get_or_create_profiler(trace_dir)
self.profiler = self.profiler_controller.profiler
self.local_rank = get_world_group().local_rank
@@ -498,17 +497,6 @@ class ComposedPipelineBase(ABC):
self._stage_name_mapping[stage_name] = stage
setattr(self, stage_name, stage)
def profile(self, is_start: bool = True):
if self.profiler is None:
raise RuntimeError("Profiler is not enabled.")
if is_start:
self.profiler.start()
else:
self.profiler.stop()
# only print profiler results on rank 0
if self.local_rank == 0:
print(self.profiler.key_averages().table(sort_by="self_cuda_time_total"))
# TODO(will): don't hardcode no_grad
@torch.no_grad()
def forward(
+116 -56
View File
@@ -5,8 +5,9 @@ The profiler is shared across the process; this module adds a light-weight
controller that gates collection based on named *regions*. Regions may be
enabled through dedicated environment variables (e.g.
``FASTVIDEO_TORCH_PROFILE_MODEL_LOADING=1``) or via the consolidated
``FASTVIDEO_TORCH_PROFILE_REGIONS`` comma-separated list (e.g.
``FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit``).
``FASTVIDEO_TORCH_PROFILE_REGIONS`` comma-separated list. Short names work
(``FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_train`` resolves the
``profiler_region_`` prefix automatically).
Typical usage from client code::
@@ -34,7 +35,6 @@ from fastvideo.logger import init_logger
logger = init_logger(__name__)
_GLOBAL_PROFILER: torch.profiler.profile | None = None
_GLOBAL_CONTROLLER: TorchProfilerController | None = None
@@ -44,7 +44,6 @@ class ProfilerRegion:
name: str
description: str
default_enabled: bool = False
def __post_init__(self) -> None:
if not self.name or self.name.strip() != self.name:
@@ -63,8 +62,6 @@ def _normalize_token(token: str) -> str:
def register_profiler_region(
name: str,
description: str,
*,
default_enabled: bool = False,
) -> None:
"""Register a profiler region so configuration can validate inputs."""
@@ -75,16 +72,22 @@ def register_profiler_region(
region = ProfilerRegion(
name=canonical,
description=description,
default_enabled=bool(default_enabled),
)
_REGISTERED_REGIONS[canonical] = region
def resolve_profiler_region(name: str) -> ProfilerRegion | None:
"""Return the registered region matching ``name`` or ``None`` if absent."""
"""Return the registered region for ``name`` (long or short form).
Accepts both the canonical name and the short form without the
``profiler_region_`` prefix, so ``REGIONS=inference_denoising`` works.
"""
canonical = _normalize_token(name)
return _REGISTERED_REGIONS.get(canonical)
region = _REGISTERED_REGIONS.get(canonical)
if region is None and not canonical.startswith("profiler_region_"):
region = _REGISTERED_REGIONS.get(f"profiler_region_{canonical}")
return region
def list_profiler_regions() -> list[ProfilerRegion]:
@@ -99,17 +102,6 @@ _DEFAULT_ACTIVITIES: tuple[torch.profiler.ProfilerActivity, ...] = (
)
def get_global_profiler() -> torch.profiler.profile | None:
"""Return the global profiler instance if one was created."""
return _GLOBAL_PROFILER
def set_global_profiler(profiler: torch.profiler.profile | None) -> None:
global _GLOBAL_PROFILER
_GLOBAL_PROFILER = profiler
def get_global_controller() -> TorchProfilerController | None:
return _GLOBAL_CONTROLLER
@@ -122,7 +114,6 @@ def set_global_controller(controller: TorchProfilerController | None) -> None:
register_profiler_region(
name="profiler_region_model_loading",
description="Module/model loading during pipeline initialization.",
default_enabled=False,
)
# register_profiler_region(
# name="profiler_region_inference_pre_denoising",
@@ -149,11 +140,11 @@ register_profiler_region(
)
register_profiler_region(
name="profiler_region_training_train_one_step",
description="High-level step orchestration in the training loop.",
description="Single optimizer step including forward/backward passes.",
)
register_profiler_region(
name="profiler_region_training_train",
description="Single optimizer step including forward/backward passes.",
description="High-level step orchestration in the training loop.",
)
# distillation specific regions
@@ -204,15 +195,18 @@ def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
profile_memory=envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
with_stack=envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
with_flops=envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
schedule=torch.profiler.schedule(
wait=envs.FASTVIDEO_TORCH_PROFILER_WAIT_STEPS,
warmup=envs.FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS,
active=envs.FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS,
),
# No schedule: nothing in the codebase calls profiler.step(), so a
# wait/warmup schedule never advances and the profiler records nothing.
# Region toggling gates collection; the single trace exports at stop().
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir, use_gzip=True),
)
controller = TorchProfilerController(profiler, _DEFAULT_ACTIVITIES)
controller._trace_dir = trace_dir
controller.start()
# The trace only exports at stop(); inference paths have no shutdown hook
# that calls it, so register one. stop() is idempotent.
import atexit
atexit.register(controller.stop)
logger.info("Torch profiler started")
return controller
@@ -287,16 +281,16 @@ class TorchProfilerController:
FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit \
python fastvideo/training/wan_training_pipeline.py ...
Wrapping a code block in a custom region::
Wrapping a code block in a registered region::
controller = TorchProfilerController(profiler, activities)
with controller.region("training_validation"):
run_validation_epoch()
from fastvideo.profiler import profiler_region
Adding a new region requires three steps:
1. Define an env var in ``envs.py``.
2. Add a default entry to ``register_profiler_region`` in this module.
3. Wrap the target code in :meth:`region` using the new name.
with profiler_region("inference_denoising"):
run_denoising_loop()
Adding a new region requires two steps:
1. Register it with ``register_profiler_region`` in this module.
2. Wrap the target code in :func:`profiler_region` using the new name.
"""
def __init__(
@@ -317,10 +311,13 @@ class TorchProfilerController:
self._profiler = profiler
self._activities = activities_tuple
self._config = config or TorchProfilerConfig.from_env()
self._collection_enabled = False
# torch.profiler collects from start(); reflect that so the initial
# _set_collection(False) in start() actually toggles it off instead of
# short-circuiting (which captured everything before the first region).
self._collection_enabled = True
self._active_region_depth = 0
self._trace_dir: str | None = None
logger.info("PROFILER: TorchProfilerController initialized with config: %s", self._config)
set_global_profiler(self._profiler)
set_global_controller(self)
@property
@@ -336,7 +333,10 @@ class TorchProfilerController:
if self._profiler is None:
return False
return self._config.regions.get(region, False)
resolved = resolve_profiler_region(region)
if resolved is None:
return False
return self._config.regions.get(resolved.name, False)
def _set_collection(self, enabled: bool) -> None:
if self._profiler is None:
@@ -348,6 +348,8 @@ class TorchProfilerController:
self._profiler.toggle_collection_dynamic(enabled, self._activities)
self._collection_enabled = enabled
_warned_unregistered: set[str] = set()
@contextlib.contextmanager
def region(self, region: str):
"""Context manager that enables profiling for ``region`` if configured."""
@@ -356,24 +358,42 @@ class TorchProfilerController:
yield
return
if resolve_profiler_region(region) is None:
# a typo here would otherwise silently profile nothing, forever
if region not in self._warned_unregistered:
self._warned_unregistered.add(region)
logger.warning(
"PROFILER: region %r is not registered (typo?); available: %s", region,
", ".join(r.name for r in list_profiler_regions()))
yield
return
if not self.is_region_enabled(region):
yield
return
with torch.profiler.record_function(f"fastvideo.region::{region}"):
self._active_region_depth += 1
if self._active_region_depth == 1:
logger.info("PROFILER: Setting collection to True (depth=%s) for region %s", self._active_region_depth,
region)
self._set_collection(True)
try:
# NVTX range so the same region names are visible in nsys timelines
nvtx = torch.cuda.is_available()
if nvtx:
torch.cuda.nvtx.range_push(f"fastvideo.region::{region}")
self._active_region_depth += 1
if self._active_region_depth == 1:
logger.info("PROFILER: Setting collection to True (depth=%s) for region %s", self._active_region_depth,
region)
self._set_collection(True)
try:
# record_function opens after collection is enabled so the region
# marker itself lands in the trace.
with torch.profiler.record_function(f"fastvideo.region::{region}"):
yield
finally:
self._active_region_depth -= 1
logger.info("PROFILER: Decreasing active region depth to %s", self._active_region_depth)
if self._active_region_depth == 0:
logger.info("PROFILER: Setting collection to False upon exiting region %s", region)
self._set_collection(False)
finally:
self._active_region_depth -= 1
logger.info("PROFILER: Decreasing active region depth to %s", self._active_region_depth)
if self._active_region_depth == 0:
logger.info("PROFILER: Setting collection to False upon exiting region %s", region)
self._set_collection(False)
if nvtx:
torch.cuda.nvtx.range_pop()
def start(self) -> None:
"""Start the profiler and pause collection until a region is entered."""
@@ -388,6 +408,34 @@ class TorchProfilerController:
self._set_collection(False)
logger.info("PROFILER: Profiler started with collection disabled")
def _write_summary(self) -> None:
"""Compact per-rank op summary next to the trace: a key_averages table
and a JSON with input shapes, so operator-split analysis does not
require parsing multi-GB chrome traces."""
if self._profiler is None or not self._trace_dir:
return
try:
import json as _json
import os as _os
rank = _os.environ.get("RANK", "0")
averages = self._profiler.key_averages(group_by_input_shape=True)
stem = _os.path.join(self._trace_dir, f"summary_rank{rank}")
with open(f"{stem}.txt", "w") as fh:
fh.write(averages.table(sort_by="self_device_time_total", row_limit=60))
rows = [{
"name": e.key,
"shapes": str(e.input_shapes),
"self_device_us": e.self_device_time_total,
"device_us": e.device_time_total,
"count": e.count,
} for e in averages]
with open(f"{stem}.json", "w") as fh:
_json.dump(rows, fh)
if rank == "0":
logger.info("PROFILER: summary written to %s.txt", stem)
except Exception: # noqa: BLE001 -- summaries must never break shutdown
logger.exception("PROFILER: summary generation failed")
def stop(self) -> None:
"""Stop the profiler after disabling collection and clearing state."""
@@ -396,9 +444,10 @@ class TorchProfilerController:
logger.info("PROFILER: Stopping profiler...")
self._profiler.stop()
self._write_summary()
self._profiler = None # makes stop() idempotent (atexit may re-enter)
logger.info("PROFILER: Profiler stopped")
self._active_region_depth = 0
set_global_profiler(None)
set_global_controller(None)
@property
@@ -411,9 +460,20 @@ class TorchProfilerController:
def activities(self) -> tuple[torch.profiler.ProfilerActivity, ...]:
return tuple(self._activities)
@property
def profiler(self) -> torch.profiler.profile | None:
return self._profiler
@contextlib.contextmanager
def profiler_region(region: str):
"""Module-level region context manager bound to the global controller.
A no-op when no profiler is configured — stages can use this without the
get_global_controller()/nullcontext dance.
"""
controller = get_global_controller()
if controller is None:
yield
return
with controller.region(region):
yield
def profile_region(region: str) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
@@ -0,0 +1,107 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU contract tests for the profiler region system.
The region system's failure mode is silent: a profiler that records nothing
(the old wait/warmup schedule bug), a trace that never exports (no shutdown
hook), or a typo'd region name that no-ops forever. Each test here runs a
child process — the profiler is process-global and its export happens at
atexit, so in-process tests can't observe the contract.
"""
from __future__ import annotations
import glob
import gzip
import json
import os
import subprocess
import sys
# Five-window child: ops before any region, inside a region, between regions,
# inside a second (short-named) region, after the last region. Exits without
# calling stop() — export must happen via the atexit hook. Each window is
# labeled with a record_function marker so the parent can check which windows
# the trace actually captured.
_CHILD = r"""
import torch
from fastvideo.profiler import get_or_create_profiler, profiler_region
controller = get_or_create_profiler("{trace_dir}")
def burn(tag):
with torch.profiler.record_function(tag):
torch.mm(torch.ones(8, 8), torch.ones(8, 8))
burn("win_pre")
with profiler_region("profiler_region_inference_denoising"):
burn("win_region1")
burn("win_between")
with profiler_region("model_loading"): # short name must resolve
burn("win_region2")
for _ in range(2): # warn-once: second use must not log again
with profiler_region("definitely_not_a_region"):
burn("win_typo")
burn("win_post")
# no controller.stop(): atexit owns the export
"""
def _run_child(tmp_path):
trace_dir = str(tmp_path / "traces")
env = os.environ.copy()
env["FASTVIDEO_TORCH_PROFILER_DIR"] = trace_dir
env["FASTVIDEO_TORCH_PROFILE_REGIONS"] = "inference_denoising,model_loading"
proc = subprocess.run(
[sys.executable, "-c", _CHILD.format(trace_dir=trace_dir)],
env=env,
capture_output=True,
text=True,
timeout=300,
)
assert proc.returncode == 0, proc.stderr
return trace_dir, proc.stdout + proc.stderr
def _trace_event_names(trace_dir):
traces = glob.glob(os.path.join(trace_dir, "**", "*.json*"), recursive=True)
traces = [t for t in traces if "summary" not in os.path.basename(t)]
assert traces, f"no trace exported to {trace_dir} (atexit hook missing?)"
opener = gzip.open if traces[0].endswith(".gz") else open
with opener(traces[0], "rt") as fh:
events = json.load(fh).get("traceEvents", [])
return {e.get("name", "") for e in events}
def test_regions_gate_collection_and_atexit_exports(tmp_path):
trace_dir, output = _run_child(tmp_path)
names = _trace_event_names(trace_dir)
assert "win_region1" in names, "op inside an enabled region was not captured"
assert "win_region2" in names, "short region name did not resolve/capture"
assert "fastvideo.region::profiler_region_inference_denoising" in names
for leaked in ("win_pre", "win_between", "win_typo", "win_post"):
assert leaked not in names, f"op outside any region leaked into trace: {leaked}"
# unregistered region warns exactly once across repeated uses
assert output.count("definitely_not_a_region") >= 1
assert output.count("is not registered") == 1
# per-rank op summary written next to the trace
summaries = glob.glob(os.path.join(trace_dir, "summary_rank0.*"))
assert sorted(os.path.splitext(s)[1] for s in summaries) == [".json", ".txt"]
def test_noop_without_profiler_dir(tmp_path):
# profiler_region must be a clean no-op when profiling is not configured
child = (
"from fastvideo.profiler import profiler_region\n"
"with profiler_region('inference_denoising'):\n"
" x = 1\n"
"assert x == 1\n"
)
env = os.environ.copy()
env.pop("FASTVIDEO_TORCH_PROFILER_DIR", None)
env.pop("FASTVIDEO_TORCH_PROFILE_REGIONS", None)
proc = subprocess.run([sys.executable, "-c", child], env=env,
capture_output=True, text=True, timeout=300)
assert proc.returncode == 0, proc.stderr