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
}