Skip to main content

sva_analysis/experimental/
masking.rs

1// Concern: per-band signal-to-mask ratio, subtracting the target back out of the mix | Non-concern: onset detection (onsets.rs), the FFT (magnitudes) | IO: (target, mix) -> Masking or AnalysisError
2
3use std::ops::Range;
4
5use sva_samples::{magnitudes, third_octave_edges};
6
7use crate::AnalysisError;
8use crate::decibels::to_db;
9use crate::stable::onsets::Onset;
10
11/// The span after an onset a gated window scores: long enough for the attack and early
12/// sustain, short enough to stay that onset's own rather than sliding into the next.
13pub const GATE_WINDOW_SECS: f64 = 0.1;
14
15#[derive(Clone, Copy, Debug, PartialEq)]
16pub struct MaskingBand {
17    pub lo_hz: f64,
18    pub hi_hz: f64,
19    pub target_db: f64,
20    pub against_db: f64,
21    /// `target_db - against_db`; a silent side leaves no finite ratio and reads `null`. The
22    /// two levels beside it say which side was silent.
23    pub smr_db: f64,
24}
25
26#[derive(Clone, Debug, PartialEq)]
27pub struct Masking {
28    pub gated: bool,
29    pub frames_scored: usize,
30    pub bands: Vec<MaskingBand>,
31}
32
33/// `combined` must be the buffer `target` was mixed additively into (a bus, `song`); this
34/// subtracts `target` back out of it before measuring, so the target is never counted as its
35/// own masker. `gate` scores only the target's own onset windows rather than the whole buffer.
36pub fn analyze(
37    target: &[f32],
38    combined: &[f32],
39    sample_rate: f64,
40    start_secs: f64,
41    gate: Option<&[Onset]>,
42) -> Result<Masking, AnalysisError> {
43    if target.len() != combined.len() {
44        return Err(AnalysisError(format!(
45            "target is {} samples and combined is {}; they must share one window",
46            target.len(),
47            combined.len()
48        )));
49    }
50    let complement: Vec<f32> = target.iter().zip(combined).map(|(&t, &c)| c - t).collect();
51    let windows: Vec<Range<usize>> = match gate {
52        None => std::iter::once(0..target.len()).collect(),
53        Some(onsets) => onset_windows(onsets, sample_rate, start_secs, target.len()),
54    };
55
56    let edges = third_octave_edges(sample_rate);
57    let target_power = band_power(target, sample_rate, &edges, &windows);
58    let against_power = band_power(&complement, sample_rate, &edges, &windows);
59    let bands = edges
60        .into_iter()
61        .zip(target_power.into_iter().zip(against_power))
62        .map(|((lo_hz, hi_hz), (t, a))| {
63            let target_db = to_db(t.sqrt());
64            let against_db = to_db(a.sqrt());
65            MaskingBand {
66                lo_hz,
67                hi_hz,
68                target_db,
69                against_db,
70                smr_db: target_db - against_db,
71            }
72        })
73        .collect();
74
75    Ok(Masking {
76        gated: gate.is_some(),
77        frames_scored: windows.len(),
78        bands,
79    })
80}
81
82fn onset_windows(
83    onsets: &[Onset],
84    sample_rate: f64,
85    start_secs: f64,
86    len: usize,
87) -> Vec<Range<usize>> {
88    onsets
89        .iter()
90        .filter_map(|o| {
91            let s = (((o.t_secs - start_secs) * sample_rate).round().max(0.0) as usize).min(len);
92            let e = (s + (GATE_WINDOW_SECS * sample_rate).round() as usize).min(len);
93            (e > s).then_some(s..e)
94        })
95        .collect()
96}
97
98fn band_power(
99    samples: &[f32],
100    sample_rate: f64,
101    edges: &[(f64, f64)],
102    windows: &[Range<usize>],
103) -> Vec<f64> {
104    let mut sums = vec![0.0; edges.len()];
105    let mut scored = 0usize;
106    for w in windows {
107        let slice = &samples[w.clone()];
108        if slice.is_empty() {
109            continue;
110        }
111        let frame: Vec<f64> = slice.iter().map(|&s| f64::from(s)).collect();
112        let (mags, bin_hz, ..) = magnitudes(&frame, sample_rate, None);
113        for (b, &(lo, hi)) in edges.iter().enumerate() {
114            sums[b] += mags
115                .iter()
116                .enumerate()
117                .filter(|&(k, _)| {
118                    let hz = k as f64 * bin_hz;
119                    hz >= lo && hz < hi
120                })
121                .map(|(_, m)| m * m)
122                .sum::<f64>();
123        }
124        scored += 1;
125    }
126    if scored > 0 {
127        for s in &mut sums {
128            *s /= scored as f64;
129        }
130    }
131    sums
132}
133
134#[cfg(test)]
135mod tests {
136    use super::*;
137
138    fn tone(hz: f64, sr: f64, secs: f64, amp: f32) -> Vec<f32> {
139        (0..(secs * sr) as usize)
140            .map(|i| amp * (2.0 * std::f64::consts::PI * hz * i as f64 / sr).sin() as f32)
141            .collect()
142    }
143
144    /// `combined` here IS the target, so a correct exclusion leaves nothing to mask it with.
145    #[test]
146    fn a_target_reads_above_0_db_against_its_own_true_complement() {
147        let sr = 44100.0;
148        let target = tone(1000.0, sr, 1.0, 0.5);
149        let combined = target.clone();
150        let m = analyze(&target, &combined, sr, 0.0, None).unwrap();
151        let own_band = m
152            .bands
153            .iter()
154            .find(|b| b.lo_hz <= 1000.0 && 1000.0 < b.hi_hz)
155            .unwrap();
156        assert!(own_band.smr_db > 0.0, "{:?}", own_band);
157        assert!(
158            own_band.against_db < -100.0,
159            "silence left after subtracting the target out of itself: {:?}",
160            own_band
161        );
162    }
163
164    #[test]
165    fn a_target_buried_under_a_much_louder_complement_reads_negative_smr() {
166        let sr = 44100.0;
167        let target = tone(1000.0, sr, 1.0, 0.05);
168        let masker = tone(1000.0, sr, 1.0, 0.9);
169        let combined: Vec<f32> = target.iter().zip(&masker).map(|(&t, &m)| t + m).collect();
170        let m = analyze(&target, &combined, sr, 0.0, None).unwrap();
171        let own_band = m
172            .bands
173            .iter()
174            .find(|b| b.lo_hz <= 1000.0 && 1000.0 < b.hi_hz)
175            .unwrap();
176        assert!(own_band.smr_db < -10.0, "{:?}", own_band);
177    }
178
179    #[test]
180    fn mismatched_lengths_refuse_rather_than_panic() {
181        let err = analyze(&[0.0; 10], &[0.0; 5], 44100.0, 0.0, None).unwrap_err();
182        assert!(err.0.contains("10"));
183    }
184
185    #[test]
186    fn a_gated_window_scores_only_the_span_after_each_onset() {
187        let sr = 44100.0;
188        let target = tone(1000.0, sr, 1.0, 0.5);
189        let combined = target.clone();
190        let onsets = [Onset {
191            t_secs: 0.5,
192            strength: 1.0,
193        }];
194        let gated = analyze(&target, &combined, sr, 0.0, Some(&onsets)).unwrap();
195        assert!(gated.gated);
196        assert_eq!(gated.frames_scored, 1);
197    }
198}