ok
This commit is contained in:
@@ -8,6 +8,7 @@ from torch.nn import functional as F
|
||||
from torch.nn.utils import remove_weight_norm, spectral_norm, weight_norm
|
||||
|
||||
from style_bert_vits2.models import attentions, commons, modules, monotonic_alignment
|
||||
from style_bert_vits2.models.matcha_flow import MatchaFlow
|
||||
from style_bert_vits2.nlp.symbols import NUM_LANGUAGES, NUM_TONES, SYMBOLS
|
||||
|
||||
|
||||
@@ -938,6 +939,8 @@ class SynthesizerTrn(nn.Module):
|
||||
"use_spk_conditioned_encoder", True
|
||||
)
|
||||
self.use_sdp = use_sdp
|
||||
self.use_matcha = kwargs.get("use_matcha", False)
|
||||
self.matcha_n_timesteps = kwargs.get("matcha_n_timesteps", 8)
|
||||
self.use_noise_scaled_mas = kwargs.get("use_noise_scaled_mas", False)
|
||||
self.mas_noise_scale_initial = kwargs.get("mas_noise_scale_initial", 0.01)
|
||||
self.noise_scale_delta = kwargs.get("noise_scale_delta", 2e-6)
|
||||
@@ -1007,6 +1010,15 @@ class SynthesizerTrn(nn.Module):
|
||||
self.emb_g = nn.Embedding(n_speakers, gin_channels)
|
||||
else:
|
||||
self.ref_enc = ReferenceEncoder(spec_channels, gin_channels)
|
||||
if self.use_matcha:
|
||||
self.matcha = MatchaFlow(
|
||||
latent_channels=inter_channels,
|
||||
speaker_channels=gin_channels,
|
||||
channels=kwargs.get("matcha_channels", hidden_channels),
|
||||
num_heads=kwargs.get("matcha_num_heads", n_heads),
|
||||
dropout=kwargs.get("matcha_dropout", 0.05),
|
||||
sigma_min=kwargs.get("matcha_sigma_min", 1e-4),
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -1029,6 +1041,7 @@ class SynthesizerTrn(nn.Module):
|
||||
torch.Tensor,
|
||||
tuple[torch.Tensor, ...],
|
||||
tuple[torch.Tensor, ...],
|
||||
torch.Tensor,
|
||||
]:
|
||||
if self.n_speakers > 0:
|
||||
g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]
|
||||
@@ -1093,10 +1106,20 @@ class SynthesizerTrn(nn.Module):
|
||||
z_slice, ids_slice = commons.rand_slice_segments(
|
||||
z, y_lengths, self.segment_size
|
||||
)
|
||||
if self.use_matcha:
|
||||
matcha_target = commons.slice_segments(z_p, ids_slice, self.segment_size)
|
||||
matcha_mu = commons.slice_segments(m_p, ids_slice, self.segment_size)
|
||||
matcha_mask = commons.slice_segments(y_mask, ids_slice, self.segment_size)
|
||||
matcha_loss = self.matcha.compute_loss(
|
||||
matcha_target.detach(), matcha_mask, matcha_mu, speaker=g
|
||||
)
|
||||
else:
|
||||
matcha_loss = z.new_zeros(())
|
||||
o = self.dec(z_slice, g=g)
|
||||
return (
|
||||
o,
|
||||
l_length,
|
||||
matcha_loss,
|
||||
attn,
|
||||
ids_slice,
|
||||
x_mask,
|
||||
@@ -1151,7 +1174,16 @@ class SynthesizerTrn(nn.Module):
|
||||
1, 2
|
||||
) # [b, t', t], [b, t, d] -> [b, d, t']
|
||||
|
||||
z_p = m_p + torch.randn_like(m_p) * torch.exp(logs_p) * noise_scale
|
||||
if self.use_matcha:
|
||||
z_p = self.matcha(
|
||||
m_p,
|
||||
y_mask,
|
||||
n_timesteps=self.matcha_n_timesteps,
|
||||
temperature=noise_scale,
|
||||
speaker=g,
|
||||
)
|
||||
else:
|
||||
z_p = m_p + torch.randn_like(m_p) * torch.exp(logs_p) * noise_scale
|
||||
z = self.flow(z_p, y_mask, g=g, reverse=True)
|
||||
o = self.dec((z * y_mask)[:, :, :max_len], g=g)
|
||||
return o, attn, y_mask, (z, z_p, m_p, logs_p)
|
||||
|
||||
Reference in New Issue
Block a user