sva_analysis/experimental/
masking.rs1use 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
11pub 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 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
33pub 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 #[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}