ViuAI_TTS_200M / models /vocoder.py
ViuAI's picture
Fix: Peak audio normalization (-1 dB), robust speaker search, and tuned CFG for crystal-clear benchmark generation
d62a3fe verified
Raw History Blame Contribute Delete
8.22 kB
"""
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