Skip to content

minimax_h3_audio

Native MiniMax H3 waveform autoencoder.

Classes

fastvideo.models.vaes.minimax_h3_audio.MiniMaxH3AudioCausalAttention

MiniMaxH3AudioCausalAttention(in_dim: int, out_dim: int, num_heads: int)

Bases: Module

Causal projection from the encoder trunk width to the latent width.

Source code in fastvideo/models/vaes/minimax_h3_audio.py
def __init__(self, in_dim: int, out_dim: int, num_heads: int):
    super().__init__()
    self.out_dim = out_dim
    self.num_heads = num_heads
    self.head_dim = in_dim // num_heads
    self.qkv = ReplicatedLinear(in_dim, in_dim * 3, bias=False, params_dtype=torch.float32)
    self.q_bias = nn.Parameter(torch.zeros(in_dim))
    self.v_bias = nn.Parameter(torch.zeros(in_dim))
    self.register_buffer("zero_k_bias", torch.zeros(in_dim))
    self.proj = ReplicatedLinear(out_dim, out_dim, params_dtype=torch.float32)

fastvideo.models.vaes.minimax_h3_audio.MiniMaxH3AudioDiagonalGaussianDistribution

MiniMaxH3AudioDiagonalGaussianDistribution(mean: Tensor, logs: Tensor)

Diagonal Gaussian parameterized by mean and log standard deviation.

Source code in fastvideo/models/vaes/minimax_h3_audio.py
def __init__(self, mean: torch.Tensor, logs: torch.Tensor):
    self.mean = mean
    self.logs = logs
    self.std = torch.exp(logs)

fastvideo.models.vaes.minimax_h3_audio.MiniMaxH3AudioVAE

MiniMaxH3AudioVAE(config: MiniMaxH3AudioVAEConfig)

Bases: Module

DAC encoder plus BigVGAN decoder for mono 32 kHz waveforms.

Source code in fastvideo/models/vaes/minimax_h3_audio.py
def __init__(self, config: MiniMaxH3AudioVAEConfig):
    super().__init__()
    self.config = config
    self.fastvideo_config = config
    arch = config.arch_config

    encoder_rates = tuple(int(rate) for rate in arch.encoder_rates)
    decoder_rates = tuple(int(rate) for rate in arch.decoder_rates)
    self.hop_length = math.prod(encoder_rates)
    self.sampling_rate = int(arch.sampling_rate)
    self.latent_channels = int(arch.latent_channels)
    self.audio_channels = 1
    latents_mean = arch.latents_mean if arch.latents_mean is not None else [0.0] * self.latent_channels
    latents_std = arch.latents_std if arch.latents_std is not None else [1.0] * self.latent_channels
    self.register_buffer(
        "latents_mean",
        torch.tensor(latents_mean, dtype=torch.float32).view(1, -1, 1),
        persistent=False,
    )
    self.register_buffer(
        "latents_std",
        torch.tensor(latents_std, dtype=torch.float32).view(1, -1, 1),
        persistent=False,
    )

    if math.prod(decoder_rates) != self.hop_length:
        raise ValueError(f"`decoder_rates` must upsample by the encoder hop length {self.hop_length}, got "
                         f"{math.prod(decoder_rates)}.")
    if arch.latent_dim % arch.latent_channels != 0:
        raise ValueError(f"`latent_dim` ({arch.latent_dim}) must be a multiple of `latent_channels` "
                         f"({arch.latent_channels}).")

    self.encoder = MiniMaxH3AudioEncoder(
        d_model=arch.encoder_dim,
        strides=encoder_rates,
        d_latent=arch.latent_dim,
    )
    self.pre_block = MiniMaxH3AudioAttnProjection(
        arch.latent_dim,
        arch.latent_channels,
        num_heads=arch.num_attention_heads,
    )
    self.mean_proj = nn.Conv1d(arch.latent_channels, arch.latent_channels, 1)
    self.logs_proj = nn.Conv1d(arch.latent_channels, arch.latent_channels, 1)

    self.dec_in_proj = nn.Conv1d(arch.latent_channels, arch.latent_dim, 1)
    self.decoder = MiniMaxH3AudioBigVGANDecoder(
        in_channels=arch.latent_dim,
        upsample_initial_channel=arch.decoder_dim,
        upsample_rates=decoder_rates,
        upsample_kernel_sizes=tuple(int(kernel) for kernel in arch.decoder_kernel_sizes),
        resblock_kernel_sizes=tuple(int(kernel) for kernel in arch.resblock_kernel_sizes),
        resblock_dilation_sizes=tuple(
            tuple(int(dilation) for dilation in group) for group in arch.resblock_dilation_sizes),
    )

    # The H3 checkpoint and waveform numerics require this component to stay FP32.
    self.float()

Functions:

fastvideo.models.vaes.minimax_h3_audio.kaiser_sinc_filter1d

kaiser_sinc_filter1d(cutoff: float, half_width: float, kernel_size: int) -> Tensor

Build the persistent low-pass filter used by alias-free activations.

Source code in fastvideo/models/vaes/minimax_h3_audio.py
def kaiser_sinc_filter1d(cutoff: float, half_width: float, kernel_size: int) -> torch.Tensor:
    """Build the persistent low-pass filter used by alias-free activations."""

    half_size = kernel_size // 2
    attenuation = 2.285 * (half_size - 1) * math.pi * (4 * half_width) + 7.95
    if attenuation > 50.0:
        beta = 0.1102 * (attenuation - 8.7)
    elif attenuation >= 21.0:
        beta = 0.5842 * (attenuation - 21)**0.4 + 0.07886 * (attenuation - 21.0)
    else:
        beta = 0.0

    window = torch.kaiser_window(kernel_size, beta=beta, periodic=False, dtype=torch.float32)
    if kernel_size % 2 == 0:
        time = torch.arange(-half_size, half_size, dtype=torch.float32) + 0.5
    else:
        time = torch.arange(kernel_size, dtype=torch.float32) - half_size

    filter_ = 2 * cutoff * window * torch.sinc(2 * cutoff * time)
    filter_ /= filter_.sum()
    return filter_.view(1, 1, kernel_size)