""" Neural Audio Vocoder for ViuAI_TTS_200M. BigVGAN / HiFi-GAN inspired multi-receptive field periodic generator. Converts 80-channel mel-spectrograms into high-fidelity 24kHz raw audio waveforms. """ import torch import torch.nn as nn import torch.nn.functional as F from typing import List class ResBlock(nn.Module): def __init__(self, channels: int, kernel_size: int = 3, dilations: List[int] = [1, 3, 5]): super().__init__() self.convs1 = nn.ModuleList([ nn.Conv1d( channels, channels, kernel_size, dilation=d, padding=((kernel_size - 1) * d) // 2 ) for d in dilations ]) self.convs2 = nn.ModuleList([ nn.Conv1d( channels, channels, kernel_size, dilation=1, padding=(kernel_size - 1) // 2 ) for _ in dilations ]) def forward(self, x: torch.Tensor) -> torch.Tensor: for c1, c2 in zip(self.convs1, self.convs2): xt = F.leaky_relu(x, 0.1) xt = c1(xt) xt = F.leaky_relu(xt, 0.1) xt = c2(xt) x = xt + x return x try: from vocos import Vocos HAS_VOCOS = True except ImportError: HAS_VOCOS = False class NeuralVocoder(nn.Module): """ High-Fidelity Neural Vocoder (~18M - 20M parameters). Converts 80-dim Mel Spectrograms to raw 24,000 Hz broadcast audio. Supports studio-grade Vocos 24kHz engine with automatic internal HiFi-GAN fallback. """ def __init__( self, in_channels: int = 80, upsample_initial_channel: int = 512, upsample_rates: List[int] = [8, 8, 2, 2], upsample_kernel_sizes: List[int] = [16, 16, 4, 4], resblock_kernel_sizes: List[int] = [3, 7, 11], resblock_dilation_sizes: List[List[int]] = [[1, 3, 5], [1, 3, 5], [1, 3, 5]], use_vocos: bool = True, train_vocoder: bool = False, ): super().__init__() self.use_vocos = use_vocos self.train_vocoder = train_vocoder # Use object.__setattr__ so Vocos is NOT registered as a PyTorch submodule. # This keeps state_dict clean and ensures strict=True loading passes without missing keys! object.__setattr__(self, "_vocos_instance", None) object.__setattr__(self, "_mel_proj_matrix", None) self.num_upsamples = len(upsample_rates) self.conv_pre = nn.Conv1d(in_channels, upsample_initial_channel, 7, 1, padding=3) self.ups = nn.ModuleList() curr_channel = upsample_initial_channel for r, k in zip(upsample_rates, upsample_kernel_sizes): self.ups.append( nn.ConvTranspose1d( curr_channel, curr_channel // 2, kernel_size=k, stride=r, padding=(k - r) // 2 ) ) curr_channel = curr_channel // 2 self.resblocks = nn.ModuleList() for i in range(len(self.ups)): ch = upsample_initial_channel // (2 ** (i + 1)) for k, d in zip(resblock_kernel_sizes, resblock_dilation_sizes): self.resblocks.append(ResBlock(ch, kernel_size=k, dilations=d)) self.conv_post = nn.Conv1d(ch, 1, 7, 1, padding=3) def _get_vocos(self, device): inst = getattr(self, "_vocos_instance", None) if inst is None and HAS_VOCOS and self.use_vocos: try: v = Vocos.from_pretrained("charactr/vocos-mel-24khz") v = v.to(device) v.eval() object.__setattr__(self, "_vocos_instance", v) except Exception as e: print(f" [!] Vocos load warning: {e}. Internal vocoder fallback active.") object.__setattr__(self, "_vocos_instance", False) inst = getattr(self, "_vocos_instance", None) return inst if inst is not False else None def _get_mel_80_to_100_proj(self, device: torch.device) -> torch.Tensor: """ Computes or retrieves the mathematically exact acoustic projection matrix mapping 80-channel mel-spectrograms (0-8000Hz) to Vocos' 100-channel mel-spectrograms (0-12000Hz). P = M_100.T @ pinv(M_80.T) in R^[100, 80]. Preserves harmonic alignment and formant frequencies without image-like bilinear stretching. """ proj = getattr(self, "_mel_proj_matrix", None) if proj is None: try: import torchaudio mel_80_fb = torchaudio.functional.melscale_fbanks( n_freqs=513, f_min=0.0, f_max=8000.0, n_mels=80, sample_rate=24000, norm="slaney" ) # [513, 80] mel_100_fb = torchaudio.functional.melscale_fbanks( n_freqs=513, f_min=0.0, f_max=12000.0, n_mels=100, sample_rate=24000, norm="slaney" ) # [513, 100] P = mel_100_fb.T @ torch.pinverse(mel_80_fb.T) # [100, 80] proj = P.float() except Exception: P = torch.zeros(100, 80, dtype=torch.float32) for i in range(80): P[int(i * 100 / 80), i] = 1.0 proj = P object.__setattr__(self, "_mel_proj_matrix", proj) return getattr(self, "_mel_proj_matrix", None).to(device) def forward(self, mel: torch.Tensor) -> torch.Tensor: """ Args: mel: Mel-spectrogram [B, 80, T_mel] Returns: waveform: Raw audio [B, 1, T_audio] where T_audio = T_mel * 256 """ # If train_vocoder=True and in training mode, bypass Vocos so internal generator receives gradients! is_training_internal = self.training and self.train_vocoder if not is_training_internal: vocos = self._get_vocos(mel.device) if vocos: try: with torch.no_grad(): # Acoustically project 80-channel mel to Vocos' 100-channel pretrained standard if mel.shape[1] == 80: proj = self._get_mel_80_to_100_proj(mel.device) mel_in = torch.matmul(proj, mel) # Roll-off air-band channels (87 to 99) naturally from channel 86 # Prevents zero-energy injection which log-mel treats as loud high-frequency noise! roll_off = torch.linspace(0.4, 3.5, 13, device=mel.device).unsqueeze(0).unsqueeze(-1) mel_in[:, 87:, :] = mel_in[:, 86:87, :].expand(-1, 13, -1) - roll_off else: mel_in = mel wav = vocos.decode(mel_in) if wav.ndim == 2: wav = wav.unsqueeze(1) # Automatic DC offset subtraction and sub-bass filtering wav = wav - wav.mean(dim=-1, keepdim=True) try: import torchaudio wav = torchaudio.functional.highpass_biquad(wav, 24000, cutoff_freq=50.0) wav = wav - wav.mean(dim=-1, keepdim=True) except Exception: pass return wav except Exception as ex: print(f" [!] Vocos decode exception: {ex}") # Internal Neural Vocoder fallback x = self.conv_pre(mel) num_resblocks_per_stage = len(self.resblocks) // self.num_upsamples for i in range(self.num_upsamples): x = F.leaky_relu(x, 0.1) x = self.ups[i](x) xs = None for j in range(num_resblocks_per_stage): idx = i * num_resblocks_per_stage + j res = self.resblocks[idx](x) xs = res if xs is None else xs + res x = xs / num_resblocks_per_stage x = F.leaky_relu(x) x = self.conv_post(x) x = torch.tanh(x) # DC offset subtraction on internal vocoder output x = x - x.mean(dim=-1, keepdim=True) return x