1use 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
12pub 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
87pub 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}