Skip to content

mmaudio_validation

Validation loss and audio sampling for MMAudio feature training.

Classes

fastvideo.train.callbacks.mmaudio_validation.MMAudioValidationCallback

MMAudioValidationCallback(*, data_path: str, every_steps: int = 5000, max_batches: int = 0, batch_size: int = 8, num_data_workers: int = 2, run_at_start: bool = False, use_ema: bool = False, inference_every_steps: int = 20000, inference_model_path: str = '', inference_num_samples: int = 16, inference_num_steps: int = 25, inference_guidance_scale: float = 4.5, inference_seed: int = 14159265, inference_save_video: bool = True, inference_log_to_tracker: bool = True, output_dir: str | None = None)

Bases: Callback

Evaluate cached val features and periodically run native V2A inference.

Validation loss follows the official MMAudio val_fn: sample a VAE posterior latent, a logit-normal flow time, prior noise, and independent video/text CFG masks. The RNG is reset for every pass, making values at different training steps directly comparable.

Optional inference reuses the live FSDP transformer in :class:MMAudioPipeline, while frozen VAE/vocoder weights are loaded from inference_model_path. Precomputed CLIP, Synchformer, and text features go directly into the pipeline, so validation does not decode source video or run feature encoders again.

Source code in fastvideo/train/callbacks/mmaudio_validation.py
def __init__(
    self,
    *,
    data_path: str,
    every_steps: int = 5000,
    max_batches: int = 0,
    batch_size: int = 8,
    num_data_workers: int = 2,
    run_at_start: bool = False,
    use_ema: bool = False,
    inference_every_steps: int = 20000,
    inference_model_path: str = "",
    inference_num_samples: int = 16,
    inference_num_steps: int = 25,
    inference_guidance_scale: float = 4.5,
    inference_seed: int = 14159265,
    inference_save_video: bool = True,
    inference_log_to_tracker: bool = True,
    output_dir: str | None = None,
) -> None:
    self.data_path = str(data_path)
    self.every_steps = int(every_steps)
    self.max_batches = max(0, int(max_batches))
    self.batch_size = max(1, int(batch_size))
    self.num_data_workers = max(0, int(num_data_workers))
    self.run_at_start = bool(run_at_start)
    self.use_ema = bool(use_ema)
    self.inference_every_steps = max(0, int(inference_every_steps))
    self.inference_model_path = str(inference_model_path)
    self.inference_num_samples = max(1, int(inference_num_samples))
    self.inference_num_steps = max(1, int(inference_num_steps))
    self.inference_guidance_scale = float(inference_guidance_scale)
    self.inference_seed = int(inference_seed)
    self.inference_save_video = bool(inference_save_video)
    self.inference_log_to_tracker = bool(inference_log_to_tracker)
    self.output_dir = output_dir

    self._dataloader: Any = None
    self._pipeline: Any = None
    self._rank = 0

Functions: