Download Modules/codec_decoder_mel_speaker.py from FashionFlora/SFlowTTS: direct link, hf CLI and curl.
- Browser
- Download file 58.3 kB
-
https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/codec_decoder_mel_speaker.py
- Command line
-
hf download hf://FashionFlora/SFlowTTS/Modules/codec_decoder_mel_speaker.py
-
curl -L -o codec_decoder_mel_speaker.py https://huggingface.co/FashionFlora/SFlowTTS/resolve/main/Modules/codec_decoder_mel_speaker.py
58.3 kB
| import math | |
| import random | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from torch.nn.utils import weight_norm, remove_weight_norm | |
| from scipy.signal import get_window | |
| from einops import rearrange | |
| from typing import Tuple, Optional, List, Dict, Union | |
| from .conformer import Conformer | |
| from .utils import init_weights, get_padding | |
| class TorchSTFT(nn.Module): | |
| def __init__(self, filter_length=800, hop_length=200, win_length=800, window="hann"): | |
| super().__init__() | |
| self.filter_length = filter_length | |
| self.hop_length = hop_length | |
| self.win_length = win_length | |
| self.window = torch.from_numpy( | |
| get_window(window, win_length, fftbins=True).astype(np.float32) | |
| ) | |
| def transform(self, input_data): | |
| forward_transform = torch.stft( | |
| input_data, | |
| self.filter_length, | |
| self.hop_length, | |
| self.win_length, | |
| window=self.window.to(input_data.device), | |
| return_complex=True, | |
| ) | |
| return torch.abs(forward_transform), torch.angle(forward_transform) | |
| def inverse(self, magnitude, phase): | |
| inverse_transform = torch.istft( | |
| magnitude * torch.exp(phase * 1j), | |
| self.filter_length, | |
| self.hop_length, | |
| self.win_length, | |
| window=self.window.to(magnitude.device), | |
| ) | |
| return inverse_transform.unsqueeze(-2) | |
| class Snake1d(nn.Module): | |
| """Learned periodic activation from BigVGAN.""" | |
| def __init__(self, in_features): | |
| super().__init__() | |
| self.alpha = nn.Parameter(torch.ones(1, in_features, 1)) | |
| def forward(self, x): | |
| return x + (1.0 / (self.alpha + 1e-9)) * (torch.sin(self.alpha * x) ** 2) | |
| class AdaIN1d(nn.Module): | |
| """ | |
| Adaptive Instance Normalization for 1D signals. | |
| Takes a style vector [B, style_dim] and applies affine transformation | |
| to normalized features [B, C, T]. | |
| """ | |
| def __init__(self, style_dim, num_features): | |
| super().__init__() | |
| self.norm = nn.InstanceNorm1d(num_features, affine=False) | |
| self.fc = nn.Linear(style_dim, num_features * 2) | |
| def forward(self, x, s): | |
| """ | |
| Args: | |
| x: [B, C, T] input features | |
| s: [B, style_dim] style/speaker embedding | |
| Returns: | |
| [B, C, T] AdaIN-transformed features | |
| """ | |
| h = self.fc(s) | |
| h = h.view(h.size(0), h.size(1), 1) | |
| gamma, beta = torch.chunk(h, chunks=2, dim=1) | |
| return (1 + gamma) * self.norm(x) + beta | |
| class SpeakerAdaINResBlock1(nn.Module): | |
| """ | |
| Residual block with AdaIN speaker conditioning. | |
| Uses global speaker embedding [B, speaker_dim] for style. | |
| """ | |
| def __init__(self, channels, kernel_size=3, dilation=(1, 3, 5), speaker_dim=128): | |
| super().__init__() | |
| self.convs1 = nn.ModuleList([ | |
| weight_norm(nn.Conv1d(channels, channels, kernel_size, 1, | |
| dilation=d, padding=get_padding(kernel_size, d))) | |
| for d in dilation | |
| ]) | |
| self.convs1.apply(init_weights) | |
| self.convs2 = nn.ModuleList([ | |
| weight_norm(nn.Conv1d(channels, channels, kernel_size, 1, | |
| dilation=1, padding=get_padding(kernel_size, 1))) | |
| for _ in dilation | |
| ]) | |
| self.convs2.apply(init_weights) | |
| self.adain1 = nn.ModuleList([AdaIN1d(speaker_dim, channels) for _ in dilation]) | |
| self.adain2 = nn.ModuleList([AdaIN1d(speaker_dim, channels) for _ in dilation]) | |
| self.snakes1 = nn.ModuleList([Snake1d(channels) for _ in dilation]) | |
| self.snakes2 = nn.ModuleList([Snake1d(channels) for _ in dilation]) | |
| def forward(self, x, speaker_emb): | |
| """ | |
| Args: | |
| x: [B, C, T] input features | |
| speaker_emb: [B, speaker_dim] speaker embedding | |
| """ | |
| for c1, c2, n1, n2, s1, s2 in zip( | |
| self.convs1, self.convs2, self.adain1, self.adain2, self.snakes1, self.snakes2 | |
| ): | |
| xt = n1(x, speaker_emb) | |
| xt = s1(xt) | |
| xt = c1(xt) | |
| xt = n2(xt, speaker_emb) | |
| xt = s2(xt) | |
| xt = c2(xt) | |
| x = xt + x | |
| return x | |
| # ============================================================================== | |
| # Harmonic Source Module | |
| # ============================================================================== | |
| class SineGen(nn.Module): | |
| """Sine generator for F0-based harmonic source with phase caching.""" | |
| def __init__(self, samp_rate, upsample_scale, harmonic_num=0, | |
| sine_amp=0.1, noise_std=0.003, voiced_threshold=0, | |
| flag_for_pulse=False): | |
| super().__init__() | |
| self.sine_amp = sine_amp | |
| self.noise_std = noise_std | |
| self.harmonic_num = harmonic_num | |
| self.dim = harmonic_num + 1 | |
| self.sampling_rate = samp_rate | |
| self.voiced_threshold = voiced_threshold | |
| self.upsample_scale = upsample_scale | |
| self.flag_for_pulse = flag_for_pulse | |
| def _f02uv(self, f0): | |
| return (f0 > self.voiced_threshold).float() | |
| def _f02sine(self, f0_values, initial_phase=None): | |
| rad_values = (f0_values / self.sampling_rate) % 1 | |
| rand_ini = torch.rand(f0_values.shape[0], f0_values.shape[2], device=f0_values.device) | |
| rand_ini[:, 0] = 0 | |
| rad_values[:, 0, :] = rad_values[:, 0, :] + rand_ini | |
| rad_values = F.interpolate( | |
| rad_values.transpose(1, 2), | |
| scale_factor=1 / self.upsample_scale, | |
| mode="linear", | |
| ).transpose(1, 2) | |
| phase = torch.cumsum(rad_values, dim=1) * 2 * np.pi | |
| if initial_phase is not None: | |
| phase = phase + initial_phase | |
| phase = F.interpolate( | |
| phase.transpose(1, 2) * self.upsample_scale, | |
| scale_factor=self.upsample_scale, | |
| mode="linear", | |
| ).transpose(1, 2) | |
| last_phase = phase[:, -1:, :] | |
| if self.flag_for_pulse: | |
| sines = torch.cos(phase) | |
| else: | |
| sines = torch.sin(phase) | |
| return sines, last_phase | |
| def forward(self, f0, initial_phase=None): | |
| f0_buf = torch.zeros(f0.shape[0], f0.shape[1], self.dim, device=f0.device) | |
| fn = torch.multiply( | |
| f0, torch.FloatTensor([[range(1, self.harmonic_num + 2)]]).to(f0.device) | |
| ) | |
| sine_waves, next_phase = self._f02sine(fn, initial_phase) | |
| sine_waves = sine_waves * self.sine_amp | |
| uv = self._f02uv(f0) | |
| noise_amp = uv * self.noise_std + (1 - uv) * self.sine_amp / 3 | |
| noise = noise_amp * torch.randn_like(sine_waves) | |
| sine_waves = sine_waves * uv + noise | |
| return sine_waves, uv, noise, next_phase | |
| class SourceModuleHnNSF(nn.Module): | |
| """Source module for harmonic-plus-noise synthesis.""" | |
| def __init__(self, sampling_rate, upsample_scale, harmonic_num=0, | |
| sine_amp=0.1, add_noise_std=0.003, voiced_threshold=0): | |
| super().__init__() | |
| self.sine_amp = sine_amp | |
| self.noise_std = add_noise_std | |
| self.l_sin_gen = SineGen( | |
| sampling_rate, upsample_scale, harmonic_num, | |
| sine_amp, add_noise_std, voiced_threshold, | |
| flag_for_pulse=False | |
| ) | |
| self.l_linear = nn.Linear(harmonic_num + 1, 1) | |
| self.l_tanh = nn.Tanh() | |
| def forward(self, x, cache=None): | |
| initial_phase = cache | |
| with torch.no_grad(): | |
| sine_wavs, uv, _, next_phase = self.l_sin_gen(x, initial_phase=initial_phase) | |
| sine_merge = self.l_tanh(self.l_linear(sine_wavs)) | |
| noise = torch.randn_like(uv) * self.sine_amp / 3 | |
| return sine_merge, noise, uv, next_phase | |
| # ============================================================================== | |
| # Pixel Shuffle Upsampling | |
| # ============================================================================== | |
| def pixel_shuffle_1d(x: torch.Tensor, r: int) -> torch.Tensor: | |
| B, Cr, L = x.size() | |
| C = Cr // r | |
| x = x.view(B, C, r, L).permute(0, 1, 3, 2) | |
| return x.reshape(B, C, L * r) | |
| class UpsamplePixelShuffle1D(nn.Module): | |
| def __init__(self, in_ch: int, out_ch: int, kernel_size: int, r: int): | |
| super().__init__() | |
| self.r = r | |
| pad_l, pad_r = (kernel_size - 1) // 2, kernel_size // 2 | |
| self.pad = nn.ReflectionPad1d((pad_l, pad_r)) | |
| self.conv = weight_norm(nn.Conv1d(in_ch, out_ch * r, kernel_size, padding=0)) | |
| self._init_icnr(in_ch, out_ch, r, kernel_size) | |
| def _init_icnr(self, in_ch, out_ch, r, kernel_size): | |
| """ICNR initialization for smooth upsampling.""" | |
| weight = self.conv.weight.data | |
| kernel = torch.zeros(out_ch, in_ch, kernel_size) | |
| nn.init.kaiming_normal_(kernel) | |
| weight.copy_(kernel.repeat(r, 1, 1)) | |
| if self.conv.bias is not None: | |
| self.conv.bias.data.fill_(0) | |
| def forward(self, x): | |
| x = self.pad(x) | |
| x = self.conv(x) | |
| return pixel_shuffle_1d(x, self.r) | |
| class EncoderResBlock1d(nn.Module): | |
| """Lightweight residual block used in the mel/pitch encoder.""" | |
| def __init__(self, channels: int, kernel_size: int = 3, dilation: int = 1): | |
| super().__init__() | |
| padding = get_padding(kernel_size, dilation) | |
| self.block = nn.Sequential( | |
| weight_norm(nn.Conv1d(channels, channels, kernel_size, padding=padding, dilation=dilation)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(channels, channels, kernel_size, padding=padding, dilation=dilation)), | |
| ) | |
| self.act = nn.SiLU() | |
| def forward(self, x): | |
| return self.act(x + self.block(x)) | |
| # ============================================================================== | |
| # Mel + Pitch Encoder (replaces text_emb + energy encoder) | |
| # ============================================================================== | |
| class MelPitchEncoderSpeaker(nn.Module): | |
| """ | |
| Encoder that takes mel spectrogram and pitch as input. | |
| Produces a latent representation suitable for quantization. | |
| Supports variable compression ratios via strides parameter: | |
| - strides=[2]: 2x compression (~50Hz tokens) | |
| - strides=[4]: 4x compression (~25Hz tokens) | |
| - strides=[2,2]: 4x compression (~25Hz tokens) | |
| """ | |
| def __init__( | |
| self, | |
| n_mels: int = 80, | |
| speaker_dim: int = 128, | |
| latent_dim: int = 256, | |
| hidden_dim: int = 256, | |
| strides: List[int] = [2], | |
| ): | |
| super().__init__() | |
| self.latent_dim = latent_dim | |
| self.speaker_dim = speaker_dim | |
| self.n_mels = n_mels | |
| self.compression_ratio = int(np.prod(strides)) | |
| self.strides = strides | |
| # Build mel encoder from EXACT stride list for predictable token rate. | |
| mel_layers = [] | |
| in_ch = n_mels | |
| for stage_idx, stride in enumerate(self.strides): | |
| k = 2 * stride + 1 | |
| p = stride | |
| mel_layers.extend([ | |
| weight_norm(nn.Conv1d(in_ch, hidden_dim, kernel_size=k, stride=stride, padding=p)), | |
| nn.SiLU(), | |
| EncoderResBlock1d(hidden_dim, kernel_size=3, dilation=1), | |
| ]) | |
| in_ch = hidden_dim | |
| mel_layers.extend([ | |
| weight_norm(nn.Conv1d(hidden_dim, hidden_dim, 3, stride=1, padding=1)), | |
| nn.SiLU(), | |
| EncoderResBlock1d(hidden_dim, kernel_size=3, dilation=1), | |
| weight_norm(nn.Conv1d(hidden_dim, hidden_dim * 2, 3, stride=1, padding=1)), | |
| nn.SiLU(), | |
| ]) | |
| self.mel_encoder = nn.Sequential(*mel_layers) | |
| # Build pitch encoder from EXACT stride list. | |
| # Input is [logf0, uv] -> 2 channels. | |
| pitch_layers = [] | |
| in_ch = 2 | |
| pitch_hidden = hidden_dim // 2 | |
| for stage_idx, stride in enumerate(self.strides): | |
| k = 2 * stride + 1 | |
| p = stride | |
| pitch_layers.extend([ | |
| weight_norm(nn.Conv1d(in_ch, pitch_hidden, kernel_size=k, stride=stride, padding=p)), | |
| nn.SiLU(), | |
| EncoderResBlock1d(pitch_hidden, kernel_size=3, dilation=1), | |
| ]) | |
| in_ch = pitch_hidden | |
| pitch_layers.extend([ | |
| weight_norm(nn.Conv1d(pitch_hidden, hidden_dim, 3, stride=1, padding=1)), | |
| nn.SiLU(), | |
| EncoderResBlock1d(hidden_dim, kernel_size=3, dilation=1), | |
| ]) | |
| self.pitch_encoder = nn.Sequential(*pitch_layers) | |
| # Fusion of mel and pitch features | |
| input_dim = hidden_dim * 2 + hidden_dim # mel features + pitch features | |
| self.fusion = nn.Sequential( | |
| weight_norm(nn.Conv1d(input_dim, hidden_dim * 2, 7, padding=3)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim * 2, hidden_dim * 2, 5, padding=2)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim * 2, hidden_dim * 2, 3, padding=1)), | |
| nn.SiLU(), | |
| ) | |
| # Refinement layers | |
| self.refine = nn.Sequential( | |
| weight_norm(nn.Conv1d(hidden_dim * 2, hidden_dim * 2, 7, padding=3)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim * 2, hidden_dim * 2, 5, padding=2)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim * 2, hidden_dim * 2, 3, padding=1)), | |
| nn.SiLU(), | |
| ) | |
| # Project to latent dimension | |
| self.to_latent = nn.Sequential( | |
| weight_norm(nn.Conv1d(hidden_dim * 2, hidden_dim * 2, 5, padding=2)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim * 2, latent_dim, 1)), | |
| ) | |
| def forward(self, mel, pitch): | |
| """ | |
| Encode mel spectrogram + pitch into latent. | |
| Args: | |
| mel: [B, n_mels, T] - mel spectrogram | |
| pitch: [B, T] - pitch contour (log F0) | |
| Returns: | |
| latent: [B, latent_dim, T / compression_ratio] | |
| """ | |
| # Encode mel spectrogram | |
| mel_feat = self.mel_encoder(mel) # [B, hidden_dim * 2, T/compression] | |
| # Encode pitch + voiced/unvoiced cue | |
| pitch_in = pitch.unsqueeze(1) | |
| uv = (pitch > 0).float().unsqueeze(1) | |
| pitch_feat = self.pitch_encoder(torch.cat([pitch_in, uv], dim=1)) # [B, hidden_dim, T/compression] | |
| # Align lengths | |
| min_len = min(mel_feat.shape[-1], pitch_feat.shape[-1]) | |
| mel_feat = mel_feat[..., :min_len] | |
| pitch_feat = pitch_feat[..., :min_len] | |
| # Concatenate and fuse | |
| x = torch.cat([mel_feat, pitch_feat], dim=1) | |
| x = self.fusion(x) | |
| x = self.refine(x) | |
| return self.to_latent(x) | |
| # ============================================================================== | |
| # Learnable Upsampling with Anti-Aliasing | |
| # ============================================================================== | |
| class TokenRatePhonemePredictor(nn.Module): | |
| """ | |
| Predict phoneme logits at TOKEN RATE directly from quantized_latent. | |
| No temporal conv (kernel_size=1 only) to prevent the head from learning | |
| phonetic structure by itself. | |
| Inputs: | |
| z_q: [B, latent_dim, T_tok] | |
| lang_emb (optional): [B, lang_dim] | |
| Output: | |
| logits: [B, n_phonemes, T_tok] | |
| """ | |
| def __init__( | |
| self, | |
| latent_dim: int, | |
| n_phonemes: int, | |
| language_emb_dim: int = 64, | |
| hidden_dim: int = 256, | |
| dropout_p: float = 0.1, | |
| ): | |
| super().__init__() | |
| self.n_phonemes = n_phonemes | |
| self.language_emb_dim = language_emb_dim | |
| in_dim = latent_dim + (language_emb_dim if language_emb_dim > 0 else 0) | |
| self.net = nn.Sequential( | |
| weight_norm(nn.Conv1d(in_dim, hidden_dim, kernel_size=1)), | |
| nn.SiLU(), | |
| nn.Dropout(dropout_p), | |
| nn.Conv1d(hidden_dim, n_phonemes, kernel_size=1), | |
| ) | |
| def durations_frames_to_tokens( | |
| durations_frames: torch.Tensor, | |
| compression: int, | |
| t_tok: int, | |
| ) -> torch.Tensor: | |
| """ | |
| Convert mel-frame durations -> token durations using boundary rounding. | |
| durations_frames: [B, T_ph] (mel frames) | |
| returns durations_tokens: [B, T_ph] (tokens) | |
| """ | |
| d = durations_frames.long().clamp(min=0) | |
| end_f = torch.cumsum(d, dim=1) # [B, T_ph] in frames | |
| # Round boundaries to token grid | |
| end_t = torch.round(end_f.float() / float(compression)).long() | |
| end_t = torch.clamp(end_t, 0, t_tok) | |
| start_t = torch.cat( | |
| [torch.zeros_like(end_t[:, :1]), end_t[:, :-1]], | |
| dim=1, | |
| ) | |
| d_tok = (end_t - start_t).clamp(min=0) | |
| return d_tok | |
| def durations_frames_to_tokens_by_length(durations_frames, T_mel, t_tok): | |
| """ | |
| Convert mel-frame durations -> token durations using mapping by length. | |
| Aligns boundaries proportionally between T_mel and t_tok. | |
| """ | |
| d = durations_frames.long().clamp(min=0) | |
| end_f = torch.cumsum(d, dim=1) # [B, T_ph] frame boundaries | |
| # Map frame boundary positions -> token boundary positions | |
| # end_t in [0, t_tok] | |
| end_t = torch.round(end_f.float() * float(t_tok) / float(T_mel)).long() | |
| end_t = end_t.clamp(0, t_tok) | |
| start_t = torch.cat([torch.zeros_like(end_t[:, :1]), end_t[:, :-1]], dim=1) | |
| return (end_t - start_t).clamp(min=0) | |
| def expand_phonemes_by_durations( | |
| phonemes: torch.Tensor, | |
| durations: torch.Tensor, | |
| target_length: int, | |
| ): | |
| """ | |
| Expand phoneme ids by durations to token-level targets. | |
| phonemes: [B, T_ph] | |
| durations: [B, T_ph] in tokens (int) | |
| returns: | |
| expanded: [B, target_length] | |
| valid_mask: [B, target_length] bool | |
| """ | |
| durations = durations.long().clamp(min=0) | |
| end_idxs = torch.cumsum(durations, dim=1) | |
| start_idxs = end_idxs - durations | |
| t = torch.arange(target_length, device=phonemes.device).view(1, 1, -1) | |
| starts = start_idxs.unsqueeze(2) | |
| ends = end_idxs.unsqueeze(2) | |
| mask = (t >= starts) & (t < ends) # [B, T_ph, T_target] | |
| expanded = (phonemes.unsqueeze(2) * mask.long()).sum(dim=1) | |
| valid_mask = mask.sum(dim=1) > 0 | |
| return expanded, valid_mask | |
| def forward(self, z_q: torch.Tensor, lang_emb: Optional[torch.Tensor] = None): | |
| if self.language_emb_dim > 0: | |
| if lang_emb is None: | |
| raise ValueError("lang_emb is required (language_emb_dim > 0).") | |
| # [B, D] -> [B, D, T] | |
| lang = lang_emb.unsqueeze(-1).expand(-1, -1, z_q.size(-1)) | |
| x = torch.cat([z_q, lang], dim=1) | |
| else: | |
| x = z_q | |
| return self.net(x) | |
| class LearnableUpsample1d(nn.Module): | |
| """ | |
| Learnable upsampling using transposed convolution with anti-aliasing. | |
| Better than nn.Upsample for preserving high-frequency details. | |
| """ | |
| def __init__(self, in_channels: int, out_channels: int, scale_factor: int = 2, kernel_size: int = None): | |
| super().__init__() | |
| self.scale_factor = scale_factor | |
| # Kernel size should be 2x scale factor for good coverage | |
| kernel_size = kernel_size or scale_factor * 4 | |
| padding = (kernel_size - scale_factor) // 2 | |
| # Main upsampling via transposed conv | |
| self.upsample = weight_norm(nn.ConvTranspose1d( | |
| in_channels, out_channels, | |
| kernel_size=kernel_size, | |
| stride=scale_factor, | |
| padding=padding, | |
| )) | |
| # Anti-aliasing low-pass filter (learnable) | |
| self.antialiasing = nn.Sequential( | |
| weight_norm(nn.Conv1d(out_channels, out_channels, kernel_size=5, padding=2, groups=out_channels)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(out_channels, out_channels, kernel_size=3, padding=1)), | |
| ) | |
| # Initialize for smooth upsampling | |
| nn.init.kaiming_normal_(self.upsample.weight) | |
| if self.upsample.bias is not None: | |
| nn.init.zeros_(self.upsample.bias) | |
| def forward(self, x): | |
| x = self.upsample(x) | |
| x = self.antialiasing(x) | |
| return x | |
| class LearnableUpsampleBlock(nn.Module): | |
| """ | |
| Multi-stage learnable upsampling block. | |
| Replaces nn.Upsample with learnable transposed convolutions. | |
| """ | |
| def __init__(self, channels: int, total_upsample: int): | |
| super().__init__() | |
| self.total_upsample = total_upsample | |
| # Decompose into 2x upsamples | |
| layers = [] | |
| remaining = total_upsample | |
| while remaining > 1: | |
| factor = min(2, remaining) | |
| layers.append(LearnableUpsample1d(channels, channels, scale_factor=factor)) | |
| remaining //= factor | |
| self.layers = nn.ModuleList(layers) | |
| def forward(self, x): | |
| for layer in self.layers: | |
| x = layer(x) | |
| return x | |
| # ============================================================================== | |
| # Finite Scalar Quantization | |
| # ============================================================================== | |
| class FiniteScalarQuantization(nn.Module): | |
| def __init__(self, input_dim=256, levels: List[int] = [4]*6): | |
| super().__init__() | |
| self.input_dim = input_dim | |
| self.levels = levels | |
| self.dims = len(levels) | |
| self.codebook_size = math.prod(levels) | |
| self.in_proj = nn.Sequential( | |
| nn.Linear(input_dim, input_dim // 2), | |
| nn.SiLU(), | |
| nn.Linear(input_dim // 2, self.dims), | |
| ) | |
| self.out_proj = nn.Sequential( | |
| nn.Linear(self.dims, input_dim // 2), | |
| nn.SiLU(), | |
| nn.Linear(input_dim // 2, input_dim), | |
| ) | |
| self.scale = nn.Parameter(torch.ones(self.dims) * 1.5) | |
| self.bias = nn.Parameter(torch.zeros(self.dims)) | |
| for m in self.in_proj.modules(): | |
| if isinstance(m, nn.Linear): | |
| nn.init.xavier_uniform_(m.weight, gain=2.0) | |
| if m.bias is not None: | |
| nn.init.zeros_(m.bias) | |
| for m in self.out_proj.modules(): | |
| if isinstance(m, nn.Linear): | |
| nn.init.xavier_uniform_(m.weight, gain=1.0) | |
| if m.bias is not None: | |
| nn.init.zeros_(m.bias) | |
| self.register_buffer('levels_tensor', torch.tensor(levels, dtype=torch.float32)) | |
| _basis = torch.cumprod(torch.tensor([1] + levels[:-1]), dim=0) | |
| self.register_buffer('basis', _basis) | |
| self.register_buffer('num_steps', torch.tensor(0)) | |
| self.warmup_steps = 5000 | |
| def forward(self, x, n_quantizers=None): | |
| x = x.transpose(1, 2) | |
| z = self.in_proj(x) | |
| z = z * self.scale + self.bias | |
| z_bound = torch.tanh(z) | |
| if self.training: | |
| self.num_steps += 1 | |
| noise_scale = max(0.3 * (1 - self.num_steps.float() / self.warmup_steps), 0.05) | |
| noise = (torch.rand_like(z_bound) - 0.5) * 2 * noise_scale | |
| z_bound_noisy = z_bound + noise | |
| z_bound_noisy = torch.clamp(z_bound_noisy, -1, 1) | |
| else: | |
| z_bound_noisy = z_bound | |
| levels = self.levels_tensor.to(z.device) | |
| half_l = (levels - 1) / 2 | |
| z_scaled = z_bound_noisy * half_l | |
| z_shifted = z_scaled + half_l | |
| z_ind = z_shifted.round() | |
| z_ind = torch.clamp(z_ind, torch.zeros_like(levels), levels - 1) | |
| z_q_target = z_ind - half_l | |
| z_q = z_scaled + (z_q_target - z_scaled).detach() | |
| out = self.out_proj(z_q) | |
| z_ind_long = z_ind.long() | |
| indices = (z_ind_long * self.basis).sum(dim=-1) | |
| out = out.transpose(1, 2) | |
| aux_loss = self._entropy_loss(z_shifted, levels) | |
| return out, indices.unsqueeze(1), aux_loss | |
| def _entropy_loss(self, z_shifted, levels): | |
| B, T, D = z_shifted.shape | |
| total_entropy_loss = torch.tensor(0.0, device=z_shifted.device) | |
| for d in range(D): | |
| vals = z_shifted[..., d].reshape(-1) | |
| num_levels = int(levels[d].item()) | |
| centers = torch.arange(num_levels, device=z_shifted.device, dtype=torch.float32) | |
| dist = (vals.unsqueeze(1) - centers.unsqueeze(0)).pow(2) | |
| probs = F.softmax(-dist / 0.5, dim=1) | |
| avg_probs = probs.mean(dim=0) | |
| uniform = torch.ones_like(avg_probs) / num_levels | |
| kl_div = (avg_probs * (torch.log(avg_probs + 1e-7) - torch.log(uniform + 1e-7))).sum() | |
| total_entropy_loss = total_entropy_loss + kl_div | |
| return 0.1 * total_entropy_loss / D | |
| def decode(self, indices): | |
| if indices.dim() == 3: | |
| indices = indices.squeeze(1) | |
| z_q = [] | |
| remainder = indices | |
| for i in range(self.dims): | |
| val = remainder % self.levels[i] | |
| remainder = remainder // self.levels[i] | |
| z_q.append(val) | |
| z_q = torch.stack(z_q, dim=-1).float().to(indices.device) | |
| levels = self.levels_tensor.to(indices.device) | |
| half_l = (levels - 1) / 2 | |
| z_q = z_q - half_l | |
| out = self.out_proj(z_q) | |
| return out.transpose(1, 2) | |
| # ========================================================================= | |
| # FLOW MATCHING SUPPORT METHODS | |
| # ========================================================================= | |
| def embed_dim(self) -> int: | |
| """Dimension of continuous FSQ embeddings (for flow matching).""" | |
| return self.dims # 6 for [4,4,4,4,4,4] | |
| def encode_continuous(self, x: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Encode latent to continuous FSQ space (pre-quantization). | |
| This is the TARGET for flow matching training. | |
| Args: | |
| x: [B, input_dim, T] - latent from encoder | |
| Returns: | |
| z_continuous: [B, T, dims] - continuous embeddings in [-half_l, half_l] | |
| For levels=[4,4,4,4,4,4], range is [-1.5, 1.5] per dim | |
| """ | |
| x = x.transpose(1, 2) # [B, T, input_dim] | |
| z = self.in_proj(x) # [B, T, dims] | |
| z = z * self.scale + self.bias | |
| z_bound = torch.tanh(z) # [-1, 1] | |
| levels = self.levels_tensor.to(z.device) | |
| half_l = (levels - 1) / 2 | |
| z_continuous = z_bound * half_l # [-half_l, half_l] per dim | |
| return z_continuous # [B, T, dims] | |
| def quantize_continuous(self, z_continuous: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Quantize continuous FSQ embeddings to token indices. | |
| Use this after flow matching generates z_continuous. | |
| Args: | |
| z_continuous: [B, T, dims] - continuous in [-half_l, half_l] | |
| Returns: | |
| indices: [B, T] - token indices | |
| """ | |
| levels = self.levels_tensor.to(z_continuous.device) | |
| half_l = (levels - 1) / 2 | |
| # Shift to [0, L-1] range and round | |
| z_shifted = z_continuous + half_l | |
| z_ind = z_shifted.round() | |
| z_ind = torch.clamp(z_ind, torch.zeros_like(levels), levels - 1) | |
| # Convert to single index | |
| z_ind_long = z_ind.long() | |
| indices = (z_ind_long * self.basis).sum(dim=-1) | |
| return indices # [B, T] | |
| def continuous_to_latent(self, z_continuous: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Convert continuous FSQ embeddings to decoder-ready latent. | |
| Quantizes and projects back to input_dim. | |
| Args: | |
| z_continuous: [B, T, dims] - from flow matching prediction | |
| Returns: | |
| latent: [B, input_dim, T] - ready for decoder | |
| """ | |
| levels = self.levels_tensor.to(z_continuous.device) | |
| half_l = (levels - 1) / 2 | |
| # Quantize (round to nearest level) | |
| z_shifted = z_continuous + half_l | |
| z_ind = z_shifted.round() | |
| z_ind = torch.clamp(z_ind, torch.zeros_like(levels), levels - 1) | |
| z_q = z_ind - half_l | |
| # Project back to latent dim | |
| out = self.out_proj(z_q) | |
| return out.transpose(1, 2) # [B, input_dim, T] | |
| def indices_to_continuous(self, indices: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Convert token indices to continuous FSQ embeddings. | |
| Useful for getting GT targets from precomputed tokens. | |
| Args: | |
| indices: [B, T] or [B, 1, T] - token indices | |
| Returns: | |
| z_continuous: [B, T, dims] - continuous embeddings | |
| """ | |
| if indices.dim() == 3: | |
| indices = indices.squeeze(1) | |
| # Decompose index into per-dimension values | |
| z_q = [] | |
| remainder = indices | |
| for i in range(self.dims): | |
| val = remainder % self.levels[i] | |
| remainder = remainder // self.levels[i] | |
| z_q.append(val) | |
| z_q = torch.stack(z_q, dim=-1).float() # [B, T, dims] | |
| levels = self.levels_tensor.to(indices.device) | |
| half_l = (levels - 1) / 2 | |
| z_continuous = z_q - half_l # Center around 0 | |
| return z_continuous # [B, T, dims] | |
| # ============================================================================== | |
| # Speaker-Conditioned Fusion Module with AdaIN1d | |
| # ============================================================================== | |
| class SpeakerFusionResBlock(nn.Module): | |
| """ | |
| Fusion ResBlock conditioned on speaker embedding via AdaIN1d. | |
| """ | |
| def __init__( | |
| self, | |
| dim_in, | |
| dim_out, | |
| speaker_dim=128, | |
| actv=nn.LeakyReLU(0.2), | |
| dropout_p=0.0, | |
| ): | |
| super().__init__() | |
| self.actv = actv | |
| self.learned_sc = dim_in != dim_out | |
| self.dropout = nn.Dropout(dropout_p) | |
| self.conv1 = weight_norm(nn.Conv1d(dim_in, dim_out, 3, 1, 1)) | |
| self.conv2 = weight_norm(nn.Conv1d(dim_out, dim_out, 3, 1, 1)) | |
| self.norm1 = AdaIN1d(speaker_dim, dim_in) | |
| self.norm2 = AdaIN1d(speaker_dim, dim_out) | |
| if self.learned_sc: | |
| self.conv1x1 = weight_norm(nn.Conv1d(dim_in, dim_out, 1, 1, 0, bias=False)) | |
| def _shortcut(self, x): | |
| if self.learned_sc: | |
| x = self.conv1x1(x) | |
| return x | |
| def _residual(self, x, speaker_emb): | |
| x = self.norm1(x, speaker_emb) | |
| x = self.actv(x) | |
| x = self.conv1(self.dropout(x)) | |
| x = self.norm2(x, speaker_emb) | |
| x = self.actv(x) | |
| x = self.conv2(self.dropout(x)) | |
| return x | |
| def forward(self, x, speaker_emb): | |
| out = self._residual(x, speaker_emb) | |
| out = (out + self._shortcut(x)) / math.sqrt(2) | |
| return out | |
| class SpeakerLatentFusionModule(nn.Module): | |
| """ | |
| ResNet-style fusion module for latent only (no text), with speaker conditioning. | |
| Takes quantized latent and processes it with speaker AdaIN conditioning. | |
| """ | |
| def __init__(self, latent_dim, hidden_dim, speaker_dim=128): | |
| super().__init__() | |
| self.input_mix = SpeakerFusionResBlock(latent_dim, hidden_dim, speaker_dim) | |
| self.decode = nn.ModuleList() | |
| concat_dim = hidden_dim + latent_dim | |
| self.decode.append(SpeakerFusionResBlock(concat_dim, hidden_dim, speaker_dim)) | |
| self.decode.append(SpeakerFusionResBlock(concat_dim, hidden_dim, speaker_dim)) | |
| self.decode.append(SpeakerFusionResBlock(concat_dim, hidden_dim, speaker_dim)) | |
| def forward(self, latent, speaker_emb): | |
| """ | |
| Args: | |
| latent: [B, latent_dim, T] | |
| speaker_emb: [B, speaker_dim] - global speaker embedding | |
| """ | |
| x = self.input_mix(latent, speaker_emb) | |
| for block in self.decode: | |
| x = torch.cat([x, latent], dim=1) | |
| x = block(x, speaker_emb) | |
| return x | |
| # ============================================================================== | |
| # Waveform Decoder with Speaker Conditioning (for Mel codec) | |
| # ============================================================================== | |
| class MelWaveformDecoderSpeaker(nn.Module): | |
| """ | |
| Waveform decoder for mel codec, conditioned on learnable speaker embeddings. | |
| Takes quantized latent and decodes to waveform. | |
| """ | |
| def __init__( | |
| self, | |
| latent_dim: int = 512, | |
| speaker_dim: int = 128, | |
| hidden_dim: int = 512, | |
| upsample_rates: List[int] = [12, 10], | |
| resblock_kernel_sizes: List[int] = [3, 7, 11], | |
| resblock_dilation_sizes: List[List[int]] = [[1, 3, 5], [1, 3, 5], [1, 3, 5]], | |
| gen_istft_n_fft: int = 30, | |
| gen_istft_hop_size: int = 5, | |
| sample_rate: int = 44100, | |
| source_upsample_rate: Optional[int] = None, | |
| codec_strides: Optional[List[int]] = None, | |
| ): | |
| super().__init__() | |
| self.num_upsamples = len(upsample_rates) | |
| self.num_kernels = len(resblock_kernel_sizes) | |
| self.gen_istft_n_fft = gen_istft_n_fft | |
| self.gen_istft_hop_size = gen_istft_hop_size | |
| self.codec_strides = codec_strides or [1] | |
| self.codec_compression = int(np.prod(self.codec_strides)) | |
| self.speaker_dim = speaker_dim | |
| total_upsample = int(np.prod(upsample_rates)) * gen_istft_hop_size | |
| self.source_upsample_rate = source_upsample_rate or total_upsample | |
| # Learnable upsampler for latent (replaces nn.Upsample) | |
| # Uses transposed convolutions with anti-aliasing for better quality | |
| self.latent_upsampler = LearnableUpsampleBlock(latent_dim, self.codec_compression) | |
| # Simple linear F0 upsampler - just interpolate + smooth with conv1d | |
| # No hidden dims, just direct linear processing | |
| self.f0_upsample_factor = self.codec_compression | |
| self.f0_smooth = nn.Sequential( | |
| weight_norm(nn.Conv1d(1, 1, kernel_size=5, padding=2)), # Smooth after interpolation | |
| weight_norm(nn.Conv1d(1, 1, kernel_size=3, padding=1)), # Final refinement | |
| ) | |
| # F0 predictor from latent | |
| self.f0_predictor = nn.Sequential( | |
| weight_norm(nn.Conv1d(latent_dim, hidden_dim, 3, padding=1)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim, hidden_dim, 3, padding=1)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim, hidden_dim // 2, 3, padding=1)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim // 2, hidden_dim // 4, 3, padding=1)), | |
| nn.SiLU(), | |
| weight_norm(nn.Conv1d(hidden_dim // 4, 1, 3, padding=1)) | |
| ) | |
| # Harmonic source module | |
| self.m_source = SourceModuleHnNSF( | |
| sampling_rate=sample_rate, | |
| upsample_scale=self.source_upsample_rate, | |
| harmonic_num=14, | |
| voiced_threshold=1, | |
| ) | |
| self.f0_upsamp = nn.Upsample(scale_factor=self.source_upsample_rate) | |
| # Speaker-conditioned pre-decoder | |
| self.pre_decoder = SpeakerLatentFusionModule( | |
| latent_dim=latent_dim, | |
| hidden_dim=hidden_dim, | |
| speaker_dim=speaker_dim | |
| ) | |
| # Conformer layers | |
| self.conformers = nn.ModuleList() | |
| for i in range(len(upsample_rates)): | |
| ch = hidden_dim // (2 ** i) | |
| self.conformers.append( | |
| Conformer( | |
| dim=ch, | |
| depth=4, | |
| dim_head=64, | |
| heads=8, | |
| ff_mult=4, | |
| conv_expansion_factor=2, | |
| conv_kernel_size=31, | |
| attn_dropout=0.1, | |
| ff_dropout=0.1, | |
| conv_dropout=0.1, | |
| ) | |
| ) | |
| # Snake activations | |
| self.snakes = nn.ModuleList() | |
| self.snakes.append(Snake1d(hidden_dim)) | |
| # Upsampling layers | |
| self.ups = nn.ModuleList() | |
| upsample_kernel_sizes = [2 * u for u in upsample_rates] | |
| for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)): | |
| in_ch = hidden_dim // (2 ** i) | |
| out_ch = hidden_dim // (2 ** (i + 1)) | |
| self.ups.append(UpsamplePixelShuffle1D(in_ch, out_ch, kernel_size=k, r=u)) | |
| self.snakes.append(Snake1d(out_ch)) | |
| # Noise injection layers | |
| self.noise_convs = nn.ModuleList() | |
| self.noise_res = nn.ModuleList() | |
| for i in range(len(upsample_rates)): | |
| c_cur = hidden_dim // (2 ** (i + 1)) | |
| if i + 1 < len(upsample_rates): | |
| stride_f0 = int(np.prod(upsample_rates[i + 1:])) | |
| self.noise_convs.append( | |
| weight_norm(nn.Conv1d( | |
| gen_istft_n_fft + 2, c_cur, | |
| kernel_size=stride_f0 * 2, | |
| stride=stride_f0, | |
| padding=(stride_f0 + 1) // 2, | |
| )) | |
| ) | |
| self.noise_res.append(SpeakerAdaINResBlock1(c_cur, 7, [1, 3, 5], speaker_dim)) | |
| else: | |
| self.noise_convs.append( | |
| weight_norm(nn.Conv1d(gen_istft_n_fft + 2, c_cur, kernel_size=1)) | |
| ) | |
| self.noise_res.append(SpeakerAdaINResBlock1(c_cur, 11, [1, 3, 5], speaker_dim)) | |
| # ResBlocks with speaker AdaIN conditioning | |
| self.resblocks = nn.ModuleList() | |
| for i in range(len(upsample_rates)): | |
| ch = hidden_dim // (2 ** (i + 1)) | |
| for k, d in zip(resblock_kernel_sizes, resblock_dilation_sizes): | |
| self.resblocks.append(SpeakerAdaINResBlock1(ch, k, d, speaker_dim)) | |
| # Post convolution for STFT | |
| self.post_n_fft = gen_istft_n_fft | |
| final_ch = hidden_dim // (2 ** len(upsample_rates)) | |
| self.conv_post = weight_norm(nn.Conv1d(final_ch, self.post_n_fft + 2, 7, padding=3)) | |
| # STFT for inverse transform | |
| self.stft = TorchSTFT( | |
| filter_length=gen_istft_n_fft, | |
| hop_length=gen_istft_hop_size, | |
| win_length=gen_istft_n_fft, | |
| ) | |
| self.reflection_pad = nn.ReflectionPad1d((1, 0)) | |
| def forward(self, latent, speaker_emb, f0_gt=None, cache=None): | |
| """ | |
| Args: | |
| latent: [B, latent_dim, T_comp] - quantized latent | |
| speaker_emb: [B, speaker_dim] - global speaker embedding | |
| f0_gt: [B, T] optional ground truth F0 | |
| cache: dict for streaming inference | |
| Returns: | |
| wav, spec, phase, f0_pred, (new_cache if streaming) | |
| """ | |
| B = latent.shape[0] | |
| # Predict F0 from latent | |
| f0_pred_latent = self.f0_predictor(latent) # [B, 1, T_comp] | |
| # Simple linear upsample - just interpolate and smooth | |
| f0_pred = F.interpolate(f0_pred_latent, scale_factor=self.f0_upsample_factor, mode='linear', align_corners=False) | |
| f0_pred = self.f0_smooth(f0_pred) # Smooth with conv1d | |
| # Match F0 to target length | |
| if f0_gt is not None: | |
| f0_pred = F.interpolate(f0_pred, size=f0_gt.shape[-1], mode='linear') | |
| else: | |
| target_len = int(latent.shape[-1] * self.codec_compression) | |
| if f0_pred.shape[-1] != target_len: | |
| f0_pred = F.interpolate(f0_pred, size=target_len, mode='linear') | |
| f0_pred = f0_pred.squeeze(1) | |
| f0_to_use = f0_gt if f0_gt is not None else f0_pred.detach() | |
| # Generate harmonic source | |
| uv = (f0_to_use > 0).float() # f0_to_use is log10(F0), 0 = unvoiced | |
| f0_log_up = self.f0_upsamp(f0_to_use[:, None]).transpose(1, 2) | |
| uv_up = self.f0_upsamp(uv[:, None]).transpose(1, 2) | |
| uv_up = (uv_up > 0.5).to(f0_log_up.dtype) | |
| f0_lin = (10.0 ** f0_log_up.float()).to(f0_log_up.dtype) * uv_up | |
| source_phase_cache = cache.get("source_phase") if cache is not None else None | |
| har_source, noi_source, uv, next_source_phase = self.m_source(f0_lin, cache=source_phase_cache) | |
| har_source = har_source.transpose(1, 2).squeeze(1) | |
| har_spec, har_phase = self.stft.transform(har_source) | |
| har = torch.cat([har_spec, har_phase], dim=1) | |
| # Upsample latent | |
| latent_up = self.latent_upsampler(latent) | |
| # Speaker-conditioned fusion | |
| x = self.pre_decoder(latent_up, speaker_emb) | |
| # Upsampling with conformers and residual blocks | |
| for i in range(self.num_upsamples): | |
| x = self.snakes[i](x) | |
| x = rearrange(x, "b f t -> b t f") | |
| x = self.conformers[i](x) | |
| x = rearrange(x, "b t f -> b f t") | |
| x = self.ups[i](x) | |
| x_source = self.noise_convs[i](har) | |
| x_source = self.noise_res[i](x_source, speaker_emb) | |
| if i == self.num_upsamples - 1: | |
| x = self.reflection_pad(x) | |
| if x.shape[-1] != x_source.shape[-1]: | |
| min_len_add = min(x.shape[-1], x_source.shape[-1]) | |
| x = x[..., :min_len_add] | |
| x_source = x_source[..., :min_len_add] | |
| x = x + x_source | |
| xs = None | |
| for j in range(self.num_kernels): | |
| if xs is None: | |
| xs = self.resblocks[i * self.num_kernels + j](x, speaker_emb) | |
| else: | |
| xs += self.resblocks[i * self.num_kernels + j](x, speaker_emb) | |
| x = xs / self.num_kernels | |
| x = self.snakes[-1](x) | |
| x = self.conv_post(x) | |
| spec = torch.exp(x[:, :self.post_n_fft // 2 + 1, :]) | |
| phase = torch.sin(x[:, self.post_n_fft // 2 + 1:, :]) | |
| out = self.stft.inverse(spec, phase) | |
| if cache is not None: | |
| new_cache = { | |
| "source_phase": next_source_phase | |
| } | |
| return out, spec, phase, f0_pred, new_cache | |
| return out, spec, phase, f0_pred | |
| # ============================================================================== | |
| # Main Mel Codec with Learnable Speaker Embeddings | |
| # ============================================================================== | |
| class MelCodecVocoderSpeaker(nn.Module): | |
| """ | |
| Mel Spectrogram Codec with LEARNABLE SPEAKER EMBEDDINGS. | |
| Key features: | |
| - Input: mel spectrogram + pitch (instead of text_emb + energy) | |
| - Learnable speaker embedding: nn.Embedding(num_speakers, speaker_dim) | |
| - Speaker conditioning via AdaIN1d throughout the decoder | |
| - F0 prediction from latent | |
| Usage: | |
| model = MelCodecVocoderSpeaker(num_speakers=11, speaker_dim=128, ...) | |
| output = model(mel, pitch, speaker_ids=speaker_ids) | |
| """ | |
| def __init__( | |
| self, | |
| num_speakers: int = 11, | |
| speaker_dim: int = 128, | |
| n_mels: int = 80, | |
| latent_dim: int = 512, | |
| hidden_dim: int = 512, | |
| codec_strides: List[int] = [2, 2], | |
| codebook_size: int = 4096, | |
| upsample_rates: List[int] = [9, 7], | |
| gen_istft_n_fft: int = 30, | |
| gen_istft_hop_size: int = 5, | |
| sample_rate: int = 44100, | |
| source_upsample_rate: int = 441, | |
| fsq_levels: Optional[List[int]] = None, | |
| # Phoneme predictor settings | |
| n_phonemes: int = 178, | |
| num_languages: int = 10, | |
| language_dim: int = 64, | |
| ): | |
| super().__init__() | |
| self.num_speakers = num_speakers | |
| self.speaker_dim = speaker_dim | |
| self.n_mels = n_mels | |
| self.latent_dim = latent_dim | |
| self.codec_compression = math.prod(codec_strides) | |
| self.use_fsq = fsq_levels is not None | |
| self.fsq_levels = fsq_levels or [4] * 6 | |
| self.speaker_emb_dim = 64 | |
| self.language_emb_dim = 64 | |
| self.speaker_embedding = nn.Embedding( | |
| num_embeddings=num_speakers, | |
| embedding_dim=self.speaker_emb_dim | |
| ) | |
| nn.init.normal_(self.speaker_embedding.weight, mean=0, std=0.5) | |
| self.language_embedding = nn.Embedding( | |
| num_embeddings=num_languages, | |
| embedding_dim=self.language_emb_dim | |
| ) | |
| nn.init.normal_(self.language_embedding.weight, mean=0, std=0.5) | |
| # Mel + Pitch encoder | |
| self.encoder = MelPitchEncoderSpeaker( | |
| n_mels=n_mels, | |
| speaker_dim=self.speaker_emb_dim + self.language_emb_dim, # Combined dim (128) | |
| latent_dim=latent_dim, | |
| hidden_dim=hidden_dim, | |
| strides=codec_strides, | |
| ) | |
| # FSQ Quantizer | |
| self.quantizer = FiniteScalarQuantization( | |
| input_dim=latent_dim, | |
| levels=self.fsq_levels, | |
| ) | |
| # Decoder | |
| self.decoder = MelWaveformDecoderSpeaker( | |
| latent_dim=latent_dim, | |
| speaker_dim=self.speaker_emb_dim + self.language_emb_dim, # Combined dim (128) | |
| hidden_dim=hidden_dim, | |
| upsample_rates=upsample_rates, | |
| gen_istft_n_fft=gen_istft_n_fft, | |
| gen_istft_hop_size=gen_istft_hop_size, | |
| sample_rate=sample_rate, | |
| source_upsample_rate=source_upsample_rate, | |
| codec_strides=codec_strides, | |
| ) | |
| # ===================================================================== | |
| # PHONEME PREDICTOR for linguistic supervision | |
| # ===================================================================== | |
| self.phoneme_predictor = TokenRatePhonemePredictor( | |
| latent_dim=latent_dim, | |
| n_phonemes=n_phonemes, | |
| language_emb_dim=self.language_emb_dim, # 64 | |
| hidden_dim=hidden_dim // 2, | |
| dropout_p=0.1, | |
| ) | |
| def forward(self, mel, pitch, speaker_ids, language_ids=None, n_quantizers=None, use_predicted_f0=False, | |
| target_phonemes=None, target_durations=None): | |
| """ | |
| Training forward pass. | |
| Args: | |
| mel: [B, n_mels, T] - mel spectrogram | |
| pitch: [B, T] - pitch contour (log F0) | |
| speaker_ids: [B] - speaker IDs | |
| language_ids: [B] - language IDs (for phoneme prediction) | |
| n_quantizers: unused, for compatibility | |
| use_predicted_f0: bool - whether to use predicted F0 | |
| target_phonemes: [B, T_ph] - optional, target phoneme IDs for loss | |
| target_durations: [B, T_ph] - optional, target durations for loss | |
| Returns: | |
| dict with wav, tokens, speaker_emb, f0_pred, phoneme_logits, etc. | |
| """ | |
| # Get speaker embedding from ID | |
| s_emb = self.speaker_embedding(speaker_ids) # [B, 64] | |
| if language_ids is None: | |
| # Try to infer from device/shape if possible or default to 0? | |
| # Ideally we raise error, but to be robust: | |
| # raise ValueError("language_ids must be provided for style construction") | |
| # Assuming language_ids are required. | |
| raise ValueError("language_ids are required for concatenation.") | |
| l_emb = self.language_embedding(language_ids) # [B, 64] | |
| # Concatenate for Style [B, 128] | |
| speaker_emb = torch.cat([s_emb, l_emb], dim=1) | |
| # Encode mel + pitch | |
| latent = self.encoder(mel, pitch) | |
| # Quantize | |
| quantized_latent, tokens, commitment_loss = self.quantizer(latent) | |
| # Decide whether to use GT or predicted F0 | |
| decoder_f0 = None if use_predicted_f0 else pitch | |
| # Decode with speaker conditioning | |
| wav, mag, phase, f0_pred = self.decoder( | |
| quantized_latent, | |
| speaker_emb, | |
| f0_gt=decoder_f0, | |
| cache=None, | |
| ) | |
| # Phoneme prediction from latent (for linguistic supervision) | |
| phoneme_logits = None | |
| loss_phoneme = None | |
| if language_ids is not None: | |
| # Predict at frame rate (upsampled latent) | |
| # quantized_latent: [B, C, T_tok] | |
| T_mel = pitch.size(-1) | |
| # Upsample latent to frame resolution (nearest neighbor - non-learnable) | |
| quantized_latent_upsampled = F.interpolate(quantized_latent, size=T_mel, mode='nearest') | |
| phoneme_logits = self.phoneme_predictor(quantized_latent_upsampled, l_emb.detach()) # [B, P, T_mel] | |
| if target_phonemes is not None and target_durations is not None: | |
| # Expand phonemes to frame-level targets using frame durations directly | |
| targets_frames, valid_mask_frames = TokenRatePhonemePredictor.expand_phonemes_by_durations( | |
| phonemes=target_phonemes.long(), | |
| durations=target_durations, | |
| target_length=T_mel, | |
| ) | |
| # Masked CE over valid frames | |
| # NOTE: do NOT ignore_index=0 because 0 is a real phoneme in your vocab. | |
| logits_flat = phoneme_logits.transpose(1, 2).reshape(-1, phoneme_logits.size(1)) | |
| targets_flat = targets_frames.reshape(-1) | |
| mask_flat = valid_mask_frames.reshape(-1) | |
| if mask_flat.any(): | |
| loss_phoneme = F.cross_entropy(logits_flat[mask_flat], targets_flat[mask_flat]) | |
| else: | |
| loss_phoneme = torch.tensor(0.0, device=mel.device) | |
| return { | |
| "wav": wav, | |
| "mag": mag, | |
| "phase": phase, | |
| "tokens": tokens, | |
| "latent": latent, | |
| "quantized_latent": quantized_latent, | |
| "speaker_emb": speaker_emb, | |
| "commitment_loss": commitment_loss, | |
| "f0_pred": f0_pred, | |
| "f0_gt": pitch, | |
| "phoneme_logits": phoneme_logits, | |
| "phoneme_loss": loss_phoneme, | |
| } | |
| def get_speaker_embedding(self, speaker_ids): | |
| """Get speaker embedding from IDs (only speaker part).""" | |
| return self.speaker_embedding(speaker_ids) | |
| def get_style_embedding(self, speaker_ids, language_ids): | |
| """Get full style embedding (speaker + language).""" | |
| s_emb = self.speaker_embedding(speaker_ids) | |
| l_emb = self.language_embedding(language_ids) | |
| return torch.cat([s_emb, l_emb], dim=1) | |
| def tokenize(self, mel, pitch, speaker_ids, language_ids=None, n_quantizers=None): | |
| """Tokenize mel + pitch.""" | |
| s_emb = self.speaker_embedding(speaker_ids) | |
| if language_ids is not None: | |
| l_emb = self.language_embedding(language_ids) | |
| style_emb = torch.cat([s_emb, l_emb], dim=1) | |
| else: | |
| # Return partial if language not provided (might break if used for decoding) | |
| style_emb = s_emb | |
| latent = self.encoder(mel, pitch) | |
| _, tokens, _ = self.quantizer(latent) | |
| return tokens, style_emb | |
| # ========================================================================= | |
| # FLOW MATCHING SUPPORT | |
| # ========================================================================= | |
| def encode_for_flow_matching(self, mel, pitch): | |
| """ | |
| Encode mel+pitch to continuous FSQ embeddings for flow matching training. | |
| Args: | |
| mel: [B, n_mels, T] - mel spectrogram | |
| pitch: [B, T] - pitch contour (log F0) | |
| Returns: | |
| z_continuous: [B, T_codec, fsq_dims] - continuous FSQ embeddings | |
| These are the TARGETS for flow matching. | |
| fsq_dims = 6 for levels=[4,4,4,4,4,4] | |
| Range: [-1.5, 1.5] per dimension | |
| """ | |
| latent = self.encoder(mel, pitch) # [B, latent_dim, T_codec] | |
| z_continuous = self.quantizer.encode_continuous(latent) # [B, T_codec, fsq_dims] | |
| return z_continuous | |
| def tokenize_with_continuous(self, mel, pitch, speaker_ids, language_ids=None): | |
| """ | |
| Tokenize and also return continuous FSQ embeddings. | |
| Returns: | |
| tokens: [B, 1, T_codec] - discrete token indices | |
| z_continuous: [B, T_codec, fsq_dims] - continuous FSQ embeddings | |
| style_emb: [B, style_dim] - combined speaker+language embedding | |
| """ | |
| s_emb = self.speaker_embedding(speaker_ids) | |
| if language_ids is not None: | |
| l_emb = self.language_embedding(language_ids) | |
| style_emb = torch.cat([s_emb, l_emb], dim=1) | |
| else: | |
| style_emb = s_emb | |
| latent = self.encoder(mel, pitch) | |
| z_continuous = self.quantizer.encode_continuous(latent) | |
| _, tokens, _ = self.quantizer(latent) | |
| return tokens, z_continuous, style_emb | |
| def decode_from_continuous(self, z_continuous, speaker_ids, language_ids, f0=None): | |
| """ | |
| Decode from continuous FSQ embeddings (flow matching output). | |
| Args: | |
| z_continuous: [B, T_codec, fsq_dims] - from flow matching | |
| speaker_ids: [B] - speaker IDs | |
| language_ids: [B] - language IDs | |
| f0: [B, T] optional F0 | |
| Returns: | |
| wav: [B, 1, T_audio] - waveform | |
| tokens: [B, T_codec] - quantized token indices | |
| f0_pred: [B, T] - predicted F0 | |
| """ | |
| s_emb = self.speaker_embedding(speaker_ids) | |
| l_emb = self.language_embedding(language_ids) | |
| speaker_emb = torch.cat([s_emb, l_emb], dim=1) | |
| # Quantize continuous to tokens | |
| tokens = self.quantizer.quantize_continuous(z_continuous) | |
| # Convert to decoder latent | |
| quantized_latent = self.quantizer.continuous_to_latent(z_continuous) | |
| wav, _, _, f0_pred = self.decoder( | |
| quantized_latent, | |
| speaker_emb, | |
| f0_gt=f0, | |
| cache=None, | |
| ) | |
| return wav, tokens, f0_pred | |
| def fsq_embed_dim(self) -> int: | |
| """Dimension of continuous FSQ embeddings (for flow matching).""" | |
| return self.quantizer.embed_dim | |
| def decode_tokens(self, tokens, speaker_ids, language_ids, f0=None): | |
| """ | |
| Decode tokens with speaker ID and Language ID. | |
| Args: | |
| tokens: [B, 1, T_comp] - tokens | |
| speaker_ids: [B] - speaker IDs | |
| language_ids: [B] - language IDs | |
| f0: [B, T] optional F0 to condition on | |
| """ | |
| s_emb = self.speaker_embedding(speaker_ids) | |
| l_emb = self.language_embedding(language_ids) | |
| speaker_emb = torch.cat([s_emb, l_emb], dim=1) | |
| quantized_latent = self.quantizer.decode(tokens) | |
| wav, _, _, f0_pred = self.decoder( | |
| quantized_latent, | |
| speaker_emb, | |
| f0_gt=f0, | |
| cache=None, | |
| ) | |
| return wav, f0_pred | |
| def decode_tokens_with_speaker_emb(self, tokens, speaker_emb, f0=None): | |
| """ | |
| Decode tokens with pre-computed speaker embedding. | |
| Useful for speaker interpolation. | |
| Args: | |
| tokens: [B, 1, T_comp] - tokens | |
| speaker_emb: [B, speaker_dim] - speaker embedding | |
| f0: [B, T] optional F0 | |
| """ | |
| quantized_latent = self.quantizer.decode(tokens) | |
| wav, _, _, f0_pred = self.decoder( | |
| quantized_latent, | |
| speaker_emb, | |
| f0_gt=f0, | |
| cache=None, | |
| ) | |
| return wav, f0_pred | |
| def decode_chunk(self, tokens, speaker_ids, language_ids, cache=None, f0=None): | |
| """ | |
| Streaming inference by chunk. | |
| Args: | |
| tokens: Chunk of tokens | |
| speaker_ids: [B] speaker IDs | |
| language_ids: [B] language IDs | |
| cache: Dictionary from previous chunk call | |
| f0: [B, T] optional F0 | |
| Returns: | |
| wav_chunk, f0_pred, new_cache | |
| """ | |
| if cache is None: | |
| cache = {} | |
| s_emb = self.speaker_embedding(speaker_ids) | |
| l_emb = self.language_embedding(language_ids) | |
| speaker_emb = torch.cat([s_emb, l_emb], dim=1) | |
| quantized_latent = self.quantizer.decode(tokens) | |
| wav, _, _, f0_pred, new_cache = self.decoder( | |
| quantized_latent, | |
| speaker_emb, | |
| f0_gt=f0, | |
| cache=cache, | |
| ) | |
| return wav, f0_pred, new_cache | |
| def interpolate_speakers(self, speaker_id_1, speaker_id_2, alpha=0.5): | |
| """ | |
| Interpolate between two speaker embeddings. | |
| Args: | |
| speaker_id_1: int - first speaker ID | |
| speaker_id_2: int - second speaker ID | |
| alpha: float - interpolation weight (0 = speaker_1, 1 = speaker_2) | |
| Returns: | |
| [1, speaker_dim] interpolated embedding | |
| """ | |
| emb1 = self.speaker_embedding(torch.tensor([speaker_id_1], device=self.speaker_embedding.weight.device)) | |
| emb2 = self.speaker_embedding(torch.tensor([speaker_id_2], device=self.speaker_embedding.weight.device)) | |
| return (1 - alpha) * emb1 + alpha * emb2 | |
| def encode_and_reconstruct(self, mel, pitch, speaker_ids, language_ids, use_predicted_f0=False): | |
| """ | |
| Encode mel+pitch and reconstruct waveform. | |
| Useful for testing reconstruction quality. | |
| """ | |
| output = self.forward(mel, pitch, speaker_ids, language_ids, use_predicted_f0=use_predicted_f0) | |
| return output['wav'], output['f0_pred'], output['tokens'] | |