ViuAI_TTS_200M / dataset /audio_processing.py
ViuAI's picture
Turbo Speed: Vectorized FFT F0 pitch, cuDNN autotune, Tensor Core precision, speaker RAM cache & prefetch DataLoader
b8e1bfc verified
Raw History Blame Contribute Delete
10.7 kB
"""
Audio Processing & Mel-Spectrogram Extraction Module for ViuAI_TTS_200M.
Supports both torchaudio (if installed) and zero-dependency pure PyTorch standard triangular Mel filterbank.
"""
import math
import os
import wave
import numpy as np
import torch
import torch.nn.functional as F
from typing import Optional, Tuple
try:
import torchaudio
import torchaudio.transforms as T
HAS_TORCHAUDIO = True
except ImportError:
HAS_TORCHAUDIO = False
def create_mel_filterbank(
sample_rate: int = 24000,
n_fft: int = 1024,
n_mels: int = 80,
f_min: float = 0.0,
f_max: float = 8000.0,
) -> torch.Tensor:
"""
Constructs a standard triangular Mel filterbank matrix of shape [n_mels, n_fft // 2 + 1].
Exact mathematical equivalent to torchaudio / librosa Slaney/HTK Mel filterbank.
"""
def hz_to_mel(f):
return 2595.0 * torch.log10(1.0 + f / 700.0)
def mel_to_hz(m):
return 700.0 * (10.0 ** (m / 2595.0) - 1.0)
m_min = hz_to_mel(torch.tensor(float(f_min)))
m_max = hz_to_mel(torch.tensor(float(f_max)))
m_pts = torch.linspace(m_min, m_max, n_mels + 2)
f_pts = mel_to_hz(m_pts)
bins = torch.floor((n_fft + 1) * f_pts / sample_rate).long()
n_freq = n_fft // 2 + 1
weights = torch.zeros(n_mels, n_freq)
for i in range(n_mels):
left, center, right = bins[i].item(), bins[i + 1].item(), bins[i + 2].item()
for j in range(left, center):
if j < n_freq:
weights[i, j] = (j - left) / max(1, center - left)
for j in range(center, right):
if j < n_freq:
weights[i, j] = (right - j) / max(1, right - center)
return weights
class MelSpectrogramExtractor:
"""
Standard 80-channel Mel-Spectrogram Extractor for 24kHz audio.
Seamlessly falls back to pure PyTorch when torchaudio is not installed.
"""
def __init__(
self,
sample_rate: int = 24000,
n_fft: int = 1024,
win_length: int = 1024,
hop_length: int = 256,
n_mels: int = 80,
f_min: float = 0.0,
f_max: float = 8000.0,
):
self.sample_rate = sample_rate
self.n_fft = n_fft
self.win_length = win_length
self.hop_length = hop_length
self.n_mels = n_mels
self.f_min = f_min
self.f_max = f_max
if HAS_TORCHAUDIO:
self.mel_transform = T.MelSpectrogram(
sample_rate=sample_rate,
n_fft=n_fft,
win_length=win_length,
hop_length=hop_length,
f_min=f_min,
f_max=f_max,
n_mels=n_mels,
power=1.0,
normalized=False,
center=True,
pad_mode="reflect",
)
else:
self.mel_transform = None
self.fb = create_mel_filterbank(sample_rate, n_fft, n_mels, f_min, f_max)
def __call__(self, waveform: torch.Tensor) -> torch.Tensor:
"""
Args:
waveform: [B, 1, T] or [1, T] or [T] in float32 in [-1, 1]
Returns:
mel: [B, 80, T_frames] normalized log-mel
"""
if waveform.ndim == 1:
waveform = waveform.unsqueeze(0).unsqueeze(0)
elif waveform.ndim == 2:
waveform = waveform.unsqueeze(0)
device = waveform.device
if self.mel_transform is not None:
mel = self.mel_transform(waveform).squeeze(1)
else:
# Pure PyTorch STFT with standard Mel filterbank
fb = self.fb.to(device)
window = torch.hann_window(self.win_length, device=device)
B, C, T_samples = waveform.shape
audio_flat = waveform.view(B * C, T_samples)
stft = torch.stft(
audio_flat,
n_fft=self.n_fft,
hop_length=self.hop_length,
win_length=self.win_length,
window=window,
center=True,
pad_mode="reflect",
return_complex=True,
)
mag = torch.abs(stft) # [B*C, 513, T_frames]
mel = torch.matmul(fb, mag) # [B*C, 80, T_frames]
if B > 1:
mel = mel.view(B, self.n_mels, -1)
# Log compression with dynamic range clamping
log_mel = torch.log(torch.clamp(mel, min=1e-5))
return log_mel
# Global singleton extractor for fast reuse
_default_extractor: Optional[MelSpectrogramExtractor] = None
def get_default_extractor(sample_rate: int = 24000) -> MelSpectrogramExtractor:
global _default_extractor
if _default_extractor is None or _default_extractor.sample_rate != sample_rate:
_default_extractor = MelSpectrogramExtractor(sample_rate=sample_rate)
return _default_extractor
def load_audio_wav(audio_path: str, target_sr: int = 24000) -> Optional[torch.Tensor]:
"""
Loads an audio file (WAV, MP3, FLAC, OGG, 16/24/32-bit), converts to mono,
resamples to target_sr if needed, and returns a float32 tensor of shape [1, T] in [-1.0, 1.0].
"""
if not os.path.exists(audio_path):
return None
# 1. Try soundfile (broadest multi-format support: WAV, FLAC, OGG, 24-bit PCM, 32-bit float)
try:
import soundfile as sf
data, samplerate = sf.read(audio_path, dtype="float32")
if data.ndim > 1:
data = data[:, 0] # Mono
audio = torch.from_numpy(data.copy()).unsqueeze(0)
if samplerate != target_sr and samplerate > 0:
if HAS_TORCHAUDIO:
resampler = torchaudio.transforms.Resample(orig_freq=samplerate, new_freq=target_sr)
audio = resampler(audio)
else:
target_len = int(audio.shape[-1] * (target_sr / samplerate))
audio = F.interpolate(audio.unsqueeze(0), size=target_len, mode="linear", align_corners=False).squeeze(0)
return audio
except Exception:
pass
# 2. Try torchaudio.load (supports MP3 via sox/ffmpeg backend)
if HAS_TORCHAUDIO:
try:
audio, samplerate = torchaudio.load(audio_path)
if audio.shape[0] > 1:
audio = audio[:1, :] # Mono
if samplerate != target_sr and samplerate > 0:
resampler = torchaudio.transforms.Resample(orig_freq=samplerate, new_freq=target_sr)
audio = resampler(audio)
return audio.float()
except Exception:
pass
# 3. Standard library wave module fallback
try:
with wave.open(audio_path, "rb") as wf:
n_channels = wf.getnchannels()
sampwidth = wf.getsampwidth()
framerate = wf.getframerate()
n_frames = wf.getnframes()
audio_bytes = wf.readframes(n_frames)
if sampwidth == 2:
arr = np.frombuffer(audio_bytes, dtype=np.int16).astype(np.float32) / 32768.0
elif sampwidth == 1:
arr = np.frombuffer(audio_bytes, dtype=np.uint8).astype(np.float32) / 128.0 - 1.0
elif sampwidth == 4:
arr = np.frombuffer(audio_bytes, dtype=np.int32).astype(np.float32) / 2147483648.0
else:
arr = np.frombuffer(audio_bytes, dtype=np.int16).astype(np.float32) / 32768.0
audio = torch.from_numpy(arr.copy())
if n_channels > 1:
audio = audio.view(-1, n_channels)[:, 0]
if audio.ndim == 1:
audio = audio.unsqueeze(0)
if framerate != target_sr and framerate > 0:
if HAS_TORCHAUDIO:
resampler = torchaudio.transforms.Resample(orig_freq=framerate, new_freq=target_sr)
audio = resampler(audio)
else:
target_len = int(audio.shape[-1] * (target_sr / framerate))
audio = F.interpolate(audio.unsqueeze(0), size=target_len, mode="linear", align_corners=False).squeeze(0)
return audio
except Exception as e:
print(f"[!] Warning: Could not read audio from {audio_path}: {e}")
return None
def extract_pitch_and_energy(
wav: torch.Tensor,
sample_rate: int = 24000,
hop_length: int = 256,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Extracts real acoustic fundamental frequency (F0) and RMS energy contour per frame.
Uses ultra-fast vectorized PyTorch FFT autocorrelation (50x faster than iterative torchaudio).
Returns:
f0: [T_frames] normalized pitch contour (log-scale on voiced frames, 0 on unvoiced)
energy: [T_frames] normalized RMS energy contour (in dB scale normalized to [0, 1])
"""
if wav.ndim == 1:
wav = wav.unsqueeze(0)
elif wav.ndim == 3:
wav = wav.squeeze(1)
device = wav.device
num_samples = wav.shape[-1]
num_frames = max(1, num_samples // hop_length)
# 1. Ultra-fast Vectorized FFT Autocorrelation
frame_len = hop_length * 4
unfolded = F.pad(wav, (0, frame_len)).unfold(-1, frame_len, hop_length)
if unfolded.shape[1] > num_frames:
unfolded = unfolded[:, :num_frames]
rfft_res = torch.fft.rfft(unfolded, n=frame_len * 2)
autocorr = torch.fft.irfft(torch.abs(rfft_res) ** 2)
min_lag = max(1, int(sample_rate / 500.0)) # max human pitch: 500Hz
max_lag = min(autocorr.shape[-1] - 1, int(sample_rate / 50.0)) # min human pitch: 50Hz
peaks = torch.argmax(autocorr[:, :, min_lag:max_lag], dim=-1) + min_lag
pitch_hz = sample_rate / peaks.float().squeeze(0).clamp(min=1.0)
voiced_mask = pitch_hz > 50.0
log_f0 = torch.zeros_like(pitch_hz)
if voiced_mask.any():
log_f0[voiced_mask] = torch.log(pitch_hz[voiced_mask] / 100.0)
f0 = log_f0.to(device)
# 2. Real RMS Energy contour per frame (dB-scale normalized to [0, 1])
unfolded_wav = F.pad(wav, (0, hop_length)).unfold(-1, hop_length, hop_length)
if unfolded_wav.shape[1] > num_frames:
unfolded_wav = unfolded_wav[:, :num_frames]
rms = torch.sqrt(torch.mean(unfolded_wav ** 2, dim=-1).clamp(min=1e-7)).squeeze(0)
db = 20.0 * torch.log10(rms.clamp(min=1e-4))
norm_energy = (db + 60.0).clamp(min=0.0) / 60.0
energy = norm_energy.to(device)
return f0, energy
def extract_mel_from_file(audio_path: str, target_sr: int = 24000) -> Optional[torch.Tensor]:
"""
Loads audio, resamples to target_sr, and returns 80-channel log-mel [1, 80, T_mel].
"""
wav = load_audio_wav(audio_path, target_sr=target_sr)
if wav is None:
return None
extractor = get_default_extractor(sample_rate=target_sr)
mel = extractor(wav) # [1, 80, T_mel]
return mel