diff --git a/models.py b/models.py index 2364d62..9a4e1e7 100644 --- a/models.py +++ b/models.py @@ -267,7 +267,8 @@ class TextEncoder(nn.Module): self.language_emb = nn.Embedding(num_languages, hidden_channels) nn.init.normal_(self.language_emb.weight, 0.0, hidden_channels ** -0.5) self.bert_proj = nn.Conv1d(1024, hidden_channels, 1) - + self.ja_bert_proj = nn.Conv1d(768, hidden_channels, 1) + self.encoder = attentions.Encoder( hidden_channels, filter_channels, @@ -278,8 +279,11 @@ class TextEncoder(nn.Module): gin_channels=self.gin_channels) self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) - def forward(self, x, x_lengths, tone, language, bert, g=None): - x = (self.emb(x)+ self.tone_emb(tone)+ self.language_emb(language)+self.bert_proj(bert).transpose(1,2)) * math.sqrt(self.hidden_channels) # [b, t, h] + def forward(self, x, x_lengths, tone, language, bert, ja_bert, g=None): + bert_emb = self.zh_bert_proj(bert).transpose(1,2) + ja_bert_emb = = self.ja_bert_proj(ja_bert).transpose(1,2) + x = (self.emb(x)+ self.tone_emb(tone)+ self.language_emb(language) + + bert_emb +ja_bert_emb) * math.sqrt(self.hidden_channels) # [b, t, h] x = torch.transpose(x, 1, -1) # [b, h, t] x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype) @@ -637,12 +641,12 @@ class SynthesizerTrn(nn.Module): else: self.ref_enc = ReferenceEncoder(spec_channels, gin_channels) - def forward(self, x, x_lengths, y, y_lengths, sid, tone, language, bert): + def forward(self, x, x_lengths, y, y_lengths, sid, tone, language, bert, ja_bert): 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,g=g) + x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert, ja_bert, 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) @@ -681,14 +685,14 @@ class SynthesizerTrn(nn.Module): o = self.dec(z_slice, g=g) return o, l_length, attn, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q), (x, logw, logw_) - def infer(self, x, x_lengths, sid, tone, language, bert, noise_scale=.667, length_scale=1, noise_scale_w=0.8, max_len=None, sdp_ratio=0,y=None): + def infer(self, x, x_lengths, sid, tone, language, bert, ja_bert, noise_scale=.667, length_scale=1, noise_scale_w=0.8, max_len=None, sdp_ratio=0,y=None): #x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert) # g = self.gst(y) 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,g=g) + x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, tone, language, bert, ja_bert, g=g) logw = self.sdp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w) * (sdp_ratio) + self.dp(x, x_mask, g=g) * (1 - sdp_ratio) w = torch.exp(logw) * x_mask * length_scale w_ceil = torch.ceil(w)