Skip to content

mmaudio

MMAudio multimodal flow-prediction transformer.

The model operates on one-dimensional audio latent sequences and jointly attends to semantic video, synchronization-video, and text conditions. The initial implementation intentionally uses torch SDPA so its single-GPU numeric contract remains explicit; sequence/tensor parallel support is a later, separately verified optimization.

Classes

fastvideo.models.dits.mmaudio.MMAudioTransformer

MMAudioTransformer(config: MMAudioTransformerConfig, hf_config: dict[str, Any], **kwargs)

Bases: BaseDiT

Source code in fastvideo/models/dits/mmaudio.py
def __init__(self, config: MMAudioTransformerConfig, hf_config: dict[str, Any], **kwargs) -> None:
    del kwargs
    super().__init__(config=config, hf_config=hf_config)
    arch = config.arch_config
    self.v2 = arch.v2
    self.latent_dim = arch.latent_dim
    self._latent_seq_len = arch.latent_seq_len
    self._clip_seq_len = arch.clip_seq_len
    self._sync_seq_len = arch.sync_seq_len
    self._text_seq_len = arch.text_seq_len
    self.hidden_dim = arch.hidden_dim
    self.num_heads = arch.num_heads
    self.hidden_size = arch.hidden_size
    self.num_attention_heads = arch.num_attention_heads
    self.num_channels_latents = arch.num_channels_latents

    activation = nn.SiLU if arch.v2 else nn.SELU
    self.audio_input_proj = nn.Sequential(
        ChannelLastConv1d(arch.latent_dim, arch.hidden_dim, kernel_size=7, padding=3),
        activation(),
        MMAudioConvMLP(arch.hidden_dim, arch.hidden_dim * 4, kernel_size=7, padding=3),
    )
    clip_layers: list[nn.Module] = [nn.Linear(arch.clip_dim, arch.hidden_dim)]
    if arch.v2:
        clip_layers.append(nn.SiLU())
    clip_layers.append(MMAudioConvMLP(arch.hidden_dim, arch.hidden_dim * 4, kernel_size=3, padding=1))
    self.clip_input_proj = nn.Sequential(*clip_layers)
    self.sync_input_proj = nn.Sequential(
        ChannelLastConv1d(arch.sync_dim, arch.hidden_dim, kernel_size=7, padding=3),
        activation(),
        MMAudioConvMLP(arch.hidden_dim, arch.hidden_dim * 4, kernel_size=3, padding=1),
    )
    text_layers: list[nn.Module] = [nn.Linear(arch.text_dim, arch.hidden_dim)]
    if arch.v2:
        text_layers.append(nn.SiLU())
    text_layers.append(MMAudioMLP(arch.hidden_dim, arch.hidden_dim * 4))
    self.text_input_proj = nn.Sequential(*text_layers)

    self.clip_cond_proj = nn.Linear(arch.hidden_dim, arch.hidden_dim)
    self.text_cond_proj = nn.Linear(arch.hidden_dim, arch.hidden_dim)
    self.global_cond_mlp = MMAudioMLP(arch.hidden_dim, arch.hidden_dim * 4)
    self.sync_pos_emb = nn.Parameter(torch.zeros((1, 1, 8, arch.sync_dim)))
    self.final_layer = FinalBlock(arch.hidden_dim, arch.latent_dim)
    self.t_embed = TimestepEmbedder(
        arch.hidden_dim,
        frequency_embedding_size=(arch.hidden_dim if arch.v2 else 256),
        max_period=(1 if arch.v2 else 10000),
    )
    self.joint_blocks = nn.ModuleList(
        [
            JointBlock(
                arch.hidden_dim,
                arch.num_heads,
                mlp_ratio=arch.mlp_ratio,
                pre_only=(index == arch.depth - arch.fused_depth - 1),
            )
            for index in range(arch.depth - arch.fused_depth)
        ]
    )
    self.fused_blocks = nn.ModuleList(
        [
            MMDitSingleBlock(arch.hidden_dim, arch.num_heads, mlp_ratio=arch.mlp_ratio, kernel_size=3, padding=1)
            for _ in range(arch.fused_depth)
        ]
    )

    self.latent_mean = nn.Parameter(torch.full((1, 1, arch.latent_dim), float("nan")), requires_grad=False)
    self.latent_std = nn.Parameter(torch.full((1, 1, arch.latent_dim), float("nan")), requires_grad=False)
    self.empty_string_feat = nn.Parameter(torch.zeros((arch.text_seq_len, arch.text_dim)), requires_grad=False)
    self.empty_clip_feat = nn.Parameter(torch.zeros(1, arch.clip_dim), requires_grad=True)
    self.empty_sync_feat = nn.Parameter(torch.zeros(1, arch.sync_dim), requires_grad=True)

    self.initialize_weights()
    self.initialize_rotations()
    self.__post_init__()

Methods:

fastvideo.models.dits.mmaudio.MMAudioTransformer.materialize_non_persistent_buffers
materialize_non_persistent_buffers(device: device, dtype: dtype | None = None) -> None

Rebuild derived buffers after meta-device production loading.

Source code in fastvideo/models/dits/mmaudio.py
def materialize_non_persistent_buffers(
    self,
    device: torch.device,
    dtype: torch.dtype | None = None,
) -> None:
    """Rebuild derived buffers after meta-device production loading."""
    if self.t_embed.freqs.is_meta:
        frequency_dim = self.t_embed.mlp[0].in_features
        freqs = 1.0 / (
            10000
            ** (
                torch.arange(0, frequency_dim, 2, dtype=torch.float32, device=device)
                / frequency_dim
            )
        )
        freqs = (10000 / self.t_embed.max_period) * freqs
        self.t_embed._buffers["freqs"] = freqs.to(dtype=dtype or torch.float32)
    if self.latent_rot.is_meta or self.clip_rot.is_meta:
        head_dim = self.hidden_dim // self.num_heads
        self._buffers["latent_rot"] = compute_rope_rotations(
            self._latent_seq_len, head_dim, 10000, device=device
        )
        self._buffers["clip_rot"] = compute_rope_rotations(
            self._clip_seq_len,
            head_dim,
            10000,
            freq_scaling=self._latent_seq_len / self._clip_seq_len,
            device=device,
        )