"""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)