285 lines
9.8 KiB
Python
285 lines
9.8 KiB
Python
"""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)
|