diff --git a/models.py b/models.py index d53b430..5e98754 100644 --- a/models.py +++ b/models.py @@ -10,6 +10,7 @@ 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 @@ -321,6 +322,7 @@ class TextEncoder(nn.Module): n_layers, kernel_size, p_dropout, + n_speakers, gin_channels=0, ): super().__init__() @@ -342,6 +344,18 @@ 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, @@ -354,10 +368,33 @@ 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, g=None): + def forward( + self, x, x_lengths, tone, language, bert, ja_bert, en_bert, emo, 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) @@ -365,6 +402,7 @@ 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] @@ -377,7 +415,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 + return x, m, logs, x_mask, emo_commit_loss class ResidualCouplingBlock(nn.Module): @@ -810,6 +848,7 @@ class SynthesizerTrn(nn.Module): n_layers, kernel_size, p_dropout, + self.n_speakers, gin_channels=self.enc_gin_channels, ) self.dec = Generator( @@ -860,7 +899,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) @@ -877,13 +916,14 @@ class SynthesizerTrn(nn.Module): bert, ja_bert, en_bert, + emo=None, ): if self.n_speakers > 0: 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 = self.enc_p( - x, x_lengths, tone, language, bert, ja_bert, en_bert, g=g + 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 ) z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g) z_p = self.flow(z, y_mask, g=g) @@ -949,6 +989,7 @@ class SynthesizerTrn(nn.Module): y_mask, (z, z_p, m_p, logs_p, m_q, logs_q), (x, logw, logw_), + loss_commit, ) def infer( @@ -961,6 +1002,7 @@ class SynthesizerTrn(nn.Module): bert, ja_bert, en_bert, + emo=None, noise_scale=0.667, length_scale=1, noise_scale_w=0.8, @@ -974,8 +1016,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 = self.enc_p( - x, x_lengths, tone, language, bert, ja_bert, en_bert, g=g + x, m_p, logs_p, x_mask, _ = self.enc_p( + x, x_lengths, tone, language, bert, ja_bert, en_bert, emo, sid, g=g ) logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * ( sdp_ratio