Download dataset/audio_processing.py from ViuAI/ViuAI_TTS_200M: direct link, hf CLI and curl.
- Browser
- Download file 10.7 kB
-
https://huggingface.co/ViuAI/ViuAI_TTS_200M/resolve/main/dataset/audio_processing.py
- Command line
-
hf download hf://ViuAI/ViuAI_TTS_200M/dataset/audio_processing.py
-
curl -L -o audio_processing.py https://huggingface.co/ViuAI/ViuAI_TTS_200M/resolve/main/dataset/audio_processing.py
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 | |