Files
sbv2-v2/style_bert_vits2/models/mel_synthesizer.py
tuna2134 8125666e22 split
2026-07-24 16:24:50 +09:00

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)