# Copyright (c) 2024 NVIDIA CORPORATION. # Licensed under the MIT license. # Adapted from https://github.com/NVIDIA/BigVGAN under the MIT license. # LICENSE is in incl_licenses directory. from typing import * import torch import torch.nn as nn from torch.nn import Conv1d, ConvTranspose1d from torch.nn.utils import weight_norm, remove_weight_norm from . import bigvgan_activations as activations from .alias_free_act import Activation1d as TorchActivation1d from .temporal_adapter import TemporalAdapterBlock def init_weights(m, mean=0.0, std=0.01): classname = m.__class__.__name__ if classname.find("Conv") != -1: m.weight.data.normal_(mean, std) def apply_weight_norm(m): classname = m.__class__.__name__ if classname.find("Conv") != -1: weight_norm(m) def get_padding(kernel_size, dilation=1): return int((kernel_size * dilation - dilation) / 2) class FiLMBlock(nn.Module): """Feature-wise Linear Modulation for watermark conditioning. Maps a binary watermark vector to per-channel scale (γ) and shift (β) parameters that modulate intermediate feature maps. """ def __init__(self, watermark_bits: int, channels: int, hidden_dim: int = 128): super().__init__() self.scale_net = nn.Sequential( nn.Linear(watermark_bits, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, channels), ) self.shift_net = nn.Sequential( nn.Linear(watermark_bits, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, channels), ) # Initialize close to identity: γ≈1, β≈0 nn.init.zeros_(self.scale_net[-1].weight) nn.init.zeros_(self.scale_net[-1].bias) nn.init.zeros_(self.shift_net[-1].weight) nn.init.zeros_(self.shift_net[-1].bias) def forward(self, x, watermark): """ Args: x: [B, C, T] feature tensor watermark: [B, watermark_bits] float tensor Returns: [B, C, T] modulated features """ gamma = self.scale_net(watermark).unsqueeze(-1) # [B, C, 1] beta = self.shift_net(watermark).unsqueeze(-1) # [B, C, 1] return (1 + gamma) * x + beta # γ≈0 → identity class AMPBlock1(torch.nn.Module): """ AMPBlock applies Snake / SnakeBeta activation functions with trainable parameters that control periodicity, defined for each layer. AMPBlock1 has additional self.convs2 that contains additional Conv1d layers with a fixed dilation=1 followed by each layer in self.convs1 Args: h (AttrDict): Hyperparameters. channels (int): Number of convolution channels. kernel_size (int): Size of the convolution kernel. Default is 3. dilation (tuple): Dilation rates for the convolutions. Each dilation layer has two convolutions. Default is (1, 3, 5). activation (str): Activation function type. Should be either 'snake' or 'snakebeta'. Default is None. """ def __init__( self, channels: int, kernel_size: int = 3, dilation: tuple = (1, 3, 5), activation: str = None, snake_logscale: bool = True, ): super().__init__() self.convs1 = nn.ModuleList( [ weight_norm( Conv1d( channels, channels, kernel_size, stride=1, dilation=d, padding=get_padding(kernel_size, d), ) ) for d in dilation ] ) self.convs1.apply(init_weights) self.convs2 = nn.ModuleList( [ weight_norm( Conv1d( channels, channels, kernel_size, stride=1, dilation=1, padding=get_padding(kernel_size, 1), ) ) for _ in range(len(dilation)) ] ) self.convs2.apply(init_weights) self.num_layers = len(self.convs1) + len( self.convs2 ) # Total number of conv layers Activation1d = TorchActivation1d # Activation functions if activation == "snake": self.activations = nn.ModuleList( [ Activation1d( activation=activations.Snake( channels, alpha_logscale=snake_logscale ) ) for _ in range(self.num_layers) ] ) elif activation == "snakebeta": self.activations = nn.ModuleList( [ Activation1d( activation=activations.SnakeBeta( channels, alpha_logscale=snake_logscale ) ) for _ in range(self.num_layers) ] ) else: raise NotImplementedError( "activation incorrectly specified. check the config file and look for 'activation'." ) def forward(self, x): acts1, acts2 = self.activations[::2], self.activations[1::2] for c1, c2, a1, a2 in zip(self.convs1, self.convs2, acts1, acts2): xt = a1(x) xt = c1(xt) xt = a2(xt) xt = c2(xt) x = xt + x return x def remove_weight_norm(self): for l in self.convs1: remove_weight_norm(l) for l in self.convs2: remove_weight_norm(l) class BigVGAN(torch.nn.Module): """ BigVGAN is a neural vocoder model that applies anti-aliased periodic activation for residual blocks (resblocks). New in BigVGAN-v2: it can optionally use optimized CUDA kernels for AMP (anti-aliased multi-periodicity) blocks. Args: h (AttrDict): Hyperparameters. use_cuda_kernel (bool): If set to True, loads optimized CUDA kernels for AMP. This should be used for inference only, as training is not supported with CUDA kernels. Note: - The `use_cuda_kernel` parameter should be used for inference only, as training with CUDA kernels is not supported. - Ensure that the activation function is correctly specified in the hyperparameters (h.activation). """ def __init__( self, num_mels: int = 96, global_channels: int = -1, upsample_initial_channel: int = 1536, resblock_kernel_sizes: List[int] = [3, 7, 11], resblock_dilation_sizes: List[Tuple[int]] = [(1, 3, 5), (1, 3, 5), (1, 3, 5)], upsample_rates: List[int] = [4,4,2,2,2,2], upsample_kernel_sizes: List[int] = [8,8,4,4,4,4], snake_logscale: bool = True, activation: str = 'snakebeta', use_bias_at_final:bool = False, use_tanh_at_final:bool = False, # Watermark FiLM conditioning watermark_bits: int = 0, watermark_film_hidden: int = 128, # VocBulwark Temporal Adapter vocbulwark_bits: int = 0, vocbulwark_ta_hidden: int = 128, vocbulwark_zero_conv_std: float = 0.02, ): super().__init__() Activation1d = TorchActivation1d self.num_kernels = len(resblock_kernel_sizes) self.num_upsamples = len(upsample_rates) # Pre-conv self.conv_pre = weight_norm( Conv1d(num_mels, upsample_initial_channel, 7, 1, padding=3) ) if global_channels > 0: self.global_conv = torch.nn.Conv1d(global_channels, upsample_initial_channel, 1) # Define which AMPBlock to use. BigVGAN uses AMPBlock1 as default resblock_class = AMPBlock1 # Transposed conv-based upsamplers. does not apply anti-aliasing self.ups = nn.ModuleList() for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): self.ups.append( nn.ModuleList( [ weight_norm( ConvTranspose1d( upsample_initial_channel // (2 ** i), upsample_initial_channel // (2 ** (i + 1)), k, u, padding=(k - u) // 2, ) ) ] ) ) # Residual blocks using anti-aliased multi-periodicity composition modules (AMP) self.resblocks = nn.ModuleList() for i in range(len(self.ups)): ch = upsample_initial_channel // (2 ** (i + 1)) for j, (k, d) in enumerate( zip(resblock_kernel_sizes, resblock_dilation_sizes) ): self.resblocks.append( resblock_class(ch, k, d, activation=activation, snake_logscale=snake_logscale) ) # Watermark FiLM layers (one per upsample stage) self.watermark_bits = watermark_bits if watermark_bits > 0: self.film_layers = nn.ModuleList() for i in range(len(self.ups)): ch_film = upsample_initial_channel // (2 ** (i + 1)) self.film_layers.append( FiLMBlock(watermark_bits, ch_film, watermark_film_hidden) ) # VocBulwark Temporal Adapter layers (one per upsample stage) self.vocbulwark_bits = vocbulwark_bits if vocbulwark_bits > 0: self.ta_layers = nn.ModuleList() for i in range(len(self.ups)): ch_ta = upsample_initial_channel // (2 ** (i + 1)) self.ta_layers.append( TemporalAdapterBlock(vocbulwark_bits, ch_ta, vocbulwark_ta_hidden, zero_conv_std=vocbulwark_zero_conv_std) ) # Post-conv activation_post = ( activations.Snake(ch, alpha_logscale=snake_logscale) if activation == "snake" else ( activations.SnakeBeta(ch, alpha_logscale=snake_logscale) if activation == "snakebeta" else None ) ) if activation_post is None: raise NotImplementedError( "activation incorrectly specified. check the config file and look for 'activation'." ) self.activation_post = Activation1d(activation=activation_post) # Whether to use bias for the final conv_post. Default to True for backward compatibility self.use_bias_at_final = use_bias_at_final self.conv_post = weight_norm( Conv1d(ch, 1, 7, 1, padding=3, bias=self.use_bias_at_final) ) # Weight initialization for i in range(len(self.ups)): self.ups[i].apply(init_weights) self.conv_post.apply(init_weights) # Final tanh activation. Defaults to True for backward compatibility self.use_tanh_at_final = use_tanh_at_final def forward( self, x: torch.Tensor, g: Optional[torch.Tensor] = None, watermark: Optional[torch.Tensor] = None, ): # Pre-conv x = self.conv_pre(x) if g is not None: x = x + self.global_conv(g) for i in range(self.num_upsamples): # Upsampling for i_up in range(len(self.ups[i])): x = self.ups[i][i_up](x) # AMP blocks xs = None for j in range(self.num_kernels): if xs is None: xs = self.resblocks[i * self.num_kernels + j](x) else: xs += self.resblocks[i * self.num_kernels + j](x) x = xs / self.num_kernels # Watermark FiLM conditioning if watermark is not None and self.watermark_bits > 0: x = self.film_layers[i](x, watermark) # VocBulwark Temporal Adapter conditioning if watermark is not None and self.vocbulwark_bits > 0: x = self.ta_layers[i](x, watermark) # Post-conv x = self.activation_post(x) x = self.conv_post(x) # Final tanh activation if self.use_tanh_at_final: x = torch.tanh(x) else: x = torch.clamp(x, min=-1.0, max=1.0) # Bound the output to [-1, 1] return x