vocoder-small / bigvgan_model.py
mlr2000's picture
Upload folder using huggingface_hub
1fdc661 verified
Raw
History Blame Contribute Delete
12.8 kB
# 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