Download tests/test_spectrogram.py from replicate/mel_spectrogram: direct link, hf CLI and curl.
- Browser
- Download file 3.63 kB
-
https://huggingface.co/replicate/mel_spectrogram/resolve/main/tests/test_spectrogram.py
- Command line
-
hf download hf://replicate/mel_spectrogram/tests/test_spectrogram.py
-
curl -L -o test_spectrogram.py https://huggingface.co/replicate/mel_spectrogram/resolve/main/tests/test_spectrogram.py
3.63 kB
| from typing import Tuple | |
| import pytest | |
| import torch | |
| import mel_spectrogram | |
| def torch_mel_spectrogram( | |
| samples: torch.Tensor, | |
| filters: torch.Tensor, | |
| fft_size: int, | |
| fft_step: int, | |
| n_frames: int, | |
| ) -> torch.Tensor: | |
| """Native PyTorch implementation of mel spectrogram generation.""" | |
| # Calculate frames based on CUDA implementation's frame count | |
| frames = [] | |
| # Frame the audio to match CUDA implementation | |
| for i in range(n_frames): | |
| start = i * fft_step | |
| end = min(start + fft_size, samples.size(0)) | |
| if end - start < fft_size: | |
| # Zero-pad if needed | |
| frame = torch.zeros(fft_size, device=samples.device) | |
| frame[: end - start] = samples[start:end] | |
| frames.append(frame) | |
| else: | |
| frames.append(samples[start:end]) | |
| # Stack frames | |
| frames = torch.stack(frames) # [n_frames, fft_size] | |
| # Apply Hanning window | |
| window = torch.hann_window(fft_size, device=samples.device) | |
| windowed = frames * window | |
| # Compute FFT | |
| fft_complex = torch.fft.rfft(windowed, dim=1) # [n_frames, n_fft] | |
| fft_magnitudes = torch.abs(fft_complex) # [n_frames, n_fft] | |
| # Apply mel filterbank | |
| mel_spec = torch.matmul(fft_magnitudes, filters.T) # [n_frames, n_mel] | |
| # Log-scale and normalization | |
| eps = 1e-10 | |
| mel_spec = torch.log10(torch.clamp(mel_spec, min=eps)) | |
| # Normalize like the CUDA implementation | |
| max_val = mel_spec.max() | |
| min_val = max_val - 8.0 | |
| mel_spec = torch.clamp(mel_spec, min=min_val) / 4.0 + 1.0 | |
| # Transpose to match CUDA implementation output shape [n_mel, n_frames] | |
| return mel_spec.T | |
| def test_mel_spectrogram( | |
| n_samples: int, n_mel: int, fft_size: int, fft_step: int, seed: int | |
| ) -> None: | |
| """Test the CUDA mel_spectrogram function against a native PyTorch implementation.""" | |
| if not torch.cuda.is_available(): | |
| pytest.skip("CUDA not available") | |
| # Set seed for reproducibility | |
| torch.manual_seed(seed) | |
| torch.cuda.manual_seed(seed) | |
| device = "cuda" | |
| # Calculate the number of frames and FFT bins | |
| n_frames = n_samples // fft_step | |
| n_fft = 1 + fft_size // 2 # Number of unique FFT bins | |
| # Create tensors with the correct shapes | |
| samples = torch.randn(n_samples, dtype=torch.float32, device=device) | |
| filters = torch.abs(torch.randn(n_mel, n_fft, dtype=torch.float32, device=device)) | |
| filters = filters / filters.sum(dim=1, keepdim=True) # Normalize filters | |
| # Run the CUDA implementation | |
| cuda_output = torch.zeros(n_mel, n_frames, dtype=torch.float32, device=device) | |
| mel_spectrogram.mel_spectrogram( | |
| cuda_output, | |
| samples, | |
| filters, | |
| fft_size, | |
| fft_step, | |
| ) | |
| # Get actual frame count from CUDA output | |
| actual_frames = cuda_output.shape[1] | |
| # Run the PyTorch implementation with matched frame count | |
| torch_output = torch_mel_spectrogram( | |
| samples, filters, fft_size, fft_step, actual_frames | |
| ) | |
| # Check shapes match | |
| assert cuda_output.shape == torch_output.shape | |
| # Compare results (allowing for numerical differences) | |
| mae = torch.abs(cuda_output - torch_output).mean().item() | |
| print(f"Mean Absolute Error: {mae}") | |
| # Simple validation | |
| assert not torch.isnan(cuda_output).any(), "CUDA output contains NaN values" | |
| assert not torch.allclose(cuda_output, torch.zeros_like(cuda_output)) | |