Skip to main content

combs_media/
audio.rs

1//! Audio preprocessing for speech-to-text models: WAV decoding, linear
2//! resampling to 16 kHz, and Whisper-style log-mel spectrograms.
3//!
4//! The mel pipeline mirrors the Whisper reference exactly: 400-sample
5//! **periodic** Hann window, hop 160, centered STFT with reflect padding
6//! and the final frame dropped, power spectrum, 80 **slaney-scale /
7//! slaney-normalized** mel filters over 0–8 kHz, `log10` clamped at 1e-10,
8//! a global dynamic-range floor of `max − 8`, then `(x + 4) / 4`. Any
9//! deviation shifts every downstream encoder activation, so the constants
10//! here are contract, not preference.
11
12use realfft::num_complex::Complex;
13use realfft::{RealFftPlanner, RealToComplex};
14use std::sync::Arc;
15
16use crate::{MediaError, Result};
17
18/// Sample rate every speech model input is resampled to.
19pub const SAMPLE_RATE: usize = 16_000;
20/// STFT window length (25 ms at 16 kHz).
21pub const N_FFT: usize = 400;
22/// STFT hop (10 ms at 16 kHz).
23pub const HOP_LENGTH: usize = 160;
24/// Mel bins.
25pub const N_MELS: usize = 80;
26/// Samples in one 30 s model window.
27pub const CHUNK_SAMPLES: usize = 30 * SAMPLE_RATE;
28/// Frames one 30 s window produces (CHUNK_SAMPLES / HOP_LENGTH).
29pub const CHUNK_FRAMES: usize = CHUNK_SAMPLES / HOP_LENGTH;
30
31/// Decodes a WAV payload to mono f32 samples in [-1, 1].
32///
33/// v1 scope: 16-bit PCM (format tag 1), mono or stereo (stereo is averaged
34/// to mono), any sample rate (resample separately). Unknown RIFF chunks
35/// are skipped, including the padding byte after odd-sized chunks.
36pub 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; // (tag, channels, rate, bits)
46    let mut data: Option<(usize, usize)> = None; // (offset, len)
47    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        // Chunks are word-aligned; odd sizes carry one padding byte.
72        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
102/// Linear-interpolation resampler. Documented v1 simplification: no
103/// low-pass filter, which is adequate for speech into a 16 kHz pipeline;
104/// a windowed-sinc resampler can replace this without changing callers.
105pub 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
123/// Zero-pads or truncates to exactly `len` samples (the 30 s model window).
124pub 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
130/// Whisper log-mel extractor. Owns the FFT plan, window, and filterbank;
131/// build once, reuse per utterance.
132pub struct LogMel {
133    fft: Arc<dyn RealToComplex<f32>>,
134    window: Vec<f32>,
135    /// Row-major `[N_MELS, N_FFT/2 + 1]` slaney filterbank.
136    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        // Periodic Hann: divide by N, not N-1 (torch.hann_window default).
150        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    /// Computes the log-mel spectrogram of `samples` (16 kHz mono).
164    /// Returns row-major `[N_MELS, n_frames]` with
165    /// `n_frames = samples.len() / HOP_LENGTH` (centered STFT, last frame
166    /// dropped — the Whisper convention).
167    pub fn compute(&self, samples: &[f32]) -> (Vec<f32>, usize) {
168        let n = samples.len();
169        let half = N_FFT / 2;
170        // Reflect padding (no edge repeat): [s[half]..s[1]] + s + [s[n-2]..s[n-half-1]]
171        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        // Whisper drops the final STFT frame.
186        let n_frames = n_frames_full.saturating_sub(1);
187        let n_bins = half + 1;
188
189        // Power spectrum per frame, then mel projection.
190        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        // log10 clamp, global dynamic-range floor, (x + 4) / 4.
222        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
237/// Slaney mel scale (librosa `htk=False`): linear below 1 kHz, log above.
238fn 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
254/// Builds the `[N_MELS, N_FFT/2 + 1]` slaney-normalized triangular
255/// filterbank over 0–8 kHz, matching librosa/transformers for Whisper.
256fn 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    // N_MELS + 2 corner points, uniform in mel space.
261    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    /// Builds a minimal WAV byte stream around raw i16 frames.
284    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()); // size (unchecked)
289        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); // word-alignment padding for the odd size
295        }
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()); // PCM
299        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        // Frames: (1000, 3000) -> 2000; (-2000, -4000) -> -3000.
327        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; // format tag -> IEEE float
339        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        // A linear ramp stays a linear ramp under linear interpolation.
349        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        // Filter peaks must be strictly ordered in frequency.
387        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        // After (x+4)/4 with a max-8 floor, the span is exactly ≤ 2.
425        assert!(max - min <= 2.0 + 1e-5);
426        assert!(max < 3.0);
427    }
428}