Skip to content

mmaudio

Classes

fastvideo.train.models.mmaudio.MMAudioModel

MMAudioModel(*, init_from: str | None = None, training_config: TrainingConfig, variant: str | None = None, from_scratch: bool = False, empty_string_features_path: str | None = None, latent_statistics_path: str | None = None, latent_statistics_chunk_size: int = 32, trainable: bool = True, allow_v2_training: bool = False, lora: LoraConfig | dict[str, Any] | None = None, attention_backend: AttentionBackendEnum | str | None = None, transformer: Module | None = None)

Bases: ModelBase

Official-compatible conditional flow-matching training adapter.

Only the MMAudio transformer participates in the training graph. Audio VAE, DFN5B, and Synchformer outputs are read from an offline feature cache.

Source code in fastvideo/train/models/mmaudio/mmaudio.py
def __init__(
    self,
    *,
    init_from: str | None = None,
    training_config: TrainingConfig,
    variant: str | None = None,
    from_scratch: bool = False,
    empty_string_features_path: str | None = None,
    latent_statistics_path: str | None = None,
    latent_statistics_chunk_size: int = 32,
    trainable: bool = True,
    allow_v2_training: bool = False,
    lora: LoraConfig | dict[str, Any] | None = None,
    attention_backend: AttentionBackendEnum | str | None = None,
    transformer: torch.nn.Module | None = None,
) -> None:
    super().__init__(
        trainable=trainable,
        lora=lora,
        attention_backend=attention_backend,
    )
    if int(training_config.distributed.sp_size or 1) != 1:
        raise ValueError("MMAudio training does not yet support sequence parallelism; set sp_size=1")
    if int(training_config.distributed.tp_size or 1) != 1:
        raise ValueError("MMAudio training does not yet support tensor parallelism; set tp_size=1")
    use_ddp = is_ddp_strategy(training_config)
    if use_ddp and (int(training_config.distributed.hsdp_replicate_dim) != 1
                    or int(training_config.distributed.hsdp_shard_dim) not in {-1, 1}):
        raise ValueError("MMAudio DDP strategy does not use HSDP dimensions; "
                         "set hsdp_replicate_dim=1 and hsdp_shard_dim=1")

    self.training_config = training_config
    if variant is not None:
        try:
            pipeline_config_cls = MMAUDIO_PIPELINE_CONFIGS[variant]
        except KeyError as exc:
            supported = ", ".join(MMAUDIO_PIPELINE_CONFIGS)
            raise ValueError(f"Unknown MMAudio variant {variant!r}; expected one of: {supported}") from exc
        self.training_config.pipeline_config = pipeline_config_cls()
    elif self.training_config.pipeline_config is None:
        self.training_config.pipeline_config = MMAudioV2AConfig()

    if transformer is None:
        if from_scratch:
            if init_from:
                raise ValueError("MMAudio from_scratch=true cannot also set init_from; use "
                                 "empty_string_features_path for the fixed CLIP embedding")
            if variant is None:
                raise ValueError("MMAudio from-scratch training requires variant")
            if variant not in MMAUDIO_TRAINING_VARIANTS:
                supported = ", ".join(MMAUDIO_TRAINING_VARIANTS)
                raise ValueError(f"The official MMAudio recipe does not train {variant!r}; "
                                 f"supported from-scratch variants: {supported}")
            if not empty_string_features_path:
                raise ValueError("MMAudio from-scratch training requires "
                                 "empty_string_features_path")
            latent_mean, latent_std = _load_or_compute_latent_stats(
                self.training_config,
                variant=variant,
                cache_path=latent_statistics_path,
                chunk_size=latent_statistics_chunk_size,
            )
            empty_string_feat = _load_empty_string_features(empty_string_features_path)
            architecture = MMAUDIO_VARIANT_ARCHITECTURES[variant]
            expected_text_shape = (77, 1024)
            if tuple(empty_string_feat.shape) != expected_text_shape:
                raise ValueError("MMAudio empty-string features must have shape "
                                 f"{expected_text_shape}, got {tuple(empty_string_feat.shape)}")
            model_config = get_mmaudio_transformer_config(variant)
            hf_config = {
                "_class_name": self._transformer_cls_name,
                **architecture,
                "clip_dim": 1024,
                "clip_seq_len": 64,
                "sync_dim": 768,
                "sync_seq_len": 192,
                "text_dim": 1024,
                "text_seq_len": 77,
                "mlp_ratio": 4.0,
            }
            default_dtype = PRECISION_TO_TYPE[self.training_config.dit_precision]
            init_params = {
                "config": model_config,
                "hf_config": hf_config,
                "latent_mean": latent_mean,
                "latent_std": latent_std,
                "empty_string_feat": empty_string_feat,
            }
            if use_ddp:
                transformer = build_replicated_model_from_scratch(
                    MMAudioTransformer,
                    init_params,
                    device=get_local_torch_device(),
                    default_dtype=default_dtype,
                    seed=self.training_config.data.seed,
                )
            else:
                transformer = build_fsdp_model_from_scratch(
                    model_cls=MMAudioTransformer,
                    init_params=init_params,
                    device=get_local_torch_device(),
                    hsdp_replicate_dim=self.training_config.distributed.hsdp_replicate_dim,
                    hsdp_shard_dim=self.training_config.distributed.hsdp_shard_dim,
                    default_dtype=default_dtype,
                    param_dtype=torch.bfloat16,
                    reduce_dtype=torch.float32,
                    seed=self.training_config.data.seed,
                    pin_cpu_memory=self.training_config.distributed.pin_cpu_memory,
                )
            self._init_from = f"scratch:{variant}"
        else:
            if not init_from:
                raise ValueError("MMAudio pretrained training requires init_from, or set "
                                 "from_scratch=true with a variant")
            if use_ddp:
                raise NotImplementedError("MMAudio DDP currently supports from_scratch=true only; "
                                          "pretrained component loading remains on the FSDP path")
            self._init_from = str(init_from)
            transformer = load_module_from_path(
                model_path=self._init_from,
                module_type="transformer",
                training_config=self.training_config,
                override_transformer_cls_name=self._transformer_cls_name,
                attention_backend=self.attention_backend,
            )
    else:
        self._init_from = str(init_from or f"provided:{variant or 'unknown'}")
    self.transformer = transformer
    if bool(getattr(self.transformer, "v2", False)) and not allow_v2_training:
        raise ValueError("The official MMAudio training recipe does not support `_v2` "
                         "checkpoints. Use small_44k, medium_44k, or large_44k.")

    if not self._enable_lora_if_configured(self.transformer):
        self.transformer = apply_trainable(self.transformer, trainable=self._trainable)
    if self._trainable:
        # These are checkpoint statistics/fixed text features in the
        # official model, not optimization variables. ``apply_trainable``
        # intentionally enables a whole module, so restore their contract.
        for name in ("latent_mean", "latent_std", "empty_string_feat"):
            parameter = getattr(self.transformer, name, None)
            if isinstance(parameter, torch.nn.Parameter):
                parameter.requires_grad_(False)
        # The two learned null-video tokens are trained by MMAudio.
        for name in ("empty_clip_feat", "empty_sync_feat"):
            parameter = getattr(self.transformer, name, None)
            if isinstance(parameter, torch.nn.Parameter):
                parameter.requires_grad_(True)
    if use_ddp:
        # Official MMAudio wraps the fully initialized FP32 model after
        # selecting trainable parameters. It disables buffer broadcasts;
        # the fixed latent statistics and positional buffers are identical
        # because every rank uses the same initialization seed.
        self.transformer = wrap_module_ddp(
            self.transformer,
            device=get_local_torch_device(),
            broadcast_buffers=False,
        )

    self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
        shift=1.0,
        invert_sigmas=True,
        sigma_min=0.0,
        use_reference_discrete_timesteps=True,
    )
    self.dataloader: Any = None
    self.start_step = 0

