use super::ids::{SpeakerId, SpeakerIdRemap};
use super::measures::TimeRange;
use serde::{Deserialize, Serialize};
pub fn remap_segments(segments: &mut [Segment], remap: &SpeakerIdRemap) {
for seg in segments.iter_mut() {
if let Some(spk) = seg.speaker {
seg.speaker = Some(remap.remap(spk));
}
}
}
pub fn remap_turns(turns: &mut [SpeakerTurn], remap: &SpeakerIdRemap) {
for turn in turns.iter_mut() {
turn.speaker = remap.remap(turn.speaker);
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Segment {
pub time: TimeRange,
pub speaker: Option<SpeakerId>,
pub confidence: Option<f32>,
}
fn default_turn_stable() -> bool {
true
}
fn is_true(v: &bool) -> bool {
*v
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SpeakerTurn {
pub speaker: SpeakerId,
pub time: TimeRange,
#[serde(skip_serializing_if = "Option::is_none")]
pub text: Option<String>,
#[serde(default = "default_turn_stable", skip_serializing_if = "is_true")]
pub stable: bool,
}
impl SpeakerTurn {
pub fn new(speaker: SpeakerId, time: TimeRange) -> Self {
Self {
speaker,
time,
text: None,
stable: true,
}
}
pub fn with_stability(speaker: SpeakerId, time: TimeRange, stable: bool) -> Self {
Self {
speaker,
time,
text: None,
stable,
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct WordAlignment {
pub word: String,
pub time: TimeRange,
pub speaker: Option<SpeakerId>,
pub confidence: f32,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub interpolated: bool,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Word {
pub word: String,
pub time: TimeRange,
pub confidence: f32,
}
#[derive(Debug, Clone, PartialEq, Default, Serialize, Deserialize)]
pub struct Transcript {
pub words: Vec<Word>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
pub struct AudioMeta {
pub duration_secs: f64,
pub sample_rate: u32,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
pub struct Provenance {
pub version: String,
pub profile: String,
pub segmenter: String,
pub embedder: String,
pub clusterer: String,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct SpeakerSummary {
pub label: String,
pub id: u32,
pub total_speech_s: f64,
pub turn_count: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub embedding: Option<Vec<f32>>,
}
fn default_schema_version() -> String {
"diarization-result-v1".to_owned()
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct DiarizationResult {
pub segments: Vec<Segment>,
pub turns: Vec<SpeakerTurn>,
pub num_speakers: usize,
#[serde(default = "default_schema_version")]
pub schema_version: String,
#[serde(default)]
pub audio: AudioMeta,
#[serde(default)]
pub provenance: Provenance,
#[serde(default)]
pub speakers: Vec<SpeakerSummary>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub exclusive_turns: Vec<SpeakerTurn>,
}
impl DiarizationResult {
pub fn new(segments: Vec<Segment>, turns: Vec<SpeakerTurn>, num_speakers: usize) -> Self {
let speakers = speaker_summaries(&turns);
Self {
segments,
turns,
num_speakers,
schema_version: default_schema_version(),
audio: AudioMeta::default(),
provenance: Provenance {
version: env!("CARGO_PKG_VERSION").to_owned(),
..Provenance::default()
},
speakers,
exclusive_turns: Vec::new(),
}
}
pub fn with_audio(mut self, duration_secs: f64, sample_rate: u32) -> Self {
self.audio = AudioMeta {
duration_secs,
sample_rate,
};
self
}
pub fn with_provenance(mut self, provenance: Provenance) -> Self {
let version = if provenance.version.is_empty() {
self.provenance.version.clone()
} else {
provenance.version.clone()
};
self.provenance = Provenance {
version,
..provenance
};
self
}
pub fn with_exclusive(mut self) -> Self {
self.exclusive_turns = exclusive_turns(&self.turns);
self
}
pub fn with_speaker_embeddings(mut self, embeddings: &[(SpeakerId, Vec<f32>)]) -> Self {
for sp in &mut self.speakers {
if let Some((_, emb)) = embeddings.iter().find(|(id, _)| id.0 == sp.id) {
let mut v = emb.clone();
crate::utils::l2_normalize(&mut v);
sp.embedding = Some(v);
}
}
self
}
}
fn speaker_summaries(turns: &[SpeakerTurn]) -> Vec<SpeakerSummary> {
use std::collections::BTreeMap;
let mut agg: BTreeMap<u32, (f64, usize)> = BTreeMap::new();
for t in turns {
let e = agg.entry(t.speaker.0).or_insert((0.0, 0));
e.0 += t.time.duration();
e.1 += 1;
}
agg.into_iter()
.map(|(id, (total, count))| SpeakerSummary {
label: SpeakerId(id).to_string(),
id,
total_speech_s: total,
turn_count: count,
embedding: None,
})
.collect()
}
impl TimeRange {
pub(crate) const FRAME_GRID_RESOLUTION_SECS: f64 = 0.01;
pub(crate) const MAX_GRID_FRAMES: usize = 24 * 3600 * 100;
pub(crate) fn grid_frame_count(max_time: f64) -> usize {
((max_time / Self::FRAME_GRID_RESOLUTION_SECS).ceil() as usize + 1)
.min(Self::MAX_GRID_FRAMES)
}
pub(crate) fn grid_frame_range(&self) -> (usize, usize) {
(
(self.start / Self::FRAME_GRID_RESOLUTION_SECS) as usize,
(self.end / Self::FRAME_GRID_RESOLUTION_SECS).ceil() as usize,
)
}
}
pub(crate) const EXCLUSIVE_FRAME_SECS: f64 = TimeRange::FRAME_GRID_RESOLUTION_SECS;
pub fn exclusive_turns(turns: &[SpeakerTurn]) -> Vec<SpeakerTurn> {
if turns.is_empty() {
return Vec::new();
}
let max_time = turns.iter().map(|t| t.time.end).fold(0.0f64, f64::max);
if !max_time.is_finite() || max_time <= 0.0 {
return Vec::new();
}
let n_frames = TimeRange::grid_frame_count(max_time);
let mut best: Vec<Option<(u32, f64)>> = vec![None; n_frames];
for turn in turns {
if !turn.time.start.is_finite()
|| !turn.time.end.is_finite()
|| turn.time.end <= turn.time.start
{
continue;
}
let dur = turn.time.duration();
let (start_f, end_f) = turn.time.grid_frame_range();
for frame in best.iter_mut().take(end_f.min(n_frames)).skip(start_f) {
match frame {
None => *frame = Some((turn.speaker.0, dur)),
Some((spk, best_dur)) => {
if dur > *best_dur + f64::EPSILON
|| ((dur - *best_dur).abs() <= f64::EPSILON && turn.speaker.0 < *spk)
{
*frame = Some((turn.speaker.0, dur));
}
}
}
}
}
let mut out: Vec<SpeakerTurn> = Vec::new();
let mut i = 0usize;
while i < n_frames {
let Some((spk, _)) = best[i] else {
i += 1;
continue;
};
let start = i;
i += 1;
while i < n_frames {
match best[i] {
Some((s, _)) if s == spk => i += 1,
_ => break,
}
}
out.push(SpeakerTurn::new(
SpeakerId(spk),
TimeRange {
start: start as f64 * EXCLUSIVE_FRAME_SECS,
end: i as f64 * EXCLUSIVE_FRAME_SECS,
},
));
}
out
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod tests {
use super::*;
fn seg(start: f64, end: f64, speaker: Option<u32>) -> Segment {
Segment {
time: TimeRange { start, end },
speaker: speaker.map(SpeakerId),
confidence: Some(0.9),
}
}
fn turn_range(id: u32, start: f64, end: f64) -> SpeakerTurn {
SpeakerTurn::new(SpeakerId(id), TimeRange { start, end })
}
fn sample_result() -> DiarizationResult {
let turns = vec![
turn_range(1, 2.0, 3.5),
turn_range(0, 0.0, 1.0),
turn_range(0, 1.0, 2.0),
];
DiarizationResult::new(vec![seg(0.0, 1.0, Some(0))], turns, 2)
}
#[test]
fn remap_segments_and_turns_apply_mapping() {
let remap = SpeakerIdRemap::from_mapping(vec![(SpeakerId(1), SpeakerId(0))]).unwrap();
let mut segments = vec![
seg(0.0, 1.0, Some(1)),
seg(1.0, 2.0, None),
seg(2.0, 3.0, Some(5)),
];
remap_segments(&mut segments, &remap);
assert_eq!(segments[0].speaker, Some(SpeakerId(0)));
assert_eq!(segments[1].speaker, None); assert_eq!(segments[2].speaker, Some(SpeakerId(5)));
let mut turns = vec![turn_range(1, 0.0, 1.0), turn_range(5, 1.0, 2.0)];
remap_turns(&mut turns, &remap);
assert_eq!(turns[0].speaker, SpeakerId(0));
assert_eq!(turns[1].speaker, SpeakerId(5));
}
#[test]
fn speaker_turn_constructors() {
let t = SpeakerTurn::new(
SpeakerId(1),
TimeRange {
start: 0.0,
end: 1.0,
},
);
assert!(t.stable);
assert!(t.text.is_none());
let p = SpeakerTurn::with_stability(
SpeakerId(1),
TimeRange {
start: 0.0,
end: 1.0,
},
false,
);
assert!(!p.stable);
}
#[test]
fn speaker_turn_serde_stability_compatibility() {
let t = turn_range(1, 0.0, 1.0);
let json = serde_json::to_string(&t).unwrap();
assert!(!json.contains("stable"), "{json}");
assert!(!json.contains("text"), "{json}");
assert_eq!(serde_json::from_str::<SpeakerTurn>(&json).unwrap(), t);
let p = SpeakerTurn::with_stability(
SpeakerId(1),
TimeRange {
start: 0.0,
end: 1.0,
},
false,
);
let json = serde_json::to_string(&p).unwrap();
assert!(json.contains("\"stable\":false"), "{json}");
assert_eq!(serde_json::from_str::<SpeakerTurn>(&json).unwrap(), p);
let legacy = r#"{"speaker":2,"time":{"start":1.0,"end":2.0}}"#;
let back: SpeakerTurn = serde_json::from_str(legacy).unwrap();
assert!(back.stable);
assert_eq!(back.speaker, SpeakerId(2));
assert!(back.text.is_none());
}
#[test]
fn word_alignment_interpolated_omitted_when_false() {
let wa = WordAlignment {
word: "hi".into(),
time: TimeRange {
start: 0.0,
end: 0.4,
},
speaker: Some(SpeakerId(0)),
confidence: 0.9,
interpolated: false,
};
let json = serde_json::to_string(&wa).unwrap();
assert!(!json.contains("interpolated"), "{json}");
assert_eq!(serde_json::from_str::<WordAlignment>(&json).unwrap(), wa);
let wa = WordAlignment {
interpolated: true,
..wa
};
let json = serde_json::to_string(&wa).unwrap();
assert!(json.contains("\"interpolated\":true"), "{json}");
assert_eq!(serde_json::from_str::<WordAlignment>(&json).unwrap(), wa);
}
#[test]
fn segment_and_transcript_serde_roundtrips() {
let s = seg(0.1, 0.9, Some(3));
let json = serde_json::to_string(&s).unwrap();
assert_eq!(serde_json::from_str::<Segment>(&json).unwrap(), s);
let s = Segment {
speaker: None,
confidence: None,
..s
};
let json = serde_json::to_string(&s).unwrap();
assert_eq!(serde_json::from_str::<Segment>(&json).unwrap(), s);
let t = Transcript {
words: vec![Word {
word: "hi".into(),
time: TimeRange {
start: 0.0,
end: 0.3,
},
confidence: 0.8,
}],
};
let json = serde_json::to_string(&t).unwrap();
assert_eq!(serde_json::from_str::<Transcript>(&json).unwrap(), t);
assert!(Transcript::default().words.is_empty());
}
#[test]
fn result_new_builds_sorted_speaker_rollup() {
let r = sample_result();
assert_eq!(r.schema_version, "diarization-result-v1");
assert!(!r.provenance.version.is_empty());
assert_eq!(r.speakers.len(), 2);
assert_eq!(r.speakers[0].id, 0);
assert_eq!(r.speakers[0].label, "SPEAKER_00");
assert!((r.speakers[0].total_speech_s - 2.0).abs() < 1e-9);
assert_eq!(r.speakers[0].turn_count, 2);
assert_eq!(r.speakers[1].id, 1);
assert_eq!(r.speakers[1].label, "SPEAKER_01");
assert!((r.speakers[1].total_speech_s - 1.5).abs() < 1e-9);
assert_eq!(r.speakers[1].turn_count, 1);
assert!(r.speakers.iter().all(|s| s.embedding.is_none()));
}
#[test]
fn result_builders_attach_audio_and_provenance() {
let r = sample_result().with_audio(10.0, 16000);
assert_eq!(r.audio.duration_secs, 10.0);
assert_eq!(r.audio.sample_rate, 16000);
let keep = sample_result().with_provenance(Provenance {
profile: "balanced".into(),
embedder: "wespeaker".into(),
..Provenance::default()
});
assert_eq!(keep.provenance.profile, "balanced");
assert_eq!(keep.provenance.embedder, "wespeaker");
assert!(!keep.provenance.version.is_empty());
let over = sample_result().with_provenance(Provenance {
version: "9.9.9".into(),
..Provenance::default()
});
assert_eq!(over.provenance.version, "9.9.9");
}
#[test]
fn speaker_summary_embedding_omitted_when_absent() {
let r = sample_result();
let json = serde_json::to_string(&r.speakers[0]).unwrap();
assert!(!json.contains("embedding"), "{json}");
}
#[test]
fn with_speaker_embeddings_normalizes_and_matches_by_id() {
let r = sample_result().with_speaker_embeddings(&[
(SpeakerId(0), vec![3.0, 4.0]), (SpeakerId(9), vec![1.0, 0.0]), ]);
let e0 = r.speakers[0].embedding.as_ref().unwrap();
let norm = e0.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-5, "norm={norm}");
assert!((e0[0] - 0.6).abs() < 1e-6);
assert!(r.speakers[1].embedding.is_none());
}
#[test]
fn result_serde_roundtrip_and_legacy_defaults() {
let r = sample_result().with_audio(5.0, 16000).with_exclusive();
let json = serde_json::to_string(&r).unwrap();
assert_eq!(serde_json::from_str::<DiarizationResult>(&json).unwrap(), r);
let legacy = r#"{"segments":[],"turns":[],"num_speakers":0}"#;
let back: DiarizationResult = serde_json::from_str(legacy).unwrap();
assert_eq!(back.schema_version, "diarization-result-v1");
assert_eq!(back.audio, AudioMeta::default());
assert_eq!(back.provenance, Provenance::default());
assert!(back.speakers.is_empty());
assert!(back.exclusive_turns.is_empty());
}
#[test]
fn grid_frame_helpers() {
assert_eq!(TimeRange::grid_frame_count(1.0), 101);
assert_eq!(
TimeRange::grid_frame_count(1e12),
TimeRange::MAX_GRID_FRAMES
);
let tr = TimeRange {
start: 0.0,
end: 0.03,
};
assert_eq!(tr.grid_frame_range(), (0, 3));
let neg = TimeRange {
start: -1.0,
end: 0.02,
};
assert_eq!(neg.grid_frame_range().0, 0);
}
#[test]
fn exclusive_turns_empty_and_degenerate_inputs() {
assert!(exclusive_turns(&[]).is_empty());
let inf = turn_range(0, 0.0, f64::INFINITY);
assert!(exclusive_turns(&[inf]).is_empty());
let zero = turn_range(0, 1.0, 1.0);
let neg = turn_range(1, 2.0, 1.0);
let nan = turn_range(2, f64::NAN, 3.0);
assert!(exclusive_turns(&[zero, neg, nan]).is_empty());
}
#[test]
fn exclusive_turns_single_turn_collapses_frames() {
let out = exclusive_turns(&[turn_range(0, 0.0, 0.03)]);
assert_eq!(out.len(), 1);
assert_eq!(out[0].speaker, SpeakerId(0));
assert!(out[0].time.start.abs() < 1e-9);
assert!((out[0].time.end - 0.03).abs() <= EXCLUSIVE_FRAME_SECS + 1e-9);
}
#[test]
fn exclusive_turns_overlap_picks_longer_covering_turn() {
let turns = vec![
turn_range(0, 0.0, 1.0), turn_range(1, 0.2, 0.4), ];
let out = exclusive_turns(&turns);
assert_eq!(out.len(), 1);
assert_eq!(out[0].speaker, SpeakerId(0));
}
#[test]
fn exclusive_turns_tie_breaks_to_smaller_speaker_id() {
let turns = vec![turn_range(2, 0.0, 0.5), turn_range(1, 0.0, 0.5)];
let out = exclusive_turns(&turns);
assert_eq!(out.len(), 1);
assert_eq!(out[0].speaker, SpeakerId(1));
}
#[test]
fn exclusive_turns_silence_splits_same_speaker() {
let turns = vec![turn_range(0, 0.0, 0.1), turn_range(0, 0.5, 0.7)];
let out = exclusive_turns(&turns);
assert_eq!(out.len(), 2);
assert_eq!(out[0].speaker, SpeakerId(0));
assert_eq!(out[1].speaker, SpeakerId(0));
assert!((out[1].time.start - 0.5).abs() <= EXCLUSIVE_FRAME_SECS + 1e-9);
}
#[test]
fn exclusive_turns_adjacent_speakers_stay_contiguous() {
let turns = vec![turn_range(0, 0.0, 0.2), turn_range(1, 0.2, 0.4)];
let out = exclusive_turns(&turns);
assert_eq!(out.len(), 2);
assert_eq!(out[0].speaker, SpeakerId(0));
assert_eq!(out[1].speaker, SpeakerId(1));
assert!((out[0].time.end - out[1].time.start).abs() < 1e-9);
}
#[test]
fn with_exclusive_fills_and_is_idempotent() {
let r = sample_result().with_exclusive();
assert!(!r.exclusive_turns.is_empty());
assert_eq!(r.turns.len(), 3);
let again = r.clone().with_exclusive();
assert_eq!(r.exclusive_turns, again.exclusive_turns);
}
}