Use only one vq is enough

This commit is contained in:
Sihan Wang
2023-12-07 22:07:51 +08:00
parent f550f86caa
commit 7ddc560617

View File

@@ -345,10 +345,7 @@ class TextEncoder(nn.Module):
self.ja_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.en_bert_proj = nn.Conv1d(1024, hidden_channels, 1)
self.emo_proj = nn.Linear(1024, 1024) self.emo_proj = nn.Linear(1024, 1024)
self.emo_quantizer = nn.ModuleList() self.emo_quantizer = VectorQuantize(
for i in range(0, n_speakers):
self.emo_quantizer.append(
VectorQuantize(
dim=1024, dim=1024,
codebook_size=10, codebook_size=10,
decay=0.8, decay=0.8,
@@ -356,7 +353,6 @@ class TextEncoder(nn.Module):
learnable_codebook=True, learnable_codebook=True,
ema_update=False, ema_update=False,
) )
)
self.emo_q_proj = nn.Linear(1024, hidden_channels) self.emo_q_proj = nn.Linear(1024, hidden_channels)
self.encoder = attentions.Encoder( self.encoder = attentions.Encoder(
@@ -373,17 +369,16 @@ class TextEncoder(nn.Module):
def forward( 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, emo, sid, g=None
): ):
sid = sid.cpu()
bert_emb = self.bert_proj(bert).transpose(1, 2) bert_emb = self.bert_proj(bert).transpose(1, 2)
ja_bert_emb = self.ja_bert_proj(ja_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) en_bert_emb = self.en_bert_proj(en_bert).transpose(1, 2)
if emo.size(-1) == 1024: if emo.size(-1) == 1024:
emo_emb = self.emo_proj(emo.unsqueeze(1)) emo_emb = self.emo_proj(emo.unsqueeze(1))
emo_commit_loss = torch.zeros(1) emo_commit_loss = torch.zeros(1).to(emo_emb.device)
emo_emb_ = [] emo_emb_ = []
for i in range(emo_emb.size(0)): for i in range(emo_emb.size(0)):
temp_emo_emb, _, temp_emo_commit_loss = self.emo_quantizer[sid[i]]( temp_emo_emb, _, temp_emo_commit_loss = self.emo_quantizer(
emo_emb[i].unsqueeze(0).cpu() emo_emb[i].unsqueeze(0)
) )
emo_commit_loss += temp_emo_commit_loss emo_commit_loss += temp_emo_commit_loss
emo_emb_.append(temp_emo_emb) emo_emb_.append(temp_emo_emb)
@@ -391,8 +386,7 @@ class TextEncoder(nn.Module):
emo_commit_loss = emo_commit_loss.to(emo_emb.device) emo_commit_loss = emo_commit_loss.to(emo_emb.device)
else: else:
emo_emb = ( emo_emb = (
self.emo_quantizer[sid[0]] self.emo_quantizer.get_output_from_indices(emo.to(torch.int))
.get_output_from_indices(emo.to(torch.int).cpu())
.unsqueeze(0) .unsqueeze(0)
.to(emo.device) .to(emo.device)
) )