use crate::types::{SpeakerId, SpeakerTurn, TimeRange};
#[cfg(feature = "segmentation")]
const TIME_RANGE_EPS_SECS: f64 = 1e-6;
pub trait Resegmenter: Send + Sync {
fn resegment(&self, inputs: ResegmentInputs<'_>) -> Result<Vec<SpeakerTurn>, ResegmentError>;
}
#[derive(Debug, Clone)]
pub struct ResegmentInputs<'a> {
pub primary_turns: &'a [SpeakerTurn],
pub speaker_centroids: &'a [SpeakerCentroid],
pub overlap_regions: &'a [OverlapRegionInput],
}
#[derive(Debug, Clone, PartialEq)]
pub struct SpeakerCentroid {
pub speaker: SpeakerId,
pub embedding: Vec<f32>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct OverlapRegionInput {
pub time: TimeRange,
pub primary_speaker: SpeakerId,
pub secondary_speaker: Option<SpeakerId>,
pub embedding: Vec<f32>,
}
#[derive(Debug, thiserror::Error)]
pub enum ResegmentError {
#[error("centroid dim mismatch at index {index}: expected {expected}, got {actual}")]
CentroidDimMismatch {
index: usize,
expected: usize,
actual: usize,
},
#[error("overlap embedding dim mismatch at index {index}: expected {expected}, got {actual}")]
OverlapDimMismatch {
index: usize,
expected: usize,
actual: usize,
},
#[error("primary speaker {primary} for overlap region {index} not present in centroids")]
MissingPrimaryCentroid { index: usize, primary: SpeakerId },
}
pub fn compute_centroids(embeddings: &[Vec<f32>], labels: &[usize]) -> Vec<SpeakerCentroid> {
if embeddings.len() != labels.len() || embeddings.is_empty() {
return Vec::new();
}
let mut buckets: std::collections::BTreeMap<usize, Vec<&Vec<f32>>> =
std::collections::BTreeMap::new();
for (emb, &lbl) in embeddings.iter().zip(labels.iter()) {
buckets.entry(lbl).or_default().push(emb);
}
let mut out = Vec::with_capacity(buckets.len());
for (lbl, members) in buckets {
let owned: Vec<Vec<f32>> = members.iter().map(|e| (*e).clone()).collect();
if let Some(mut mean) = crate::utils::mean_vector(&owned) {
crate::utils::l2_normalize(&mut mean);
let id = SpeakerId(lbl as u32);
out.push(SpeakerCentroid {
speaker: id,
embedding: mean,
});
}
}
out.sort_by_key(|c| c.speaker.0);
out
}
#[cfg(feature = "segmentation")]
pub fn extract_overlap_time_ranges(
segments: &[crate::segmentation::RawSegment],
) -> Vec<(TimeRange, u8, u8)> {
let mut pairs: Vec<(TimeRange, u8, u8)> = Vec::new();
for (i, a) in segments.iter().enumerate() {
if !a.is_overlap {
continue;
}
for b in segments.iter().skip(i + 1) {
if !b.is_overlap {
continue;
}
if a.local_speaker_idx == b.local_speaker_idx {
continue;
}
if (a.time.start - b.time.start).abs() > TIME_RANGE_EPS_SECS
|| (a.time.end - b.time.end).abs() > TIME_RANGE_EPS_SECS
{
continue;
}
let (lo, hi) = if a.local_speaker_idx < b.local_speaker_idx {
(a.local_speaker_idx, b.local_speaker_idx)
} else {
(b.local_speaker_idx, a.local_speaker_idx)
};
pairs.push((a.time, lo, hi));
}
}
pairs
}
#[derive(Debug, Clone, Copy)]
pub struct OverlapResegmenter {
threshold: f32,
min_overlap_secs: f32,
}
impl OverlapResegmenter {
pub fn new(threshold: f32, min_overlap_secs: f32) -> Self {
Self {
threshold,
min_overlap_secs: min_overlap_secs.max(0.0),
}
}
pub fn threshold(&self) -> f32 {
self.threshold
}
pub fn min_overlap_secs(&self) -> f32 {
self.min_overlap_secs
}
}
impl Default for OverlapResegmenter {
fn default() -> Self {
Self::new(0.0, 0.1)
}
}
impl Resegmenter for OverlapResegmenter {
fn resegment(&self, inputs: ResegmentInputs<'_>) -> Result<Vec<SpeakerTurn>, ResegmentError> {
let mut out: Vec<SpeakerTurn> = inputs.primary_turns.to_vec();
if inputs.speaker_centroids.len() < 2 || inputs.overlap_regions.is_empty() {
out.sort_by(|a, b| a.time.start.total_cmp(&b.time.start));
return Ok(out);
}
let expected_dim = inputs.speaker_centroids[0].embedding.len();
for (i, c) in inputs.speaker_centroids.iter().enumerate() {
if c.embedding.len() != expected_dim {
return Err(ResegmentError::CentroidDimMismatch {
index: i,
expected: expected_dim,
actual: c.embedding.len(),
});
}
}
for (i, region) in inputs.overlap_regions.iter().enumerate() {
match region.secondary_speaker {
Some(secondary) => {
if region.time.duration() < f64::from(self.min_overlap_secs) {
continue;
}
out.push(SpeakerTurn {
speaker: region.primary_speaker,
time: region.time,
text: None,
stable: true,
});
if secondary != region.primary_speaker {
out.push(SpeakerTurn {
speaker: secondary,
time: region.time,
text: None,
stable: true,
});
}
}
None => {
if region.embedding.len() != expected_dim {
return Err(ResegmentError::OverlapDimMismatch {
index: i,
expected: expected_dim,
actual: region.embedding.len(),
});
}
if !inputs
.speaker_centroids
.iter()
.any(|c| c.speaker == region.primary_speaker)
{
return Err(ResegmentError::MissingPrimaryCentroid {
index: i,
primary: region.primary_speaker,
});
}
if region.time.duration() < f64::from(self.min_overlap_secs) {
continue;
}
let mut best: Option<(SpeakerId, f32)> = None;
for c in inputs.speaker_centroids.iter() {
if c.speaker == region.primary_speaker {
continue;
}
let s = crate::utils::cosine_similarity(®ion.embedding, &c.embedding);
let take = match best {
None => true,
Some((_, b)) => s > b,
};
if take {
best = Some((c.speaker, s));
}
}
if let Some((id, score)) = best
&& score > self.threshold
{
out.push(SpeakerTurn {
speaker: id,
time: region.time,
text: None,
stable: true,
});
}
}
}
}
out.sort_by(|a, b| a.time.start.total_cmp(&b.time.start));
Ok(out)
}
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[path = "trait_tests.rs"]
mod trait_tests;
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[path = "centroid_tests.rs"]
mod centroid_tests;
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[cfg(feature = "segmentation")]
#[path = "overlap_extract_tests.rs"]
mod overlap_extract_tests;
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[path = "resegmenter_tests.rs"]
mod resegmenter_tests;