1use std::cmp::Ordering;
4
5use crate::frame::{SpectralFrame, spectral_frames};
6use crate::{AnalysisError, TempoGrid};
7
8pub const FRAME_SECS: f64 = 0.023;
10pub const HOP_SECS: f64 = FRAME_SECS / 2.0;
12pub const MIN_ONSET_GAP_SECS: f64 = 0.05;
14pub const ADAPTIVE_WINDOW_FRAMES: usize = 10;
16pub const THRESHOLD_MULTIPLIER: f64 = 1.5;
17pub const THRESHOLD_DELTA: f64 = 1e-6;
18pub const IOI_BUCKET_SECS: f64 = 0.025;
19pub const IOI_MAX_SECS: f64 = 2.0;
20pub const LEVEL_FRACTION: f64 = 0.03;
22pub const STRIKE_SPAN_SECS: f64 = 0.012;
23pub const JUMP_ENERGY_RATIO: f64 = 4.0;
25
26#[derive(Clone, Copy, Debug, PartialEq)]
27pub struct Onset {
28 pub t_secs: f64,
29 pub strength: f64,
30}
31
32#[derive(Clone, Copy, Debug, PartialEq)]
33pub struct IoiBucket {
34 pub lo_secs: f64,
35 pub hi_secs: f64,
36 pub count: usize,
37}
38
39#[derive(Clone, Debug, PartialEq)]
40pub struct Onsets {
41 pub onsets: Vec<Onset>,
42 pub ioi_histogram: Vec<IoiBucket>,
43 pub onsets_per_bar: Option<Vec<usize>>,
44 pub syncopation_index: Option<f64>,
45 pub resolution_secs: f64,
46}
47
48pub fn detect(
49 samples: &[f32],
50 sample_rate: f64,
51 start_secs: f64,
52 tempo: Option<TempoGrid>,
53) -> Result<Onsets, AnalysisError> {
54 let duration_secs = samples.len() as f64 / sample_rate;
55 if let Some(t) = tempo {
56 if t.seconds_per_bar.partial_cmp(&0.0) != Some(Ordering::Greater) {
57 return Err(AnalysisError(format!(
58 "a bar of {} seconds is no tempo grid; `onsets` needs a bar of positive length",
59 t.seconds_per_bar
60 )));
61 }
62 if duration_secs / t.seconds_per_bar > samples.len() as f64 {
64 return Err(AnalysisError(format!(
65 "a bar of {:e} seconds divides {} samples into more bars than there are \
66 samples; `onsets` needs a bar at least one sample long",
67 t.seconds_per_bar,
68 samples.len()
69 )));
70 }
71 }
72 let frames = spectral_frames(samples, sample_rate, start_secs, FRAME_SECS, HOP_SECS);
73 let floor = LEVEL_FRACTION * frames.iter().map(|f| total(&f.mags)).fold(0.0, f64::max);
74 let flux = flux_of(&frames, floor);
75 let onsets = pick_peaks(
76 &flux,
77 &frames,
78 &Energy::of(samples, sample_rate, start_secs),
79 floor,
80 );
81 Ok(Onsets {
82 ioi_histogram: ioi_histogram(&onsets),
83 onsets_per_bar: tempo
84 .map(|t| per_bar(&onsets, start_secs, duration_secs, t.seconds_per_bar)),
85 syncopation_index: tempo.map(|t| syncopation(&onsets, t)),
86 onsets,
87 resolution_secs: HOP_SECS,
88 })
89}
90
91fn total(mags: &[f64]) -> f64 {
92 mags.iter().sum()
93}
94
95fn flux_of(frames: &[SpectralFrame], floor: f64) -> Vec<f64> {
98 let opens_on_sound = frames
99 .first()
100 .is_some_and(|f| floor > 0.0 && total(&f.mags) > floor);
101 (0..frames.len())
102 .map(|i| match i {
103 0 => match opens_on_sound {
104 true => total(&frames[0].mags),
105 false => 0.0,
106 },
107 _ => spectral_flux(&frames[i - 1].mags, &frames[i].mags),
108 })
109 .collect()
110}
111
112fn spectral_flux(prev: &[f64], curr: &[f64]) -> f64 {
113 prev.iter().zip(curr).map(|(&p, &c)| (c - p).max(0.0)).sum()
114}
115
116fn pick_peaks(flux: &[f64], frames: &[SpectralFrame], energy: &Energy, floor: f64) -> Vec<Onset> {
119 let mut onsets: Vec<Onset> = Vec::new();
120 for i in 0..flux.len() {
121 let lo = i.saturating_sub(ADAPTIVE_WINDOW_FRAMES);
122 let hi = (i + ADAPTIVE_WINDOW_FRAMES + 1).min(flux.len());
123 let local = &flux[lo..hi];
124 let mean = local.iter().sum::<f64>() / local.len() as f64;
125 let threshold = THRESHOLD_DELTA + THRESHOLD_MULTIPLIER * mean;
126 let reach = (MIN_ONSET_GAP_SECS / HOP_SECS).round().max(1.0) as usize;
128 let near = &flux[i.saturating_sub(reach)..(i + reach + 1).min(flux.len())];
129 if flux[i] <= threshold || flux[i] < floor || near.iter().any(|&v| v > flux[i]) {
130 continue;
131 }
132 let Some(t) = energy.strike(frames[i].t_secs, frames[i].span_secs) else {
133 continue;
134 };
135 if onsets
136 .last()
137 .is_some_and(|o| t - o.t_secs < MIN_ONSET_GAP_SECS)
138 {
139 continue;
140 }
141 onsets.push(Onset {
142 t_secs: t,
143 strength: flux[i],
144 });
145 }
146 onsets
147}
148
149struct Energy {
151 running: Vec<f64>,
152 sample_rate: f64,
153 start_secs: f64,
154 reach: usize,
155 floor: f64,
156}
157
158impl Energy {
159 fn of(samples: &[f32], sample_rate: f64, start_secs: f64) -> Energy {
160 let mut running = Vec::with_capacity(samples.len() + 1);
161 running.push(0.0);
162 for s in samples {
163 let held = running[running.len() - 1];
164 running.push(held + f64::from(*s) * f64::from(*s));
165 }
166 let reach = (STRIKE_SPAN_SECS * sample_rate).round().max(1.0) as usize;
167 let loudest = (0..samples.len())
168 .map(|n| running[(n + reach).min(samples.len())] - running[n])
169 .fold(0.0, f64::max);
170 Energy {
171 running,
172 sample_rate,
173 start_secs,
174 reach,
175 floor: LEVEL_FRACTION * LEVEL_FRACTION * loudest,
176 }
177 }
178
179 fn over(&self, from: usize, to: usize) -> f64 {
181 let last = self.running.len() - 1;
182 self.running[to.min(last)] - self.running[from.min(last)]
183 }
184
185 fn strike(&self, from_secs: f64, span_secs: f64) -> Option<f64> {
188 let at = |secs: f64| {
189 ((secs - self.start_secs) * self.sample_rate)
190 .round()
191 .max(0.0) as usize
192 };
193 let (from, to) = (at(from_secs), at(from_secs + span_secs));
194 let mut best: Option<(f64, usize)> = None;
195 for n in from..to.min(self.running.len() - 1) {
196 let before = self.over(n.saturating_sub(self.reach), n) + self.floor;
197 let after = self.over(n, n + self.reach);
198 if before == 0.0 {
199 continue;
200 }
201 let ratio = after / before;
202 if best.is_none_or(|(r, _)| ratio >= r) {
203 best = Some((ratio, n));
204 }
205 }
206 let (ratio, n) = best?;
207 (ratio >= JUMP_ENERGY_RATIO).then(|| self.start_secs + n as f64 / self.sample_rate)
208 }
209}
210
211fn ioi_histogram(onsets: &[Onset]) -> Vec<IoiBucket> {
212 let n = (IOI_MAX_SECS / IOI_BUCKET_SECS).ceil() as usize + 1;
213 let mut counts = vec![0usize; n];
214 for w in onsets.windows(2) {
215 let ioi = w[1].t_secs - w[0].t_secs;
216 counts[((ioi / IOI_BUCKET_SECS) as usize).min(n - 1)] += 1;
217 }
218 counts
219 .into_iter()
220 .enumerate()
221 .map(|(i, count)| IoiBucket {
222 lo_secs: i as f64 * IOI_BUCKET_SECS,
223 hi_secs: if i + 1 == n {
224 f64::INFINITY
225 } else {
226 (i + 1) as f64 * IOI_BUCKET_SECS
227 },
228 count,
229 })
230 .collect()
231}
232
233fn per_bar(
234 onsets: &[Onset],
235 start_secs: f64,
236 duration_secs: f64,
237 seconds_per_bar: f64,
238) -> Vec<usize> {
239 let bars = ((duration_secs / seconds_per_bar).ceil() as usize).max(1);
240 let mut counts = vec![0usize; bars];
241 for o in onsets {
242 let bar = (((o.t_secs - start_secs) / seconds_per_bar) as usize).min(bars - 1);
243 counts[bar] += 1;
244 }
245 counts
246}
247
248fn metrical_weight(k: u32, subdivisions: u32, beats_per_bar: u32) -> f64 {
252 if k == 0 {
253 return 1.0;
254 }
255 if subdivisions.is_power_of_two() {
256 return (k.trailing_zeros() + 1) as f64 / (subdivisions.trailing_zeros() + 1) as f64;
257 }
258 let per_beat = (subdivisions / beats_per_bar).max(1);
259 if k.is_multiple_of(per_beat) { 0.5 } else { 0.0 }
260}
261
262fn syncopation(onsets: &[Onset], tempo: TempoGrid) -> f64 {
265 if onsets.is_empty() {
266 return 0.0;
267 }
268 let beats_per_bar = tempo.beats_per_bar.round().max(1.0) as u32;
269 let subdivisions = (beats_per_bar * 4).max(1);
270 let total: f64 = onsets
271 .iter()
272 .map(|o| {
273 let phase = o.t_secs.rem_euclid(tempo.seconds_per_bar) / tempo.seconds_per_bar;
274 let k = (phase * subdivisions as f64).round() as u32 % subdivisions;
275 1.0 - metrical_weight(k, subdivisions, beats_per_bar)
276 })
277 .sum();
278 total / onsets.len() as f64
279}
280
281#[cfg(test)]
282mod tests {
283 use super::*;
284
285 fn click_track(sr: f64, secs: f64, gap_secs: f64) -> Vec<f32> {
286 let n = (secs * sr) as usize;
287 let gap = (gap_secs * sr) as usize;
288 let mut out = vec![0.0f32; n];
289 let mut at = 0;
290 while at + 8 < n {
291 for (i, s) in out[at..at + 8].iter_mut().enumerate() {
292 *s = (1.0 - i as f32 / 8.0) * if i % 2 == 0 { 1.0 } else { -1.0 };
293 }
294 at += gap;
295 }
296 out
297 }
298
299 #[test]
300 fn evenly_spaced_clicks_are_found_at_roughly_their_own_spacing() {
301 let sr = 44100.0;
302 let gap = 0.25;
303 let samples = click_track(sr, 4.0, gap);
304 let found = detect(&samples, sr, 0.0, None).expect("no tempo, no grid");
305 assert!(found.onsets.len() >= 12, "{}", found.onsets.len());
306 for w in found.onsets.windows(2) {
307 let ioi = w[1].t_secs - w[0].t_secs;
308 assert!((ioi - gap).abs() < 0.03, "ioi {ioi} far from {gap}");
309 }
310 }
311
312 #[test]
313 fn silence_holds_no_onsets() {
314 let found = detect(&vec![0.0f32; 44100], 44100.0, 0.0, None).expect("no tempo");
315 assert!(found.onsets.is_empty());
316 assert!(found.onsets_per_bar.is_none());
317 assert!(found.syncopation_index.is_none());
318 }
319
320 #[test]
321 fn a_close_double_trigger_on_one_transient_is_suppressed() {
322 let sr = 44100.0;
323 let mut samples = vec![0.0f32; (sr * 0.2) as usize];
324 for (i, s) in samples.iter_mut().enumerate().take(200) {
325 *s = (1.0 - i as f32 / 200.0) * if i % 2 == 0 { 1.0 } else { -1.0 };
326 }
327 let found = detect(&samples, sr, 0.0, None).expect("no tempo, no grid");
328 assert!(found.onsets.len() <= 1, "{:?}", found.onsets);
329 }
330
331 #[test]
332 fn onsets_squarely_on_the_downbeat_read_as_barely_syncopated() {
333 let sr = 44100.0;
334 let tempo = TempoGrid {
335 seconds_per_bar: 2.0,
336 beats_per_bar: 4.0,
337 };
338 let samples = click_track(sr, 8.0, 2.0);
339 let found = detect(&samples, sr, 0.0, Some(tempo)).expect("a two-second bar");
340 assert!(
341 found.syncopation_index.unwrap() < 0.2,
342 "{:?}",
343 found.syncopation_index
344 );
345 assert_eq!(found.onsets_per_bar.as_ref().unwrap().len(), 4);
346 }
347
348 #[test]
349 fn metrical_weight_favours_the_downbeat_over_an_off_sixteenth() {
350 assert!(metrical_weight(0, 16, 4) > metrical_weight(1, 16, 4));
351 assert!(metrical_weight(8, 16, 4) > metrical_weight(1, 16, 4));
352 }
353}