kwau/sovits-isla
0
1import torch2from torch import nn3from torch.nn import Conv1d, Conv2d4from torch.nn import functional as F5from torch.nn.utils import spectral_norm, weight_norm6 7import modules.attentions as attentions8import modules.commons as commons9import modules.modules as modules10import utils11from modules.commons import get_padding12from utils import f0_to_coarse13 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 share_parameter=False25 ):26 super().__init__()27 self.channels = channels28 self.hidden_channels = hidden_channels29 self.kernel_size = kernel_size30 self.dilation_rate = dilation_rate31 self.n_layers = n_layers32 self.n_flows = n_flows33 self.gin_channels = gin_channels34 35 self.flows = nn.ModuleList()36 37 self.wn = modules.WN(hidden_channels, kernel_size, dilation_rate, n_layers, p_dropout=0, gin_channels=gin_channels) if share_parameter else None38 39 for i in range(n_flows):40 self.flows.append(41 modules.ResidualCouplingLayer(channels, hidden_channels, kernel_size, dilation_rate, n_layers,42 gin_channels=gin_channels, mean_only=True, wn_sharing_parameter=self.wn))43 self.flows.append(modules.Flip())44 45 def forward(self, x, x_mask, g=None, reverse=False):46 if not reverse:47 for flow in self.flows:48 x, _ = flow(x, x_mask, g=g, reverse=reverse)49 else:50 for flow in reversed(self.flows):51 x = flow(x, x_mask, g=g, reverse=reverse)52 return x53 54 55class Encoder(nn.Module):56 def __init__(self,57 in_channels,58 out_channels,59 hidden_channels,60 kernel_size,61 dilation_rate,62 n_layers,63 gin_channels=0):64 super().__init__()65 self.in_channels = in_channels66 self.out_channels = out_channels67 self.hidden_channels = hidden_channels68 self.kernel_size = kernel_size69 self.dilation_rate = dilation_rate70 self.n_layers = n_layers71 self.gin_channels = gin_channels72 73 self.pre = nn.Conv1d(in_channels, hidden_channels, 1)74 self.enc = modules.WN(hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels)75 self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)76 77 def forward(self, x, x_lengths, g=None):78 # print(x.shape,x_lengths.shape)79 x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)80 x = self.pre(x) * x_mask81 x = self.enc(x, x_mask, g=g)82 stats = self.proj(x) * x_mask83 m, logs = torch.split(stats, self.out_channels, dim=1)84 z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask85 return z, m, logs, x_mask86 87 88class TextEncoder(nn.Module):89 def __init__(self,90 out_channels,91 hidden_channels,92 kernel_size,93 n_layers,94 gin_channels=0,95 filter_channels=None,96 n_heads=None,97 p_dropout=None):98 super().__init__()99 self.out_channels = out_channels100 self.hidden_channels = hidden_channels101 self.kernel_size = kernel_size102 self.n_layers = n_layers103 self.gin_channels = gin_channels104 self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)105 self.f0_emb = nn.Embedding(256, hidden_channels)106 107 self.enc_ = attentions.Encoder(108 hidden_channels,109 filter_channels,110 n_heads,111 n_layers,112 kernel_size,113 p_dropout)114 115 def forward(self, x, x_mask, f0=None, noice_scale=1):116 x = x + self.f0_emb(f0).transpose(1, 2)117 x = self.enc_(x * x_mask, x_mask)118 stats = self.proj(x) * x_mask119 m, logs = torch.split(stats, self.out_channels, dim=1)120 z = (m + torch.randn_like(m) * torch.exp(logs) * noice_scale) * x_mask121 122 return z, m, logs, x_mask123 124 125class DiscriminatorP(torch.nn.Module):126 def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):127 super(DiscriminatorP, self).__init__()128 self.period = period129 self.use_spectral_norm = use_spectral_norm130 norm_f = weight_norm if use_spectral_norm is False else spectral_norm131 self.convs = nn.ModuleList([132 norm_f(Conv2d(1, 32, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),133 norm_f(Conv2d(32, 128, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),134 norm_f(Conv2d(128, 512, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),135 norm_f(Conv2d(512, 1024, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),136 norm_f(Conv2d(1024, 1024, (kernel_size, 1), 1, padding=(get_padding(kernel_size, 1), 0))),137 ])138 self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))139 140 def forward(self, x):141 fmap = []142 143 # 1d to 2d144 b, c, t = x.shape145 if t % self.period != 0: # pad first146 n_pad = self.period - (t % self.period)147 x = F.pad(x, (0, n_pad), "reflect")148 t = t + n_pad149 x = x.view(b, c, t // self.period, self.period)150 151 for l in self.convs:152 x = l(x)153 x = F.leaky_relu(x, modules.LRELU_SLOPE)154 fmap.append(x)155 x = self.conv_post(x)156 fmap.append(x)157 x = torch.flatten(x, 1, -1)158 159 return x, fmap160 161 162class DiscriminatorS(torch.nn.Module):163 def __init__(self, use_spectral_norm=False):164 super(DiscriminatorS, self).__init__()165 norm_f = weight_norm if use_spectral_norm is False else spectral_norm166 self.convs = nn.ModuleList([167 norm_f(Conv1d(1, 16, 15, 1, padding=7)),168 norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),169 norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),170 norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),171 norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),172 norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),173 ])174 self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))175 176 def forward(self, x):177 fmap = []178 179 for l in self.convs:180 x = l(x)181 x = F.leaky_relu(x, modules.LRELU_SLOPE)182 fmap.append(x)183 x = self.conv_post(x)184 fmap.append(x)185 x = torch.flatten(x, 1, -1)186 187 return x, fmap188 189 190class MultiPeriodDiscriminator(torch.nn.Module):191 def __init__(self, use_spectral_norm=False):192 super(MultiPeriodDiscriminator, self).__init__()193 periods = [2, 3, 5, 7, 11]194 195 discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]196 discs = discs + [DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods]197 self.discriminators = nn.ModuleList(discs)198 199 def forward(self, y, y_hat):200 y_d_rs = []201 y_d_gs = []202 fmap_rs = []203 fmap_gs = []204 for i, d in enumerate(self.discriminators):205 y_d_r, fmap_r = d(y)206 y_d_g, fmap_g = d(y_hat)207 y_d_rs.append(y_d_r)208 y_d_gs.append(y_d_g)209 fmap_rs.append(fmap_r)210 fmap_gs.append(fmap_g)211 212 return y_d_rs, y_d_gs, fmap_rs, fmap_gs213 214 215class SpeakerEncoder(torch.nn.Module):216 def __init__(self, mel_n_channels=80, model_num_layers=3, model_hidden_size=256, model_embedding_size=256):217 super(SpeakerEncoder, self).__init__()218 self.lstm = nn.LSTM(mel_n_channels, model_hidden_size, model_num_layers, batch_first=True)219 self.linear = nn.Linear(model_hidden_size, model_embedding_size)220 self.relu = nn.ReLU()221 222 def forward(self, mels):223 self.lstm.flatten_parameters()224 _, (hidden, _) = self.lstm(mels)225 embeds_raw = self.relu(self.linear(hidden[-1]))226 return embeds_raw / torch.norm(embeds_raw, dim=1, keepdim=True)227 228 def compute_partial_slices(self, total_frames, partial_frames, partial_hop):229 mel_slices = []230 for i in range(0, total_frames - partial_frames, partial_hop):231 mel_range = torch.arange(i, i + partial_frames)232 mel_slices.append(mel_range)233 234 return mel_slices235 236 def embed_utterance(self, mel, partial_frames=128, partial_hop=64):237 mel_len = mel.size(1)238 last_mel = mel[:, -partial_frames:]239 240 if mel_len > partial_frames:241 mel_slices = self.compute_partial_slices(mel_len, partial_frames, partial_hop)242 mels = list(mel[:, s] for s in mel_slices)243 mels.append(last_mel)244 mels = torch.stack(tuple(mels), 0).squeeze(1)245 246 with torch.no_grad():247 partial_embeds = self(mels)248 embed = torch.mean(partial_embeds, axis=0).unsqueeze(0)249 # embed = embed / torch.linalg.norm(embed, 2)250 else:251 with torch.no_grad():252 embed = self(last_mel)253 254 return embed255 256class F0Decoder(nn.Module):257 def __init__(self,258 out_channels,259 hidden_channels,260 filter_channels,261 n_heads,262 n_layers,263 kernel_size,264 p_dropout,265 spk_channels=0):266 super().__init__()267 self.out_channels = out_channels268 self.hidden_channels = hidden_channels269 self.filter_channels = filter_channels270 self.n_heads = n_heads271 self.n_layers = n_layers272 self.kernel_size = kernel_size273 self.p_dropout = p_dropout274 self.spk_channels = spk_channels275 276 self.prenet = nn.Conv1d(hidden_channels, hidden_channels, 3, padding=1)277 self.decoder = attentions.FFT(278 hidden_channels,279 filter_channels,280 n_heads,281 n_layers,282 kernel_size,283 p_dropout)284 self.proj = nn.Conv1d(hidden_channels, out_channels, 1)285 self.f0_prenet = nn.Conv1d(1, hidden_channels, 3, padding=1)286 self.cond = nn.Conv1d(spk_channels, hidden_channels, 1)287 288 def forward(self, x, norm_f0, x_mask, spk_emb=None):289 x = torch.detach(x)290 if (spk_emb is not None):291 x = x + self.cond(spk_emb)292 x += self.f0_prenet(norm_f0)293 x = self.prenet(x) * x_mask294 x = self.decoder(x * x_mask, x_mask)295 x = self.proj(x) * x_mask296 return x297 298 299class SynthesizerTrn(nn.Module):300 """301 Synthesizer for Training302 """303 304 def __init__(self,305 spec_channels,306 segment_size,307 inter_channels,308 hidden_channels,309 filter_channels,310 n_heads,311 n_layers,312 kernel_size,313 p_dropout,314 resblock,315 resblock_kernel_sizes,316 resblock_dilation_sizes,317 upsample_rates,318 upsample_initial_channel,319 upsample_kernel_sizes,320 gin_channels,321 ssl_dim,322 n_speakers,323 sampling_rate=44100,324 vol_embedding=False,325 vocoder_name = "nsf-hifigan",326 use_depthwise_conv = False,327 use_automatic_f0_prediction = True,328 flow_share_parameter = False,329 n_flow_layer = 4,330 **kwargs):331 332 super().__init__()333 self.spec_channels = spec_channels334 self.inter_channels = inter_channels335 self.hidden_channels = hidden_channels336 self.filter_channels = filter_channels337 self.n_heads = n_heads338 self.n_layers = n_layers339 self.kernel_size = kernel_size340 self.p_dropout = p_dropout341 self.resblock = resblock342 self.resblock_kernel_sizes = resblock_kernel_sizes343 self.resblock_dilation_sizes = resblock_dilation_sizes344 self.upsample_rates = upsample_rates345 self.upsample_initial_channel = upsample_initial_channel346 self.upsample_kernel_sizes = upsample_kernel_sizes347 self.segment_size = segment_size348 self.gin_channels = gin_channels349 self.ssl_dim = ssl_dim350 self.vol_embedding = vol_embedding351 self.emb_g = nn.Embedding(n_speakers, gin_channels)352 self.use_depthwise_conv = use_depthwise_conv353 self.use_automatic_f0_prediction = use_automatic_f0_prediction354 if vol_embedding:355 self.emb_vol = nn.Linear(1, hidden_channels)356 357 self.pre = nn.Conv1d(ssl_dim, hidden_channels, kernel_size=5, padding=2)358 359 self.enc_p = TextEncoder(360 inter_channels,361 hidden_channels,362 filter_channels=filter_channels,363 n_heads=n_heads,364 n_layers=n_layers,365 kernel_size=kernel_size,366 p_dropout=p_dropout367 )368 hps = {369 "sampling_rate": sampling_rate,370 "inter_channels": inter_channels,371 "resblock": resblock,372 "resblock_kernel_sizes": resblock_kernel_sizes,373 "resblock_dilation_sizes": resblock_dilation_sizes,374 "upsample_rates": upsample_rates,375 "upsample_initial_channel": upsample_initial_channel,376 "upsample_kernel_sizes": upsample_kernel_sizes,377 "gin_channels": gin_channels,378 "use_depthwise_conv":use_depthwise_conv379 }380 381 modules.set_Conv1dModel(self.use_depthwise_conv)382 383 if vocoder_name == "nsf-hifigan":384 from vdecoder.hifigan.models import Generator385 self.dec = Generator(h=hps)386 elif vocoder_name == "nsf-snake-hifigan":387 from vdecoder.hifiganwithsnake.models import Generator388 self.dec = Generator(h=hps)389 else:390 print("[?] Unkown vocoder: use default(nsf-hifigan)")391 from vdecoder.hifigan.models import Generator392 self.dec = Generator(h=hps)393 394 self.enc_q = Encoder(spec_channels, inter_channels, hidden_channels, 5, 1, 16, gin_channels=gin_channels)395 self.flow = ResidualCouplingBlock(inter_channels, hidden_channels, 5, 1, n_flow_layer, gin_channels=gin_channels, share_parameter= flow_share_parameter)396 if self.use_automatic_f0_prediction:397 self.f0_decoder = F0Decoder(398 1,399 hidden_channels,400 filter_channels,401 n_heads,402 n_layers,403 kernel_size,404 p_dropout,405 spk_channels=gin_channels406 )407 self.emb_uv = nn.Embedding(2, hidden_channels)408 self.character_mix = False409 410 def EnableCharacterMix(self, n_speakers_map, device):411 self.speaker_map = torch.zeros((n_speakers_map, 1, 1, self.gin_channels)).to(device)412 for i in range(n_speakers_map):413 self.speaker_map[i] = self.emb_g(torch.LongTensor([[i]]).to(device))414 self.speaker_map = self.speaker_map.unsqueeze(0).to(device)415 self.character_mix = True416 417 def forward(self, c, f0, uv, spec, g=None, c_lengths=None, spec_lengths=None, vol = None):418 g = self.emb_g(g).transpose(1,2)419 420 # vol proj421 vol = self.emb_vol(vol[:,:,None]).transpose(1,2) if vol is not None and self.vol_embedding else 0422 423 # ssl prenet424 x_mask = torch.unsqueeze(commons.sequence_mask(c_lengths, c.size(2)), 1).to(c.dtype)425 x = self.pre(c) * x_mask + self.emb_uv(uv.long()).transpose(1,2) + vol426 427 # f0 predict428 if self.use_automatic_f0_prediction:429 lf0 = 2595. * torch.log10(1. + f0.unsqueeze(1) / 700.) / 500430 norm_lf0 = utils.normalize_f0(lf0, x_mask, uv)431 pred_lf0 = self.f0_decoder(x, norm_lf0, x_mask, spk_emb=g)432 else:433 lf0 = 0434 norm_lf0 = 0435 pred_lf0 = 0436 # encoder437 z_ptemp, m_p, logs_p, _ = self.enc_p(x, x_mask, f0=f0_to_coarse(f0))438 z, m_q, logs_q, spec_mask = self.enc_q(spec, spec_lengths, g=g)439 440 # flow441 z_p = self.flow(z, spec_mask, g=g)442 z_slice, pitch_slice, ids_slice = commons.rand_slice_segments_with_pitch(z, f0, spec_lengths, self.segment_size)443 444 # nsf decoder445 o = self.dec(z_slice, g=g, f0=pitch_slice)446 447 return o, ids_slice, spec_mask, (z, z_p, m_p, logs_p, m_q, logs_q), pred_lf0, norm_lf0, lf0448 449 @torch.no_grad()450 def infer(self, c, f0, uv, g=None, noice_scale=0.35, seed=52468, predict_f0=False, vol = None):451 452 if c.device == torch.device("cuda"):453 torch.cuda.manual_seed_all(seed)454 else:455 torch.manual_seed(seed)456 457 c_lengths = (torch.ones(c.size(0)) * c.size(-1)).to(c.device)458 459 if self.character_mix and len(g) > 1: # [N, S] * [S, B, 1, H]460 g = g.reshape((g.shape[0], g.shape[1], 1, 1, 1)) # [N, S, B, 1, 1]461 g = g * self.speaker_map # [N, S, B, 1, H]462 g = torch.sum(g, dim=1) # [N, 1, B, 1, H]463 g = g.transpose(0, -1).transpose(0, -2).squeeze(0) # [B, H, N]464 else:465 if g.dim() == 1:466 g = g.unsqueeze(0)467 g = self.emb_g(g).transpose(1, 2)468 469 x_mask = torch.unsqueeze(commons.sequence_mask(c_lengths, c.size(2)), 1).to(c.dtype)470 # vol proj471 472 vol = self.emb_vol(vol[:,:,None]).transpose(1,2) if vol is not None and self.vol_embedding else 0473 474 x = self.pre(c) * x_mask + self.emb_uv(uv.long()).transpose(1, 2) + vol475 476 477 if self.use_automatic_f0_prediction and predict_f0:478 lf0 = 2595. * torch.log10(1. + f0.unsqueeze(1) / 700.) / 500479 norm_lf0 = utils.normalize_f0(lf0, x_mask, uv, random_scale=False)480 pred_lf0 = self.f0_decoder(x, norm_lf0, x_mask, spk_emb=g)481 f0 = (700 * (torch.pow(10, pred_lf0 * 500 / 2595) - 1)).squeeze(1)482 483 z_p, m_p, logs_p, c_mask = self.enc_p(x, x_mask, f0=f0_to_coarse(f0), noice_scale=noice_scale)484 z = self.flow(z_p, c_mask, g=g, reverse=True)485 o = self.dec(z * c_mask, g=g, f0=f0)486 return o,f0487 488 