Methods:

fastvideo.train.models.mmaudio.MMAudioModel.compile_training_forward
compile_training_forward(compile_kwargs: dict[str, Any]) -> None

Compile transformer math while leaving DDP control flow eager.

Source code in fastvideo/train/models/mmaudio/mmaudio.py
def compile_training_forward(
    self,
    compile_kwargs: dict[str, Any],
) -> None:
    """Compile transformer math while leaving DDP control flow eager."""
    module = unwrap_ddp_module(self.transformer)
    compiled_forward = torch.compile(module.forward, **compile_kwargs)
    module.forward = compiled_forward
    logger.info(
        "Compiled inner MMAudio transformer forward with kwargs=%s",
        compile_kwargs,
    )
fastvideo.train.models.mmaudio.MMAudioModel.flow_matching_forward
flow_matching_forward(noisy_latents: Tensor, timestep: Tensor, clip_features: Tensor, sync_features: Tensor, text_features: Tensor) -> Tensor

Tensor-only adapter around the DDP-wrapped training forward.

Source code in fastvideo/train/models/mmaudio/mmaudio.py
def flow_matching_forward(
    self,
    noisy_latents: torch.Tensor,
    timestep: torch.Tensor,
    clip_features: torch.Tensor,
    sync_features: torch.Tensor,
    text_features: torch.Tensor,
) -> torch.Tensor:
    """Tensor-only adapter around the DDP-wrapped training forward."""
    device_type = noisy_latents.device.type
    with torch.autocast(
            device_type=device_type,
            dtype=torch.bfloat16,
            enabled=device_type == "cuda",
    ):
        return self.transformer(
            hidden_states=noisy_latents,
            encoder_hidden_states={
                "clip_features": clip_features,
                "sync_features": sync_features,
                "text_features": text_features,
            },
            timestep=timestep,
        )