Skip to content

performance

Low-overhead performance accounting for modular training.

The monitor counts logical Transformer invocations at the role boundary. This is important for distillation: one optimizer step may contain several student rollouts plus critic and teacher forwards. FLOPs are an analytic estimate of the Transformer core; measured step and sample throughput remain exact wall clock metrics.

Classes

fastvideo.train.utils.performance.TrainingPerformanceMonitor

TrainingPerformanceMonitor()

Count per-role Transformer work for one optimizer step at a time.

Source code in fastvideo/train/utils/performance.py
def __init__(self) -> None:
    self._handles: list[Any] = []
    self._work: list[_ForwardWork] = []
    self._forward_calls: dict[str, int] = defaultdict(int)
    self._backward_forwards: dict[str, int] = defaultdict(int)

Functions:

fastvideo.train.utils.performance.estimate_transformer_forward

estimate_transformer_forward(module: Module, kwargs: dict[str, Any], output: Any, *, role: str, attention_metadata: Any | None, cross_attention_cached: bool = False) -> _ForwardWork | None

Estimate one Wan-style Transformer invocation from its logical input.

Source code in fastvideo/train/utils/performance.py
def estimate_transformer_forward(
    module: torch.nn.Module,
    kwargs: dict[str, Any],
    output: Any,
    *,
    role: str,
    attention_metadata: Any | None,
    cross_attention_cached: bool = False,
) -> _ForwardWork | None:
    """Estimate one Wan-style Transformer invocation from its logical input."""
    hidden_states = _first_tensor(kwargs.get("hidden_states"))
    if hidden_states is None or hidden_states.ndim != 5:
        return None

    arch = _arch_config(module)
    if arch is None:
        return None
    required = ("hidden_size", "ffn_dim", "num_layers", "patch_size")
    if any(not hasattr(arch, name) for name in required):
        return None

    batch_size = int(hidden_states.shape[0])
    raw_frames = int(hidden_states.shape[2])
    raw_height = int(hidden_states.shape[3])
    raw_width = int(hidden_states.shape[4])
    patch_size = arch.patch_size
    patch: tuple[int,
                 ...] = ((1, patch_size,
                          patch_size) if isinstance(patch_size, int) else tuple(int(value) for value in patch_size))
    if len(patch) != 3 or any(value <= 0 for value in patch):
        return None

    frames = raw_frames // patch[0]
    spatial_tokens = (raw_height // patch[1]) * (raw_width // patch[2])
    seq_len = frames * spatial_tokens
    if seq_len <= 0:
        return None

    hidden_size = int(arch.hidden_size)
    ffn_dim = int(arch.ffn_dim)
    num_layers = int(arch.num_layers)

    context = _first_tensor(kwargs.get("encoder_hidden_states"))
    context_tokens = int(context.shape[-2]) if context is not None and context.ndim >= 2 else 0
    image_context = _first_tensor(kwargs.get("encoder_hidden_states_image"))
    if image_context is not None and image_context.ndim >= 2:
        context_tokens += int(image_context.shape[-2])

    is_causal = hasattr(module, "num_frame_per_block")
    teacher_forcing = is_causal and _first_tensor(kwargs.get("clean_x")) is not None
    query_tokens = seq_len
    dense_pairs = float(seq_len * seq_len)
    vsa_tiles = 0
    causal_chunks = 0

    if is_causal:
        frames_per_block = int(getattr(module, "num_frame_per_block", 1))
        causal_chunks = math.ceil(frames / max(1, frames_per_block))
        local_attn_size = int(getattr(module, "local_attn_size", -1))
        kv_cache = kwargs.get("kv_cache")
        if kv_cache is not None:
            current_start = int(kwargs.get("current_start", 0))
            # Causal Wan caps the KV window at
            # GLOBAL_ATTN_COMPAT_MAX_LATENT_FRAMES (21) frames when
            # local_attn_size is unset; sliding_window_num_frames only
            # sizes the streaming KV cache. MatrixGame2's 15-frame
            # compatibility window is not modeled here.
            max_frames = local_attn_size if local_attn_size >= 0 else 21
            key_tokens = min(current_start + seq_len, max_frames * spatial_tokens)
            attention_pairs = float(seq_len * key_tokens)
            dense_pairs = float(seq_len * (current_start + seq_len))
        elif teacher_forcing:
            query_tokens = 2 * seq_len
            attention_pairs = float(_teacher_forcing_frame_pairs(frames, frames_per_block) * spatial_tokens**2)
            dense_pairs = float(query_tokens * query_tokens)
        else:
            attention_pairs = float(
                _blockwise_causal_frame_pairs(frames, frames_per_block, local_attn_size) * spatial_tokens**2)
    else:
        attention_pairs, vsa_tiles = _vsa_attention_pairs(seq_len, attention_metadata)

    # Causal Wan pads text to ``text_len`` before every cross-attention block.
    if is_causal:
        context_tokens = max(context_tokens, int(getattr(module, "text_len", 0)))

    b = batch_size
    seq_length = query_tokens
    dim = hidden_size
    # Four self-attention projections, two MLP projections, cross-attention
    # Q/K/V/out projections, and the two attention matmuls.
    per_layer = (8.0 * b * seq_length * dim * dim + 4.0 * b * seq_length * dim * ffn_dim +
                 4.0 * b * seq_length * dim * dim +
                 (0.0 if cross_attention_cached else 4.0 * b * context_tokens * dim * dim) +
                 4.0 * b * seq_length * context_tokens * dim + 4.0 * b * attention_pairs * dim)

    if vsa_tiles:
        # Wan VSA adds one gate projection and dense pooled QK/AV matmuls.
        per_layer += 2.0 * b * seq_length * dim * dim
        per_layer += 4.0 * b * vsa_tiles * vsa_tiles * dim

    forward_flops = per_layer * num_layers
    backward_expected = _output_requires_grad(output)
    # Standard MFU counts model-algorithm FLOPs, excluding activation
    # checkpoint recompute: training is approximately F + 2B = 3F.
    useful_flops = forward_flops * (3.0 if backward_expected else 1.0)
    return _ForwardWork(
        role=role,
        forward_flops=forward_flops,
        useful_flops=useful_flops,
        causal_chunks=causal_chunks,
        query_frames=b * (2 * frames if teacher_forcing else frames),
        query_tokens=b * query_tokens,
        attention_pairs=b * attention_pairs,
        dense_attention_pairs=b * dense_pairs,
    )

fastvideo.train.utils.performance.infer_peak_bf16_tflops

infer_peak_bf16_tflops(device_name: str) -> float | None

Infer dense BF16 tensor-core peak TFLOP/s from a CUDA device name.

Values intentionally exclude structured sparsity. Explicit configuration should be used when board form factor or clocks differ from the common data-center variant.

Source code in fastvideo/train/utils/performance.py
def infer_peak_bf16_tflops(device_name: str) -> float | None:
    """Infer dense BF16 tensor-core peak TFLOP/s from a CUDA device name.

    Values intentionally exclude structured sparsity. Explicit configuration
    should be used when board form factor or clocks differ from the common
    data-center variant.
    """
    name = device_name.upper()
    # Order matters: GB200 contains B200.
    peaks = (
        ("GB200", 2500.0),
        ("B200", 2250.0),
        ("H200", 989.5),
        ("H100", 989.5),
        ("L40S", 362.05),
        ("A100", 312.0),
        ("A40", 149.7),
    )
    return next((peak for token, peak in peaks if token in name), None)