CoolFace
Apppublic

kwau/sovits-isla

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
models.py488 linesDownload Raw Back to root
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