1use realfft::num_complex::Complex;
13use realfft::{RealFftPlanner, RealToComplex};
14use std::sync::Arc;
15
16use crate::{MediaError, Result};
17
18pub const SAMPLE_RATE: usize = 16_000;
20pub const N_FFT: usize = 400;
22pub const HOP_LENGTH: usize = 160;
24pub const N_MELS: usize = 80;
26pub const CHUNK_SAMPLES: usize = 30 * SAMPLE_RATE;
28pub const CHUNK_FRAMES: usize = CHUNK_SAMPLES / HOP_LENGTH;
30
31pub fn decode_wav(bytes: &[u8]) -> Result<(Vec<f32>, u32)> {
37 let err = |m: &str| MediaError::AudioDecode(m.to_string());
38 if bytes.len() < 12 || &bytes[0..4] != b"RIFF" || &bytes[8..12] != b"WAVE" {
39 return Err(err("not a RIFF/WAVE payload"));
40 }
41 let u16_at = |o: usize| u16::from_le_bytes([bytes[o], bytes[o + 1]]);
42 let u32_at = |o: usize| u32::from_le_bytes([bytes[o], bytes[o + 1], bytes[o + 2], bytes[o + 3]]);
43
44 let mut pos = 12usize;
45 let mut fmt: Option<(u16, u16, u32, u16)> = None; let mut data: Option<(usize, usize)> = None; while pos + 8 <= bytes.len() {
48 let id = &bytes[pos..pos + 4];
49 let size = u32_at(pos + 4) as usize;
50 let body = pos + 8;
51 if body + size > bytes.len() {
52 return Err(err("chunk overruns file"));
53 }
54 match id {
55 b"fmt " => {
56 if size < 16 {
57 return Err(err("fmt chunk too short"));
58 }
59 fmt = Some((
60 u16_at(body),
61 u16_at(body + 2),
62 u32_at(body + 4),
63 u16_at(body + 14),
64 ));
65 }
66 b"data" => {
67 data = Some((body, size));
68 }
69 _ => {}
70 }
71 pos = body + size + (size & 1);
73 }
74
75 let (tag, channels, rate, bits) = fmt.ok_or_else(|| err("missing fmt chunk"))?;
76 let (off, len) = data.ok_or_else(|| err("missing data chunk"))?;
77 if tag != 1 {
78 return Err(err("only PCM (format tag 1) is supported"));
79 }
80 if bits != 16 {
81 return Err(err("only 16-bit samples are supported"));
82 }
83 if channels == 0 || channels > 2 {
84 return Err(err("only mono or stereo is supported"));
85 }
86 let ch = channels as usize;
87 let frame_bytes = 2 * ch;
88 let n_frames = len / frame_bytes;
89 let mut out = Vec::with_capacity(n_frames);
90 for f in 0..n_frames {
91 let base = off + f * frame_bytes;
92 let mut acc = 0.0f32;
93 for c in 0..ch {
94 let s = i16::from_le_bytes([bytes[base + 2 * c], bytes[base + 2 * c + 1]]);
95 acc += f32::from(s) / 32768.0;
96 }
97 out.push(acc / ch as f32);
98 }
99 Ok((out, rate))
100}
101
102pub fn resample_linear(samples: &[f32], from_rate: u32, to_rate: u32) -> Vec<f32> {
106 if from_rate == to_rate || samples.is_empty() {
107 return samples.to_vec();
108 }
109 let n_out = ((samples.len() as u64 * to_rate as u64) / from_rate as u64) as usize;
110 let step = from_rate as f64 / to_rate as f64;
111 let mut out = Vec::with_capacity(n_out);
112 for i in 0..n_out {
113 let pos = i as f64 * step;
114 let i0 = pos as usize;
115 let frac = (pos - i0 as f64) as f32;
116 let a = samples[i0.min(samples.len() - 1)];
117 let b = samples[(i0 + 1).min(samples.len() - 1)];
118 out.push(a + (b - a) * frac);
119 }
120 out
121}
122
123pub fn pad_or_trim(samples: &[f32], len: usize) -> Vec<f32> {
125 let mut out = samples.to_vec();
126 out.resize(len, 0.0);
127 out
128}
129
130pub struct LogMel {
133 fft: Arc<dyn RealToComplex<f32>>,
134 window: Vec<f32>,
135 filters: Vec<f32>,
137}
138
139impl Default for LogMel {
140 fn default() -> Self {
141 Self::new()
142 }
143}
144
145impl LogMel {
146 pub fn new() -> Self {
147 let mut planner = RealFftPlanner::<f32>::new();
148 let fft = planner.plan_fft_forward(N_FFT);
149 let window: Vec<f32> = (0..N_FFT)
151 .map(|i| {
152 let x = core::f32::consts::TAU * i as f32 / N_FFT as f32;
153 0.5 * (1.0 - x.cos())
154 })
155 .collect();
156 LogMel {
157 fft,
158 window,
159 filters: mel_filterbank(),
160 }
161 }
162
163 pub fn compute(&self, samples: &[f32]) -> (Vec<f32>, usize) {
168 let n = samples.len();
169 let half = N_FFT / 2;
170 let mut padded = Vec::with_capacity(n + N_FFT);
172 for i in (1..=half).rev() {
173 padded.push(samples[i.min(n.saturating_sub(1))]);
174 }
175 padded.extend_from_slice(samples);
176 for i in 2..=(half + 1) {
177 padded.push(samples[n.saturating_sub(i)]);
178 }
179
180 let n_frames_full = if padded.len() >= N_FFT {
181 1 + (padded.len() - N_FFT) / HOP_LENGTH
182 } else {
183 0
184 };
185 let n_frames = n_frames_full.saturating_sub(1);
187 let n_bins = half + 1;
188
189 let mut frame = vec![0.0f32; N_FFT];
191 let mut spectrum = vec![Complex::new(0.0f32, 0.0f32); n_bins];
192 let mut power = vec![0.0f32; n_bins * n_frames];
193 let mut scratch = self.fft.make_scratch_vec();
194 for f in 0..n_frames {
195 let start = f * HOP_LENGTH;
196 for i in 0..N_FFT {
197 frame[i] = padded[start + i] * self.window[i];
198 }
199 self.fft
200 .process_with_scratch(&mut frame, &mut spectrum, &mut scratch)
201 .expect("fft length is fixed");
202 for (k, c) in spectrum.iter().enumerate() {
203 power[k * n_frames + f] = c.re * c.re + c.im * c.im;
204 }
205 }
206
207 let mut mel = vec![0.0f32; N_MELS * n_frames];
208 for m in 0..N_MELS {
209 for k in 0..n_bins {
210 let w = self.filters[m * n_bins + k];
211 if w != 0.0 {
212 let row = &power[k * n_frames..(k + 1) * n_frames];
213 let out = &mut mel[m * n_frames..(m + 1) * n_frames];
214 for f in 0..n_frames {
215 out[f] += w * row[f];
216 }
217 }
218 }
219 }
220
221 let mut max_val = f32::MIN;
223 for v in mel.iter_mut() {
224 *v = v.max(1e-10).log10();
225 if *v > max_val {
226 max_val = *v;
227 }
228 }
229 let floor = max_val - 8.0;
230 for v in mel.iter_mut() {
231 *v = (v.max(floor) + 4.0) / 4.0;
232 }
233 (mel, n_frames)
234 }
235}
236
237fn hz_to_mel(f: f32) -> f32 {
239 if f < 1000.0 {
240 f * 3.0 / 200.0
241 } else {
242 15.0 + 27.0 * (f / 1000.0).ln() / 6.4f32.ln()
243 }
244}
245
246fn mel_to_hz(m: f32) -> f32 {
247 if m < 15.0 {
248 m * 200.0 / 3.0
249 } else {
250 1000.0 * (6.4f32.ln() * (m - 15.0) / 27.0).exp()
251 }
252}
253
254fn mel_filterbank() -> Vec<f32> {
257 let n_bins = N_FFT / 2 + 1;
258 let fmax = SAMPLE_RATE as f32 / 2.0;
259 let mel_max = hz_to_mel(fmax);
260 let corners: Vec<f32> = (0..N_MELS + 2)
262 .map(|i| mel_to_hz(mel_max * i as f32 / (N_MELS + 1) as f32))
263 .collect();
264 let mut fb = vec![0.0f32; N_MELS * n_bins];
265 for m in 0..N_MELS {
266 let (lo, mid, hi) = (corners[m], corners[m + 1], corners[m + 2]);
267 let norm = 2.0 / (hi - lo);
268 for k in 0..n_bins {
269 let f = k as f32 * SAMPLE_RATE as f32 / N_FFT as f32;
270 let rising = (f - lo) / (mid - lo);
271 let falling = (hi - f) / (hi - mid);
272 let w = rising.min(falling).max(0.0);
273 fb[m * n_bins + k] = w * norm;
274 }
275 }
276 fb
277}
278
279#[cfg(test)]
280mod tests {
281 use super::*;
282
283 fn wav_bytes(channels: u16, rate: u32, samples: &[i16], extra_chunk: bool) -> Vec<u8> {
285 let data_len = samples.len() * 2;
286 let mut out = Vec::new();
287 out.extend_from_slice(b"RIFF");
288 out.extend_from_slice(&0u32.to_le_bytes()); out.extend_from_slice(b"WAVE");
290 if extra_chunk {
291 out.extend_from_slice(b"LIST");
292 out.extend_from_slice(&3u32.to_le_bytes());
293 out.extend_from_slice(b"abc");
294 out.push(0); }
296 out.extend_from_slice(b"fmt ");
297 out.extend_from_slice(&16u32.to_le_bytes());
298 out.extend_from_slice(&1u16.to_le_bytes()); out.extend_from_slice(&channels.to_le_bytes());
300 out.extend_from_slice(&rate.to_le_bytes());
301 out.extend_from_slice(&(rate * u32::from(channels) * 2).to_le_bytes());
302 out.extend_from_slice(&(channels * 2).to_le_bytes());
303 out.extend_from_slice(&16u16.to_le_bytes());
304 out.extend_from_slice(b"data");
305 out.extend_from_slice(&(data_len as u32).to_le_bytes());
306 for s in samples {
307 out.extend_from_slice(&s.to_le_bytes());
308 }
309 out
310 }
311
312 #[test]
313 fn wav_mono_roundtrip() {
314 let bytes = wav_bytes(1, 16_000, &[0, 16384, -16384, 32767], false);
315 let (samples, rate) = decode_wav(&bytes).unwrap();
316 assert_eq!(rate, 16_000);
317 assert_eq!(samples.len(), 4);
318 assert!((samples[0]).abs() < 1e-6);
319 assert!((samples[1] - 0.5).abs() < 1e-4);
320 assert!((samples[2] + 0.5).abs() < 1e-4);
321 assert!(samples[3] > 0.999);
322 }
323
324 #[test]
325 fn wav_stereo_averages_and_skips_chunks() {
326 let bytes = wav_bytes(2, 44_100, &[1000, 3000, -2000, -4000], true);
328 let (samples, rate) = decode_wav(&bytes).unwrap();
329 assert_eq!(rate, 44_100);
330 assert_eq!(samples.len(), 2);
331 assert!((samples[0] - 2000.0 / 32768.0).abs() < 1e-6);
332 assert!((samples[1] + 3000.0 / 32768.0).abs() < 1e-6);
333 }
334
335 #[test]
336 fn wav_rejects_non_pcm() {
337 let mut bytes = wav_bytes(1, 16_000, &[0, 0], false);
338 bytes[20] = 3; assert!(decode_wav(&bytes).is_err());
340 }
341
342 #[test]
343 fn resample_identity_and_halving() {
344 let s: Vec<f32> = (0..100).map(|i| i as f32).collect();
345 assert_eq!(resample_linear(&s, 16_000, 16_000), s);
346 let half = resample_linear(&s, 32_000, 16_000);
347 assert_eq!(half.len(), 50);
348 assert!((half[10] - 20.0).abs() < 1e-4);
350 }
351
352 #[test]
353 fn pad_and_trim() {
354 let s = vec![1.0f32; 10];
355 let padded = pad_or_trim(&s, 16);
356 assert_eq!(padded.len(), 16);
357 assert_eq!(padded[9], 1.0);
358 assert_eq!(padded[10], 0.0);
359 assert_eq!(pad_or_trim(&s, 4).len(), 4);
360 }
361
362 #[test]
363 fn hann_window_is_periodic() {
364 let lm = LogMel::new();
365 assert!(lm.window[0].abs() < 1e-7);
366 assert!((lm.window[N_FFT / 2] - 1.0).abs() < 1e-6);
367 for k in 1..N_FFT {
368 assert!(
369 (lm.window[k] - lm.window[N_FFT - k]).abs() < 1e-6,
370 "periodic Hann symmetry broke at {k}"
371 );
372 }
373 }
374
375 #[test]
376 fn filterbank_shape_and_coverage() {
377 let fb = mel_filterbank();
378 let n_bins = N_FFT / 2 + 1;
379 assert_eq!(fb.len(), N_MELS * n_bins);
380 for m in 0..N_MELS {
381 let row = &fb[m * n_bins..(m + 1) * n_bins];
382 let sum: f32 = row.iter().sum();
383 assert!(sum > 0.0, "mel filter {m} is empty");
384 assert!(row.iter().all(|w| *w >= 0.0));
385 }
386 let peak = |m: usize| {
388 let row = &fb[m * n_bins..(m + 1) * n_bins];
389 row.iter()
390 .enumerate()
391 .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
392 .unwrap()
393 .0
394 };
395 assert!(peak(0) < peak(N_MELS / 2));
396 assert!(peak(N_MELS / 2) < peak(N_MELS - 1));
397 }
398
399 #[test]
400 fn log_mel_frame_counts() {
401 let lm = LogMel::new();
402 let clip: Vec<f32> = (0..8000)
403 .map(|i| (core::f32::consts::TAU * 440.0 * i as f32 / 16_000.0).sin() * 0.1)
404 .collect();
405 let (mel, frames) = lm.compute(&clip);
406 assert_eq!(frames, 50);
407 assert_eq!(mel.len(), N_MELS * 50);
408 assert!(mel.iter().all(|v| v.is_finite()));
409
410 let (mel30, frames30) = lm.compute(&pad_or_trim(&clip, CHUNK_SAMPLES));
411 assert_eq!(frames30, CHUNK_FRAMES);
412 assert_eq!(mel30.len(), N_MELS * CHUNK_FRAMES);
413 }
414
415 #[test]
416 fn log_mel_range_is_normalized() {
417 let lm = LogMel::new();
418 let clip: Vec<f32> = (0..16_000)
419 .map(|i| (core::f32::consts::TAU * 1000.0 * i as f32 / 16_000.0).sin() * 0.5)
420 .collect();
421 let (mel, _) = lm.compute(&clip);
422 let max = mel.iter().cloned().fold(f32::MIN, f32::max);
423 let min = mel.iter().cloned().fold(f32::MAX, f32::min);
424 assert!(max - min <= 2.0 + 1e-5);
426 assert!(max < 3.0);
427 }
428}