use crate::segmentation::decoder::{
MAX_LOCAL_SPEAKERS, NUM_POWERSET_CLASSES, PowersetClass, PowersetDecoder,
};
use crate::vad::hysteresis::{HysteresisGate, RegionEvent, RegionTracker, TailPolicy};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct BinarizationConfig {
pub onset: f32,
pub offset: f32,
pub min_duration_on: f32,
pub min_duration_off: f32,
}
impl Default for BinarizationConfig {
fn default() -> Self {
Self {
onset: 0.5,
offset: 0.5,
min_duration_on: 0.0,
min_duration_off: 0.0,
}
}
}
#[allow(clippy::needless_range_loop)]
pub fn binarize_frames(
avg_probs: &[[f32; NUM_POWERSET_CLASSES]],
has_data: &[bool],
stride: f32,
cfg: &BinarizationConfig,
) -> (Vec<Option<PowersetClass>>, Vec<f32>) {
let n = avg_probs.len();
let mut speaker_probs = vec![[0.0_f32; MAX_LOCAL_SPEAKERS]; n];
for g in 0..n {
if !has_data[g] {
continue;
}
for c in 0..NUM_POWERSET_CLASSES {
if let Some(class) = PowersetDecoder::class_for_index(c) {
for s in class.speakers() {
speaker_probs[g][s as usize] += avg_probs[g][c];
}
}
}
}
let min_on = (cfg.min_duration_on / stride).round() as usize;
let min_off = (cfg.min_duration_off / stride).round() as usize;
let mut active = vec![[false; MAX_LOCAL_SPEAKERS]; n];
for s in 0..MAX_LOCAL_SPEAKERS {
let mut gate = HysteresisGate::new(cfg.onset, cfg.offset);
let mut tracker = RegionTracker::new(min_off, min_on, TailPolicy::Trim);
for g in 0..n {
if !has_data[g] {
gate.reset();
let event = tracker.reset();
mark_region(&mut active, s, &tracker, event);
continue;
}
let on = gate.update(speaker_probs[g][s]);
let event = tracker.advance(on, g);
mark_region(&mut active, s, &tracker, event);
}
let event = tracker.flush(n);
mark_region(&mut active, s, &tracker, event);
}
let mut classes = Vec::with_capacity(n);
let mut confidences = Vec::with_capacity(n);
for g in 0..n {
if !has_data[g] {
classes.push(None);
confidences.push(0.0);
continue;
}
let mut on: Vec<u8> = (0..MAX_LOCAL_SPEAKERS as u8)
.filter(|&s| active[g][s as usize])
.collect();
if on.len() > 2 {
on.sort_by(|a, b| {
speaker_probs[g][*b as usize].total_cmp(&speaker_probs[g][*a as usize])
});
on.truncate(2);
on.sort_unstable();
}
debug_assert!(
matches!(on.as_slice(), [] | [_] | [_, _]),
"at most two speakers survive the top-2 truncation"
);
let class = PowersetClass::from_speakers(&on);
classes.push(class);
let conf = if on.is_empty() {
avg_probs[g][0]
} else {
on.iter()
.map(|s| speaker_probs[g][*s as usize])
.sum::<f32>()
/ on.len() as f32
};
confidences.push(conf.clamp(0.0, 1.0));
}
(classes, confidences)
}
fn mark_region(
active: &mut [[bool; MAX_LOCAL_SPEAKERS]],
speaker: usize,
tracker: &RegionTracker,
event: Option<RegionEvent>,
) {
if let Some(RegionEvent::End {
start_frame,
end_frame,
}) = event
&& tracker.keeps(start_frame, end_frame)
{
for frame in active.iter_mut().take(end_frame).skip(start_frame) {
frame[speaker] = true;
}
}
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod tests {
use super::*;
fn solo(s: usize, p: f32) -> [f32; NUM_POWERSET_CLASSES] {
let mut row = [0.0; NUM_POWERSET_CLASSES];
row[0] = 1.0 - p;
row[1 + s] = p;
row
}
fn silence() -> [f32; NUM_POWERSET_CLASSES] {
solo(0, 0.0)
}
#[test]
fn default_config_is_plain_thresholding() {
let cfg = BinarizationConfig::default();
assert!((cfg.onset - 0.5).abs() < 1e-6);
assert!((cfg.offset - 0.5).abs() < 1e-6);
assert!((cfg.min_duration_on - 0.0).abs() < 1e-6);
assert!((cfg.min_duration_off - 0.0).abs() < 1e-6);
}
#[test]
fn uncovered_frames_emit_none_and_zero_confidence() {
let avg = vec![silence(), silence(), silence()];
let has_data = vec![true, false, true];
let (classes, confs) =
binarize_frames(&avg, &has_data, 0.01, &BinarizationConfig::default());
assert_eq!(classes.len(), 3);
assert_eq!(confs.len(), 3);
assert_eq!(classes[0], Some(PowersetClass::Silence));
assert_eq!(classes[1], None);
assert_eq!(classes[2], Some(PowersetClass::Silence));
assert!((confs[1] - 0.0).abs() < 1e-6);
}
#[test]
fn empty_input_returns_empty_tracks() {
let (classes, confs) = binarize_frames(&[], &[], 0.01, &BinarizationConfig::default());
assert!(classes.is_empty());
assert!(confs.is_empty());
}
#[test]
fn silence_frame_confidence_is_silence_probability() {
let avg = vec![silence()];
let has_data = vec![true];
let (classes, confs) =
binarize_frames(&avg, &has_data, 0.01, &BinarizationConfig::default());
assert_eq!(classes[0], Some(PowersetClass::Silence));
assert!((confs[0] - 1.0).abs() < 1e-6, "got {}", confs[0]);
}
#[test]
fn single_active_speaker_maps_to_solo_class_with_mean_confidence() {
let avg = vec![silence(), solo(1, 0.9), solo(1, 0.7), silence()];
let has_data = vec![true; 4];
let (classes, confs) =
binarize_frames(&avg, &has_data, 0.01, &BinarizationConfig::default());
assert_eq!(
classes,
vec![
Some(PowersetClass::Silence),
Some(PowersetClass::Speaker(1)),
Some(PowersetClass::Speaker(1)),
Some(PowersetClass::Silence),
]
);
assert!((confs[1] - 0.9).abs() < 1e-6, "got {}", confs[1]);
assert!((confs[2] - 0.7).abs() < 1e-6, "got {}", confs[2]);
}
#[test]
fn two_active_speakers_map_to_pair_class() {
let mut frame = [0.0; NUM_POWERSET_CLASSES];
frame[0] = 0.1;
frame[1] = 0.5; frame[2] = 0.4; let avg = vec![silence(), frame, silence()];
let has_data = vec![true; 3];
let cfg = BinarizationConfig {
onset: 0.2,
..BinarizationConfig::default()
};
let (classes, confs) = binarize_frames(&avg, &has_data, 0.01, &cfg);
assert_eq!(classes[1], Some(PowersetClass::Pair(0, 1)));
assert!((confs[1] - 0.45).abs() < 1e-6, "got {}", confs[1]);
}
#[test]
fn three_active_speakers_truncate_to_top_two_by_probability() {
let mut frame = [0.0; NUM_POWERSET_CLASSES];
frame[1] = 0.4; frame[2] = 0.35; frame[3] = 0.25; let avg = vec![frame];
let has_data = vec![true];
let cfg = BinarizationConfig {
onset: 0.2,
offset: 0.2,
..Default::default()
};
let (classes, confs) = binarize_frames(&avg, &has_data, 0.01, &cfg);
assert_eq!(classes[0], Some(PowersetClass::Pair(0, 1)));
assert!((confs[0] - 0.375).abs() < 1e-6, "got {}", confs[0]);
}
#[test]
fn hysteresis_holds_speaker_on_through_dip_above_offset() {
let avg = vec![
silence(),
solo(0, 0.7), solo(0, 0.45), solo(0, 0.1), ];
let has_data = vec![true; 4];
let cfg = BinarizationConfig {
onset: 0.6,
offset: 0.4,
..Default::default()
};
let (classes, _) = binarize_frames(&avg, &has_data, 1.0, &cfg);
assert_eq!(
classes,
vec![
Some(PowersetClass::Silence),
Some(PowersetClass::Speaker(0)),
Some(PowersetClass::Speaker(0)),
Some(PowersetClass::Silence),
]
);
}
#[test]
fn short_gap_is_bridged_by_min_duration_off() {
let avg = vec![
silence(),
solo(0, 0.9),
solo(0, 0.0),
solo(0, 0.9),
silence(),
];
let has_data = vec![true; 5];
let cfg = BinarizationConfig {
min_duration_off: 2.0, ..Default::default()
};
let (classes, _) = binarize_frames(&avg, &has_data, 1.0, &cfg);
assert_eq!(
classes,
vec![
Some(PowersetClass::Silence),
Some(PowersetClass::Speaker(0)),
Some(PowersetClass::Speaker(0)), Some(PowersetClass::Speaker(0)),
Some(PowersetClass::Silence),
]
);
}
#[test]
fn short_active_blip_is_dropped_by_min_duration_on() {
let avg = vec![silence(), solo(0, 0.9), solo(0, 0.9), silence()];
let has_data = vec![true; 4];
let cfg = BinarizationConfig {
min_duration_on: 3.0, ..Default::default()
};
let (classes, _) = binarize_frames(&avg, &has_data, 1.0, &cfg);
assert_eq!(classes, vec![Some(PowersetClass::Silence); 4]);
}
#[test]
fn coverage_hole_hard_closes_region_instead_of_bridging() {
let avg = vec![
solo(0, 0.9),
solo(0, 0.9),
silence(),
solo(1, 0.9),
solo(1, 0.9),
];
let has_data = vec![true, true, false, true, true];
let cfg = BinarizationConfig {
min_duration_off: 10.0, ..Default::default()
};
let (classes, confs) = binarize_frames(&avg, &has_data, 1.0, &cfg);
assert_eq!(
classes,
vec![
Some(PowersetClass::Speaker(0)),
Some(PowersetClass::Speaker(0)),
None,
Some(PowersetClass::Speaker(1)),
Some(PowersetClass::Speaker(1)),
]
);
assert!((confs[2] - 0.0).abs() < 1e-6);
}
#[test]
fn speakers_are_independent_tracks() {
let avg = vec![silence(), solo(0, 0.9), solo(2, 0.9), silence()];
let has_data = vec![true; 4];
let (classes, _) = binarize_frames(&avg, &has_data, 0.01, &BinarizationConfig::default());
assert_eq!(
classes,
vec![
Some(PowersetClass::Silence),
Some(PowersetClass::Speaker(0)),
Some(PowersetClass::Speaker(2)),
Some(PowersetClass::Silence),
]
);
}
}