Compare commits

...
17 changed files with 1002 additions and 262 deletions
+68 -14
View File
@@ -9,33 +9,82 @@ FastVideo exposes a process-wide torch profiler that you can enable via environm
```bash
FASTVIDEO_TORCH_PROFILER_DIR=/mnt/traces/fastvideo \
FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_model_loading,profiler_region_training_step"
FASTVIDEO_TORCH_PROFILE_REGIONS="model_loading,training_train_one_step" \
bash examples/train/run.sh /path/to/config.yaml
```
All profiled regions must be registered in `fastvideo.profiler`; the current list includes:
The `profiler_region_` prefix is optional. All profiled regions must be
registered in `fastvideo.profiler`; the current list includes:
- `profiler_region_model_loading` — pipeline/module loading
- `profiler_region_inference_pre_denoising`
- `profiler_region_inference_denoising`
- `profiler_region_inference_post_denoising`
- `profiler_region_training_checkpoint_saving`
- `profiler_region_training_dit`
- `profiler_region_training_train` — the complete training run
- `profiler_region_training_train_one_step` — one complete optimizer step
- `profiler_region_training_dataloader` — fetch the next batch in the trainer process
- `profiler_region_training_forward` — method forward and loss computation
- `profiler_region_training_backward` — backward pass
- `profiler_region_training_optimizer` — gradient clipping, optimizer/scheduler step, and zeroing gradients
- `profiler_region_training_callbacks` — end-of-step callbacks such as EMA updates
- `profiler_region_training_validation`
- `profiler_region_training_epoch`
- `profiler_region_training_step`
- `profiler_region_training_backward`
- `profiler_region_training_optimizer`
- `profiler_region_training_save_checkpoint`
- `profiler_region_distillation_teacher_forward`
- `profiler_region_distillation_student_forward`
- `profiler_region_distillation_loss`
- `profiler_region_distillation_update`
- `profiler_region_dmd2_student_rollout` — one DMD2 student rollout, including simulated prefix steps
- `profiler_region_dmd2_generator_loss` — teacher/critic scoring for the generator loss
- `profiler_region_dmd2_critic_loss` — critic flow-matching loss, including its student rollout
### Profiling modular training
The YAML-driven trainer under `fastvideo/train/` initializes and flushes the
profiler automatically. For example, this captures six DMD2 steps: one warmup
step, four steady-state critic updates, and the configured 1-in-5 generator
update, without changing the method's update cadence:
```bash
TRACE_DIR="$(pwd)/profiler_traces/wan_dmd2"
mkdir -p "$TRACE_DIR"
NUM_GPUS=4 \
WANDB_MODE=offline \
FASTVIDEO_TORCH_PROFILER_DIR="$TRACE_DIR" \
FASTVIDEO_TORCH_PROFILE_REGIONS="training_train,training_train_one_step,training_dataloader,training_forward,training_backward,training_optimizer,training_callbacks,dmd2_student_rollout,dmd2_generator_loss,dmd2_critic_loss" \
bash examples/train/run.sh \
examples/train/configs/distribution_matching/wan/dmd2_t2v.yaml \
--training.loop.max_train_steps 6 \
--training.checkpoint.training_state_checkpointing_steps 0 \
--callbacks.validation.every_steps 0
```
The two overrides disable validation and checkpoint writes so their I/O does
not contaminate training-step measurements. Remove them when profiling those
regions. Methods that manage optimization internally emit the enclosing
`training_train_one_step` region but do not emit the generic
dataloader/forward/backward/optimizer child regions because those boundaries
belong to the method. With `dataloader_num_workers > 0`, the
`training_dataloader` region measures the training process waiting for and
receiving the next batch; work performed inside dataloader worker subprocesses
does not appear in that process's torch-profiler trace.
The modular trainer also logs `dataloader_time_sec` on every ordinary training
step (including DMD2), so routine runs can monitor rank 0's
`next(dataloader)` wait without enabling the torch profiler. It is
intentionally not reduced across ranks; cross-rank synchronization would
perturb the hot path being measured.
While profiling is enabled, FastVideo records additional annotations:
- `fastvideo.region::<name>` spans are emitted when entering a region.
- `fastvideo.profiler.enable_collection` / `fastvideo.profiler.disable_collection` events mark when torch profiler collection is toggled on or off.
Only one profiler instance is created per process; subsequent pipelines reuse the same controller. If you set `FASTVIDEO_TORCH_PROFILE_REGIONS` incorrectly (e.g. misspelled name), FastVideo logs a warning and ignores that entry.
Only one profiler controller is created per process; subsequent pipelines
reuse it. Each outermost enabled region invocation produces one complete
CPU/CUDA trace segment. Enabled regions nested inside it are annotations in
that segment. Include an enclosing region such as `training_train_one_step`
when selecting its child regions so they share one trace and profiler startup
does not perturb each phase independently. If you set
`FASTVIDEO_TORCH_PROFILE_REGIONS` incorrectly (e.g. misspelled name), FastVideo
logs a warning and ignores that entry.
Additional knobs:
@@ -44,13 +93,18 @@ Additional knobs:
- `FASTVIDEO_TORCH_PROFILER_WITH_STACK`
- `FASTVIDEO_TORCH_PROFILER_WITH_FLOPS`
Traces can be visualized using <https://ui.perfetto.dev/>.
Traces can be visualized using <https://ui.perfetto.dev/>. Each rank also
writes `summary_rank<N>_segment<MMMM>_<region>.txt` and the corresponding JSON
beside every trace segment for a quick operator-level view without loading the
full timeline. The JSON contains CPU and device totals; input-shape grouping is
only enabled with `FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1` because it makes
summary generation materially more expensive on large traces.
### Best Practices
- Keep the profiled step count small; traces can be large and slow down job shutdown while the profiler flushes data.
- After profiling, clean up trace directories to avoid filling disk storage.
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `with self.profiler_controller.region("your_region"):` or the `@profile_region` decorator.
- When adding new regions, register them in `fastvideo.profiler` and wrap the corresponding code block with `profiler_region("your_region")` or the `@profile_region` decorator.
## Related: Activation Trace Mode
+3 -1
View File
@@ -431,7 +431,9 @@ callbacks:
```
The EMA callback owns its own state and checkpoints independently — EMA weights
are saved and restored automatically on resume.
are saved and restored automatically on resume. It advances only after the
training method reports that the student optimizer updated; alternating methods
such as DMD2 therefore skip EMA work on critic-only steps.
### ValidationCallback
+1 -6
View File
@@ -250,12 +250,7 @@ environment_variables: dict[str, Callable[[], Any]] = {
# not profile flops.
"FASTVIDEO_TORCH_PROFILER_WITH_FLOPS":
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.
# Warmup steps per profiling cycle (torch.profiler.schedule warmup parameter)
# Defaults to 1 if not set.
# Active steps per profiling cycle (torch.profiler.schedule active parameter)
# Defaults to 2 if not set.
# Comma-separated names of registered profiler regions to capture.
"FASTVIDEO_TORCH_PROFILE_REGIONS":
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
+185 -84
View File
@@ -2,18 +2,17 @@
"""Utilities for managing the PyTorch profiler within FastVideo.
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. Short names work
controller that gates collection based on named *regions*. Regions are enabled
through the ``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::
controller = TorchProfilerController(profiler, activities)
with controller.region("training_dit"):
controller = get_or_create_profiler("/tmp/fastvideo-traces")
with controller.region("training_train"):
run_training_step()
controller.stop()
To introduce a new region, register it via :func:`register_profiler_region`
and wrap the corresponding code in :meth:`TorchProfilerController.region`.
@@ -22,11 +21,11 @@ and wrap the corresponding code in :meth:`TorchProfilerController.region`.
from __future__ import annotations
import contextlib
import functools
import os
from collections.abc import Callable, Iterable
from dataclasses import dataclass
from typing import Any
from collections.abc import Callable
import functools
from collections.abc import Iterable
import torch
@@ -142,6 +141,26 @@ register_profiler_region(
name="profiler_region_training_train_one_step",
description="Single optimizer step including forward/backward passes.",
)
register_profiler_region(
name="profiler_region_training_forward",
description="Training method forward pass and loss computation.",
)
register_profiler_region(
name="profiler_region_training_dataloader",
description="Fetch the next training batch in the trainer process.",
)
register_profiler_region(
name="profiler_region_training_backward",
description="Training backward pass.",
)
register_profiler_region(
name="profiler_region_training_optimizer",
description="Gradient clipping, optimizer/scheduler steps, and zero_grad.",
)
register_profiler_region(
name="profiler_region_training_callbacks",
description="End-of-step training callbacks such as EMA updates.",
)
register_profiler_region(
name="profiler_region_training_train",
description="High-level step orchestration in the training loop.",
@@ -165,6 +184,21 @@ register_profiler_region(
description="Parameter updates specific to distillation workflows.",
)
# DMD2 method regions. These sit inside ``training_forward`` and make the
# method's multi-model forward path distinguishable in a single trace.
register_profiler_region(
name="profiler_region_dmd2_student_rollout",
description="DMD2 student rollout, including its simulated prefix steps.",
)
register_profiler_region(
name="profiler_region_dmd2_generator_loss",
description="DMD2 generator loss, including teacher and critic scoring.",
)
register_profiler_region(
name="profiler_region_dmd2_critic_loss",
description="DMD2 critic flow-matching loss, including its student rollout.",
)
def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
"""Create or reuse the process-wide torch profiler controller."""
@@ -189,25 +223,28 @@ def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
)
logger.info("FASTVIDEO_TORCH_PROFILE_REGIONS=%s", envs.FASTVIDEO_TORCH_PROFILE_REGIONS)
profiler = torch.profiler.profile(
activities=_DEFAULT_ACTIVITIES,
record_shapes=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
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,
# 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),
def profiler_factory() -> Any:
return torch.profiler.profile(
activities=_DEFAULT_ACTIVITIES,
record_shapes=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES,
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,
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir, use_gzip=True),
)
controller = TorchProfilerController(
None,
_DEFAULT_ACTIVITIES,
profiler_factory=profiler_factory,
trace_dir=trace_dir,
)
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.
# Region exit normally exports each trace segment. Keep an atexit hook for
# exceptions or process shutdown while a region is still active.
import atexit
atexit.register(controller.stop)
logger.info("Torch profiler started")
logger.info("Torch profiler armed; collection starts at the first enabled region")
return controller
@@ -260,7 +297,14 @@ class TorchProfilerConfig:
class TorchProfilerController:
"""Helper that toggles torch profiler collection for named regions.
"""Create complete torch-profiler trace segments for named regions.
PyTorch's dynamic CUDA collection toggle can fail to re-enable CUPTI on
some supported stacks. In that failure mode it emits CPU operators while
silently dropping every CUDA kernel. This controller therefore starts a
fresh profiler at each outermost enabled region and stops it when that
region exits. Nested enabled regions become annotations in the same
CPU/CUDA trace segment.
Parameters
----------
@@ -273,12 +317,17 @@ class TorchProfilerController:
config:
Optional :class:`TorchProfilerConfig`. If omitted, :meth:`from_env`
constructs one during initialization.
profiler_factory:
Factory for fresh profiler instances. Required to profile more than
one outermost region invocation.
trace_dir:
Directory for per-segment summaries.
Examples
--------
Enabling an existing region from the command line::
FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_dit \
FASTVIDEO_TORCH_PROFILE_REGIONS=model_loading,training_train \
python fastvideo/training/wan_training_pipeline.py ...
Wrapping a code block in a registered region::
@@ -299,24 +348,29 @@ class TorchProfilerController:
activities: Iterable[torch.profiler.ProfilerActivity],
config: TorchProfilerConfig | None = None,
disabled: bool = False,
profiler_factory: Callable[[], Any] | None = None,
trace_dir: str | None = None,
) -> None:
activities_tuple = tuple(activities)
existing = get_global_controller()
if existing is not None and not disabled:
raise RuntimeError("TorchProfilerController already initialized globally. Use get_global_controller().")
self._profiler = profiler
self._profiler_factory = profiler_factory
self._activities = activities_tuple
self._trace_dir = trace_dir
self._segment_index = 0
self._segment_region: str | None = None
self._active_region_depth = 0
self._collection_enabled = False
if disabled:
self._profiler = None
self._configured = False
self._armed = False
return
self._profiler = profiler
self._activities = activities_tuple
self._config = config or TorchProfilerConfig.from_env()
# 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
self._configured = True
self._armed = False
logger.info("PROFILER: TorchProfilerController initialized with config: %s", self._config)
set_global_controller(self)
@@ -324,29 +378,61 @@ class TorchProfilerController:
def is_enabled(self) -> bool:
"""Return ``True`` when the underlying profiler is collecting."""
if self._profiler is None:
return False
return self._collection_enabled
def is_region_enabled(self, region: str) -> bool:
"""Return ``True`` if ``region`` should be collected."""
if self._profiler is None:
if not self.has_profiler:
return 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:
def _new_profiler(self) -> Any:
if self._profiler is not None:
profiler = self._profiler
self._profiler = None
return profiler
if self._profiler_factory is None:
raise RuntimeError("Torch profiler cannot start another trace segment without a profiler_factory")
return self._profiler_factory()
def _start_segment(self, region: str) -> None:
if self._collection_enabled:
return
if self._collection_enabled == enabled:
self._profiler = self._new_profiler()
logger.info(
"PROFILER: Starting segment %d for region %s",
self._segment_index,
region,
)
self._profiler.start()
self._segment_region = region
self._collection_enabled = True
def _finish_segment(self) -> None:
if self._profiler is None or not self._collection_enabled:
return
event = ("fastvideo.profiler.enable_collection" if enabled else "fastvideo.profiler.disable_collection")
with torch.profiler.record_function(event):
self._profiler.toggle_collection_dynamic(enabled, self._activities)
self._collection_enabled = enabled
profiler = self._profiler
segment_index = self._segment_index
segment_region = self._segment_region or "unknown"
logger.info(
"PROFILER: Stopping segment %d for region %s",
segment_index,
segment_region,
)
profiler.stop()
self._write_summary(
profiler,
segment_index=segment_index,
segment_region=segment_region,
)
self._profiler = None
self._collection_enabled = False
self._segment_region = None
self._segment_index += 1
_warned_unregistered: set[str] = set()
@@ -354,7 +440,7 @@ class TorchProfilerController:
def region(self, region: str):
"""Context manager that enables profiling for ``region`` if configured."""
if self._profiler is None:
if not self.has_profiler:
yield
return
@@ -371,64 +457,70 @@ class TorchProfilerController:
yield
return
# NVTX range so the same region names are visible in nsys timelines
if self._active_region_depth == 0:
self._start_segment(region)
# NVTX range so the same region names are visible in nsys timelines.
# Push after profiler startup so Kineto also records the annotation.
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)
if nvtx:
torch.cuda.nvtx.range_pop()
try:
if nvtx:
torch.cuda.nvtx.range_pop()
finally:
# Close NVTX before stopping Kineto so both profilers see a
# balanced outermost range in the exported segment.
if self._active_region_depth == 0:
self._finish_segment()
def start(self) -> None:
"""Start the profiler and pause collection until a region is entered."""
"""Arm the controller; collection begins at an enabled region."""
logger.info("PROFILER: Starting profiler...")
if self._profiler is None:
if not self._configured:
return
self._profiler.start()
logger.info("PROFILER: Profiler started")
# Profiler starts with collection disabled by default.
logger.info("PROFILER: Setting collection to False")
self._set_collection(False)
logger.info("PROFILER: Profiler started with collection disabled")
self._armed = True
logger.info("PROFILER: Controller armed")
def _write_summary(self) -> None:
def _write_summary(
self,
profiler: Any,
*,
segment_index: int,
segment_region: str,
) -> 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:
if 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:
rank = os.environ.get("RANK", "0")
short_region = segment_region.removeprefix("profiler_region_")
stem = os.path.join(
self._trace_dir,
f"summary_rank{rank}_segment{segment_index:04d}_{short_region}",
)
averages = profiler.key_averages(group_by_input_shape=envs.FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES, )
with open(f"{stem}.txt", "w", encoding="utf-8") 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_cpu_us": e.self_cpu_time_total,
"cpu_us": e.cpu_time_total,
"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:
with open(f"{stem}.json", "w", encoding="utf-8") as fh:
_json.dump(rows, fh)
if rank == "0":
logger.info("PROFILER: summary written to %s.txt", stem)
@@ -436,24 +528,25 @@ class TorchProfilerController:
logger.exception("PROFILER: summary generation failed")
def stop(self) -> None:
"""Stop the profiler after disabling collection and clearing state."""
"""Flush any active segment and disable this controller."""
if self._profiler is None:
if not self._configured:
return
logger.info("PROFILER: Stopping profiler...")
self._profiler.stop()
self._write_summary()
self._profiler = None # makes stop() idempotent (atexit may re-enter)
self._finish_segment()
self._profiler = None
self._configured = False
self._armed = False
logger.info("PROFILER: Profiler stopped")
self._active_region_depth = 0
set_global_controller(None)
@property
def has_profiler(self) -> bool:
"""Return ``True`` when a profiler instance is available."""
"""Return ``True`` when this controller is configured and armed."""
return self._profiler is not None
return self._configured and self._armed
@property
def activities(self) -> tuple[torch.profiler.ProfilerActivity, ...]:
@@ -476,13 +569,21 @@ def profiler_region(region: str):
def profile_region(region: str) -> Callable[[Callable[..., Any]], Callable[..., Any]]:
"""Wrap a bound method so it runs inside a profiler region if available."""
"""Wrap a bound method so it runs inside a profiler region if available.
Prefer a controller attached to the instance, then fall back to the
process-wide controller. The fallback lets lightweight owners such as the
modular trainer and its callbacks add regions without threading profiler
plumbing through their public constructors.
"""
def decorator(fn: Callable[..., Any]) -> Callable[..., Any]:
@functools.wraps(fn)
def wrapped(self, *args, **kwargs):
controller = getattr(self, "profiler_controller", None)
if controller is None:
controller = get_global_controller()
if controller is None or not controller.has_profiler:
return fn(self, *args, **kwargs)
with controller.region(region):
@@ -16,6 +16,9 @@ import os
import subprocess
import sys
import pytest
import torch
# 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
@@ -61,19 +64,21 @@ def _run_child(tmp_path):
return trace_dir, proc.stdout + proc.stderr
def _trace_event_names(trace_dir):
def _trace_events(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}
events = []
for trace in traces:
opener = gzip.open if trace.endswith(".gz") else open
with opener(trace, "rt") as fh:
events.extend(json.load(fh).get("traceEvents", []))
return events
def test_regions_gate_collection_and_atexit_exports(tmp_path):
trace_dir, output = _run_child(tmp_path)
names = _trace_event_names(trace_dir)
names = {event.get("name", "") for event in _trace_events(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"
@@ -87,8 +92,61 @@ def test_regions_gate_collection_and_atexit_exports(tmp_path):
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"]
summaries = glob.glob(os.path.join(trace_dir, "summary_rank0_segment*.*"))
assert sorted(os.path.splitext(s)[1] for s in summaries) == [
".json",
".json",
".txt",
".txt",
]
with open(next(s for s in summaries if s.endswith(".json")), encoding="utf-8") as fh:
summary_rows = json.load(fh)
assert summary_rows
assert {
"name",
"shapes",
"self_cpu_us",
"cpu_us",
"self_device_us",
"device_us",
"count",
} <= summary_rows[0].keys()
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required")
def test_cuda_region_exports_kernel_events(tmp_path):
trace_dir = str(tmp_path / "cuda_traces")
child = r"""
import torch
from fastvideo.profiler import get_or_create_profiler, profiler_region
controller = get_or_create_profiler({trace_dir!r})
x = torch.randn(1024, 1024, device="cuda")
y = torch.randn(1024, 1024, device="cuda")
torch.cuda.synchronize()
with profiler_region("training_forward"):
torch.mm(x, y)
torch.cuda.synchronize()
controller.stop()
""".format(trace_dir=trace_dir)
env = os.environ.copy()
env["FASTVIDEO_TORCH_PROFILER_DIR"] = trace_dir
env["FASTVIDEO_TORCH_PROFILE_REGIONS"] = "training_forward"
proc = subprocess.run(
[sys.executable, "-c", child],
env=env,
capture_output=True,
text=True,
timeout=300,
)
assert proc.returncode == 0, proc.stderr
events = _trace_events(trace_dir)
categories = {event.get("cat", "") for event in events}
names = {event.get("name", "") for event in events}
assert "kernel" in categories
assert "cuda_runtime" in categories
assert "fastvideo.region::training_forward" in names
def test_noop_without_profiler_dir(tmp_path):
@@ -42,9 +42,15 @@ class _Method:
self,
transformer: torch.nn.Module | None,
tracker: Any | None = None,
updated_iterations: set[int] | None = None,
) -> None:
self.student = _Student(transformer)
self.tracker = tracker
self.updated_iterations = updated_iterations
def did_update_role(self, role: str, iteration: int) -> bool:
assert role == "student"
return (self.updated_iterations is None or iteration in self.updated_iterations)
def _tiny_transformer(*, fill: float = 0.0) -> torch.nn.Module:
@@ -155,6 +161,44 @@ class TestOnTrainingStepEnd:
assert any(payload.get("ema/decay") == 0.99 and step == 0 for payload, step in tracker.entries)
def test_only_updates_after_student_optimizer_step(self) -> None:
transformer = _tiny_transformer(fill=1.0)
tracker = _RecordingTracker()
method = _Method(
transformer,
tracker=tracker,
updated_iterations={5},
)
cb = EMACallback(decay=0.9, start_iter=0)
cb.on_train_start(method, iteration=0)
with torch.no_grad():
transformer.weight.fill_(7.0)
cb.on_training_step_end(method, loss_dict={}, iteration=4)
assert not cb._ema_started
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 1.0),
)
assert tracker.entries == []
cb.on_training_step_end(method, loss_dict={}, iteration=5)
assert cb._ema_started
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 7.0),
)
assert tracker.entries == [({"ema/decay": 0.9}, 5)]
with torch.no_grad():
transformer.weight.fill_(11.0)
cb.on_training_step_end(method, loss_dict={}, iteration=6)
assert torch.allclose(
cb.student_ema.shadow["weight"],
torch.full((2, 4), 7.0),
)
assert tracker.entries == [({"ema/decay": 0.9}, 5)]
# ---------------------------------------------------------------------------
# C. ema_context
@@ -0,0 +1,189 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from types import SimpleNamespace
from typing import Any
import pytest
import torch
from fastvideo.train.methods.distribution_matching.dmd2 import DMD2Method
class _RecordingStudent:
def __init__(self) -> None:
self.predict_calls: list[torch.Tensor] = []
self.add_noise_calls: list[torch.Tensor] = []
def predict_x0(
self,
noisy_latents: torch.Tensor,
timestep: torch.Tensor,
batch: Any,
**kwargs: Any,
) -> torch.Tensor:
del batch, kwargs
self.predict_calls.append(timestep.detach().clone())
return noisy_latents + timestep.to(noisy_latents.dtype)
def add_noise(
self,
clean_latents: torch.Tensor,
noise: torch.Tensor,
timestep: torch.Tensor,
) -> torch.Tensor:
self.add_noise_calls.append(timestep.detach().clone())
return clean_latents + noise
def _rollout_method(seed: int = 1234) -> tuple[DMD2Method, _RecordingStudent]:
method = object.__new__(DMD2Method)
torch.nn.Module.__init__(method)
student = _RecordingStudent()
method.student = student
method._rollout_mode = "simulate"
method._cfg_uncond = None
method._denoising_step_list = torch.tensor([1000, 750, 500, 250])
method.cuda_generator = torch.Generator(device="cpu").manual_seed(seed)
return method, student
def _legacy_full_rollout_reference(
*,
seed: int,
target_idx: int,
shape: tuple[int, ...],
) -> tuple[torch.Tensor, torch.Tensor]:
"""Reproduce the pre-optimization simulate rollout and RNG state."""
step_list = torch.tensor([1000, 750, 500, 250])
generator = torch.Generator(device="cpu").manual_seed(seed)
current = torch.randn(shape, generator=generator)
initial = current.clone()
noise_latents: list[torch.Tensor] = []
for step_idx in range(len(step_list) - 1):
pred_clean = current + step_list[step_idx].to(current.dtype)
noise = torch.randn(shape, generator=generator)
current = pred_clean + noise
noise_latents.append(current.clone())
noisy_input: torch.Tensor
if target_idx == 0:
noisy_input = initial
else:
noisy_input = noise_latents[target_idx - 1]
output = noisy_input + step_list[target_idx].to(noisy_input.dtype)
return output, generator.get_state()
@pytest.mark.parametrize("target_idx", range(4))
def test_simulate_rollout_only_runs_required_prefix_forwards(
monkeypatch: pytest.MonkeyPatch,
target_idx: int,
) -> None:
method, student = _rollout_method()
batch = SimpleNamespace(
latents=torch.zeros((1, 2)),
dmd_latent_vis_dict={},
)
def _fixed_target(*args: Any, **kwargs: Any) -> torch.Tensor:
del args, kwargs
return torch.tensor([target_idx], dtype=torch.long)
monkeypatch.setattr(torch, "randint", _fixed_target)
method._student_rollout(batch, with_grad=True)
# One prediction per required prefix step, plus the differentiable target
# prediction. Prefix noising only happens for the required prefix steps.
assert len(student.predict_calls) == target_idx + 1
assert len(student.add_noise_calls) == target_idx
@pytest.mark.parametrize("target_idx", range(4))
def test_simulate_rollout_preserves_method_generator_progress(
monkeypatch: pytest.MonkeyPatch,
target_idx: int,
) -> None:
seed = 4321
method, _ = _rollout_method(seed)
shape = (1, 2)
batch = SimpleNamespace(
latents=torch.zeros(shape),
dmd_latent_vis_dict={},
)
def _fixed_target(*args: Any, **kwargs: Any) -> torch.Tensor:
del args, kwargs
return torch.tensor([target_idx], dtype=torch.long)
monkeypatch.setattr(torch, "randint", _fixed_target)
output = method._student_rollout(batch, with_grad=False)
# The previous implementation drew the initial latent and one noise tensor
# for each of the three possible prefix transitions. Keep consuming those
# draws so subsequent DMD2 randomness remains aligned across the change.
reference_output, reference_state = _legacy_full_rollout_reference(
seed=seed,
target_idx=target_idx,
shape=shape,
)
assert torch.equal(output, reference_output)
assert torch.equal(
method.cuda_generator.get_state(),
reference_state,
)
@pytest.mark.parametrize("target_idx", range(4))
def test_simulate_rollout_uses_global_max_prefix_without_changing_local_output(
monkeypatch: pytest.MonkeyPatch,
target_idx: int,
) -> None:
seed = 9876
method, student = _rollout_method(seed)
shape = (1, 2)
batch = SimpleNamespace(
latents=torch.zeros(shape),
dmd_latent_vis_dict={},
)
monkeypatch.setattr(
torch,
"randint",
lambda *args, **kwargs: torch.tensor([target_idx], dtype=torch.long),
)
monkeypatch.setattr(
method,
"_max_rollout_target_idx_across_ranks",
lambda sampled_idx: 3,
)
output = method._student_rollout(batch, with_grad=False)
reference_output, reference_state = _legacy_full_rollout_reference(
seed=seed,
target_idx=target_idx,
shape=shape,
)
# Every rank participates in the globally required three prefix forwards,
# then evaluates its own target. The local result and method-owned RNG
# sequence still match the legacy full rollout exactly.
assert len(student.predict_calls) == 4
assert len(student.add_noise_calls) == 3
assert torch.equal(output, reference_output)
assert torch.equal(method.cuda_generator.get_state(), reference_state)
def test_dmd2_role_update_cadence_follows_selected_optimizers() -> None:
method = object.__new__(DMD2Method)
torch.nn.Module.__init__(method)
method.method_config = {"generator_update_interval": 5}
method._student_optimizer = object()
method._critic_optimizer = object()
assert not method.did_update_role("student", iteration=4)
assert method.did_update_role("student", iteration=5)
assert method.did_update_role("critic", iteration=4)
assert not method.did_update_role("teacher", iteration=5)
@@ -0,0 +1,61 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only tests for modular training profiler lifecycle ownership."""
from __future__ import annotations
from contextlib import contextmanager
from types import SimpleNamespace
from fastvideo.train.entrypoint.train import run_training_from_config
class _RecordingProfiler:
def __init__(self) -> None:
self.events: list[tuple[str, str] | tuple[str]] = []
@contextmanager
def region(self, name: str):
self.events.append(("enter", name))
try:
yield
finally:
self.events.append(("exit", name))
def stop(self) -> None:
self.events.append(("stop", ))
def test_dry_run_profiles_model_build_and_flushes(monkeypatch) -> None:
profiler = _RecordingProfiler()
training = SimpleNamespace(
distributed=SimpleNamespace(tp_size=1, sp_size=1),
vsa_sparsity=0.0,
model_path="model",
)
cfg = SimpleNamespace(training=training)
monkeypatch.setattr(
"fastvideo.train.utils.config.load_run_config",
lambda *args, **kwargs: cfg,
)
monkeypatch.setattr(
"fastvideo.distributed.maybe_init_distributed_environment_and_model_parallel",
lambda *args, **kwargs: None,
)
monkeypatch.setattr(
"fastvideo.train.utils.builder.build_from_config",
lambda loaded_cfg: (loaded_cfg.training, object(), object(), 0),
)
monkeypatch.setattr(
"fastvideo.train.entrypoint.train.get_or_create_profiler",
lambda trace_dir: profiler,
)
run_training_from_config("unused.yaml", dry_run=True)
assert profiler.events == [
("enter", "profiler_region_model_loading"),
("exit", "profiler_region_model_loading"),
("stop", ),
]
@@ -0,0 +1,137 @@
# SPDX-License-Identifier: Apache-2.0
"""CPU-only tests for modular Trainer profiler boundaries."""
from __future__ import annotations
from contextlib import contextmanager
from types import SimpleNamespace
from typing import Any
import torch
from fastvideo.profiler import list_profiler_regions
from fastvideo.train.trainer import Trainer
from fastvideo.train.utils.training_config import TrainingConfig
class _RecordingProfiler:
def __init__(self) -> None:
self.events: list[tuple[str, str]] = []
@property
def has_profiler(self) -> bool:
return True
@contextmanager
def region(self, name: str):
self.events.append(("enter", name))
try:
yield
finally:
self.events.append(("exit", name))
class _DummyTracker:
def log(self, metrics: dict[str, float], step: int) -> None:
del metrics, step
def finish(self) -> None:
pass
class _DummyMethod:
def __init__(self) -> None:
self.weight = torch.nn.Parameter(torch.tensor(1.0))
def set_tracker(self, tracker: Any) -> None:
del tracker
def on_train_start(self) -> None:
pass
def manages_optimization(self) -> bool:
return False
def single_train_step(
self,
batch: dict[str, Any],
iteration: int,
) -> tuple[dict[str, torch.Tensor], dict[str, Any], dict[str, float]]:
del batch, iteration
return {"total_loss": self.weight.square()}, {}, {}
def backward(
self,
loss_map: dict[str, torch.Tensor],
outputs: dict[str, Any],
*,
grad_accum_rounds: int,
) -> None:
del outputs
(loss_map["total_loss"] / grad_accum_rounds).backward()
def optimizers_schedulers_step(self, iteration: int) -> None:
del iteration
def optimizers_zero_grad(self, iteration: int) -> None:
del iteration
self.weight.grad = None
def test_modular_trainer_emits_nested_step_regions(monkeypatch) -> None:
group = SimpleNamespace(rank=0, local_rank=0, rank_in_group=0, world_size=1)
profiler = _RecordingProfiler()
monkeypatch.setattr("fastvideo.train.trainer.get_world_group", lambda: group)
monkeypatch.setattr("fastvideo.train.trainer.get_sp_group", lambda: group)
monkeypatch.setattr(
"fastvideo.train.trainer.build_tracker",
lambda *args, **kwargs: _DummyTracker(),
)
monkeypatch.setattr("fastvideo.profiler._GLOBAL_CONTROLLER", profiler)
trainer = Trainer(TrainingConfig())
trainer.run(
_DummyMethod(),
dataloader=[{}],
max_steps=1,
)
assert profiler.events == [
("enter", "profiler_region_training_train"),
("enter", "profiler_region_training_train_one_step"),
("enter", "profiler_region_training_dataloader"),
("exit", "profiler_region_training_dataloader"),
("enter", "profiler_region_training_forward"),
("exit", "profiler_region_training_forward"),
("enter", "profiler_region_training_backward"),
("exit", "profiler_region_training_backward"),
("enter", "profiler_region_training_optimizer"),
("exit", "profiler_region_training_optimizer"),
("enter", "profiler_region_training_callbacks"),
("exit", "profiler_region_training_callbacks"),
("exit", "profiler_region_training_train_one_step"),
("exit", "profiler_region_training_train"),
]
def test_modular_training_regions_are_registered() -> None:
names = {region.name for region in list_profiler_regions()}
assert {
"profiler_region_model_loading",
"profiler_region_training_train",
"profiler_region_training_train_one_step",
"profiler_region_training_dataloader",
"profiler_region_training_forward",
"profiler_region_training_backward",
"profiler_region_training_optimizer",
"profiler_region_training_callbacks",
"profiler_region_training_save_checkpoint",
"profiler_region_training_validation",
"profiler_region_dmd2_student_rollout",
"profiler_region_dmd2_generator_loss",
"profiler_region_dmd2_critic_loss",
} <= names
@@ -136,4 +136,5 @@ def test_trainer_runs_validation_callback_during_training(monkeypatch, ) -> None
assert method.zero_grad_steps == [0, 1, 2, 3]
assert method.optimizer_steps == [1, 2, 3]
assert [step for _, step in tracker.logs] == [1, 2, 3]
assert all("dataloader_time_sec" in metrics for metrics, _ in tracker.logs)
assert tracker.finished is True
+2
View File
@@ -94,6 +94,8 @@ class EMACallback(Callback):
if iteration < self._start_iter:
return
if not method.did_update_role("student", iteration):
return
if not self._ema_started:
logger.info(
"Starting EMA updates at iteration %d "
+2
View File
@@ -31,6 +31,7 @@ from fastvideo.distributed import (
get_world_group,
)
from fastvideo.logger import init_logger
from fastvideo.profiler import profile_region
from fastvideo.pipelines import ForwardBatch
from fastvideo.train.callbacks.callback import Callback
from fastvideo.train.utils.instantiate import resolve_target
@@ -285,6 +286,7 @@ class ValidationCallback(Callback):
# Core validation logic
# ----------------------------------------------------------
@profile_region("profiler_region_training_validation")
def _run_validation(
self,
method: TrainingMethod,
+45 -35
View File
@@ -24,7 +24,9 @@ from typing import Any
import torch
import fastvideo.envs as envs
from fastvideo.logger import init_logger
from fastvideo.profiler import get_or_create_profiler
logger = init_logger(__name__)
@@ -74,47 +76,55 @@ def run_training_from_config(
tc.distributed.sp_size,
)
_, method, dataloader, start_step = build_from_config(cfg)
profiler_controller = get_or_create_profiler(envs.FASTVIDEO_TORCH_PROFILER_DIR, )
try:
with profiler_controller.region("profiler_region_model_loading"):
_, method, dataloader, start_step = build_from_config(cfg)
if dry_run:
logger.info("Dry-run: config parsed and "
"build_from_config succeeded.")
return
if dry_run:
logger.info("Dry-run: config parsed and "
"build_from_config succeeded.")
return
trainer = Trainer(
tc,
config=cfg.resolved_config(),
callback_configs=cfg.callbacks,
)
trainer = Trainer(
tc,
config=cfg.resolved_config(),
callback_configs=cfg.callbacks,
)
# Attach the exact YAML used for this run to the
# tracker (e.g., W&B Files).
trainer.tracker.log_file(
os.path.abspath(os.path.expanduser(config_path)),
name="run.yaml",
)
# Attach the exact YAML used for this run to the
# tracker (e.g., W&B Files).
trainer.tracker.log_file(
os.path.abspath(os.path.expanduser(config_path)),
name="run.yaml",
)
ckpt_config = CheckpointConfig(
save_steps=int(tc.checkpoint.training_state_checkpointing_steps or 0),
keep_last=int(tc.checkpoint.checkpoints_total_limit or 0),
)
ckpt_config = CheckpointConfig(
save_steps=int(tc.checkpoint.training_state_checkpointing_steps or 0),
keep_last=int(tc.checkpoint.checkpoints_total_limit or 0),
)
checkpoint_manager = CheckpointManager(
method=method,
dataloader=dataloader,
output_dir=tc.checkpoint.output_dir,
config=ckpt_config,
callbacks=trainer.callbacks,
raw_config=cfg.raw,
)
checkpoint_manager = CheckpointManager(
method=method,
dataloader=dataloader,
output_dir=tc.checkpoint.output_dir,
config=ckpt_config,
callbacks=trainer.callbacks,
raw_config=cfg.raw,
)
trainer.run(
method,
dataloader=dataloader,
max_steps=tc.loop.max_train_steps,
start_step=start_step,
checkpoint_manager=checkpoint_manager,
)
trainer.run(
method,
dataloader=dataloader,
max_steps=tc.loop.max_train_steps,
start_step=start_step,
checkpoint_manager=checkpoint_manager,
)
finally:
# torch.profiler exports its trace at stop(). Keep atexit as a fallback
# for abrupt exits, but flush here so normal training returns only after
# all per-rank traces and summaries are complete.
profiler_controller.stop()
def main(
+17
View File
@@ -208,6 +208,23 @@ class TrainingMethod(torch.nn.Module, ABC):
"""
return False
def did_update_role(
self,
role: str,
iteration: int,
) -> bool:
"""Whether ``role`` completed an optimizer update this iteration.
Callbacks run after optimization and use this hook to follow the
method's actual update cadence. The default implementation matches
the role's optimizer against the optimizers selected for this
iteration; methods that step optimizers internally can override it.
"""
role_optimizer = self._optimizer_dict.get(role)
if role_optimizer is None:
return False
return any(optimizer is role_optimizer for optimizer in self.get_optimizers(iteration))
def managed_train_step(
self,
data_stream: Any,
@@ -6,8 +6,10 @@ from __future__ import annotations
from typing import Any, Literal
import torch
import torch.distributed as dist
import torch.nn.functional as F
from fastvideo.profiler import profile_region
from fastvideo.train.methods.base import TrainingMethod, LogScalar
from fastvideo.train.models.base import ModelBase
from fastvideo.train.utils.optimizer import (
@@ -431,6 +433,23 @@ class DMD2Method(TrainingMethod):
)
return step_list[index]
@staticmethod
def _max_rollout_target_idx_across_ranks(target_timestep_idx: torch.Tensor, ) -> int:
"""Return the largest sampled rollout index across all data ranks.
Each rank intentionally samples its own DMD2 timestep, but FSDP ranks
must execute the same number of transformer forwards so their
collectives stay ordered. The largest local target is therefore the
minimum prefix length every rank must execute this iteration.
"""
max_target_timestep_idx = target_timestep_idx.detach().clone()
if dist.is_available() and dist.is_initialized():
dist.all_reduce(
max_target_timestep_idx,
op=dist.ReduceOp.MAX,
)
return int(max_target_timestep_idx.item())
def _parse_score_timestep_bounds(self) -> tuple[int, int]:
"""Resolve the score-model timestep window used by legacy DMD.
@@ -476,6 +495,7 @@ class DMD2Method(TrainingMethod):
self._score_max_timestep,
)
@profile_region("profiler_region_dmd2_student_rollout")
def _student_rollout(
self,
batch: Any,
@@ -516,6 +536,7 @@ class DMD2Method(TrainingMethod):
generator=self.cuda_generator,
)
target_timestep_idx_int = int(target_timestep_idx.item())
synchronized_target_idx = self._max_rollout_target_idx_across_ranks(target_timestep_idx, )
target_timestep = step_list[target_timestep_idx]
current_noise_latents = torch.randn(
@@ -524,56 +545,62 @@ class DMD2Method(TrainingMethod):
dtype=dtype,
generator=self.cuda_generator,
)
current_noise_latents_copy = (current_noise_latents.clone())
max_target_idx = len(step_list) - 1
noise_latents: list[torch.Tensor] = []
noise_latent_index = target_timestep_idx_int - 1
noisy_input = current_noise_latents
if max_target_idx > 0:
with torch.no_grad():
for step_idx in range(max_target_idx):
current_timestep = step_list[step_idx]
current_timestep_tensor = (current_timestep * torch.ones(
1,
device=device,
dtype=torch.long,
))
# FSDP ranks must run the same number of forwards. Ranks
# whose local target is earlier keep advancing a throwaway
# trajectory until the largest target sampled globally.
needs_prefix_step = step_idx < synchronized_target_idx
pred_clean: torch.Tensor | None = None
noise_dtype = dtype
if needs_prefix_step:
current_timestep = step_list[step_idx]
current_timestep_tensor = (current_timestep * torch.ones(
1,
device=device,
dtype=torch.long,
))
pred_clean = self.student.predict_x0(
current_noise_latents,
current_timestep_tensor,
batch,
conditional=True,
cfg_uncond=self._cfg_uncond,
attn_kind="vsa",
)
pred_clean = self.student.predict_x0(
current_noise_latents,
current_timestep_tensor,
batch,
conditional=True,
cfg_uncond=self._cfg_uncond,
attn_kind="vsa",
)
noise_dtype = pred_clean.dtype
next_timestep = step_list[step_idx + 1]
next_timestep_tensor = (next_timestep * torch.ones(
1,
device=device,
dtype=torch.long,
))
# Preserve the method-owned generator sequence even when
# the corresponding prefix forward is unnecessary. This
# keeps all later DMD2 random draws aligned with the old
# full-rollout implementation.
noise = torch.randn(
latents.shape,
device=device,
dtype=pred_clean.dtype,
dtype=noise_dtype,
generator=self.cuda_generator,
)
current_noise_latents = (self.student.add_noise(
pred_clean,
noise,
next_timestep_tensor,
))
noise_latents.append(current_noise_latents.clone())
if noise_latent_index >= 0:
if noise_latent_index >= len(noise_latents):
raise RuntimeError("noise_latent_index is out of bounds")
noisy_input = noise_latents[noise_latent_index]
else:
noisy_input = current_noise_latents_copy
if needs_prefix_step:
assert pred_clean is not None
next_timestep = step_list[step_idx + 1]
next_timestep_tensor = (next_timestep * torch.ones(
1,
device=device,
dtype=torch.long,
))
current_noise_latents = (self.student.add_noise(
pred_clean,
noise,
next_timestep_tensor,
))
if step_idx + 1 == target_timestep_idx_int:
noisy_input = current_noise_latents
if with_grad:
pred_x0 = self.student.predict_x0(
@@ -598,6 +625,7 @@ class DMD2Method(TrainingMethod):
batch.dmd_latent_vis_dict["generator_timestep"] = target_timestep.float().detach()
return pred_x0
@profile_region("profiler_region_dmd2_critic_loss")
def _critic_flow_matching_loss(
self,
batch: Any,
@@ -638,6 +666,7 @@ class DMD2Method(TrainingMethod):
outputs,
)
@profile_region("profiler_region_dmd2_generator_loss")
def _dmd_loss(
self,
generator_pred_x0: torch.Tensor,
+113 -77
View File
@@ -11,6 +11,7 @@ import torch
from tqdm.auto import tqdm
from fastvideo.distributed import get_sp_group, get_world_group
from fastvideo.profiler import profile_region, profiler_region
from fastvideo.train.callbacks.callback import CallbackDict
from fastvideo.train.methods.base import LogScalar, TrainingMethod
from fastvideo.train.utils.tracking import build_tracker
@@ -98,6 +99,113 @@ class Trainer:
if self.global_rank == 0 and validation_metrics:
self.tracker.log(validation_metrics, iteration)
@profile_region("profiler_region_training_train_one_step")
def _run_train_step(
self,
method: TrainingMethod,
*,
data_stream: Iterator[dict[str, Any]],
step: int,
grad_accum: int,
method_manages_optimization: bool,
) -> None:
t0 = time.perf_counter()
# Accumulate on GPU during grad-accum; materialise to CPU once per
# step right before logging.
loss_sums: dict[str, float | torch.Tensor] = {}
metric_sums: dict[str, float | torch.Tensor] = {}
dataloader_time_sec = 0.0
if method_manages_optimization:
# Managed methods own their forward/backward/optimizer boundaries,
# so the enclosing training_train_one_step region is the truthful
# granularity available to the trainer.
loss_map, outputs, step_metrics = method.managed_train_step(
data_stream,
step,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
loss_sums[k] = v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
metric_sums[k] = _coerce_log_scalar(
v,
where=("method.managed_train_step()"
f".metrics[{k!r}]"),
)
else:
for _ in range(grad_accum):
dataloader_t0 = time.perf_counter()
with profiler_region("profiler_region_training_dataloader"):
batch = next(data_stream)
dataloader_time_sec += (time.perf_counter() - dataloader_t0)
with profiler_region("profiler_region_training_forward"):
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
step,
))
with profiler_region("profiler_region_training_backward"):
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
v,
where=("method.single_train_step()"
f".metrics[{k!r}]"),
))
if not method_manages_optimization:
with profiler_region("profiler_region_training_optimizer"):
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
# Single CPU sync point: materialise GPU tensors to float right before
# logging.
divisor = 1 if method_manages_optimization else grad_accum
metrics = {k: float(v) / divisor for k, v in loss_sums.items()}
metrics.update({k: float(v) / divisor for k, v in metric_sums.items()})
metrics["step_time_sec"] = (time.perf_counter() - t0)
if not method_manages_optimization:
# This is the local training process's wait for next(data_stream).
# Track it without a cross-rank reduction to avoid adding a
# synchronization to every training step.
metrics["dataloader_time_sec"] = dataloader_time_sec
metrics["vsa_sparsity"] = float(self.training_config.vsa_sparsity)
if self.global_rank == 0 and metrics:
self.tracker.log(metrics, step)
with profiler_region("profiler_region_training_callbacks"):
self.callbacks.on_training_step_end(
method,
metrics,
iteration=step,
)
@profile_region("profiler_region_training_train")
def run(
self,
method: TrainingMethod,
@@ -150,84 +258,12 @@ class Trainer:
# Allow method-specific optimization flow (e.g. DiffusionNFT).
method_manages_optimization = bool(method.manages_optimization())
for step in progress:
t0 = time.perf_counter()
# Accumulate on GPU during grad-accum; materialise
# to CPU once per step right before logging.
loss_sums: dict[str, float | torch.Tensor] = {}
metric_sums: dict[str, float | torch.Tensor] = {}
if method_manages_optimization:
loss_map, outputs, step_metrics = method.managed_train_step(
data_stream,
step,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
loss_sums[k] = v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
metric_sums[k] = _coerce_log_scalar(
v,
where=("method.managed_train_step()"
f".metrics[{k!r}]"),
)
else:
for accum_iter in range(grad_accum):
batch = next(data_stream)
loss_map, outputs, step_metrics = (method.single_train_step(
batch,
step,
))
method.backward(
loss_map,
outputs,
grad_accum_rounds=grad_accum,
)
for k, v in loss_map.items():
if isinstance(v, torch.Tensor):
prev = loss_sums.get(k, 0.0)
loss_sums[k] = prev + v.detach()
for k, v in step_metrics.items():
if k in loss_sums:
raise ValueError(f"Metric key {k!r} collides "
"with loss key. Use a "
"different name (e.g. prefix "
"with 'train/').")
prev = metric_sums.get(k, 0.0)
metric_sums[k] = (prev + _coerce_log_scalar(
v,
where=("method.single_train_step()"
f".metrics[{k!r}]"),
))
if not method_manages_optimization:
self.callbacks.on_before_optimizer_step(
method,
iteration=step,
)
method.optimizers_schedulers_step(step)
method.optimizers_zero_grad(step)
# Single CPU sync point: materialise GPU tensors
# to float right before logging.
divisor = 1 if method_manages_optimization else grad_accum
metrics = {k: float(v) / divisor for k, v in loss_sums.items()}
metrics.update({k: float(v) / divisor for k, v in metric_sums.items()})
metrics["step_time_sec"] = (time.perf_counter() - t0)
metrics["vsa_sparsity"] = float(tc.vsa_sparsity)
if self.global_rank == 0 and metrics:
self.tracker.log(metrics, step)
self.callbacks.on_training_step_end(
self._run_train_step(
method,
metrics,
iteration=step,
data_stream=data_stream,
step=step,
grad_accum=grad_accum,
method_manages_optimization=method_manages_optimization,
)
if checkpoint_manager is not None:
+2
View File
@@ -23,6 +23,7 @@ from torch.distributed.checkpoint.state_dict import (
from torch.distributed.checkpoint.stateful import Stateful
from fastvideo.logger import init_logger
from fastvideo.profiler import profile_region
logger = init_logger(__name__)
@@ -246,6 +247,7 @@ class CheckpointManager:
return
self.save(step)
@profile_region("profiler_region_training_save_checkpoint")
def save(self, step: int) -> None:
checkpoint_dir = self._checkpoint_dir(step)
dcp_dir = self._dcp_dir(step)