Files
WhispAssist/src-tauri/src/diarization/voiceprint.rs
T

264 lines
9.4 KiB
Rust

//! Voiceprint matching: identifies which diarized speaker cluster is the
//! meeting's own microphone, so it can be auto-labeled "You" instead of a
//! clustered "S1"/"S2" (bug: the mic speaker wasn't reliably first/labeled).
//! Mic and system audio are already summed into one mono stream before
//! diarization ever runs, so the only way to tell them apart afterwards is a
//! voiceprint: a short mic-only sample, captured live, compared by embedding
//! similarity against each cluster's own audio from the finished recording.
//! Runs once per meeting, entirely offline via the same sherpa-onnx
//! speaker-embedding model diarization already uses (ADR-0005).
use crate::models::SpeakerSpan;
use std::collections::HashMap;
use std::path::Path;
/// At least this much clean audio (mic sample or candidate cluster) before an
/// embedding computed from it is trusted at all — a fragment of a word gives
/// an unstable embedding that's as likely to mismatch as match.
const MIN_VOICEPRINT_SAMPLES: usize = 16_000; // 1s @ 16kHz
/// Per-candidate audio is capped so one very long-talking speaker doesn't
/// blow up embedding compute time; a few seconds is already stable.
const MAX_CANDIDATE_SAMPLES: usize = 16_000 * 10;
/// sherpa's own default "is this a match" similarity threshold
/// (`speaker_id::DEFAULT_SIMILARITY_THRESHOLD`) — kept as a local constant so
/// this module doesn't need the `diarization` feature just to state its
/// policy (used by both the real and no-op builds' doc comments/tests).
const SIMILARITY_THRESHOLD: f32 = 0.5;
#[derive(Debug, thiserror::Error)]
pub enum VoiceprintError {
#[error("model load failed: {0}")]
Load(String),
#[error("embedding failed: {0}")]
Embed(String),
#[error("failed to read the recording: {0}")]
Read(String),
}
/// The label -> display-name map to auto-apply after diarization: whichever
/// speaker's audio matches `mic_samples` best -> `"You"`; every other label,
/// in first-appearance order, -> `"Speaker 2"`, `"Speaker 3"`, … An empty map
/// means "couldn't tell" (too little mic audio, no cluster cleared the
/// similarity threshold, embedding model unavailable) — callers leave the
/// existing "S1"/"S2" labels alone rather than guess (FR-SPK-5).
#[cfg(feature = "diarization")]
pub fn match_mic_speaker(
embedding_model: &Path,
mic_samples: &[f32],
wav_path: &Path,
spans: &[SpeakerSpan],
) -> Result<HashMap<String, String>, VoiceprintError> {
if mic_samples.len() < MIN_VOICEPRINT_SAMPLES || spans.is_empty() {
return Ok(HashMap::new());
}
let labels_in_order = first_appearance_order(spans);
let wav_samples = crate::audio::read_wav_mono_16k(wav_path)
.map_err(|e| VoiceprintError::Read(e.to_string()))?;
let mut extractor =
sherpa_rs::speaker_id::EmbeddingExtractor::new(sherpa_rs::speaker_id::ExtractorConfig {
model: embedding_model.to_string_lossy().to_string(),
..Default::default()
})
.map_err(|e| VoiceprintError::Load(e.to_string()))?;
let mic_embedding = extractor
.compute_speaker_embedding(mic_samples.to_vec(), 16_000)
.map_err(|e| VoiceprintError::Embed(e.to_string()))?;
let mut best: Option<(&str, f32)> = None;
for label in &labels_in_order {
let candidate_samples = candidate_audio(&wav_samples, spans, label);
if candidate_samples.len() < MIN_VOICEPRINT_SAMPLES {
continue;
}
let embedding = extractor
.compute_speaker_embedding(candidate_samples, 16_000)
.map_err(|e| VoiceprintError::Embed(e.to_string()))?;
let score = cosine_similarity(&mic_embedding, &embedding);
let is_better = match best {
Some((_, best_score)) => score > best_score,
None => true,
};
if is_better {
best = Some((label, score));
}
}
let Some((mic_label, score)) = best else {
return Ok(HashMap::new());
};
if score < SIMILARITY_THRESHOLD {
return Ok(HashMap::new());
}
Ok(build_name_map(&labels_in_order, mic_label))
}
#[cfg(not(feature = "diarization"))]
pub fn match_mic_speaker(
_embedding_model: &Path,
_mic_samples: &[f32],
_wav_path: &Path,
_spans: &[SpeakerSpan],
) -> Result<HashMap<String, String>, VoiceprintError> {
Ok(HashMap::new())
}
/// Distinct speaker labels in first-appearance order — spans come back from
/// the diarizer already sorted by start time. `pub(crate)` so the Phase 3
/// per-stream path can reuse it to name its merged You + far-side spans.
pub(crate) fn first_appearance_order(spans: &[SpeakerSpan]) -> Vec<String> {
let mut seen = std::collections::HashSet::new();
spans
.iter()
.filter(|s| seen.insert(s.speaker.clone()))
.map(|s| s.speaker.clone())
.collect()
}
/// Concatenates up to `MAX_CANDIDATE_SAMPLES` of `label`'s audio out of the
/// full 16kHz-mono recording, using each span's millisecond range.
fn candidate_audio(wav_samples: &[f32], spans: &[SpeakerSpan], label: &str) -> Vec<f32> {
const SAMPLES_PER_MS: u64 = 16; // 16_000 Hz / 1000
let mut out = Vec::new();
for span in spans.iter().filter(|s| s.speaker == label) {
if out.len() >= MAX_CANDIDATE_SAMPLES {
break;
}
let start = (span.start_ms * SAMPLES_PER_MS) as usize;
let end = ((span.end_ms * SAMPLES_PER_MS) as usize).min(wav_samples.len());
if start < end {
out.extend_from_slice(&wav_samples[start..end]);
}
}
out.truncate(MAX_CANDIDATE_SAMPLES);
out
}
fn cosine_similarity(a: &[f32], b: &[f32]) -> f32 {
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let norm_a = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let norm_b = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm_a == 0.0 || norm_b == 0.0 {
0.0
} else {
dot / (norm_a * norm_b)
}
}
/// `mic_label` -> "You"; every other label, in first-appearance order ->
/// "Speaker 2", "Speaker 3", … (numbering starts at 2 — "You" stands in for
/// "Speaker 1" without ever being called that). `pub(crate)` so Phase 3 reuses
/// it for its merged You + far-side spans, keeping naming uniform app-wide.
pub(crate) fn build_name_map(
labels_in_order: &[String],
mic_label: &str,
) -> HashMap<String, String> {
let mut names = HashMap::new();
let mut next_speaker_number = 2;
for label in labels_in_order {
if label == mic_label {
names.insert(label.clone(), "You".to_string());
} else {
names.insert(label.clone(), format!("Speaker {next_speaker_number}"));
next_speaker_number += 1;
}
}
names
}
#[cfg(test)]
mod tests {
use super::*;
fn span(start_ms: u64, end_ms: u64, speaker: &str) -> SpeakerSpan {
SpeakerSpan {
start_ms,
end_ms,
speaker: speaker.to_string(),
}
}
#[test]
fn first_appearance_order_dedupes_in_encounter_order() {
let spans = vec![
span(0, 1000, "S2"),
span(1000, 2000, "S1"),
span(2000, 3000, "S2"),
];
assert_eq!(first_appearance_order(&spans), vec!["S2", "S1"]);
}
#[test]
fn candidate_audio_concatenates_only_that_speakers_spans() {
let wav: Vec<f32> = (0..32_000).map(|i| i as f32).collect(); // 2s @16kHz
let spans = vec![
span(0, 500, "S1"),
span(500, 1000, "S2"),
span(1000, 1500, "S1"),
];
let s1 = candidate_audio(&wav, &spans, "S1");
// 500ms + 500ms of S1 = 1s = 16_000 samples, taken from [0,8000) and [16000,24000).
assert_eq!(s1.len(), 16_000);
assert_eq!(s1[0], 0.0);
assert_eq!(s1[8000], 16_000.0);
}
#[test]
fn candidate_audio_caps_at_the_maximum() {
let wav: Vec<f32> = vec![0.0; MAX_CANDIDATE_SAMPLES + 10_000];
let spans = vec![span(0, (MAX_CANDIDATE_SAMPLES as u64 + 10_000) / 16, "S1")];
assert_eq!(
candidate_audio(&wav, &spans, "S1").len(),
MAX_CANDIDATE_SAMPLES
);
}
#[test]
fn cosine_similarity_of_identical_vectors_is_one() {
let v = [1.0, 2.0, 3.0];
assert!((cosine_similarity(&v, &v) - 1.0).abs() < 1e-6);
}
#[test]
fn cosine_similarity_of_opposite_vectors_is_negative_one() {
let a = [1.0, 0.0];
let b = [-1.0, 0.0];
assert!((cosine_similarity(&a, &b) + 1.0).abs() < 1e-6);
}
#[test]
fn cosine_similarity_handles_a_zero_vector_without_dividing_by_zero() {
let a = [0.0, 0.0];
let b = [1.0, 1.0];
assert_eq!(cosine_similarity(&a, &b), 0.0);
}
#[test]
fn build_name_map_labels_the_mic_you_and_numbers_the_rest_from_two() {
let labels = vec!["S2".to_string(), "S1".to_string(), "S3".to_string()];
let names = build_name_map(&labels, "S1");
assert_eq!(names.get("S1"), Some(&"You".to_string()));
assert_eq!(names.get("S2"), Some(&"Speaker 2".to_string()));
assert_eq!(names.get("S3"), Some(&"Speaker 3".to_string()));
}
#[test]
fn match_mic_speaker_returns_empty_when_mic_sample_is_too_short() {
let spans = vec![span(0, 1000, "S1")];
let names = match_mic_speaker(
Path::new("model.onnx"),
&[0.0; 100],
Path::new("audio.wav"),
&spans,
)
.unwrap();
assert!(names.is_empty());
}
}