split
This commit is contained in:
@@ -270,8 +270,7 @@ class MatchaFlow(nn.Module):
|
||||
error = (prediction - velocity).square() * mask
|
||||
return error.sum() / (mask.sum().clamp_min(1) * target.shape[1])
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
def sample(
|
||||
self,
|
||||
mu: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
@@ -289,3 +288,15 @@ class MatchaFlow(nn.Module):
|
||||
)
|
||||
x = x + dt * self.estimator(x, mask, mu, time, speaker=speaker)
|
||||
return x * mask
|
||||
|
||||
@torch.inference_mode()
|
||||
def forward(
|
||||
self,
|
||||
mu: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
n_timesteps: int,
|
||||
temperature: float = 1.0,
|
||||
speaker: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""Sample without building a graph (normal inference path)."""
|
||||
return self.sample(mu, mask, n_timesteps, temperature, speaker)
|
||||
|
||||
284
style_bert_vits2/models/mel_synthesizer.py
Normal file
284
style_bert_vits2/models/mel_synthesizer.py
Normal file
@@ -0,0 +1,284 @@
|
||||
"""Separated mel acoustic model and mel-to-wave vocoder composition.
|
||||
|
||||
Unlike :class:`SynthesizerTrn`, the acoustic model in this module ends at a
|
||||
log-mel spectrogram. The vocoder is therefore independently pretrainable and
|
||||
replaceable, while :class:`JointMelSynthesizer` keeps the boundary
|
||||
differentiable for end-to-end fine-tuning.
|
||||
"""
|
||||
|
||||
import math
|
||||
from typing import Any, Optional
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from style_bert_vits2.models import commons
|
||||
from style_bert_vits2.models.matcha_flow import MatchaFlow
|
||||
|
||||
|
||||
class MelAcousticModel(nn.Module):
|
||||
"""Text-to-mel model shared by the standard and JP-Extra front ends."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoder: nn.Module,
|
||||
stochastic_duration_predictor: nn.Module,
|
||||
duration_predictor: nn.Module,
|
||||
n_mel_channels: int,
|
||||
n_speakers: int,
|
||||
gin_channels: int,
|
||||
hidden_channels: int,
|
||||
matcha_channels: int,
|
||||
matcha_num_heads: int,
|
||||
matcha_dropout: float,
|
||||
matcha_sigma_min: float,
|
||||
matcha_n_timesteps: int,
|
||||
matcha_use_diff_attention: bool,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.enc_p = text_encoder
|
||||
self.sdp = stochastic_duration_predictor
|
||||
self.dp = duration_predictor
|
||||
self.n_mel_channels = n_mel_channels
|
||||
self.n_speakers = n_speakers
|
||||
self.gin_channels = gin_channels
|
||||
self.matcha_n_timesteps = matcha_n_timesteps
|
||||
if n_speakers < 1:
|
||||
raise ValueError("Separated mel models require at least one speaker")
|
||||
self.emb_g = nn.Embedding(n_speakers, gin_channels)
|
||||
self.matcha = MatchaFlow(
|
||||
latent_channels=n_mel_channels,
|
||||
speaker_channels=gin_channels,
|
||||
channels=matcha_channels or hidden_channels,
|
||||
num_heads=matcha_num_heads,
|
||||
dropout=matcha_dropout,
|
||||
sigma_min=matcha_sigma_min,
|
||||
use_diff_attention=matcha_use_diff_attention,
|
||||
)
|
||||
|
||||
def _encode(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
x_lengths: torch.Tensor,
|
||||
sid: torch.Tensor,
|
||||
encoder_args: tuple[torch.Tensor, ...],
|
||||
) -> tuple[torch.Tensor, ...]:
|
||||
g = self.emb_g(sid).unsqueeze(-1)
|
||||
hidden, mu, logs, x_mask = self.enc_p(
|
||||
x, x_lengths, *encoder_args, g=g
|
||||
)
|
||||
return hidden, mu, logs, x_mask, g
|
||||
|
||||
@staticmethod
|
||||
def _align(
|
||||
mel: torch.Tensor,
|
||||
mu: torch.Tensor,
|
||||
logs: torch.Tensor,
|
||||
x_mask: torch.Tensor,
|
||||
mel_mask: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Run MAS using the encoder's diagonal Gaussian mel prior."""
|
||||
from style_bert_vits2.models import monotonic_alignment
|
||||
|
||||
with torch.no_grad():
|
||||
inv_variance = torch.exp(-2 * logs)
|
||||
neg_cent1 = torch.sum(
|
||||
-0.5 * math.log(2 * math.pi) - logs, 1, keepdim=True
|
||||
)
|
||||
neg_cent2 = torch.matmul(
|
||||
-0.5 * (mel**2).transpose(1, 2), inv_variance
|
||||
)
|
||||
neg_cent3 = torch.matmul(
|
||||
mel.transpose(1, 2), mu * inv_variance
|
||||
)
|
||||
neg_cent4 = torch.sum(
|
||||
-0.5 * (mu**2) * inv_variance, 1, keepdim=True
|
||||
)
|
||||
score = neg_cent1 + neg_cent2 + neg_cent3 + neg_cent4
|
||||
attn_mask = x_mask.unsqueeze(2) * mel_mask.unsqueeze(-1)
|
||||
return (
|
||||
monotonic_alignment.maximum_path(score, attn_mask.squeeze(1))
|
||||
.unsqueeze(1)
|
||||
.detach()
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
x_lengths: torch.Tensor,
|
||||
mel: torch.Tensor,
|
||||
mel_lengths: torch.Tensor,
|
||||
sid: torch.Tensor,
|
||||
*encoder_args: torch.Tensor,
|
||||
out_size: Optional[int] = None,
|
||||
generate: bool = False,
|
||||
n_timesteps: Optional[int] = None,
|
||||
temperature: float = 1.0,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
"""Compute acoustic losses and optionally a differentiable mel sample."""
|
||||
hidden, mu, logs, x_mask, g = self._encode(
|
||||
x, x_lengths, sid, encoder_args
|
||||
)
|
||||
mel_mask = commons.sequence_mask(mel_lengths, mel.shape[-1]).unsqueeze(1)
|
||||
mel_mask = mel_mask.to(dtype=mu.dtype, device=mu.device)
|
||||
attn = self._align(mel, mu, logs, x_mask, mel_mask)
|
||||
durations = attn.sum(2)
|
||||
log_durations = torch.log(durations + 1e-6) * x_mask
|
||||
predicted_log_durations = self.dp(hidden, x_mask, g=g)
|
||||
duration_loss = torch.sum(
|
||||
(predicted_log_durations - log_durations) ** 2
|
||||
) / x_mask.sum().clamp_min(1)
|
||||
duration_loss = duration_loss + (
|
||||
self.sdp(hidden, x_mask, durations, g=g).sum()
|
||||
/ x_mask.sum().clamp_min(1)
|
||||
)
|
||||
|
||||
aligned_mu = torch.matmul(
|
||||
attn.squeeze(1), mu.transpose(1, 2)
|
||||
).transpose(1, 2)
|
||||
aligned_logs = torch.matmul(
|
||||
attn.squeeze(1), logs.transpose(1, 2)
|
||||
).transpose(1, 2)
|
||||
|
||||
ids_slice: Optional[torch.Tensor] = None
|
||||
if out_size is not None:
|
||||
mel, ids_slice = commons.rand_slice_segments(
|
||||
mel, mel_lengths, out_size
|
||||
)
|
||||
aligned_mu = commons.slice_segments(aligned_mu, ids_slice, out_size)
|
||||
aligned_logs = commons.slice_segments(
|
||||
aligned_logs, ids_slice, out_size
|
||||
)
|
||||
mel_mask = commons.slice_segments(mel_mask, ids_slice, out_size)
|
||||
|
||||
flow_loss = self.matcha.compute_loss(mel, mel_mask, aligned_mu, g)
|
||||
prior = (
|
||||
aligned_logs
|
||||
+ 0.5 * ((mel - aligned_mu) ** 2) * torch.exp(-2 * aligned_logs)
|
||||
+ 0.5 * math.log(2 * math.pi)
|
||||
)
|
||||
prior_loss = (prior * mel_mask).sum() / (
|
||||
mel_mask.sum().clamp_min(1) * self.n_mel_channels
|
||||
)
|
||||
result = {
|
||||
"duration_loss": duration_loss,
|
||||
"prior_loss": prior_loss,
|
||||
"flow_loss": flow_loss,
|
||||
"attn": attn,
|
||||
"mel_mask": mel_mask,
|
||||
"target_mel": mel,
|
||||
}
|
||||
if ids_slice is not None:
|
||||
result["ids_slice"] = ids_slice
|
||||
if generate:
|
||||
result["generated_mel"] = self.matcha.sample(
|
||||
aligned_mu,
|
||||
mel_mask,
|
||||
n_timesteps or self.matcha_n_timesteps,
|
||||
temperature,
|
||||
g,
|
||||
)
|
||||
return result
|
||||
|
||||
@torch.inference_mode()
|
||||
def infer_mel(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
x_lengths: torch.Tensor,
|
||||
sid: torch.Tensor,
|
||||
*encoder_args: torch.Tensor,
|
||||
noise_scale: float = 1.0,
|
||||
length_scale: float = 1.0,
|
||||
noise_scale_w: float = 0.8,
|
||||
sdp_ratio: float = 0.0,
|
||||
n_timesteps: Optional[int] = None,
|
||||
max_len: Optional[int] = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
hidden, mu, _, x_mask, g = self._encode(
|
||||
x, x_lengths, sid, encoder_args
|
||||
)
|
||||
logw = self.sdp(
|
||||
hidden, x_mask, g=g, reverse=True, noise_scale=noise_scale_w
|
||||
) * sdp_ratio + self.dp(hidden, x_mask, g=g) * (1 - sdp_ratio)
|
||||
durations = torch.ceil(torch.exp(logw) * x_mask * length_scale)
|
||||
mel_lengths = torch.clamp_min(durations.sum((1, 2)), 1).long()
|
||||
mel_mask = commons.sequence_mask(mel_lengths, None).unsqueeze(1).to(x_mask)
|
||||
attn_mask = x_mask.unsqueeze(2) * mel_mask.unsqueeze(-1)
|
||||
attn = commons.generate_path(durations, attn_mask)
|
||||
aligned_mu = torch.matmul(
|
||||
attn.squeeze(1), mu.transpose(1, 2)
|
||||
).transpose(1, 2)
|
||||
if max_len is not None:
|
||||
aligned_mu = aligned_mu[:, :, :max_len]
|
||||
mel_mask = mel_mask[:, :, :max_len]
|
||||
attn = attn[:, :, :max_len]
|
||||
generated_mel = self.matcha(
|
||||
aligned_mu,
|
||||
mel_mask,
|
||||
n_timesteps or self.matcha_n_timesteps,
|
||||
noise_scale,
|
||||
g,
|
||||
)
|
||||
return generated_mel, attn, mel_mask
|
||||
|
||||
|
||||
class JointMelSynthesizer(nn.Module):
|
||||
"""Differentiable composition of a pretrained acoustic model and vocoder."""
|
||||
|
||||
def __init__(self, acoustic_model: MelAcousticModel, vocoder: nn.Module) -> None:
|
||||
super().__init__()
|
||||
self.acoustic_model = acoustic_model
|
||||
self.vocoder = vocoder
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
x_lengths: torch.Tensor,
|
||||
mel: torch.Tensor,
|
||||
mel_lengths: torch.Tensor,
|
||||
sid: torch.Tensor,
|
||||
*encoder_args: torch.Tensor,
|
||||
out_size: int,
|
||||
n_timesteps: Optional[int] = None,
|
||||
temperature: float = 1.0,
|
||||
) -> tuple[torch.Tensor, dict[str, torch.Tensor]]:
|
||||
acoustic = self.acoustic_model(
|
||||
x,
|
||||
x_lengths,
|
||||
mel,
|
||||
mel_lengths,
|
||||
sid,
|
||||
*encoder_args,
|
||||
out_size=out_size,
|
||||
generate=True,
|
||||
n_timesteps=n_timesteps,
|
||||
temperature=temperature,
|
||||
)
|
||||
waveform = self.vocoder(acoustic["generated_mel"], sid)
|
||||
return waveform, acoustic
|
||||
|
||||
@torch.inference_mode()
|
||||
def infer(self, *args: Any, **kwargs: Any) -> tuple[torch.Tensor, ...]:
|
||||
mel, attn, mel_mask = self.acoustic_model.infer_mel(*args, **kwargs)
|
||||
sid = args[2]
|
||||
return self.vocoder(mel, sid), mel, attn, mel_mask
|
||||
|
||||
|
||||
class MelVocoder(nn.Module):
|
||||
"""Standalone mel-to-wave Generator with its own speaker embedding."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
generator: nn.Module,
|
||||
n_speakers: int,
|
||||
gin_channels: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
if n_speakers < 1:
|
||||
raise ValueError("Separated mel vocoders require at least one speaker")
|
||||
self.generator = generator
|
||||
self.emb_g = nn.Embedding(n_speakers, gin_channels)
|
||||
|
||||
def forward(self, mel: torch.Tensor, sid: torch.Tensor) -> torch.Tensor:
|
||||
speaker = self.emb_g(sid).unsqueeze(-1)
|
||||
return self.generator(mel, g=speaker)
|
||||
@@ -1133,3 +1133,89 @@ class SynthesizerTrn(nn.Module):
|
||||
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)
|
||||
|
||||
|
||||
def build_mel_synthesizer(
|
||||
n_vocab: int,
|
||||
n_mel_channels: int,
|
||||
hidden_channels: int,
|
||||
filter_channels: int,
|
||||
n_heads: int,
|
||||
n_layers: int,
|
||||
kernel_size: int,
|
||||
p_dropout: float,
|
||||
n_speakers: int,
|
||||
gin_channels: int,
|
||||
matcha_channels: int = 192,
|
||||
matcha_num_heads: int = 2,
|
||||
matcha_dropout: float = 0.05,
|
||||
matcha_sigma_min: float = 1e-4,
|
||||
matcha_n_timesteps: int = 8,
|
||||
matcha_use_diff_attention: bool = False,
|
||||
) -> "MelAcousticModel":
|
||||
"""Build the multilingual text-to-mel model without a waveform generator."""
|
||||
from style_bert_vits2.models.mel_synthesizer import MelAcousticModel
|
||||
|
||||
encoder = TextEncoder(
|
||||
n_vocab,
|
||||
n_mel_channels,
|
||||
hidden_channels,
|
||||
filter_channels,
|
||||
n_heads,
|
||||
n_layers,
|
||||
kernel_size,
|
||||
p_dropout,
|
||||
n_speakers,
|
||||
gin_channels=gin_channels,
|
||||
)
|
||||
sdp = StochasticDurationPredictor(
|
||||
hidden_channels, 192, 3, 0.5, 4, gin_channels=gin_channels
|
||||
)
|
||||
dp = DurationPredictor(
|
||||
hidden_channels, 256, 3, 0.5, gin_channels=gin_channels
|
||||
)
|
||||
return MelAcousticModel(
|
||||
encoder,
|
||||
sdp,
|
||||
dp,
|
||||
n_mel_channels,
|
||||
n_speakers,
|
||||
gin_channels,
|
||||
hidden_channels,
|
||||
matcha_channels,
|
||||
matcha_num_heads,
|
||||
matcha_dropout,
|
||||
matcha_sigma_min,
|
||||
matcha_n_timesteps,
|
||||
matcha_use_diff_attention,
|
||||
)
|
||||
|
||||
|
||||
def build_mel_vocoder(
|
||||
n_mel_channels: int,
|
||||
resblock: str,
|
||||
resblock_kernel_sizes: list[int],
|
||||
resblock_dilation_sizes: list[list[int]],
|
||||
upsample_rates: list[int],
|
||||
upsample_initial_channel: int,
|
||||
upsample_kernel_sizes: list[int],
|
||||
n_speakers: int,
|
||||
gin_channels: int,
|
||||
) -> "MelVocoder":
|
||||
"""Build a standalone HiFi-GAN-style mel-to-wave Generator."""
|
||||
from style_bert_vits2.models.mel_synthesizer import MelVocoder
|
||||
|
||||
return MelVocoder(
|
||||
Generator(
|
||||
n_mel_channels,
|
||||
resblock,
|
||||
resblock_kernel_sizes,
|
||||
resblock_dilation_sizes,
|
||||
upsample_rates,
|
||||
upsample_initial_channel,
|
||||
upsample_kernel_sizes,
|
||||
gin_channels=gin_channels,
|
||||
),
|
||||
n_speakers=n_speakers,
|
||||
gin_channels=gin_channels,
|
||||
)
|
||||
|
||||
@@ -1188,3 +1188,88 @@ class SynthesizerTrn(nn.Module):
|
||||
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)
|
||||
|
||||
|
||||
def build_mel_synthesizer(
|
||||
n_vocab: int,
|
||||
n_mel_channels: int,
|
||||
hidden_channels: int,
|
||||
filter_channels: int,
|
||||
n_heads: int,
|
||||
n_layers: int,
|
||||
kernel_size: int,
|
||||
p_dropout: float,
|
||||
n_speakers: int,
|
||||
gin_channels: int,
|
||||
matcha_channels: int = 192,
|
||||
matcha_num_heads: int = 2,
|
||||
matcha_dropout: float = 0.05,
|
||||
matcha_sigma_min: float = 1e-4,
|
||||
matcha_n_timesteps: int = 8,
|
||||
matcha_use_diff_attention: bool = False,
|
||||
) -> "MelAcousticModel":
|
||||
"""Build the JP-Extra text-to-mel model without a waveform generator."""
|
||||
from style_bert_vits2.models.mel_synthesizer import MelAcousticModel
|
||||
|
||||
encoder = TextEncoder(
|
||||
n_vocab,
|
||||
n_mel_channels,
|
||||
hidden_channels,
|
||||
filter_channels,
|
||||
n_heads,
|
||||
n_layers,
|
||||
kernel_size,
|
||||
p_dropout,
|
||||
gin_channels=gin_channels,
|
||||
)
|
||||
sdp = StochasticDurationPredictor(
|
||||
hidden_channels, 192, 3, 0.5, 4, gin_channels=gin_channels
|
||||
)
|
||||
dp = DurationPredictor(
|
||||
hidden_channels, 256, 3, 0.5, gin_channels=gin_channels
|
||||
)
|
||||
return MelAcousticModel(
|
||||
encoder,
|
||||
sdp,
|
||||
dp,
|
||||
n_mel_channels,
|
||||
n_speakers,
|
||||
gin_channels,
|
||||
hidden_channels,
|
||||
matcha_channels,
|
||||
matcha_num_heads,
|
||||
matcha_dropout,
|
||||
matcha_sigma_min,
|
||||
matcha_n_timesteps,
|
||||
matcha_use_diff_attention,
|
||||
)
|
||||
|
||||
|
||||
def build_mel_vocoder(
|
||||
n_mel_channels: int,
|
||||
resblock: str,
|
||||
resblock_kernel_sizes: list[int],
|
||||
resblock_dilation_sizes: list[list[int]],
|
||||
upsample_rates: list[int],
|
||||
upsample_initial_channel: int,
|
||||
upsample_kernel_sizes: list[int],
|
||||
n_speakers: int,
|
||||
gin_channels: int,
|
||||
) -> "MelVocoder":
|
||||
"""Build a standalone HiFi-GAN-style mel-to-wave Generator."""
|
||||
from style_bert_vits2.models.mel_synthesizer import MelVocoder
|
||||
|
||||
return MelVocoder(
|
||||
Generator(
|
||||
n_mel_channels,
|
||||
resblock,
|
||||
resblock_kernel_sizes,
|
||||
resblock_dilation_sizes,
|
||||
upsample_rates,
|
||||
upsample_initial_channel,
|
||||
upsample_kernel_sizes,
|
||||
gin_channels=gin_channels,
|
||||
),
|
||||
n_speakers=n_speakers,
|
||||
gin_channels=gin_channels,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user