Skip to main content

sva_samples/
stft.rs

1// Concern: the forward and inverse short-time transform between a buffer and frames | Non-concern: holding either (buffer.rs, frames.rs) | IO: (Buffer, w, h) <-> Frames
2
3use std::f64::consts::TAU;
4
5use crate::buffer::Buffer;
6use crate::error::SampleError;
7use crate::fft::{fft, irfft};
8use crate::frames::Frames;
9use crate::label::{Detail, Label, Rule, Source};
10use crate::profile::Profile;
11
12/// Periodic, not symmetric: the symmetric window's duplicated endpoint is what breaks
13/// constant overlap-add at a hop that divides the window.
14pub fn hann_periodic(n: usize) -> Vec<f64> {
15    (0..n)
16        .map(|k| 0.5 - 0.5 * (TAU * k as f64 / n as f64).cos())
17        .collect()
18}
19
20const COLA_TOL: f64 = 1e-12;
21
22pub fn cola_ok(window: &[f64], hop: usize) -> bool {
23    if hop == 0 || hop > window.len() {
24        return false;
25    }
26    let sums: Vec<f64> = (0..hop)
27        .map(|phase| {
28            window
29                .iter()
30                .skip(phase)
31                .step_by(hop)
32                .map(|w| w * w)
33                .sum::<f64>()
34        })
35        .collect();
36    let first = sums[0];
37    first > COLA_TOL
38        && sums
39            .iter()
40            .all(|s| (s - first).abs() <= COLA_TOL * first.max(1.0))
41}
42
43fn frame_count(samples: usize, window: usize, hop: usize) -> usize {
44    let span = samples + 2 * (window - hop);
45    span.div_ceil(hop)
46}
47
48fn start_of(frame: usize, window: usize, hop: usize) -> isize {
49    frame as isize * hop as isize - (window as isize - hop as isize)
50}
51
52pub fn forward(x: &Buffer, window: usize, hop: usize) -> Result<Frames, SampleError> {
53    if !window.is_power_of_two() {
54        return Err(SampleError::WindowNotPowerOfTwo { window });
55    }
56    let w = hann_periodic(window);
57    if !cola_ok(&w, hop) {
58        return Err(SampleError::HopOutsideCola { window, hop });
59    }
60    let samples = x.len();
61    let count = frame_count(samples, window, hop);
62    let mut out = Frames::silence(x.rate, window, hop, x.width, count, samples);
63    out.start = x.start;
64    for c in 0..x.width {
65        let plane = x.plane(c);
66        for frame in 0..count {
67            let start = start_of(frame, window, hop);
68            let mut re = vec![0.0; window];
69            let mut im = vec![0.0; window];
70            for (n, slot) in re.iter_mut().enumerate() {
71                let at = start + n as isize;
72                let s = usize::try_from(at)
73                    .ok()
74                    .and_then(|i| plane.get(i).copied())
75                    .unwrap_or(0.0);
76                *slot = s * w[n];
77            }
78            fft(&mut re, &mut im);
79            for bin in 0..out.bins {
80                out.place(c, frame, bin, re[bin], im[bin]);
81            }
82        }
83    }
84    Ok(out)
85}
86
87/// Accumulates the overlap-added signal and the summed window square in one pass, then
88/// divides one by the other per sample. That makes the leading and trailing edges exact
89/// instead of assuming a steady-state overlap sum.
90pub fn inverse(fr: &Frames, profile: &Profile) -> (Buffer, Label) {
91    let w = hann_periodic(fr.window);
92    let mut planes = Vec::with_capacity(fr.width);
93    for c in 0..fr.width {
94        let mut y = vec![0.0; fr.samples];
95        let mut d = vec![0.0; fr.samples];
96        for frame in 0..fr.frames {
97            let start = start_of(frame, fr.window, fr.hop);
98            let (re, im): (Vec<f64>, Vec<f64>) =
99                (0..fr.bins).map(|bin| fr.at(c, frame, bin)).unzip();
100            let block = irfft(&re, &im, fr.window);
101            for n in 0..fr.window {
102                let Ok(at) = usize::try_from(start + n as isize) else {
103                    continue;
104                };
105                if at >= fr.samples {
106                    break;
107                }
108                y[at] += block[n] * w[n];
109                d[at] += w[n] * w[n];
110            }
111        }
112        for (sample, weight) in y.iter_mut().zip(&d) {
113            if *weight > COLA_TOL {
114                *sample /= weight;
115            }
116        }
117        planes.push(y);
118    }
119    let mut buffer = Buffer::of_planes(fr.rate, planes);
120    buffer.start = fr.start;
121    let source = if fr.edited {
122        Source::Measured
123    } else {
124        Source::Exact
125    };
126    let label = Label::new(
127        source,
128        profile.name,
129        fr.rate,
130        Detail::Roundtrip {
131            rule: Rule::Istft,
132            edited: fr.edited,
133        },
134    );
135    (buffer, label)
136}