lint code (no other modify)

This commit is contained in:
Lengyue
2023-09-05 01:21:17 -04:00
parent eabeb72e27
commit d4d0082b05

498
models.py
View File

@@ -1,4 +1,3 @@
import copy
import math import math
import torch import torch
from torch import nn from torch import nn
@@ -9,13 +8,17 @@ import modules
import attentions import attentions
import monotonic_align import monotonic_align
from torch.nn import Conv1d, ConvTranspose1d, AvgPool1d, Conv2d from torch.nn import Conv1d, ConvTranspose1d, Conv2d
from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm
from commons import init_weights, get_padding from commons import init_weights, get_padding
from text import symbols, num_tones, num_languages from text import symbols, num_tones, num_languages
class DurationDiscriminator(nn.Module): # vits2 class DurationDiscriminator(nn.Module): # vits2
def __init__(self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0): def __init__(
self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0
):
super().__init__() super().__init__()
self.in_channels = in_channels self.in_channels = in_channels
@@ -25,24 +28,29 @@ class DurationDiscriminator(nn.Module): #vits2
self.gin_channels = gin_channels self.gin_channels = gin_channels
self.drop = nn.Dropout(p_dropout) self.drop = nn.Dropout(p_dropout)
self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size, padding=kernel_size//2) self.conv_1 = nn.Conv1d(
in_channels, filter_channels, kernel_size, padding=kernel_size // 2
)
self.norm_1 = modules.LayerNorm(filter_channels) self.norm_1 = modules.LayerNorm(filter_channels)
self.conv_2 = nn.Conv1d(filter_channels, filter_channels, kernel_size, padding=kernel_size//2) self.conv_2 = nn.Conv1d(
filter_channels, filter_channels, kernel_size, padding=kernel_size // 2
)
self.norm_2 = modules.LayerNorm(filter_channels) self.norm_2 = modules.LayerNorm(filter_channels)
self.dur_proj = nn.Conv1d(1, filter_channels, 1) self.dur_proj = nn.Conv1d(1, filter_channels, 1)
self.pre_out_conv_1 = nn.Conv1d(2*filter_channels, filter_channels, kernel_size, padding=kernel_size//2) self.pre_out_conv_1 = nn.Conv1d(
2 * filter_channels, filter_channels, kernel_size, padding=kernel_size // 2
)
self.pre_out_norm_1 = modules.LayerNorm(filter_channels) self.pre_out_norm_1 = modules.LayerNorm(filter_channels)
self.pre_out_conv_2 = nn.Conv1d(filter_channels, filter_channels, kernel_size, padding=kernel_size//2) self.pre_out_conv_2 = nn.Conv1d(
filter_channels, filter_channels, kernel_size, padding=kernel_size // 2
)
self.pre_out_norm_2 = modules.LayerNorm(filter_channels) self.pre_out_norm_2 = modules.LayerNorm(filter_channels)
if gin_channels != 0: if gin_channels != 0:
self.cond = nn.Conv1d(gin_channels, in_channels, 1) self.cond = nn.Conv1d(gin_channels, in_channels, 1)
self.output_layer = nn.Sequential( self.output_layer = nn.Sequential(nn.Linear(filter_channels, 1), nn.Sigmoid())
nn.Linear(filter_channels, 1),
nn.Sigmoid()
)
def forward_probability(self, x, x_mask, dur, g=None): def forward_probability(self, x, x_mask, dur, g=None):
dur = self.dur_proj(dur) dur = self.dur_proj(dur)
@@ -81,8 +89,10 @@ class DurationDiscriminator(nn.Module): #vits2
return output_probs return output_probs
class TransformerCouplingBlock(nn.Module): class TransformerCouplingBlock(nn.Module):
def __init__(self, def __init__(
self,
channels, channels,
hidden_channels, hidden_channels,
filter_channels, filter_channels,
@@ -92,7 +102,7 @@ class TransformerCouplingBlock(nn.Module):
p_dropout, p_dropout,
n_flows=4, n_flows=4,
gin_channels=0, gin_channels=0,
share_parameter=False share_parameter=False,
): ):
super().__init__() super().__init__()
@@ -105,11 +115,36 @@ class TransformerCouplingBlock(nn.Module):
self.flows = nn.ModuleList() self.flows = nn.ModuleList()
self.wn = attentions.FFT(hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout, isflow = True, gin_channels = self.gin_channels) if share_parameter else None self.wn = (
attentions.FFT(
hidden_channels,
filter_channels,
n_heads,
n_layers,
kernel_size,
p_dropout,
isflow=True,
gin_channels=self.gin_channels,
)
if share_parameter
else None
)
for i in range(n_flows): for i in range(n_flows):
self.flows.append( self.flows.append(
modules.TransformerCouplingLayer(channels, hidden_channels, kernel_size, n_layers, n_heads, p_dropout, filter_channels, mean_only=True, wn_sharing_parameter=self.wn, gin_channels = self.gin_channels)) modules.TransformerCouplingLayer(
channels,
hidden_channels,
kernel_size,
n_layers,
n_heads,
p_dropout,
filter_channels,
mean_only=True,
wn_sharing_parameter=self.wn,
gin_channels=self.gin_channels,
)
)
self.flows.append(modules.Flip()) self.flows.append(modules.Flip())
def forward(self, x, x_mask, g=None, reverse=False): def forward(self, x, x_mask, g=None, reverse=False):
@@ -121,8 +156,17 @@ class TransformerCouplingBlock(nn.Module):
x = flow(x, x_mask, g=g, reverse=reverse) x = flow(x, x_mask, g=g, reverse=reverse)
return x return x
class StochasticDurationPredictor(nn.Module): class StochasticDurationPredictor(nn.Module):
def __init__(self, in_channels, filter_channels, kernel_size, p_dropout, n_flows=4, gin_channels=0): def __init__(
self,
in_channels,
filter_channels,
kernel_size,
p_dropout,
n_flows=4,
gin_channels=0,
):
super().__init__() super().__init__()
filter_channels = in_channels # it needs to be removed from future version. filter_channels = in_channels # it needs to be removed from future version.
self.in_channels = in_channels self.in_channels = in_channels
@@ -136,21 +180,29 @@ class StochasticDurationPredictor(nn.Module):
self.flows = nn.ModuleList() self.flows = nn.ModuleList()
self.flows.append(modules.ElementwiseAffine(2)) self.flows.append(modules.ElementwiseAffine(2))
for i in range(n_flows): for i in range(n_flows):
self.flows.append(modules.ConvFlow(2, filter_channels, kernel_size, n_layers=3)) self.flows.append(
modules.ConvFlow(2, filter_channels, kernel_size, n_layers=3)
)
self.flows.append(modules.Flip()) self.flows.append(modules.Flip())
self.post_pre = nn.Conv1d(1, filter_channels, 1) self.post_pre = nn.Conv1d(1, filter_channels, 1)
self.post_proj = nn.Conv1d(filter_channels, filter_channels, 1) self.post_proj = nn.Conv1d(filter_channels, filter_channels, 1)
self.post_convs = modules.DDSConv(filter_channels, kernel_size, n_layers=3, p_dropout=p_dropout) self.post_convs = modules.DDSConv(
filter_channels, kernel_size, n_layers=3, p_dropout=p_dropout
)
self.post_flows = nn.ModuleList() self.post_flows = nn.ModuleList()
self.post_flows.append(modules.ElementwiseAffine(2)) self.post_flows.append(modules.ElementwiseAffine(2))
for i in range(4): for i in range(4):
self.post_flows.append(modules.ConvFlow(2, filter_channels, kernel_size, n_layers=3)) self.post_flows.append(
modules.ConvFlow(2, filter_channels, kernel_size, n_layers=3)
)
self.post_flows.append(modules.Flip()) self.post_flows.append(modules.Flip())
self.pre = nn.Conv1d(in_channels, filter_channels, 1) self.pre = nn.Conv1d(in_channels, filter_channels, 1)
self.proj = nn.Conv1d(filter_channels, filter_channels, 1) self.proj = nn.Conv1d(filter_channels, filter_channels, 1)
self.convs = modules.DDSConv(filter_channels, kernel_size, n_layers=3, p_dropout=p_dropout) self.convs = modules.DDSConv(
filter_channels, kernel_size, n_layers=3, p_dropout=p_dropout
)
if gin_channels != 0: if gin_channels != 0:
self.cond = nn.Conv1d(gin_channels, filter_channels, 1) self.cond = nn.Conv1d(gin_channels, filter_channels, 1)
@@ -171,7 +223,10 @@ class StochasticDurationPredictor(nn.Module):
h_w = self.post_pre(w) h_w = self.post_pre(w)
h_w = self.post_convs(h_w, x_mask) h_w = self.post_convs(h_w, x_mask)
h_w = self.post_proj(h_w) * x_mask h_w = self.post_proj(h_w) * x_mask
e_q = torch.randn(w.size(0), 2, w.size(2)).to(device=x.device, dtype=x.dtype) * x_mask e_q = (
torch.randn(w.size(0), 2, w.size(2)).to(device=x.device, dtype=x.dtype)
* x_mask
)
z_q = e_q z_q = e_q
for flow in self.post_flows: for flow in self.post_flows:
z_q, logdet_q = flow(z_q, x_mask, g=(x + h_w)) z_q, logdet_q = flow(z_q, x_mask, g=(x + h_w))
@@ -179,8 +234,13 @@ class StochasticDurationPredictor(nn.Module):
z_u, z1 = torch.split(z_q, [1, 1], 1) z_u, z1 = torch.split(z_q, [1, 1], 1)
u = torch.sigmoid(z_u) * x_mask u = torch.sigmoid(z_u) * x_mask
z0 = (w - u) * x_mask z0 = (w - u) * x_mask
logdet_tot_q += torch.sum((F.logsigmoid(z_u) + F.logsigmoid(-z_u)) * x_mask, [1, 2]) logdet_tot_q += torch.sum(
logq = torch.sum(-0.5 * (math.log(2 * math.pi) + (e_q ** 2)) * x_mask, [1, 2]) - logdet_tot_q (F.logsigmoid(z_u) + F.logsigmoid(-z_u)) * x_mask, [1, 2]
)
logq = (
torch.sum(-0.5 * (math.log(2 * math.pi) + (e_q**2)) * x_mask, [1, 2])
- logdet_tot_q
)
logdet_tot = 0 logdet_tot = 0
z0, logdet = self.log_flow(z0, x_mask) z0, logdet = self.log_flow(z0, x_mask)
@@ -189,12 +249,18 @@ class StochasticDurationPredictor(nn.Module):
for flow in flows: for flow in flows:
z, logdet = flow(z, x_mask, g=x, reverse=reverse) z, logdet = flow(z, x_mask, g=x, reverse=reverse)
logdet_tot = logdet_tot + logdet logdet_tot = logdet_tot + logdet
nll = torch.sum(0.5 * (math.log(2 * math.pi) + (z ** 2)) * x_mask, [1, 2]) - logdet_tot nll = (
torch.sum(0.5 * (math.log(2 * math.pi) + (z**2)) * x_mask, [1, 2])
- logdet_tot
)
return nll + logq # [b] return nll + logq # [b]
else: else:
flows = list(reversed(self.flows)) flows = list(reversed(self.flows))
flows = flows[:-2] + [flows[-1]] # remove a useless vflow flows = flows[:-2] + [flows[-1]] # remove a useless vflow
z = torch.randn(x.size(0), 2, x.size(2)).to(device=x.device, dtype=x.dtype) * noise_scale z = (
torch.randn(x.size(0), 2, x.size(2)).to(device=x.device, dtype=x.dtype)
* noise_scale
)
for flow in flows: for flow in flows:
z = flow(z, x_mask, g=x, reverse=reverse) z = flow(z, x_mask, g=x, reverse=reverse)
z0, z1 = torch.split(z, [1, 1], 1) z0, z1 = torch.split(z, [1, 1], 1)
@@ -203,7 +269,9 @@ class StochasticDurationPredictor(nn.Module):
class DurationPredictor(nn.Module): class DurationPredictor(nn.Module):
def __init__(self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0): def __init__(
self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0
):
super().__init__() super().__init__()
self.in_channels = in_channels self.in_channels = in_channels
@@ -213,9 +281,13 @@ class DurationPredictor(nn.Module):
self.gin_channels = gin_channels self.gin_channels = gin_channels
self.drop = nn.Dropout(p_dropout) self.drop = nn.Dropout(p_dropout)
self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size, padding=kernel_size // 2) self.conv_1 = nn.Conv1d(
in_channels, filter_channels, kernel_size, padding=kernel_size // 2
)
self.norm_1 = modules.LayerNorm(filter_channels) self.norm_1 = modules.LayerNorm(filter_channels)
self.conv_2 = nn.Conv1d(filter_channels, filter_channels, kernel_size, padding=kernel_size // 2) self.conv_2 = nn.Conv1d(
filter_channels, filter_channels, kernel_size, padding=kernel_size // 2
)
self.norm_2 = modules.LayerNorm(filter_channels) self.norm_2 = modules.LayerNorm(filter_channels)
self.proj = nn.Conv1d(filter_channels, 1, 1) self.proj = nn.Conv1d(filter_channels, 1, 1)
@@ -240,7 +312,8 @@ class DurationPredictor(nn.Module):
class TextEncoder(nn.Module): class TextEncoder(nn.Module):
def __init__(self, def __init__(
self,
n_vocab, n_vocab,
out_channels, out_channels,
hidden_channels, hidden_channels,
@@ -249,7 +322,8 @@ class TextEncoder(nn.Module):
n_layers, n_layers,
kernel_size, kernel_size,
p_dropout, p_dropout,
gin_channels=0): gin_channels=0,
):
super().__init__() super().__init__()
self.n_vocab = n_vocab self.n_vocab = n_vocab
self.out_channels = out_channels self.out_channels = out_channels
@@ -276,16 +350,26 @@ class TextEncoder(nn.Module):
n_layers, n_layers,
kernel_size, kernel_size,
p_dropout, p_dropout,
gin_channels=self.gin_channels) gin_channels=self.gin_channels,
)
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
def forward(self, x, x_lengths, tone, language, bert, ja_bert, g=None): def forward(self, x, x_lengths, tone, language, bert, ja_bert, g=None):
bert_emb = self.zh_bert_proj(bert).transpose(1, 2) bert_emb = self.zh_bert_proj(bert).transpose(1, 2)
ja_bert_emb = = self.ja_bert_proj(ja_bert).transpose(1,2) ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2)
x = (self.emb(x)+ self.tone_emb(tone)+ self.language_emb(language) x = (
+ bert_emb +ja_bert_emb) * math.sqrt(self.hidden_channels) # [b, t, h] self.emb(x)
+ self.tone_emb(tone)
+ self.language_emb(language)
+ bert_emb
+ ja_bert_emb
) * math.sqrt(
self.hidden_channels
) # [b, t, h]
x = torch.transpose(x, 1, -1) # [b, h, t] x = torch.transpose(x, 1, -1) # [b, h, t]
x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype) x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
x.dtype
)
x = self.encoder(x * x_mask, x_mask, g=g) x = self.encoder(x * x_mask, x_mask, g=g)
stats = self.proj(x) * x_mask stats = self.proj(x) * x_mask
@@ -295,14 +379,16 @@ class TextEncoder(nn.Module):
class ResidualCouplingBlock(nn.Module): class ResidualCouplingBlock(nn.Module):
def __init__(self, def __init__(
self,
channels, channels,
hidden_channels, hidden_channels,
kernel_size, kernel_size,
dilation_rate, dilation_rate,
n_layers, n_layers,
n_flows=4, n_flows=4,
gin_channels=0): gin_channels=0,
):
super().__init__() super().__init__()
self.channels = channels self.channels = channels
self.hidden_channels = hidden_channels self.hidden_channels = hidden_channels
@@ -315,8 +401,16 @@ class ResidualCouplingBlock(nn.Module):
self.flows = nn.ModuleList() self.flows = nn.ModuleList()
for i in range(n_flows): for i in range(n_flows):
self.flows.append( self.flows.append(
modules.ResidualCouplingLayer(channels, hidden_channels, kernel_size, dilation_rate, n_layers, modules.ResidualCouplingLayer(
gin_channels=gin_channels, mean_only=True)) channels,
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
gin_channels=gin_channels,
mean_only=True,
)
)
self.flows.append(modules.Flip()) self.flows.append(modules.Flip())
def forward(self, x, x_mask, g=None, reverse=False): def forward(self, x, x_mask, g=None, reverse=False):
@@ -330,14 +424,16 @@ class ResidualCouplingBlock(nn.Module):
class PosteriorEncoder(nn.Module): class PosteriorEncoder(nn.Module):
def __init__(self, def __init__(
self,
in_channels, in_channels,
out_channels, out_channels,
hidden_channels, hidden_channels,
kernel_size, kernel_size,
dilation_rate, dilation_rate,
n_layers, n_layers,
gin_channels=0): gin_channels=0,
):
super().__init__() super().__init__()
self.in_channels = in_channels self.in_channels = in_channels
self.out_channels = out_channels self.out_channels = out_channels
@@ -348,11 +444,19 @@ class PosteriorEncoder(nn.Module):
self.gin_channels = gin_channels self.gin_channels = gin_channels
self.pre = nn.Conv1d(in_channels, hidden_channels, 1) self.pre = nn.Conv1d(in_channels, hidden_channels, 1)
self.enc = modules.WN(hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels) self.enc = modules.WN(
hidden_channels,
kernel_size,
dilation_rate,
n_layers,
gin_channels=gin_channels,
)
self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
def forward(self, x, x_lengths, g=None): def forward(self, x, x_lengths, g=None):
x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype) x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
x.dtype
)
x = self.pre(x) * x_mask x = self.pre(x) * x_mask
x = self.enc(x, x_mask, g=g) x = self.enc(x, x_mask, g=g)
stats = self.proj(x) * x_mask stats = self.proj(x) * x_mask
@@ -362,24 +466,45 @@ class PosteriorEncoder(nn.Module):
class Generator(torch.nn.Module): class Generator(torch.nn.Module):
def __init__(self, initial_channel, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, def __init__(
upsample_initial_channel, upsample_kernel_sizes, gin_channels=0): self,
initial_channel,
resblock,
resblock_kernel_sizes,
resblock_dilation_sizes,
upsample_rates,
upsample_initial_channel,
upsample_kernel_sizes,
gin_channels=0,
):
super(Generator, self).__init__() super(Generator, self).__init__()
self.num_kernels = len(resblock_kernel_sizes) self.num_kernels = len(resblock_kernel_sizes)
self.num_upsamples = len(upsample_rates) self.num_upsamples = len(upsample_rates)
self.conv_pre = Conv1d(initial_channel, upsample_initial_channel, 7, 1, padding=3) self.conv_pre = Conv1d(
resblock = modules.ResBlock1 if resblock == '1' else modules.ResBlock2 initial_channel, upsample_initial_channel, 7, 1, padding=3
)
resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2
self.ups = nn.ModuleList() self.ups = nn.ModuleList()
for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
self.ups.append(weight_norm( self.ups.append(
ConvTranspose1d(upsample_initial_channel // (2 ** i), upsample_initial_channel // (2 ** (i + 1)), weight_norm(
k, u, padding=(k - u) // 2))) ConvTranspose1d(
upsample_initial_channel // (2**i),
upsample_initial_channel // (2 ** (i + 1)),
k,
u,
padding=(k - u) // 2,
)
)
)
self.resblocks = nn.ModuleList() self.resblocks = nn.ModuleList()
for i in range(len(self.ups)): for i in range(len(self.ups)):
ch = upsample_initial_channel // (2 ** (i + 1)) ch = upsample_initial_channel // (2 ** (i + 1))
for j, (k, d) in enumerate(zip(resblock_kernel_sizes, resblock_dilation_sizes)): for j, (k, d) in enumerate(
zip(resblock_kernel_sizes, resblock_dilation_sizes)
):
self.resblocks.append(resblock(ch, k, d)) self.resblocks.append(resblock(ch, k, d))
self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False) self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
@@ -410,11 +535,11 @@ class Generator(torch.nn.Module):
return x return x
def remove_weight_norm(self): def remove_weight_norm(self):
print('Removing weight norm...') print("Removing weight norm...")
for l in self.ups: for layer in self.ups:
remove_weight_norm(l) remove_weight_norm(layer)
for l in self.resblocks: for layer in self.resblocks:
l.remove_weight_norm() layer.remove_weight_norm()
class DiscriminatorP(torch.nn.Module): class DiscriminatorP(torch.nn.Module):
@@ -422,14 +547,56 @@ class DiscriminatorP(torch.nn.Module):
super(DiscriminatorP, self).__init__() super(DiscriminatorP, self).__init__()
self.period = period self.period = period
self.use_spectral_norm = use_spectral_norm self.use_spectral_norm = use_spectral_norm
norm_f = weight_norm if use_spectral_norm == False else spectral_norm norm_f = weight_norm if use_spectral_norm is False else spectral_norm
self.convs = nn.ModuleList([ self.convs = nn.ModuleList(
norm_f(Conv2d(1, 32, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))), [
norm_f(Conv2d(32, 128, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))), norm_f(
norm_f(Conv2d(128, 512, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))), Conv2d(
norm_f(Conv2d(512, 1024, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))), 1,
norm_f(Conv2d(1024, 1024, (kernel_size, 1), 1, padding=(get_padding(kernel_size, 1), 0))), 32,
]) (kernel_size, 1),
(stride, 1),
padding=(get_padding(kernel_size, 1), 0),
)
),
norm_f(
Conv2d(
32,
128,
(kernel_size, 1),
(stride, 1),
padding=(get_padding(kernel_size, 1), 0),
)
),
norm_f(
Conv2d(
128,
512,
(kernel_size, 1),
(stride, 1),
padding=(get_padding(kernel_size, 1), 0),
)
),
norm_f(
Conv2d(
512,
1024,
(kernel_size, 1),
(stride, 1),
padding=(get_padding(kernel_size, 1), 0),
)
),
norm_f(
Conv2d(
1024,
1024,
(kernel_size, 1),
1,
padding=(get_padding(kernel_size, 1), 0),
)
),
]
)
self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0))) self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
def forward(self, x): def forward(self, x):
@@ -443,8 +610,8 @@ class DiscriminatorP(torch.nn.Module):
t = t + n_pad t = t + n_pad
x = x.view(b, c, t // self.period, self.period) x = x.view(b, c, t // self.period, self.period)
for l in self.convs: for layer in self.convs:
x = l(x) x = layer(x)
x = F.leaky_relu(x, modules.LRELU_SLOPE) x = F.leaky_relu(x, modules.LRELU_SLOPE)
fmap.append(x) fmap.append(x)
x = self.conv_post(x) x = self.conv_post(x)
@@ -457,22 +624,24 @@ class DiscriminatorP(torch.nn.Module):
class DiscriminatorS(torch.nn.Module): class DiscriminatorS(torch.nn.Module):
def __init__(self, use_spectral_norm=False): def __init__(self, use_spectral_norm=False):
super(DiscriminatorS, self).__init__() super(DiscriminatorS, self).__init__()
norm_f = weight_norm if use_spectral_norm == False else spectral_norm norm_f = weight_norm if use_spectral_norm is False else spectral_norm
self.convs = nn.ModuleList([ self.convs = nn.ModuleList(
[
norm_f(Conv1d(1, 16, 15, 1, padding=7)), norm_f(Conv1d(1, 16, 15, 1, padding=7)),
norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)), norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),
norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)), norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),
norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)), norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),
norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)), norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),
norm_f(Conv1d(1024, 1024, 5, 1, padding=2)), norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),
]) ]
)
self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1)) self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))
def forward(self, x): def forward(self, x):
fmap = [] fmap = []
for l in self.convs: for layer in self.convs:
x = l(x) x = layer(x)
x = F.leaky_relu(x, modules.LRELU_SLOPE) x = F.leaky_relu(x, modules.LRELU_SLOPE)
fmap.append(x) fmap.append(x)
x = self.conv_post(x) x = self.conv_post(x)
@@ -488,7 +657,9 @@ class MultiPeriodDiscriminator(torch.nn.Module):
periods = [2, 3, 5, 7, 11] periods = [2, 3, 5, 7, 11]
discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)] discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]
discs = discs + [DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods] discs = discs + [
DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods
]
self.discriminators = nn.ModuleList(discs) self.discriminators = nn.ModuleList(discs)
def forward(self, y, y_hat): def forward(self, y, y_hat):
@@ -506,11 +677,12 @@ class MultiPeriodDiscriminator(torch.nn.Module):
return y_d_rs, y_d_gs, fmap_rs, fmap_gs return y_d_rs, y_d_gs, fmap_rs, fmap_gs
class ReferenceEncoder(nn.Module): class ReferenceEncoder(nn.Module):
''' """
inputs --- [N, Ty/r, n_mels*r] mels inputs --- [N, Ty/r, n_mels*r] mels
outputs --- [N, ref_enc_gru_size] outputs --- [N, ref_enc_gru_size]
''' """
def __init__(self, spec_channels, gin_channels=0): def __init__(self, spec_channels, gin_channels=0):
@@ -519,18 +691,27 @@ class ReferenceEncoder(nn.Module):
ref_enc_filters = [32, 32, 64, 64, 128, 128] ref_enc_filters = [32, 32, 64, 64, 128, 128]
K = len(ref_enc_filters) K = len(ref_enc_filters)
filters = [1] + ref_enc_filters filters = [1] + ref_enc_filters
convs = [weight_norm(nn.Conv2d(in_channels=filters[i], convs = [
weight_norm(
nn.Conv2d(
in_channels=filters[i],
out_channels=filters[i + 1], out_channels=filters[i + 1],
kernel_size=(3, 3), kernel_size=(3, 3),
stride=(2, 2), stride=(2, 2),
padding=(1, 1))) for i in range(K)] padding=(1, 1),
)
)
for i in range(K)
]
self.convs = nn.ModuleList(convs) self.convs = nn.ModuleList(convs)
# self.wns = nn.ModuleList([weight_norm(num_features=ref_enc_filters[i]) for i in range(K)]) # self.wns = nn.ModuleList([weight_norm(num_features=ref_enc_filters[i]) for i in range(K)]) # noqa: E501
out_channels = self.calculate_channels(spec_channels, 3, 2, 1, K) out_channels = self.calculate_channels(spec_channels, 3, 2, 1, K)
self.gru = nn.GRU(input_size=ref_enc_filters[-1] * out_channels, self.gru = nn.GRU(
input_size=ref_enc_filters[-1] * out_channels,
hidden_size=256 // 2, hidden_size=256 // 2,
batch_first=True) batch_first=True,
)
self.proj = nn.Linear(128, gin_channels) self.proj = nn.Linear(128, gin_channels)
def forward(self, inputs, mask=None): def forward(self, inputs, mask=None):
@@ -562,7 +743,8 @@ class SynthesizerTrn(nn.Module):
Synthesizer for Training Synthesizer for Training
""" """
def __init__(self, def __init__(
self,
n_vocab, n_vocab,
spec_channels, spec_channels,
segment_size, segment_size,
@@ -586,7 +768,8 @@ class SynthesizerTrn(nn.Module):
n_layers_trans_flow=6, n_layers_trans_flow=6,
flow_share_parameter=False, flow_share_parameter=False,
use_transformer_flow=True, use_transformer_flow=True,
**kwargs): **kwargs
):
super().__init__() super().__init__()
self.n_vocab = n_vocab self.n_vocab = n_vocab
@@ -608,7 +791,9 @@ class SynthesizerTrn(nn.Module):
self.n_speakers = n_speakers self.n_speakers = n_speakers
self.gin_channels = gin_channels self.gin_channels = gin_channels
self.n_layers_trans_flow = n_layers_trans_flow self.n_layers_trans_flow = n_layers_trans_flow
self.use_spk_conditioned_encoder = kwargs.get("use_spk_conditioned_encoder", True) self.use_spk_conditioned_encoder = kwargs.get(
"use_spk_conditioned_encoder", True
)
self.use_sdp = use_sdp self.use_sdp = use_sdp
self.use_noise_scaled_mas = kwargs.get("use_noise_scaled_mas", False) self.use_noise_scaled_mas = kwargs.get("use_noise_scaled_mas", False)
self.mas_noise_scale_initial = kwargs.get("mas_noise_scale_initial", 0.01) self.mas_noise_scale_initial = kwargs.get("mas_noise_scale_initial", 0.01)
@@ -616,7 +801,8 @@ class SynthesizerTrn(nn.Module):
self.current_mas_noise_scale = self.mas_noise_scale_initial self.current_mas_noise_scale = self.mas_noise_scale_initial
if self.use_spk_conditioned_encoder and gin_channels > 0: if self.use_spk_conditioned_encoder and gin_channels > 0:
self.enc_gin_channels = gin_channels self.enc_gin_channels = gin_channels
self.enc_p = TextEncoder(n_vocab, self.enc_p = TextEncoder(
n_vocab,
inter_channels, inter_channels,
hidden_channels, hidden_channels,
filter_channels, filter_channels,
@@ -624,17 +810,55 @@ class SynthesizerTrn(nn.Module):
n_layers, n_layers,
kernel_size, kernel_size,
p_dropout, p_dropout,
gin_channels=self.enc_gin_channels) gin_channels=self.enc_gin_channels,
self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, )
upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels) self.dec = Generator(
self.enc_q = PosteriorEncoder(spec_channels, inter_channels, hidden_channels, 5, 1, 16, inter_channels,
gin_channels=gin_channels) resblock,
resblock_kernel_sizes,
resblock_dilation_sizes,
upsample_rates,
upsample_initial_channel,
upsample_kernel_sizes,
gin_channels=gin_channels,
)
self.enc_q = PosteriorEncoder(
spec_channels,
inter_channels,
hidden_channels,
5,
1,
16,
gin_channels=gin_channels,
)
if use_transformer_flow: if use_transformer_flow:
self.flow = TransformerCouplingBlock(inter_channels, hidden_channels, filter_channels, n_heads, n_layers_trans_flow, 5, p_dropout, n_flow_layer, gin_channels=gin_channels,share_parameter= flow_share_parameter) self.flow = TransformerCouplingBlock(
inter_channels,
hidden_channels,
filter_channels,
n_heads,
n_layers_trans_flow,
5,
p_dropout,
n_flow_layer,
gin_channels=gin_channels,
share_parameter=flow_share_parameter,
)
else: else:
self.flow = ResidualCouplingBlock(inter_channels, hidden_channels, 5, 1, n_flow_layer, gin_channels=gin_channels) self.flow = ResidualCouplingBlock(
self.sdp = StochasticDurationPredictor(hidden_channels, 192, 3, 0.5, 4, gin_channels=gin_channels) inter_channels,
self.dp = DurationPredictor(hidden_channels, 256, 3, 0.5, gin_channels=gin_channels) hidden_channels,
5,
1,
n_flow_layer,
gin_channels=gin_channels,
)
self.sdp = StochasticDurationPredictor(
hidden_channels, 192, 3, 0.5, 4, gin_channels=gin_channels
)
self.dp = DurationPredictor(
hidden_channels, 256, 3, 0.5, gin_channels=gin_channels
)
if n_speakers > 1: if n_speakers > 1:
self.emb_g = nn.Embedding(n_speakers, gin_channels) self.emb_g = nn.Embedding(n_speakers, gin_channels)
@@ -646,25 +870,42 @@ class SynthesizerTrn(nn.Module):
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
else: else:
g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1) g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert, ja_bert, g=g) x, m_p, logs_p, x_mask = self.enc_p(
x, x_lengths, tone, language, bert, ja_bert, g=g
)
z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g) z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
z_p = self.flow(z, y_mask, g=g) z_p = self.flow(z, y_mask, g=g)
with torch.no_grad(): with torch.no_grad():
# negative cross-entropy # negative cross-entropy
s_p_sq_r = torch.exp(-2 * logs_p) # [b, d, t] s_p_sq_r = torch.exp(-2 * logs_p) # [b, d, t]
neg_cent1 = torch.sum(-0.5 * math.log(2 * math.pi) - logs_p, [1], keepdim=True) # [b, 1, t_s] neg_cent1 = torch.sum(
neg_cent2 = torch.matmul(-0.5 * (z_p ** 2).transpose(1, 2), -0.5 * math.log(2 * math.pi) - logs_p, [1], keepdim=True
s_p_sq_r) # [b, t_t, d] x [b, d, t_s] = [b, t_t, t_s] ) # [b, 1, t_s]
neg_cent3 = torch.matmul(z_p.transpose(1, 2), (m_p * s_p_sq_r)) # [b, t_t, d] x [b, d, t_s] = [b, t_t, t_s] neg_cent2 = torch.matmul(
neg_cent4 = torch.sum(-0.5 * (m_p ** 2) * s_p_sq_r, [1], keepdim=True) # [b, 1, t_s] -0.5 * (z_p**2).transpose(1, 2), s_p_sq_r
) # [b, t_t, d] x [b, d, t_s] = [b, t_t, t_s]
neg_cent3 = torch.matmul(
z_p.transpose(1, 2), (m_p * s_p_sq_r)
) # [b, t_t, d] x [b, d, t_s] = [b, t_t, t_s]
neg_cent4 = torch.sum(
-0.5 * (m_p**2) * s_p_sq_r, [1], keepdim=True
) # [b, 1, t_s]
neg_cent = neg_cent1 + neg_cent2 + neg_cent3 + neg_cent4 neg_cent = neg_cent1 + neg_cent2 + neg_cent3 + neg_cent4
if self.use_noise_scaled_mas: if self.use_noise_scaled_mas:
epsilon = torch.std(neg_cent) * torch.randn_like(neg_cent) * self.current_mas_noise_scale epsilon = (
torch.std(neg_cent)
* torch.randn_like(neg_cent)
* self.current_mas_noise_scale
)
neg_cent = neg_cent + epsilon neg_cent = neg_cent + epsilon
attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1) attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)
attn = monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1)).unsqueeze(1).detach() attn = (
monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1))
.unsqueeze(1)
.detach()
)
w = attn.sum(2) w = attn.sum(2)
@@ -673,7 +914,9 @@ class SynthesizerTrn(nn.Module):
logw_ = torch.log(w + 1e-6) * x_mask logw_ = torch.log(w + 1e-6) * x_mask
logw = self.dp(x, x_mask, g=g) logw = self.dp(x, x_mask, g=g)
l_length_dp = torch.sum((logw - logw_) ** 2, [1, 2]) / torch.sum(x_mask) # for averaging l_length_dp = torch.sum((logw - logw_) ** 2, [1, 2]) / torch.sum(
x_mask
) # for averaging
l_length = l_length_dp + l_length_sdp l_length = l_length_dp + l_length_sdp
@@ -681,29 +924,64 @@ class SynthesizerTrn(nn.Module):
m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2) m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2)
logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1, 2) logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1, 2)
z_slice, ids_slice = commons.rand_slice_segments(z, y_lengths, self.segment_size) z_slice, ids_slice = commons.rand_slice_segments(
z, y_lengths, self.segment_size
)
o = self.dec(z_slice, g=g) o = self.dec(z_slice, g=g)
return o, l_length, attn, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q), (x, logw, logw_) return (
o,
l_length,
attn,
ids_slice,
x_mask,
y_mask,
(z, z_p, m_p, logs_p, m_q, logs_q),
(x, logw, logw_),
)
def infer(self, x, x_lengths, sid, tone, language, bert, ja_bert, noise_scale=.667, length_scale=1, noise_scale_w=0.8, max_len=None, sdp_ratio=0,y=None): def infer(
self,
x,
x_lengths,
sid,
tone,
language,
bert,
ja_bert,
noise_scale=0.667,
length_scale=1,
noise_scale_w=0.8,
max_len=None,
sdp_ratio=0,
y=None,
):
# x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert) # x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert)
# g = self.gst(y) # g = self.gst(y)
if self.n_speakers > 0: if self.n_speakers > 0:
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
else: else:
g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1) g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1)
x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert, ja_bert, g=g) x, m_p, logs_p, x_mask = self.enc_p(
logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * (sdp_ratio) + self.dp(x, x_mask, g=g) * (1 - sdp_ratio) x, x_lengths, tone, language, bert, ja_bert, g=g
)
logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * (
sdp_ratio
) + self.dp(x, x_mask, g=g) * (1 - sdp_ratio)
w = torch.exp(logw) * x_mask * length_scale w = torch.exp(logw) * x_mask * length_scale
w_ceil = torch.ceil(w) w_ceil = torch.ceil(w)
y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), 1).long() y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), 1).long()
y_mask = torch.unsqueeze(commons.sequence_mask(y_lengths, None), 1).to(x_mask.dtype) y_mask = torch.unsqueeze(commons.sequence_mask(y_lengths, None), 1).to(
x_mask.dtype
)
attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1) attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)
attn = commons.generate_path(w_ceil, attn_mask) attn = commons.generate_path(w_ceil, attn_mask)
m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2) # [b, t', t], [b, t, d] -> [b, d, t'] m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(
logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1, 1, 2
2) # [b, t', t], [b, t, d] -> [b, d, t'] ) # [b, t', t], [b, t, d] -> [b, d, t']
logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(
1, 2
) # [b, t', t], [b, t, d] -> [b, d, t']
z_p = m_p + torch.randn_like(m_p) * torch.exp(logs_p) * noise_scale z_p = m_p + torch.randn_like(m_p) * torch.exp(logs_p) * noise_scale
z = self.flow(z_p, y_mask, g=g, reverse=True) z = self.flow(z_p, y_mask, g=g, reverse=True)