add emo and vo

This commit is contained in:
Stardust·减
2023-11-07 12:51:18 +08:00
committed by GitHub
parent a752307e1e
commit 16af947b67

View File

@@ -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