import librosa import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from typing import Optional from audiotools import AudioSignal from audiotools import ml from audiotools import STFTParams from einops import rearrange from torch.nn.utils import weight_norm from rotary_embedding_torch import RotaryEmbedding from dac.nn.attend import Attend def WNConv1d(*args, **kwargs): act = kwargs.pop("act", True) conv = weight_norm(nn.Conv1d(*args, **kwargs)) if not act: return conv return nn.Sequential(conv, nn.LeakyReLU(0.1)) def WNConv2d(*args, **kwargs): act = kwargs.pop("act", True) conv = weight_norm(nn.Conv2d(*args, **kwargs)) if not act: return conv return nn.Sequential(conv, nn.LeakyReLU(0.1)) class MPD(nn.Module): def __init__(self, period): super().__init__() self.period = period self.convs = nn.ModuleList( [ WNConv2d(2, 32, (5, 1), (3, 1), padding=(2, 0)), WNConv2d(32, 128, (5, 1), (3, 1), padding=(2, 0)), WNConv2d(128, 512, (5, 1), (3, 1), padding=(2, 0)), WNConv2d(512, 1024, (5, 1), (3, 1), padding=(2, 0)), WNConv2d(1024, 1024, (5, 1), 1, padding=(2, 0)), ] ) self.conv_post = WNConv2d( 1024, 1, kernel_size=(3, 1), padding=(1, 0), act=False ) def pad_to_period(self, x): t = x.shape[-1] x = F.pad(x, (0, self.period - t % self.period), mode="reflect") return x def forward(self, x): fmap = [] x = self.pad_to_period(x) x = rearrange(x, "b c (l p) -> b c l p", p=self.period) for layer in self.convs: x = layer(x) fmap.append(x) x = self.conv_post(x) fmap.append(x) return fmap class MSD(nn.Module): def __init__(self, rate: int = 1, sample_rate: int = 44100): super().__init__() self.convs = nn.ModuleList( [ WNConv1d(2, 16, 15, 1, padding=7), WNConv1d(16, 64, 41, 4, groups=4, padding=20), WNConv1d(64, 256, 41, 4, groups=16, padding=20), WNConv1d(256, 1024, 41, 4, groups=64, padding=20), WNConv1d(1024, 1024, 41, 4, groups=256, padding=20), WNConv1d(1024, 1024, 5, 1, padding=2), ] ) self.conv_post = WNConv1d(1024, 1, 3, 1, padding=1, act=False) self.sample_rate = sample_rate self.rate = rate def forward(self, x): x = AudioSignal(x, self.sample_rate) x.resample(self.sample_rate // self.rate) x = x.audio_data fmap = [] for l in self.convs: x = l(x) fmap.append(x) x = self.conv_post(x) fmap.append(x) return fmap BANDS = [(0.0, 0.1), (0.1, 0.25), (0.25, 0.5), (0.5, 0.75), (0.75, 1.0)] class MRD(nn.Module): def __init__( self, window_length: int, hop_factor: float = 0.25, sample_rate: int = 44100, bands: list = BANDS, ): """Complex multi-band spectrogram discriminator. Parameters ---------- window_length : int Window length of STFT. hop_factor : float, optional Hop factor of the STFT, defaults to ``0.25 * window_length``. sample_rate : int, optional Sampling rate of audio in Hz, by default 44100 bands : list, optional Bands to run discriminator over. """ super().__init__() self.window_length = window_length self.hop_factor = hop_factor self.sample_rate = sample_rate self.stft_params = STFTParams( window_length=window_length, hop_length=int(window_length * hop_factor), match_stride=True, ) n_fft = window_length // 2 + 1 bands = [(int(b[0] * n_fft), int(b[1] * n_fft)) for b in bands] self.bands = bands ch = 32 convs = lambda: nn.ModuleList( [ WNConv2d(2, ch, (3, 9), (1, 1), padding=(1, 4)), WNConv2d(ch, ch, (3, 9), (1, 2), padding=(1, 4)), WNConv2d(ch, ch, (3, 9), (1, 2), padding=(1, 4)), WNConv2d(ch, ch, (3, 9), (1, 2), padding=(1, 4)), WNConv2d(ch, ch, (3, 3), (1, 1), padding=(1, 1)), ] ) self.band_convs = nn.ModuleList([convs() for _ in range(len(self.bands))]) self.conv_post = WNConv2d(ch, 1, (3, 3), (1, 1), padding=(1, 1), act=False) def spectrogram(self, x): x = AudioSignal(x, self.sample_rate, stft_params=self.stft_params) x = torch.view_as_real(x.stft()) x = rearrange(x, "b ch f t c -> (b ch) c t f", ch=2) # Split into bands x_bands = [x[..., b[0] : b[1]] for b in self.bands] return x_bands def forward(self, x): x_bands = self.spectrogram(x) fmap = [] x = [] for band, stack in zip(x_bands, self.band_convs): for layer in stack: band = layer(band) fmap.append(band) x.append(band) x = torch.cat(x, dim=-1) x = self.conv_post(x) fmap.append(x) return fmap # neural filterbank class SubbandProjection(nn.Module): def __init__(self, bandwidth, out_dim): super(SubbandProjection, self).__init__() self.layer_norm = nn.LayerNorm(bandwidth) self.fc = nn.Linear(bandwidth, out_dim) def forward(self, x): x = rearrange(x, "b f t -> b t f") x = self.layer_norm(x) x = self.fc(x) x = rearrange(x, "b t f -> b f t") return x class NeuralFilterbank(nn.Module): def __init__( self, n_fft: int, n_filterbank: int, sample_rate: int = 48000, out_dim: int = 64 ): super(NeuralFilterbank, self).__init__() self.bandwidth_indices = self.get_bandwidth_indices( sample_rate, n_fft, n_filterbank ) self.projection_modules = self.get_projection_layers(out_dim) def get_bandwidth_indices(self, sample_rate, n_fft, n_filterbank): mel_basis = librosa.filters.mel( sr=sample_rate, n_fft=n_fft, n_mels=n_filterbank ) indices_real = [np.where(row > 0)[0] for row in mel_basis] indices_imag = [np.where(row > 0)[0] + mel_basis.shape[1] for row in mel_basis] bandwidth_indices = [ list(r) + list(i) for r, i in zip(indices_real, indices_imag) ] return bandwidth_indices def get_projection_layers(self, out_dim): projection_modules = nn.ModuleList([]) for indices in self.bandwidth_indices: indices = indices[: len(indices)] projection_modules.append(SubbandProjection(len(indices), out_dim)) return projection_modules def forward(self, spec): # spec: [(batch, 2), freq, time] spec = rearrange(spec, "(b c) f t -> b (c f) t", c=2) emb = [] for indices, layer in zip(self.bandwidth_indices, self.projection_modules): emb.append(layer(spec[:, indices[: len(indices)], :])) return rearrange(torch.stack(emb), "f b d t -> b d f t") # attention def exists(val): return val is not None class RMSNorm(nn.Module): def __init__(self, dim): super().__init__() self.scale = dim**0.5 self.gamma = nn.Parameter(torch.ones(dim)) def forward(self, x): return F.normalize(x, dim=-1) * self.scale * self.gamma # attention class FeedForward(nn.Module): def __init__(self, dim, mult=4, dropout=0.0): super().__init__() dim_inner = int(dim * mult) self.net = nn.Sequential( RMSNorm(dim), nn.Linear(dim, dim_inner), nn.GELU(), nn.Dropout(dropout), nn.Linear(dim_inner, dim), nn.Dropout(dropout), ) def forward(self, x): return self.net(x) class Attention(nn.Module): def __init__( self, dim, heads=8, dim_head=64, dropout=0.0, rotary_embed=None, flash=True ): super().__init__() self.heads = heads self.scale = dim_head**-0.5 dim_inner = heads * dim_head self.rotary_embed = rotary_embed self.attend = Attend(flash=flash, dropout=dropout) self.norm = RMSNorm(dim) self.to_qkv = nn.Linear(dim, dim_inner * 3, bias=False) self.to_gates = nn.Linear(dim, heads) self.to_out = nn.Sequential( nn.Linear(dim_inner, dim, bias=False), nn.Dropout(dropout) ) def forward(self, x): x = self.norm(x) q, k, v = rearrange( self.to_qkv(x), "b n (qkv h d) -> qkv b h n d", qkv=3, h=self.heads ) if exists(self.rotary_embed): q = self.rotary_embed.rotate_queries_or_keys(q) k = self.rotary_embed.rotate_queries_or_keys(k) out = self.attend(q, k, v) gates = self.to_gates(x) out = out * rearrange(gates, "b n h -> b h n 1").sigmoid() out = rearrange(out, "b h n d -> b n (h d)") return self.to_out(out) class Transformer(nn.Module): def __init__( self, *, dim, depth, dim_head=64, heads=8, attn_dropout=0.0, ff_dropout=0.0, ff_mult=4, norm_output=True, rotary_embed=None, flash_attn=True, ): super().__init__() self.layers = nn.ModuleList([]) for _ in range(depth): attn = Attention( dim=dim, dim_head=dim_head, heads=heads, dropout=attn_dropout, rotary_embed=rotary_embed, flash=flash_attn, ) self.layers.append( nn.ModuleList( [attn, FeedForward(dim=dim, mult=ff_mult, dropout=ff_dropout)] ) ) self.norm = RMSNorm(dim) if norm_output else nn.Identity() def forward(self, x): for attn, ff in self.layers: x = attn(x) + x x = ff(x) + x return self.norm(x) class LayerNorm(nn.Module): r"""LayerNorm that supports two data formats: channels_last (default) or channels_first. The ordering of the dimensions in the inputs. channels_last corresponds to inputs with shape (batch_size, height, width, channels) while channels_first corresponds to inputs with shape (batch_size, channels, height, width). """ def __init__(self, normalized_shape, eps=1e-6, data_format="channels_last"): super().__init__() self.weight = nn.Parameter(torch.ones(normalized_shape)) self.bias = nn.Parameter(torch.zeros(normalized_shape)) self.eps = eps self.data_format = data_format if self.data_format not in ["channels_last", "channels_first"]: raise NotImplementedError self.normalized_shape = (normalized_shape,) def forward(self, x): if self.data_format == "channels_last": return F.layer_norm( x, self.normalized_shape, self.weight, self.bias, self.eps ) elif self.data_format == "channels_first": u = x.mean(1, keepdim=True) s = (x - u).pow(2).mean(1, keepdim=True) x = (x - u) / torch.sqrt(s + self.eps) x = self.weight[:, None, None] * x + self.bias[:, None, None] return x class Block(nn.Module): r"""ConvNeXt Block. There are two equivalent implementations: (1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W) (2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back We use (2) as we find it slightly faster in PyTorch Args: dim (int): Number of input channels. drop_path (float): Stochastic depth rate. Default: 0.0 layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6. """ def __init__(self, dim, drop_path=0.0, layer_scale_init_value=1e-6): super().__init__() self.dwconv = WNConv1d( dim, dim, kernel_size=7, padding=3, groups=dim ) # depthwise conv self.norm = LayerNorm(dim, eps=1e-6) self.pwconv1 = nn.Linear( dim, 4 * dim ) # pointwise/1x1 convs, implemented with linear layers self.act = nn.GELU() self.pwconv2 = nn.Linear(4 * dim, dim) self.gamma = ( nn.Parameter(layer_scale_init_value * torch.ones((dim)), requires_grad=True) if layer_scale_init_value > 0 else None ) self.drop_path = nn.Identity() def forward(self, x): input = x x = self.dwconv(x) x = x.permute(0, 2, 1) # (N, C, H) -> (N, H, C) x = self.norm(x) x = self.pwconv1(x) x = self.act(x) x = self.pwconv2(x) if self.gamma is not None: x = self.gamma * x x = x.permute(0, 2, 1) x = input + self.drop_path(x) return x class ConvNeXtSimple(nn.Module): """No downsampling, just 4 blocks of ConvNeXt""" def __init__(self, dim, depth, drop_path=0.0, layer_scale_init_value=1e-6): super().__init__() self.blocks = nn.ModuleList( [Block(dim, drop_path, layer_scale_init_value) for _ in range(depth)] ) def forward(self, x): for block in self.blocks: x = block(x) return x class ConvNeXtBlock(nn.Module): """ConvNeXt Block adapted from https://github.com/facebookresearch/ConvNeXt to 1D audio signal. Args: dim (int): Number of input channels. intermediate_dim (int): Dimensionality of the intermediate layer. layer_scale_init_value (float, optional): Initial value for the layer scale. None means no scaling. Defaults to None. adanorm_num_embeddings (int, optional): Number of embeddings for AdaLayerNorm. None means non-conditional LayerNorm. Defaults to None. """ def __init__( self, dim: int, intermediate_dim: int, layer_scale_init_value: float, adanorm_num_embeddings: Optional[int] = None, ): super().__init__() self.dwconv = nn.Conv1d( dim, dim, kernel_size=7, padding=3, groups=dim ) # depthwise conv self.adanorm = adanorm_num_embeddings is not None if adanorm_num_embeddings: self.norm = AdaLayerNorm(adanorm_num_embeddings, dim, eps=1e-6) else: self.norm = nn.LayerNorm(dim, eps=1e-6) self.pwconv1 = nn.Linear( dim, intermediate_dim ) # pointwise/1x1 convs, implemented with linear layers self.act = nn.GELU() self.pwconv2 = nn.Linear(intermediate_dim, dim) self.gamma = ( nn.Parameter(layer_scale_init_value * torch.ones(dim), requires_grad=True) if layer_scale_init_value > 0 else None ) def forward( self, x: torch.Tensor, cond_embedding_id: Optional[torch.Tensor] = None ) -> torch.Tensor: residual = x x = self.dwconv(x) x = x.transpose(1, 2) # (B, C, T) -> (B, T, C) if self.adanorm: assert cond_embedding_id is not None x = self.norm(x, cond_embedding_id) else: x = self.norm(x) x = self.pwconv1(x) x = self.act(x) x = self.pwconv2(x) if self.gamma is not None: x = self.gamma * x x = x.transpose(1, 2) # (B, T, C) -> (B, C, T) x = residual + x return x class AdaLayerNorm(nn.Module): """ Adaptive Layer Normalization module with learnable embeddings per `num_embeddings` classes Args: num_embeddings (int): Number of embeddings. embedding_dim (int): Dimension of the embeddings. """ def __init__(self, num_embeddings: int, embedding_dim: int, eps: float = 1e-6): super().__init__() self.eps = eps self.dim = embedding_dim self.scale = nn.Embedding( num_embeddings=num_embeddings, embedding_dim=embedding_dim ) self.shift = nn.Embedding( num_embeddings=num_embeddings, embedding_dim=embedding_dim ) torch.nn.init.ones_(self.scale.weight) torch.nn.init.zeros_(self.shift.weight) def forward(self, x: torch.Tensor, cond_embedding_id: torch.Tensor) -> torch.Tensor: scale = self.scale(cond_embedding_id) shift = self.shift(cond_embedding_id) x = nn.functional.layer_norm(x, (self.dim,), eps=self.eps) x = x * scale + shift return x window_to_bands = { 2048: 160, 1024: 128, 512: 64, } class BSTD(nn.Module): def __init__( self, window_length: int, hop_factor: float = 0.25, sample_rate: int = 48000, num_channels: int = 64, num_heads: int = 4, depth: int = 5, flash_attn: bool = True, ): """Band-split complex spectrogram discriminator Parameters ---------- window_length : int Window length of STFT. hop_factor : float, optional Hop factor of the STFT, defaults to ``0.25 * window_length``. sample_rate : int, optional Sampling rate of audio in Hz, by default 44100 """ super().__init__() # parameters self.window_length = window_length self.window = torch.hann_window(window_length) self.hop_factor = hop_factor self.sample_rate = sample_rate self.num_bands = window_to_bands[window_length] # neural filterbank self.neural_filterbank = NeuralFilterbank( n_fft=window_length, n_filterbank=self.num_bands, sample_rate=sample_rate, out_dim=num_channels, ) # spectral attention transformer_kwargs = dict( dim=num_channels, heads=num_heads, dim_head=num_channels // num_heads, attn_dropout=0.0, ff_dropout=0.0, flash_attn=flash_attn, norm_output=False, ) rotary_emb = RotaryEmbedding(dim=num_channels // num_heads) freq_attn = [] for _ in range(depth): freq_attn.append( Transformer( depth=1, rotary_embed=rotary_emb, **transformer_kwargs ).bfloat16() ) self.freq_attn = nn.ModuleList(freq_attn) # conv layers self.convs = nn.ModuleList( [ WNConv2d(num_channels, num_channels, (3, 9), (1, 1), padding=(1, 4)), WNConv2d(num_channels, num_channels, (3, 9), (1, 2), padding=(1, 4)), WNConv2d(num_channels, num_channels, (3, 9), (1, 2), padding=(1, 4)), WNConv2d(num_channels, num_channels, (3, 9), (1, 2), padding=(1, 4)), WNConv2d(num_channels, num_channels, (3, 3), (1, 1), padding=(1, 1)), ] ) self.conv_post = WNConv2d( num_channels, 1, (3, 3), (1, 1), padding=(1, 1), act=False ) def spectrogram(self, x): x = AudioSignal(x, self.sample_rate, stft_params=self.stft_params) x = torch.view_as_real(x.stft()) x = rearrange(x, "b ch f t c -> (b ch) c t f", ch=2) # Split into bands x_bands = [x[..., b[0] : b[1]] for b in self.bands] return x_bands def forward(self, x): """ einops b: batch f: frequency d: feature dimension s: stereo c: complex """ # reshape audio b, s, _ = x.shape x = rearrange(x, "b s t -> (b s) t") # short-time Fourier transform (b s) t -> (b s) f t c self.window = self.window.to(x.device) spec = torch.stft( x, n_fft=self.window_length, hop_length=int(self.window_length * self.hop_factor), win_length=self.window_length, window=self.window, onesided=True, return_complex=True, ) spec = torch.view_as_real(spec) t = spec.shape[-2] # neural filterbank (b s) f t c -> (b s) d f t spec = rearrange(spec, "(b s) f t c -> (b s c) f t", b=b, s=s, t=t, c=2) out = self.neural_filterbank(spec) # freq attention + conv fmap = [] for attn, conv in zip(self.freq_attn, self.convs): out = rearrange(out, "(b s) d f t -> (b s t) f d", b=b, s=s) out = attn(out.bfloat16()).float() fmap.append(out) out = rearrange(out, "(b s t) f d -> (b s) d f t", b=b, s=s) out = conv(out) fmap.append(out) out = self.conv_post(out) fmap.append(out) return fmap class Discriminator(ml.BaseModel): def __init__( self, rates: list = [], periods: list = [2, 3, 5, 7, 11], fft_sizes: list = [2048, 1024, 512], sample_rate: int = 44100, bands: list = BANDS, ): """Discriminator that combines multiple discriminators. Parameters ---------- rates : list, optional sampling rates (in Hz) to run MSD at, by default [] If empty, MSD is not used. periods : list, optional periods (of samples) to run MPD at, by default [2, 3, 5, 7, 11] fft_sizes : list, optional Window sizes of the FFT to run MRD at, by default [2048, 1024, 512] sample_rate : int, optional Sampling rate of audio in Hz, by default 44100 bands : list, optional Bands to run MRD at, by default `BANDS` """ super().__init__() discs = [] discs += [MPD(p) for p in periods] discs += [MSD(r, sample_rate=sample_rate) for r in rates] discs += [MRD(f, sample_rate=sample_rate, bands=bands) for f in fft_sizes] discs += [BSTD(f) for f in fft_sizes] self.discriminators = nn.ModuleList(discs) def preprocess(self, y): # Remove DC offset y = y - y.mean(dim=-1, keepdims=True) # Peak normalize the volume of input audio y = 0.8 * y / (y.abs().max(dim=-1, keepdim=True)[0] + 1e-9) return y def forward(self, x): x = self.preprocess(x) fmaps = [d(x) for d in self.discriminators] return fmaps if __name__ == "__main__": disc = Discriminator() x = torch.zeros(1, 1, 44100) results = disc(x) for i, result in enumerate(results): print(f"disc{i}") for i, r in enumerate(result): print(r.shape, r.mean(), r.min(), r.max()) print()