kevinwang676/FreeVC-OpenAI-TTS
0
1import copy2import math3import torch4from torch import nn5from torch.nn import functional as F6 7import commons8import modules9 10from torch.nn import Conv1d, ConvTranspose1d, AvgPool1d, Conv2d11from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm12from commons import init_weights, get_padding13 14 15class ResidualCouplingBlock(nn.Module):16 def __init__(self,17 channels,18 hidden_channels,19 kernel_size,20 dilation_rate,21 n_layers,22 n_flows=4,23 gin_channels=0):24 super().__init__()25 self.channels = channels26 self.hidden_channels = hidden_channels27 self.kernel_size = kernel_size28 self.dilation_rate = dilation_rate29 self.n_layers = n_layers30 self.n_flows = n_flows31 self.gin_channels = gin_channels32 33 self.flows = nn.ModuleList()34 for i in range(n_flows):35 self.flows.append(modules.ResidualCouplingLayer(channels, hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels, mean_only=True))36 self.flows.append(modules.Flip())37 38 def forward(self, x, x_mask, g=None, reverse=False):39 if not reverse:40 for flow in self.flows:41 x, _ = flow(x, x_mask, g=g, reverse=reverse)42 else:43 for flow in reversed(self.flows):44 x = flow(x, x_mask, g=g, reverse=reverse)45 return x46 47 48class Encoder(nn.Module):49 def __init__(self,50 in_channels,51 out_channels,52 hidden_channels,53 kernel_size,54 dilation_rate,55 n_layers,56 gin_channels=0):57 super().__init__()58 self.in_channels = in_channels59 self.out_channels = out_channels60 self.hidden_channels = hidden_channels61 self.kernel_size = kernel_size62 self.dilation_rate = dilation_rate63 self.n_layers = n_layers64 self.gin_channels = gin_channels65 66 self.pre = nn.Conv1d(in_channels, hidden_channels, 1)67 self.enc = modules.WN(hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels)68 self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)69 70 def forward(self, x, x_lengths, g=None):71 x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)72 x = self.pre(x) * x_mask73 x = self.enc(x, x_mask, g=g)74 stats = self.proj(x) * x_mask75 m, logs = torch.split(stats, self.out_channels, dim=1)76 z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask77 return z, m, logs, x_mask78 79 80class Generator(torch.nn.Module):81 def __init__(self, initial_channel, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, upsample_initial_channel, upsample_kernel_sizes, gin_channels=0):82 super(Generator, self).__init__()83 self.num_kernels = len(resblock_kernel_sizes)84 self.num_upsamples = len(upsample_rates)85 self.conv_pre = Conv1d(initial_channel, upsample_initial_channel, 7, 1, padding=3)86 resblock = modules.ResBlock1 if resblock == '1' else modules.ResBlock287 88 self.ups = nn.ModuleList()89 for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):90 self.ups.append(weight_norm(91 ConvTranspose1d(upsample_initial_channel//(2**i), upsample_initial_channel//(2**(i+1)),92 k, u, padding=(k-u)//2)))93 94 self.resblocks = nn.ModuleList()95 for i in range(len(self.ups)):96 ch = upsample_initial_channel//(2**(i+1))97 for j, (k, d) in enumerate(zip(resblock_kernel_sizes, resblock_dilation_sizes)):98 self.resblocks.append(resblock(ch, k, d))99 100 self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)101 self.ups.apply(init_weights)102 103 if gin_channels != 0:104 self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)105 106 def forward(self, x, g=None):107 x = self.conv_pre(x)108 if g is not None:109 x = x + self.cond(g)110 111 for i in range(self.num_upsamples):112 x = F.leaky_relu(x, modules.LRELU_SLOPE)113 x = self.ups[i](x)114 xs = None115 for j in range(self.num_kernels):116 if xs is None:117 xs = self.resblocks[i*self.num_kernels+j](x)118 else:119 xs += self.resblocks[i*self.num_kernels+j](x)120 x = xs / self.num_kernels121 x = F.leaky_relu(x)122 x = self.conv_post(x)123 x = torch.tanh(x)124 125 return x126 127 def remove_weight_norm(self):128 print('Removing weight norm...')129 for l in self.ups:130 remove_weight_norm(l)131 for l in self.resblocks:132 l.remove_weight_norm()133 134 135class DiscriminatorP(torch.nn.Module):136 def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):137 super(DiscriminatorP, self).__init__()138 self.period = period139 self.use_spectral_norm = use_spectral_norm140 norm_f = weight_norm if use_spectral_norm == False else spectral_norm141 self.convs = nn.ModuleList([142 norm_f(Conv2d(1, 32, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),143 norm_f(Conv2d(32, 128, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),144 norm_f(Conv2d(128, 512, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),145 norm_f(Conv2d(512, 1024, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),146 norm_f(Conv2d(1024, 1024, (kernel_size, 1), 1, padding=(get_padding(kernel_size, 1), 0))),147 ])148 self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))149 150 def forward(self, x):151 fmap = []152 153 # 1d to 2d154 b, c, t = x.shape155 if t % self.period != 0: # pad first156 n_pad = self.period - (t % self.period)157 x = F.pad(x, (0, n_pad), "reflect")158 t = t + n_pad159 x = x.view(b, c, t // self.period, self.period)160 161 for l in self.convs:162 x = l(x)163 x = F.leaky_relu(x, modules.LRELU_SLOPE)164 fmap.append(x)165 x = self.conv_post(x)166 fmap.append(x)167 x = torch.flatten(x, 1, -1)168 169 return x, fmap170 171 172class DiscriminatorS(torch.nn.Module):173 def __init__(self, use_spectral_norm=False):174 super(DiscriminatorS, self).__init__()175 norm_f = weight_norm if use_spectral_norm == False else spectral_norm176 self.convs = nn.ModuleList([177 norm_f(Conv1d(1, 16, 15, 1, padding=7)),178 norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),179 norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),180 norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),181 norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),182 norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),183 ])184 self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))185 186 def forward(self, x):187 fmap = []188 189 for l in self.convs:190 x = l(x)191 x = F.leaky_relu(x, modules.LRELU_SLOPE)192 fmap.append(x)193 x = self.conv_post(x)194 fmap.append(x)195 x = torch.flatten(x, 1, -1)196 197 return x, fmap198 199 200class MultiPeriodDiscriminator(torch.nn.Module):201 def __init__(self, use_spectral_norm=False):202 super(MultiPeriodDiscriminator, self).__init__()203 periods = [2,3,5,7,11]204 205 discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]206 discs = discs + [DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods]207 self.discriminators = nn.ModuleList(discs)208 209 def forward(self, y, y_hat):210 y_d_rs = []211 y_d_gs = []212 fmap_rs = []213 fmap_gs = []214 for i, d in enumerate(self.discriminators):215 y_d_r, fmap_r = d(y)216 y_d_g, fmap_g = d(y_hat)217 y_d_rs.append(y_d_r)218 y_d_gs.append(y_d_g)219 fmap_rs.append(fmap_r)220 fmap_gs.append(fmap_g)221 222 return y_d_rs, y_d_gs, fmap_rs, fmap_gs223 224 225class SpeakerEncoder(torch.nn.Module):226 def __init__(self, mel_n_channels=80, model_num_layers=3, model_hidden_size=256, model_embedding_size=256):227 super(SpeakerEncoder, self).__init__()228 self.lstm = nn.LSTM(mel_n_channels, model_hidden_size, model_num_layers, batch_first=True)229 self.linear = nn.Linear(model_hidden_size, model_embedding_size)230 self.relu = nn.ReLU()231 232 def forward(self, mels):233 self.lstm.flatten_parameters()234 _, (hidden, _) = self.lstm(mels)235 embeds_raw = self.relu(self.linear(hidden[-1]))236 return embeds_raw / torch.norm(embeds_raw, dim=1, keepdim=True)237 238 def compute_partial_slices(self, total_frames, partial_frames, partial_hop):239 mel_slices = []240 for i in range(0, total_frames-partial_frames, partial_hop):241 mel_range = torch.arange(i, i+partial_frames)242 mel_slices.append(mel_range)243 244 return mel_slices245 246 def embed_utterance(self, mel, partial_frames=128, partial_hop=64):247 mel_len = mel.size(1)248 last_mel = mel[:,-partial_frames:]249 250 if mel_len > partial_frames:251 mel_slices = self.compute_partial_slices(mel_len, partial_frames, partial_hop)252 mels = list(mel[:,s] for s in mel_slices)253 mels.append(last_mel)254 mels = torch.stack(tuple(mels), 0).squeeze(1)255 256 with torch.no_grad():257 partial_embeds = self(mels)258 embed = torch.mean(partial_embeds, axis=0).unsqueeze(0)259 #embed = embed / torch.linalg.norm(embed, 2)260 else:261 with torch.no_grad():262 embed = self(last_mel)263 264 return embed265 266 267class SynthesizerTrn(nn.Module):268 """269 Synthesizer for Training270 """271 272 def __init__(self, 273 spec_channels,274 segment_size,275 inter_channels,276 hidden_channels,277 filter_channels,278 n_heads,279 n_layers,280 kernel_size,281 p_dropout,282 resblock, 283 resblock_kernel_sizes, 284 resblock_dilation_sizes, 285 upsample_rates, 286 upsample_initial_channel, 287 upsample_kernel_sizes,288 gin_channels,289 ssl_dim,290 use_spk,291 **kwargs):292 293 super().__init__()294 self.spec_channels = spec_channels295 self.inter_channels = inter_channels296 self.hidden_channels = hidden_channels297 self.filter_channels = filter_channels298 self.n_heads = n_heads299 self.n_layers = n_layers300 self.kernel_size = kernel_size301 self.p_dropout = p_dropout302 self.resblock = resblock303 self.resblock_kernel_sizes = resblock_kernel_sizes304 self.resblock_dilation_sizes = resblock_dilation_sizes305 self.upsample_rates = upsample_rates306 self.upsample_initial_channel = upsample_initial_channel307 self.upsample_kernel_sizes = upsample_kernel_sizes308 self.segment_size = segment_size309 self.gin_channels = gin_channels310 self.ssl_dim = ssl_dim311 self.use_spk = use_spk312 313 self.enc_p = Encoder(ssl_dim, inter_channels, hidden_channels, 5, 1, 16)314 self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)315 self.enc_q = Encoder(spec_channels, inter_channels, hidden_channels, 5, 1, 16, gin_channels=gin_channels) 316 self.flow = ResidualCouplingBlock(inter_channels, hidden_channels, 5, 1, 4, gin_channels=gin_channels)317 318 if not self.use_spk:319 self.enc_spk = SpeakerEncoder(model_hidden_size=gin_channels, model_embedding_size=gin_channels)320 321 def forward(self, c, spec, g=None, mel=None, c_lengths=None, spec_lengths=None):322 if c_lengths == None:323 c_lengths = (torch.ones(c.size(0)) * c.size(-1)).to(c.device)324 if spec_lengths == None:325 spec_lengths = (torch.ones(spec.size(0)) * spec.size(-1)).to(spec.device)326 327 if not self.use_spk:328 g = self.enc_spk(mel.transpose(1,2))329 g = g.unsqueeze(-1)330 331 _, m_p, logs_p, _ = self.enc_p(c, c_lengths)332 z, m_q, logs_q, spec_mask = self.enc_q(spec, spec_lengths, g=g) 333 z_p = self.flow(z, spec_mask, g=g)334 335 z_slice, ids_slice = commons.rand_slice_segments(z, spec_lengths, self.segment_size)336 o = self.dec(z_slice, g=g)337 338 return o, ids_slice, spec_mask, (z, z_p, m_p, logs_p, m_q, logs_q)339 340 def infer(self, c, g=None, mel=None, c_lengths=None):341 if c_lengths == None:342 c_lengths = (torch.ones(c.size(0)) * c.size(-1)).to(c.device)343 if not self.use_spk:344 g = self.enc_spk.embed_utterance(mel.transpose(1,2))345 g = g.unsqueeze(-1)346 347 z_p, m_p, logs_p, c_mask = self.enc_p(c, c_lengths)348 z = self.flow(z_p, c_mask, g=g, reverse=True)349 o = self.dec(z * c_mask, g=g)350 351 return o352 