mirror of
https://github.com/Nighthawk42/soprano-factory.git
synced 2026-08-30 04:30:21 +00:00
144 lines
5.2 KiB
Python
144 lines
5.2 KiB
Python
import torch
|
|
from torch import nn
|
|
from torch.nn import functional as F
|
|
from torch.nn.utils.parametrizations import weight_norm, spectral_norm
|
|
|
|
class DiscriminatorP(nn.Module):
|
|
def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):
|
|
super(DiscriminatorP, self).__init__()
|
|
self.period = period
|
|
self.use_spectral_norm = use_spectral_norm
|
|
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
|
|
self.convs = nn.ModuleList([
|
|
norm_f(nn.Conv2d(1, 32, (kernel_size, 1), (stride, 1), padding=(2, 0))),
|
|
norm_f(nn.Conv2d(32, 128, (kernel_size, 1), (stride, 1), padding=(2, 0))),
|
|
norm_f(nn.Conv2d(128, 512, (kernel_size, 1), (stride, 1), padding=(2, 0))),
|
|
norm_f(nn.Conv2d(512, 1024, (kernel_size, 1), (stride, 1), padding=(2, 0))),
|
|
norm_f(nn.Conv2d(1024, 1024, (kernel_size, 1), 1, padding=(2, 0))),
|
|
])
|
|
self.conv_post = norm_f(nn.Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
|
|
|
|
def forward(self, x):
|
|
fmap = []
|
|
|
|
# 1d to 2d
|
|
b, c, t = x.shape
|
|
if t % self.period != 0: # pad valid
|
|
n_pad = self.period - (t % self.period)
|
|
x = F.pad(x, (0, n_pad), "reflect")
|
|
t = t + n_pad
|
|
x = x.view(b, c, t // self.period, self.period)
|
|
|
|
for l in self.convs:
|
|
x = l(x)
|
|
x = F.leaky_relu(x, 0.1)
|
|
fmap.append(x)
|
|
x = self.conv_post(x)
|
|
fmap.append(x)
|
|
x = torch.flatten(x, 1, -1)
|
|
|
|
return x, fmap
|
|
|
|
|
|
class DiscriminatorS(nn.Module):
|
|
def __init__(self, use_spectral_norm=False):
|
|
super(DiscriminatorS, self).__init__()
|
|
norm_f = weight_norm if use_spectral_norm == False else spectral_norm
|
|
self.convs = nn.ModuleList([
|
|
norm_f(nn.Conv1d(1, 16, 15, 1, padding=7)),
|
|
norm_f(nn.Conv1d(16, 64, 41, 4, groups=4, padding=20)),
|
|
norm_f(nn.Conv1d(64, 256, 41, 4, groups=16, padding=20)),
|
|
norm_f(nn.Conv1d(256, 1024, 41, 4, groups=64, padding=20)),
|
|
norm_f(nn.Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),
|
|
norm_f(nn.Conv1d(1024, 1024, 5, 1, padding=2)),
|
|
])
|
|
self.conv_post = norm_f(nn.Conv1d(1024, 1, 3, 1, padding=1))
|
|
|
|
def forward(self, x):
|
|
fmap = []
|
|
for l in self.convs:
|
|
x = l(x)
|
|
x = F.leaky_relu(x, 0.1)
|
|
fmap.append(x)
|
|
x = self.conv_post(x)
|
|
fmap.append(x)
|
|
x = torch.flatten(x, 1, -1)
|
|
|
|
return x, fmap
|
|
|
|
|
|
class MultiPeriodDiscriminator(nn.Module):
|
|
def __init__(self, use_spectral_norm=False):
|
|
super(MultiPeriodDiscriminator, self).__init__()
|
|
self.discriminators = nn.ModuleList([
|
|
DiscriminatorP(2, use_spectral_norm=use_spectral_norm),
|
|
DiscriminatorP(3, use_spectral_norm=use_spectral_norm),
|
|
DiscriminatorP(5, use_spectral_norm=use_spectral_norm),
|
|
DiscriminatorP(7, use_spectral_norm=use_spectral_norm),
|
|
DiscriminatorP(11, use_spectral_norm=use_spectral_norm),
|
|
])
|
|
|
|
def forward(self, y, y_hat):
|
|
y_d_rs = []
|
|
y_d_gs = []
|
|
fmap_rs = []
|
|
fmap_gs = []
|
|
for i, d in enumerate(self.discriminators):
|
|
y_d_r, fmap_r = d(y)
|
|
y_d_g, fmap_g = d(y_hat)
|
|
y_d_rs.append(y_d_r)
|
|
y_d_gs.append(y_d_g)
|
|
fmap_rs.append(fmap_r)
|
|
fmap_gs.append(fmap_g)
|
|
|
|
return y_d_rs, y_d_gs, fmap_rs, fmap_gs
|
|
|
|
|
|
class MultiScaleDiscriminator(nn.Module):
|
|
def __init__(self, use_spectral_norm=False):
|
|
super(MultiScaleDiscriminator, self).__init__()
|
|
self.discriminators = nn.ModuleList([
|
|
DiscriminatorS(use_spectral_norm=use_spectral_norm),
|
|
DiscriminatorS(use_spectral_norm=use_spectral_norm),
|
|
DiscriminatorS(use_spectral_norm=use_spectral_norm),
|
|
])
|
|
self.meanpools = nn.ModuleList([
|
|
nn.AvgPool1d(4, 2, padding=2),
|
|
nn.AvgPool1d(4, 2, padding=2)
|
|
])
|
|
|
|
def forward(self, y, y_hat):
|
|
y_d_rs = []
|
|
y_d_gs = []
|
|
fmap_rs = []
|
|
fmap_gs = []
|
|
for i, d in enumerate(self.discriminators):
|
|
if i != 0:
|
|
y = self.meanpools[i-1](y)
|
|
y_hat = self.meanpools[i-1](y_hat)
|
|
y_d_r, fmap_r = d(y)
|
|
y_d_g, fmap_g = d(y_hat)
|
|
y_d_rs.append(y_d_r)
|
|
y_d_gs.append(y_d_g)
|
|
fmap_rs.append(fmap_r)
|
|
fmap_gs.append(fmap_g)
|
|
|
|
return y_d_rs, y_d_gs, fmap_rs, fmap_gs
|
|
|
|
class Discriminator(nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.mpd = MultiPeriodDiscriminator()
|
|
self.msd = MultiScaleDiscriminator()
|
|
|
|
def forward(self, y, y_hat):
|
|
# y: real audio, y_hat: gen audio
|
|
# Unsqueeze if needed (B, T) -> (B, 1, T)
|
|
if y.ndim == 2: y = y.unsqueeze(1)
|
|
if y_hat.ndim == 2: y_hat = y_hat.unsqueeze(1)
|
|
|
|
y_d_rs_p, y_d_gs_p, fmap_rs_p, fmap_gs_p = self.mpd(y, y_hat)
|
|
y_d_rs_s, y_d_gs_s, fmap_rs_s, fmap_gs_s = self.msd(y, y_hat)
|
|
|
|
return (y_d_rs_p + y_d_rs_s, y_d_gs_p + y_d_gs_s, fmap_rs_p + fmap_rs_s, fmap_gs_p + fmap_gs_s)
|