Skip to content

mmaudio_vae

Native 1D audio VAE used by MMAudio.

Classes

fastvideo.models.audio.mmaudio_vae.MMAudioVAE

MMAudioVAE(mode: str | dict[str, Any] = '44k', need_encoder: bool = False)

Bases: Module

MMAudio mel-spectrogram VAE for 16 kHz or 44.1 kHz audio.

Source code in fastvideo/models/audio/mmaudio_vae.py
def __init__(self, mode: str | dict[str, Any] = "44k", need_encoder: bool = False) -> None:
    super().__init__()
    if isinstance(mode, dict):
        config = mode
        mode = config.get("mode", "44k")
        need_encoder = config.get("need_encoder", need_encoder)
    if mode == "16k":
        data_dim, embed_dim, hidden_dim = 80, 20, 384
        data_mean, data_std = DATA_MEAN_80D, DATA_STD_80D
    elif mode == "44k":
        data_dim, embed_dim, hidden_dim = 128, 40, 512
        data_mean, data_std = DATA_MEAN_128D, DATA_STD_128D
    else:
        raise ValueError(f"Unknown MMAudio VAE mode: {mode}")

    self.mode = mode
    self.embed_dim = embed_dim
    self._weights_normalized = False
    self.register_buffer("data_mean", torch.tensor(data_mean, dtype=torch.float32).view(1, -1, 1))
    self.register_buffer("data_std", torch.tensor(data_std, dtype=torch.float32).view(1, -1, 1))
    if need_encoder:
        self.encoder = Encoder1D(
            dim=hidden_dim,
            ch_mult=(1, 2, 4),
            num_res_blocks=2,
            attn_layers=[3],
            down_layers=[0],
            in_dim=data_dim,
            embed_dim=embed_dim,
        )
    self.decoder = Decoder1D(
        dim=hidden_dim,
        ch_mult=(1, 2, 4),
        num_res_blocks=2,
        attn_layers=[3],
        down_layers=[0],
        in_dim=data_dim,
        out_dim=data_dim,
        embed_dim=embed_dim,
    )

Functions: