From 242428f9ab1272dde16187fb9e5979d3ece8aabe Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= Date: Wed, 8 Nov 2023 20:24:37 +0800 Subject: [PATCH] del emo --- models.py | 52 ++++++---------------------------------------------- 1 file changed, 6 insertions(+), 46 deletions(-) diff --git a/models.py b/models.py index d493878..af3b4e2 100644 --- a/models.py +++ b/models.py @@ -10,8 +10,6 @@ import monotonic_align from torch.nn import Conv1d, ConvTranspose1d, Conv2d from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm -from vector_quantize_pytorch import VectorQuantize - from commons import init_weights, get_padding from text import symbols, num_tones, num_languages @@ -322,7 +320,6 @@ class TextEncoder(nn.Module): n_layers, kernel_size, p_dropout, - n_speakers, gin_channels=0, ): super().__init__() @@ -344,18 +341,6 @@ class TextEncoder(nn.Module): self.bert_proj = nn.Conv1d(1024, hidden_channels, 1) self.ja_bert_proj = nn.Conv1d(1024, hidden_channels, 1) self.en_bert_proj = nn.Conv1d(1024, hidden_channels, 1) - self.emo_proj = nn.Linear(1024, 1024) - self.emo_quantizer = [ - VectorQuantize( - dim=1024, - codebook_size=5, - decay=0.8, - commitment_weight=1.0, - learnable_codebook=True, - ema_update=False, - ) - ] * n_speakers - self.emo_q_proj = nn.Linear(1024, hidden_channels) self.encoder = attentions.Encoder( hidden_channels, @@ -369,32 +354,11 @@ class TextEncoder(nn.Module): self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) def forward( - self, x, x_lengths, tone, language, bert, ja_bert, en_bert, emo, sid, g=None + self, x, x_lengths, tone, language, bert, ja_bert, en_bert, sid, g=None ): - sid = sid.cpu() bert_emb = self.bert_proj(bert).transpose(1, 2) ja_bert_emb = self.ja_bert_proj(ja_bert).transpose(1, 2) en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2) - if emo.size(-1) == 1024: - emo_emb = self.emo_proj(emo.unsqueeze(1)) - emo_commit_loss = torch.zeros(1) - emo_emb_ = [] - for i in range(emo_emb.size(0)): - temp_emo_emb, _, temp_emo_commit_loss = self.emo_quantizer[sid[i]]( - emo_emb[i].unsqueeze(0).cpu() - ) - emo_commit_loss += temp_emo_commit_loss - emo_emb_.append(temp_emo_emb) - emo_emb = torch.cat(emo_emb_, dim=0).to(emo_emb.device) - emo_commit_loss = emo_commit_loss.to(emo_emb.device) - else: - emo_emb = ( - self.emo_quantizer[sid[0]] - .get_output_from_indices(emo.to(torch.int).cpu()) - .unsqueeze(0) - .to(emo.device) - ) - emo_commit_loss = torch.zeros(1) x = ( self.emb(x) + self.tone_emb(tone) @@ -402,7 +366,6 @@ class TextEncoder(nn.Module): + bert_emb + ja_bert_emb + en_bert_emb - + self.emo_q_proj(emo_emb) ) * math.sqrt( self.hidden_channels ) # [b, t, h] @@ -415,7 +378,7 @@ class TextEncoder(nn.Module): stats = self.proj(x) * x_mask m, logs = torch.split(stats, self.out_channels, dim=1) - return x, m, logs, x_mask, emo_commit_loss + return x, m, logs, x_mask class ResidualCouplingBlock(nn.Module): @@ -848,7 +811,6 @@ class SynthesizerTrn(nn.Module): n_layers, kernel_size, p_dropout, - self.n_speakers, gin_channels=self.enc_gin_channels, ) self.dec = Generator( @@ -899,7 +861,7 @@ class SynthesizerTrn(nn.Module): 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) else: self.ref_enc = ReferenceEncoder(spec_channels, gin_channels) @@ -922,8 +884,8 @@ class SynthesizerTrn(nn.Module): g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1] else: g = self.ref_enc(y.transpose(1, 2)).unsqueeze(-1) - x, m_p, logs_p, x_mask, loss_commit = self.enc_p( - x, x_lengths, tone, language, bert, ja_bert, en_bert, emo, sid, g=g + x, m_p, logs_p, x_mask = self.enc_p( + x, x_lengths, tone, language, bert, ja_bert, en_bert, sid, 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) @@ -989,7 +951,6 @@ class SynthesizerTrn(nn.Module): y_mask, (z, z_p, m_p, logs_p, m_q, logs_q), (x, logw, logw_), - loss_commit, ) def infer( @@ -1002,7 +963,6 @@ class SynthesizerTrn(nn.Module): bert, ja_bert, en_bert, - emo=None, noise_scale=0.667, length_scale=1, noise_scale_w=0.8, @@ -1017,7 +977,7 @@ class SynthesizerTrn(nn.Module): else: 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, en_bert, emo, sid, g=g + x, x_lengths, tone, language, bert, ja_bert, en_bert, sid, g=g ) logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * ( sdp_ratio