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 RISE_FRACTION: f64 = 0.1;
21pub const LEVEL_FRACTION: f64 = 0.03;
23pub const MAX_RISE_SECS: f64 = 0.012;
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 = whole_frames(samples, sample_rate, start_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(&flux, &frames, samples, sample_rate, start_secs, floor);
76 Ok(Onsets {
77 ioi_histogram: ioi_histogram(&onsets),
78 onsets_per_bar: tempo
79 .map(|t| per_bar(&onsets, start_secs, duration_secs, t.seconds_per_bar)),
80 syncopation_index: tempo.map(|t| syncopation(&onsets, t)),
81 onsets,
82 resolution_secs: HOP_SECS,
83 })
84}
85
86fn whole_frames(samples: &[f32], sample_rate: f64, start_secs: f64) -> Vec<SpectralFrame> {
87 let mut frames = spectral_frames(samples, sample_rate, start_secs, FRAME_SECS, HOP_SECS);
88 frames.truncate(frames.iter().take_while(|f| f.filled).count());
89 frames
90}
91
92fn total(mags: &[f64]) -> f64 {
93 mags.iter().sum()
94}
95
96fn flux_of(frames: &[SpectralFrame], floor: f64) -> Vec<f64> {
99 let opens_on_sound = frames
100 .first()
101 .is_some_and(|f| floor > 0.0 && total(&f.mags) > floor);
102 (0..frames.len())
103 .map(|i| match i {
104 0 => match opens_on_sound {
105 true => total(&frames[0].mags),
106 false => 0.0,
107 },
108 _ => spectral_flux(&frames[i - 1].mags, &frames[i].mags),
109 })
110 .collect()
111}
112
113fn spectral_flux(prev: &[f64], curr: &[f64]) -> f64 {
114 prev.iter().zip(curr).map(|(&p, &c)| (c - p).max(0.0)).sum()
115}
116
117fn pick_peaks(
119 flux: &[f64],
120 frames: &[SpectralFrame],
121 samples: &[f32],
122 sample_rate: f64,
123 start_secs: f64,
124 floor: f64,
125) -> Vec<Onset> {
126 let mut onsets: Vec<Onset> = Vec::new();
127 for i in 0..flux.len() {
128 let lo = i.saturating_sub(ADAPTIVE_WINDOW_FRAMES);
129 let hi = (i + ADAPTIVE_WINDOW_FRAMES + 1).min(flux.len());
130 let local = &flux[lo..hi];
131 let mean = local.iter().sum::<f64>() / local.len() as f64;
132 let threshold = THRESHOLD_DELTA + THRESHOLD_MULTIPLIER * mean;
133 let reach = (MIN_ONSET_GAP_SECS / HOP_SECS).round().max(1.0) as usize;
135 let near = &flux[i.saturating_sub(reach)..(i + reach + 1).min(flux.len())];
136 if flux[i] <= threshold || near.iter().any(|&v| v > flux[i]) {
137 continue;
138 }
139 let span = frames[i].span_secs;
140 if flux[i] < floor || climb_secs(flux, i, floor) > span + HOP_SECS {
141 continue;
142 }
143 let from = frames[i].t_secs - start_secs;
145 let Some((at, rise_secs)) = attack(samples, sample_rate, from, span) else {
146 continue;
147 };
148 if rise_secs > MAX_RISE_SECS {
149 continue;
150 }
151 let t = start_secs + at;
152 if onsets
153 .last()
154 .is_some_and(|o| t - o.t_secs < MIN_ONSET_GAP_SECS)
155 {
156 continue;
157 }
158 onsets.push(Onset {
159 t_secs: t,
160 strength: flux[i],
161 });
162 }
163 onsets
164}
165
166fn climb_secs(flux: &[f64], i: usize, floor: f64) -> f64 {
168 let mut from = i;
169 while from > 0 && flux[from - 1] >= floor {
170 from -= 1;
171 }
172 (i + 1 - from) as f64 * HOP_SECS
173}
174
175fn attack(samples: &[f32], sample_rate: f64, from_secs: f64, span_secs: f64) -> Option<(f64, f64)> {
177 let at = |secs: f64| ((secs * sample_rate).round().max(0.0) as usize).min(samples.len());
178 let (from, to) = (at(from_secs), at(from_secs + span_secs));
179 let held = samples.get(from..to)?;
180 let peak = held.iter().fold(0.0f64, |a, s| a.max(f64::from(*s).abs()));
181 if peak == 0.0 {
182 return None;
183 }
184 let above = |bar: f64| held.iter().position(|s| f64::from(*s).abs() >= bar);
185 let struck = above(peak * RISE_FRACTION)?;
186 let crest = above(peak * (1.0 - RISE_FRACTION)).unwrap_or(struck);
187 Some((
188 (from + struck) as f64 / sample_rate,
189 crest.saturating_sub(struck) as f64 / sample_rate,
190 ))
191}
192
193fn ioi_histogram(onsets: &[Onset]) -> Vec<IoiBucket> {
194 let n = (IOI_MAX_SECS / IOI_BUCKET_SECS).ceil() as usize + 1;
195 let mut counts = vec![0usize; n];
196 for w in onsets.windows(2) {
197 let ioi = w[1].t_secs - w[0].t_secs;
198 counts[((ioi / IOI_BUCKET_SECS) as usize).min(n - 1)] += 1;
199 }
200 counts
201 .into_iter()
202 .enumerate()
203 .map(|(i, count)| IoiBucket {
204 lo_secs: i as f64 * IOI_BUCKET_SECS,
205 hi_secs: if i + 1 == n {
206 f64::INFINITY
207 } else {
208 (i + 1) as f64 * IOI_BUCKET_SECS
209 },
210 count,
211 })
212 .collect()
213}
214
215fn per_bar(
216 onsets: &[Onset],
217 start_secs: f64,
218 duration_secs: f64,
219 seconds_per_bar: f64,
220) -> Vec<usize> {
221 let bars = ((duration_secs / seconds_per_bar).ceil() as usize).max(1);
222 let mut counts = vec![0usize; bars];
223 for o in onsets {
224 let bar = (((o.t_secs - start_secs) / seconds_per_bar) as usize).min(bars - 1);
225 counts[bar] += 1;
226 }
227 counts
228}
229
230fn metrical_weight(k: u32, subdivisions: u32, beats_per_bar: u32) -> f64 {
234 if k == 0 {
235 return 1.0;
236 }
237 if subdivisions.is_power_of_two() {
238 return (k.trailing_zeros() + 1) as f64 / (subdivisions.trailing_zeros() + 1) as f64;
239 }
240 let per_beat = (subdivisions / beats_per_bar).max(1);
241 if k.is_multiple_of(per_beat) { 0.5 } else { 0.0 }
242}
243
244fn syncopation(onsets: &[Onset], tempo: TempoGrid) -> f64 {
247 if onsets.is_empty() {
248 return 0.0;
249 }
250 let beats_per_bar = tempo.beats_per_bar.round().max(1.0) as u32;
251 let subdivisions = (beats_per_bar * 4).max(1);
252 let total: f64 = onsets
253 .iter()
254 .map(|o| {
255 let phase = o.t_secs.rem_euclid(tempo.seconds_per_bar) / tempo.seconds_per_bar;
256 let k = (phase * subdivisions as f64).round() as u32 % subdivisions;
257 1.0 - metrical_weight(k, subdivisions, beats_per_bar)
258 })
259 .sum();
260 total / onsets.len() as f64
261}
262
263#[cfg(test)]
264mod tests {
265 use super::*;
266
267 fn click_track(sr: f64, secs: f64, gap_secs: f64) -> Vec<f32> {
268 let n = (secs * sr) as usize;
269 let gap = (gap_secs * sr) as usize;
270 let mut out = vec![0.0f32; n];
271 let mut at = 0;
272 while at + 8 < n {
273 for (i, s) in out[at..at + 8].iter_mut().enumerate() {
274 *s = (1.0 - i as f32 / 8.0) * if i % 2 == 0 { 1.0 } else { -1.0 };
275 }
276 at += gap;
277 }
278 out
279 }
280
281 #[test]
282 fn evenly_spaced_clicks_are_found_at_roughly_their_own_spacing() {
283 let sr = 44100.0;
284 let gap = 0.25;
285 let samples = click_track(sr, 4.0, gap);
286 let found = detect(&samples, sr, 0.0, None).expect("no tempo, no grid");
287 assert!(found.onsets.len() >= 12, "{}", found.onsets.len());
288 for w in found.onsets.windows(2) {
289 let ioi = w[1].t_secs - w[0].t_secs;
290 assert!((ioi - gap).abs() < 0.03, "ioi {ioi} far from {gap}");
291 }
292 }
293
294 #[test]
295 fn silence_holds_no_onsets() {
296 let found = detect(&vec![0.0f32; 44100], 44100.0, 0.0, None).expect("no tempo");
297 assert!(found.onsets.is_empty());
298 assert!(found.onsets_per_bar.is_none());
299 assert!(found.syncopation_index.is_none());
300 }
301
302 #[test]
303 fn a_close_double_trigger_on_one_transient_is_suppressed() {
304 let sr = 44100.0;
305 let mut samples = vec![0.0f32; (sr * 0.2) as usize];
306 for (i, s) in samples.iter_mut().enumerate().take(200) {
307 *s = (1.0 - i as f32 / 200.0) * if i % 2 == 0 { 1.0 } else { -1.0 };
308 }
309 let found = detect(&samples, sr, 0.0, None).expect("no tempo, no grid");
310 assert!(found.onsets.len() <= 1, "{:?}", found.onsets);
311 }
312
313 #[test]
314 fn onsets_squarely_on_the_downbeat_read_as_barely_syncopated() {
315 let sr = 44100.0;
316 let tempo = TempoGrid {
317 seconds_per_bar: 2.0,
318 beats_per_bar: 4.0,
319 };
320 let samples = click_track(sr, 8.0, 2.0);
321 let found = detect(&samples, sr, 0.0, Some(tempo)).expect("a two-second bar");
322 assert!(
323 found.syncopation_index.unwrap() < 0.2,
324 "{:?}",
325 found.syncopation_index
326 );
327 assert_eq!(found.onsets_per_bar.as_ref().unwrap().len(), 4);
328 }
329
330 #[test]
331 fn metrical_weight_favours_the_downbeat_over_an_off_sixteenth() {
332 assert!(metrical_weight(0, 16, 4) > metrical_weight(1, 16, 4));
333 assert!(metrical_weight(8, 16, 4) > metrical_weight(1, 16, 4));
334 }
335}