| 11 | /// Returns tensor of shape `[1, n_mels, time_frames]`. |
| 12 | #[allow(clippy::too_many_arguments)] |
| 13 | pub fn mel_spectrogram( |
| 14 | samples: &[f32], |
| 15 | n_fft: usize, |
| 16 | hop_length: usize, |
| 17 | n_mels: usize, |
| 18 | sample_rate: usize, |
| 19 | device: &Device, |
| 20 | dtype: DType, |
| 21 | ) -> Result<Tensor> { |
| 22 | let stft = compute_stft(samples, n_fft, hop_length); |
| 23 | let n_freq = n_fft / 2 + 1; |
| 24 | let n_frames = stft.len() / n_freq; |
| 25 | |
| 26 | // Build mel filterbank |
| 27 | let mel_filters = mel_filterbank(n_mels, n_freq, sample_rate, n_fft); |
| 28 | |
| 29 | // Apply mel filterbank: [n_mels, n_freq] @ [n_freq, n_frames] -> [n_mels, n_frames] |
| 30 | let mut mel = vec![0.0f32; n_mels * n_frames]; |
| 31 | for m in 0..n_mels { |
| 32 | for t in 0..n_frames { |
| 33 | let mut sum = 0.0f32; |
| 34 | for f in 0..n_freq { |
| 35 | sum += mel_filters[m * n_freq + f] * stft[t * n_freq + f]; |
| 36 | } |
| 37 | mel[m * n_frames + t] = sum; |
| 38 | } |
| 39 | } |
| 40 | |
| 41 | // Log mel (with floor to avoid log(0)) |
| 42 | for v in &mut mel { |
| 43 | *v = (*v).max(1e-10).ln(); |
| 44 | } |
| 45 | |
| 46 | // Convert to tensor [1, n_mels, n_frames] |
| 47 | let mel_tensor = Tensor::from_vec(mel, (1, n_mels, n_frames), device)?.to_dtype(dtype)?; |
| 48 | Ok(mel_tensor) |
| 49 | } |
| 50 | |
| 51 | /// Compute STFT magnitude squared, returned as flat Vec [n_frames * n_freq]. |
| 52 | fn compute_stft(samples: &[f32], n_fft: usize, hop_length: usize) -> Vec<f32> { |