Skip to content

posthoc_ema

FSDP2-compatible post-hoc EMA used by the official MMAudio recipe.

Classes

fastvideo.train.callbacks.posthoc_ema.PostHocEMACallback

PostHocEMACallback(*, sigma_rels: list[float] | tuple[float, ...] = (0.05, 0.1), update_every: int = 1, checkpoint_every: int = 5000, checkpoint_folder: str | None = None, start_iter: int = 0, default_output_sigma: float = 0.05, step_size_correction: bool = True)

Bases: Callback

Maintain and checkpoint multiple Karras EMA profiles on local shards.

The upstream nitrous-ema implementation deep-copies a complete model on rank zero. That is appropriate for MMAudio's DDP trainer but not for a FastVideo FSDP2/HSDP transformer. This callback applies the same update and synthesis equations independently to every local FSDP shard. Together the rank-local snapshots represent the same full EMA model without gathering it on every optimizer step.

Source code in fastvideo/train/callbacks/posthoc_ema.py
def __init__(
    self,
    *,
    sigma_rels: list[float] | tuple[float, ...] = (0.05, 0.1),
    update_every: int = 1,
    checkpoint_every: int = 5000,
    checkpoint_folder: str | None = None,
    start_iter: int = 0,
    default_output_sigma: float = 0.05,
    step_size_correction: bool = True,
) -> None:
    if len(sigma_rels) < 2:
        raise ValueError("Post-hoc EMA requires at least two sigma profiles")
    self.sigma_rels = tuple(float(value) for value in sigma_rels)
    self.gammas = tuple(sigma_rel_to_gamma(value) for value in self.sigma_rels)
    if len(set(self.gammas)) != len(self.gammas):
        raise ValueError("Post-hoc EMA sigma profiles must be unique")
    self.update_every = max(1, int(update_every))
    self.checkpoint_every = max(1, int(checkpoint_every))
    self.checkpoint_folder = checkpoint_folder
    self.start_iter = int(start_iter)
    self.default_output_sigma = float(default_output_sigma)
    self.step_size_correction = bool(step_size_correction)

    self._ema_models: list[EMA_FSDP] = []
    self._calls = 0
    self._initted = False
    self._rank = 0
    self._snapshot_root: Path | None = None

Functions:

fastvideo.train.callbacks.posthoc_ema.sigma_rel_to_gamma

sigma_rel_to_gamma(sigma_rel: float) -> float

Algorithm 2 from Karras et al., matching nitrous-ema.

Source code in fastvideo/train/callbacks/posthoc_ema.py
def sigma_rel_to_gamma(sigma_rel: float) -> float:
    """Algorithm 2 from Karras et al., matching ``nitrous-ema``."""
    if sigma_rel <= 0:
        raise ValueError("Post-hoc EMA sigma_rel must be positive")
    t = sigma_rel**-2
    return float(np.roots([1, 7, 16 - t, 12 - t]).real.max())

fastvideo.train.callbacks.posthoc_ema.solve_posthoc_weights

solve_posthoc_weights(timesteps: Tensor, gammas: Tensor, target_timestep: int, target_gamma: float) -> Tensor

Algorithm 3 from Karras et al., matching nitrous-ema.

Source code in fastvideo/train/callbacks/posthoc_ema.py
def solve_posthoc_weights(
    timesteps: torch.Tensor,
    gammas: torch.Tensor,
    target_timestep: int,
    target_gamma: float,
) -> torch.Tensor:
    """Algorithm 3 from Karras et al., matching ``nitrous-ema``."""
    t_i = timesteps.double().reshape(-1, 1)
    gamma_i = gammas.double().reshape(-1, 1)
    t_j = timesteps.double().reshape(1, -1)
    gamma_j = gammas.double().reshape(1, -1)
    matrix = _p_dot_p(t_i, gamma_i, t_j, gamma_j)
    target_t = torch.tensor([[target_timestep]], dtype=torch.float64)
    target_g = torch.tensor([[target_gamma]], dtype=torch.float64)
    rhs = _p_dot_p(t_i, gamma_i, target_t, target_g)
    return torch.linalg.solve(matrix, rhs).squeeze(-1)