From a2a5789e2a429cf6aebfa6b83cd8efdca5b1f101 Mon Sep 17 00:00:00 2001 From: Jer Miller Date: Sun, 26 Jul 2026 14:09:48 -0600 Subject: [PATCH] feat(speakers): port pyannote segmentation pass to Rust Add the pyannote segmentation model pass, per-window speaker statistics, and speaker-evidence decision logic to solstone-core-speakers. The pure crate stays dependency-free and inside the iOS gate. Add the pyannote ONNX role to solstone-core-speakers-onnx, which stays outside the iOS gate because ONNX Runtime remains host-only native linkage. This deliberately diverges from the Python structure in two places. Rust has one parameterised windowing pass where Python has three near-copies: overlap.compute_overlap_fraction, overlap.compute_overlap_and_logprobs, and diarize._run_pyannote. Rust also has one session-construction helper where Python loads pyannote twice under mismatched provider policies: overlap.py:145 uses _select_onnx_providers(), while diarize.py:78 hardcodes CPUExecutionProvider. Preserve and widen the single-site assertions rather than relaxing them. Both assertions previously scanned only include_str!("lib.rs") and would have counted zero after the module split. They now recursively walk every .rs file under src/ and assert a lower bound on files visited, so a broken walk fails loudly instead of passing vacuously. The ONNX assertion still requires exactly one Session::builder()? across the whole crate, covering both model roles through one construction path, one provider policy, and one error surface. The mel-bank assertion gets the same widening. CoreML and Apple-hardware behavior remain unverified. This host had no Apple hardware. The aarch64-apple-ios cross-target check passes, but nothing was executed on Apple silicon and no accelerator path was exercised. Use f64::round_ties_even() for frame_start because Python round() is banker's rounding while Rust f64::round() is half-away-from-zero. At stride 5, window index 1 gives start_sample 80000 and exactly 294.5: Python resolves that to 294, while f64::round() would produce 295. The wrong function would shift every downstream frame boundary in that window. The test pins the exact 294.5 case and asserts both the divergent Rust round() value and the correct ties-even value. core/fixtures/speaker_stage_boundaries.json gates the pure numeric stages and was never intended to gate the model pass. It carries no per-frame log-probability arrays, no window-start sequences, and no model-derived overlap fraction. speaker_evidence.*.overlap_fraction is a hand-chosen scalar input, 0.0 or 0.049. Every committed evidence case also has exactly one window, so the rule that only speech-bearing windows enter either denominator is unexercised by fixture; a port dropping that filter would still match all four committed decisions. The consequence is that the windowing, accumulation, frame-offset, averaged-argmax mechanics, and the speech-bearing denominator filter are covered by structural tests in CI and deferred to the real-corpus bundle differential for numeric confirmation. They are not ungated and not a fixture gap. The differential already records these as first-class bundle fields and compares them element-wise Python-vs-port: pyannote.avg_log_probs at tests/verify_speaker_differential.py:87, pyannote.window_stats at :88, evidence.mean_window_overlap_share at :92, evidence.overlap_fraction at :93, per_frame_argmax_agreement_fraction gated by LOGPROB_ARGMAX_AGREEMENT_MIN at :1011-1020, and window_stats_equal as an exact array comparison at :1064. Honesty table: Expectation Source ----------------------------------- ----------------------------- start/final/short pad hand-derived from fixture constants; implementation- independent; not committed accumulation order/count floor in-test hand-computation 294.5 -> 294 genuine independent Python oracle; not committed overlap from averaged argmax in-test hand-computation speech-bearing denominator filter in-test hand-computation per-window statistics triples committed fixture value at speaker_evidence.. windows[0]; class sequence is test input decisions and both fractions committed fixture value at speaker_evidence.. decision.* eleven constants committed fixture value at identity.source_constants.* Add a cross-language calibration-drift gate for eleven constants, each asserted equal to its identity.source_constants path so a Python-side recalibration turns the Rust tests red. Keep SPEAKER_EVIDENCE_MULTI_MIN, SPEAKER_EVIDENCE_SINGLE_MAX, and DIARIZE_MIN_OVERLAP as three separate constants even though all are 0.05 today; they are independent tuning controls. Test counts increase from 5 to 21 in solstone-core-speakers and from 5 to 7 in solstone-core-speakers-onnx. No core/fixtures/*, scripts/build_core_fixtures.py, or solstone/**/*.py files changed. Co-Authored-By: Claude Opus 5 (1M context) --- .../solstone-core-speakers-onnx/src/lib.rs | 310 ++---- .../src/pyannote.rs | 161 +++ .../src/session.rs | 107 ++ .../src/wespeaker.rs | 143 +++ .../solstone-core-speakers/src/filterbank.rs | 380 +++++++ core/crates/solstone-core-speakers/src/lib.rs | 442 ++------ .../src/segmentation.rs | 958 ++++++++++++++++++ 7 files changed, 1902 insertions(+), 599 deletions(-) create mode 100644 core/crates/solstone-core-speakers-onnx/src/pyannote.rs create mode 100644 core/crates/solstone-core-speakers-onnx/src/session.rs create mode 100644 core/crates/solstone-core-speakers-onnx/src/wespeaker.rs create mode 100644 core/crates/solstone-core-speakers/src/filterbank.rs create mode 100644 core/crates/solstone-core-speakers/src/segmentation.rs diff --git a/core/crates/solstone-core-speakers-onnx/src/lib.rs b/core/crates/solstone-core-speakers-onnx/src/lib.rs index cd43298e3..bb782bace 100644 --- a/core/crates/solstone-core-speakers-onnx/src/lib.rs +++ b/core/crates/solstone-core-speakers-onnx/src/lib.rs @@ -1,17 +1,17 @@ // SPDX-License-Identifier: AGPL-3.0-only // Copyright (c) 2026 sol pbc +mod pyannote; +mod session; +mod wespeaker; + use std::error::Error; use std::fmt; -use std::path::Path; -use ort::ep::{CPU, CoreML, ExecutionProviderDispatch}; -use ort::session::Session; -use ort::value::{Tensor, TensorElementType, ValueType}; -use solstone_core_speakers::{FeatureMatrix, WESPEAKER_EMBEDDING_SIZE, WESPEAKER_MEL_BINS}; +use solstone_core_speakers::{PYANNOTE_WINDOW_S, WESPEAKER_MEL_BINS}; -const INPUT_NAME: &str = "feats"; -const OUTPUT_NAME: &str = "embs"; +pub use pyannote::PyannoteSegmenter; +pub use wespeaker::{SpeakerEmbedding, WespeakerEmbedder}; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum PlatformFamily { @@ -42,32 +42,29 @@ pub enum SpeakerExecutionProvider { Cpu, } -#[derive(Debug, Clone, PartialEq)] -pub struct SpeakerEmbedding { - values: [f32; WESPEAKER_EMBEDDING_SIZE], -} - -impl SpeakerEmbedding { - pub fn values(&self) -> &[f32; WESPEAKER_EMBEDDING_SIZE] { - &self.values - } -} - -#[derive(Debug)] -pub struct WespeakerEmbedder { - session: Session, - input_name: String, - output_name: String, -} - #[derive(Debug, Clone, PartialEq, Eq)] pub enum SpeakerOnnxError { EmptyProviderPlan, - ProviderUnavailable { provider: &'static str }, - InvalidFeatureMatrix { frames: usize, bins: usize }, - InvalidModelIo { detail: String }, - MissingOutput { name: String }, - Ort { message: String }, + ProviderUnavailable { + provider: &'static str, + }, + InvalidFeatureMatrix { + frames: usize, + bins: usize, + }, + InvalidAudioWindow { + expected_samples: usize, + actual_samples: usize, + }, + InvalidModelIo { + detail: String, + }, + MissingOutput { + name: String, + }, + Ort { + message: String, + }, } impl fmt::Display for SpeakerOnnxError { @@ -84,6 +81,13 @@ impl fmt::Display for SpeakerOnnxError { formatter, "speaker ONNX features must have at least one frame and {WESPEAKER_MEL_BINS} bins, got frames={frames} bins={bins}" ), + Self::InvalidAudioWindow { + expected_samples, + actual_samples, + } => write!( + formatter, + "pyannote ONNX audio window must have {expected_samples} samples ({PYANNOTE_WINDOW_S}s at 16 kHz), got {actual_samples}" + ), Self::InvalidModelIo { detail } => { write!(formatter, "speaker ONNX model IO mismatch: {detail}") } @@ -117,171 +121,10 @@ pub fn default_speaker_execution_providers( } } -impl WespeakerEmbedder { - pub fn open( - model_path: &Path, - providers: &[SpeakerExecutionProvider], - ) -> Result { - if providers.is_empty() { - return Err(SpeakerOnnxError::EmptyProviderPlan); - } - let dispatches = provider_dispatches(providers)?; - let session = Session::builder()? - .with_execution_providers(dispatches)? - .commit_from_file(model_path)?; - validate_session_io(&session)?; - Ok(Self { - session, - input_name: INPUT_NAME.to_string(), - output_name: OUTPUT_NAME.to_string(), - }) - } - - pub fn embed( - &mut self, - features: &FeatureMatrix, - ) -> Result { - if features.frames() == 0 || features.bins() != WESPEAKER_MEL_BINS { - return Err(SpeakerOnnxError::InvalidFeatureMatrix { - frames: features.frames(), - bins: features.bins(), - }); - } - let input = Tensor::from_array(( - [1_usize, features.frames(), WESPEAKER_MEL_BINS], - features.data().to_vec().into_boxed_slice(), - ))?; - let mut outputs = self - .session - .run(ort::inputs![self.input_name.as_str() => input])?; - let output = - outputs - .remove(&self.output_name) - .ok_or_else(|| SpeakerOnnxError::MissingOutput { - name: self.output_name.clone(), - })?; - let (shape, values) = output.try_extract_tensor::()?; - if shape[..] != [1, WESPEAKER_EMBEDDING_SIZE as i64] { - return Err(SpeakerOnnxError::InvalidModelIo { - detail: format!("output shape {shape} is not [1, {WESPEAKER_EMBEDDING_SIZE}]"), - }); - } - let mut embedding = [0.0; WESPEAKER_EMBEDDING_SIZE]; - embedding.copy_from_slice(values); - Ok(SpeakerEmbedding { values: embedding }) - } -} - -fn provider_dispatches( - providers: &[SpeakerExecutionProvider], -) -> Result, SpeakerOnnxError> { - let mut dispatches = Vec::with_capacity(providers.len()); - for provider in providers { - match provider { - SpeakerExecutionProvider::CoreMl => { - if !cfg!(target_vendor = "apple") { - return Err(SpeakerOnnxError::ProviderUnavailable { provider: "coreml" }); - } - dispatches.push(CoreML::default().build()); - } - SpeakerExecutionProvider::Cpu => { - dispatches.push(CPU::default().build()); - } - } - } - Ok(dispatches) -} - -fn validate_session_io(session: &Session) -> Result<(), SpeakerOnnxError> { - let inputs = session.inputs(); - let outputs = session.outputs(); - if inputs.len() != 1 { - return Err(SpeakerOnnxError::InvalidModelIo { - detail: format!("expected one input, got {}", inputs.len()), - }); - } - if outputs.len() != 1 { - return Err(SpeakerOnnxError::InvalidModelIo { - detail: format!("expected one output, got {}", outputs.len()), - }); - } - expect_tensor( - "input", - inputs[0].name(), - inputs[0].dtype(), - INPUT_NAME, - &[ - ExpectedDim::Any, - ExpectedDim::Any, - ExpectedDim::Exact(WESPEAKER_MEL_BINS as i64), - ], - )?; - expect_tensor( - "output", - outputs[0].name(), - outputs[0].dtype(), - OUTPUT_NAME, - &[ - ExpectedDim::Any, - ExpectedDim::Exact(WESPEAKER_EMBEDDING_SIZE as i64), - ], - )?; - Ok(()) -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -enum ExpectedDim { - Any, - Exact(i64), -} - -fn expect_tensor( - label: &str, - name: &str, - value_type: &ValueType, - expected_name: &str, - expected_shape: &[ExpectedDim], -) -> Result<(), SpeakerOnnxError> { - if name != expected_name { - return Err(SpeakerOnnxError::InvalidModelIo { - detail: format!("{label} name {name:?} is not {expected_name:?}"), - }); - } - let ValueType::Tensor { ty, shape, .. } = value_type else { - return Err(SpeakerOnnxError::InvalidModelIo { - detail: format!("{label} {name:?} is not a tensor"), - }); - }; - if *ty != TensorElementType::Float32 { - return Err(SpeakerOnnxError::InvalidModelIo { - detail: format!("{label} {name:?} is {ty}, not float32"), - }); - } - if shape.len() != expected_shape.len() { - return Err(SpeakerOnnxError::InvalidModelIo { - detail: format!("{label} {name:?} shape {shape} has wrong rank"), - }); - } - for (index, (actual, expected)) in shape.iter().zip(expected_shape).enumerate() { - match expected { - ExpectedDim::Any => {} - ExpectedDim::Exact(value) if actual == value => {} - ExpectedDim::Exact(value) => { - return Err(SpeakerOnnxError::InvalidModelIo { - detail: format!("{label} {name:?} dim {index} is {actual}, not {value}"), - }); - } - } - } - Ok(()) -} - #[cfg(test)] mod tests { use super::*; - use solstone_core_speakers::{WESPEAKER_SAMPLE_RATE_HZ, compute_wespeaker_filterbank_cmn}; - - const FIXTURE: &str = include_str!("../../../fixtures/speaker_filterbank.json"); + use crate::test_support::source_tree_needle_count; #[test] fn provider_plan_selects_coreml_then_cpu_for_synthetic_apple() { @@ -308,51 +151,34 @@ mod tests { } #[test] - fn coreml_provider_open_is_rejected_on_non_apple_builds() { - if cfg!(target_vendor = "apple") { - return; - } - let error = provider_dispatches(&[SpeakerExecutionProvider::CoreMl]).unwrap_err(); - assert_eq!( - error, - SpeakerOnnxError::ProviderUnavailable { provider: "coreml" } - ); - } + fn session_builder_has_single_production_site_scans_src_tree() { + let needle = concat!("Session", "::builder()?"); + let (visited_files, count) = source_tree_needle_count(needle); - #[test] - fn committed_wespeaker_model_accepts_fixture_features_and_returns_256_floats() { - let fixture = fixture(); - let audio = decode_waveform(&fixture); - let features = - compute_wespeaker_filterbank_cmn(&audio, WESPEAKER_SAMPLE_RATE_HZ).expect("features"); - let model_path = repo_root().join( - "packages/solstone-journal-models/solstone_journal_models/assets/wespeaker-resnet34-256.onnx", + assert!( + visited_files >= 4, + "source walk visited too few Rust files: {visited_files}" ); - let mut embedder = WespeakerEmbedder::open(&model_path, &[SpeakerExecutionProvider::Cpu]) - .expect("embedder"); - - let embedding = embedder.embed(&features).expect("embedding"); - - assert_eq!(embedding.values().len(), WESPEAKER_EMBEDDING_SIZE); - assert!(embedding.values().iter().all(|value| value.is_finite())); + assert_eq!(count, 1); } +} - #[test] - fn session_builder_has_single_production_site() { - let source = include_str!("lib.rs"); - let needle = concat!("Session", "::builder()?"); - assert_eq!(source.matches(needle).count(), 1); - } +#[cfg(test)] +pub(crate) mod test_support { + use serde_json::Value; + use std::path::{Path, PathBuf}; + + pub(crate) const FIXTURE: &str = include_str!("../../../fixtures/speaker_filterbank.json"); - fn repo_root() -> std::path::PathBuf { + pub(crate) fn repo_root() -> std::path::PathBuf { std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("../../..") } - fn fixture() -> serde_json::Value { + pub(crate) fn fixture() -> Value { serde_json::from_str(FIXTURE).expect("fixture JSON") } - fn decode_waveform(fixture: &serde_json::Value) -> Vec { + pub(crate) fn decode_waveform(fixture: &Value) -> Vec { let encoded = fixture["waveform"]["samples_f32_le_base64"] .as_str() .expect("waveform base64"); @@ -363,7 +189,7 @@ mod tests { .collect() } - fn decode_base64(input: &str) -> Vec { + pub(crate) fn decode_base64(input: &str) -> Vec { let mut out = Vec::with_capacity(input.len() / 4 * 3); let mut quartet = [0_u8; 4]; let mut len = 0; @@ -392,4 +218,34 @@ mod tests { assert_eq!(len, 0); out } + + pub(crate) fn source_tree_needle_count(needle: &str) -> (usize, usize) { + let src = Path::new(env!("CARGO_MANIFEST_DIR")).join("src"); + let mut files = Vec::new(); + collect_rust_files(&src, &mut files); + let count = files + .iter() + .map(|path| { + std::fs::read_to_string(path) + .unwrap_or_else(|error| panic!("failed to read {}: {error}", path.display())) + .matches(needle) + .count() + }) + .sum(); + (files.len(), count) + } + + fn collect_rust_files(path: &Path, files: &mut Vec) { + for entry in std::fs::read_dir(path) + .unwrap_or_else(|error| panic!("failed to read {}: {error}", path.display())) + { + let entry = entry.expect("directory entry"); + let path = entry.path(); + if path.is_dir() { + collect_rust_files(&path, files); + } else if path.extension().and_then(|value| value.to_str()) == Some("rs") { + files.push(path); + } + } + } } diff --git a/core/crates/solstone-core-speakers-onnx/src/pyannote.rs b/core/crates/solstone-core-speakers-onnx/src/pyannote.rs new file mode 100644 index 000000000..8f9671d40 --- /dev/null +++ b/core/crates/solstone-core-speakers-onnx/src/pyannote.rs @@ -0,0 +1,161 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::path::Path; + +use ort::session::Session; +use ort::value::Tensor; +use solstone_core_speakers::{ + FeatureMatrix, PYANNOTE_CLASS_COUNT, PYANNOTE_FRAMES_PER_WINDOW, PYANNOTE_SAMPLE_RATE_HZ, + PYANNOTE_WINDOW_S, +}; + +use crate::session::{ExpectedDim, expect_tensor, open_session}; +use crate::{SpeakerExecutionProvider, SpeakerOnnxError}; + +const INPUT_NAME: &str = "input_values"; +const OUTPUT_NAME: &str = "logits"; + +#[derive(Debug)] +pub struct PyannoteSegmenter { + session: Session, + input_name: String, + output_name: String, +} + +impl PyannoteSegmenter { + pub fn open( + model_path: &Path, + providers: &[SpeakerExecutionProvider], + ) -> Result { + let session = open_session(model_path, providers)?; + validate_session_io(&session)?; + Ok(Self { + session, + input_name: INPUT_NAME.to_string(), + output_name: OUTPUT_NAME.to_string(), + }) + } + + pub fn infer_window( + &mut self, + audio_window: &[f32], + ) -> Result { + validate_audio_window(audio_window)?; + let input = Tensor::from_array(( + [1_usize, 1_usize, audio_window.len()], + audio_window.to_vec().into_boxed_slice(), + ))?; + let mut outputs = self + .session + .run(ort::inputs![self.input_name.as_str() => input])?; + let output = + outputs + .remove(&self.output_name) + .ok_or_else(|| SpeakerOnnxError::MissingOutput { + name: self.output_name.clone(), + })?; + let (shape, values) = output.try_extract_tensor::()?; + if shape[..] + != [ + 1, + PYANNOTE_FRAMES_PER_WINDOW as i64, + PYANNOTE_CLASS_COUNT as i64, + ] + { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!( + "output shape {shape} is not [1, {PYANNOTE_FRAMES_PER_WINDOW}, {PYANNOTE_CLASS_COUNT}]" + ), + }); + } + FeatureMatrix::from_row_major( + PYANNOTE_FRAMES_PER_WINDOW, + PYANNOTE_CLASS_COUNT, + values.to_vec(), + ) + .map_err(|error| SpeakerOnnxError::InvalidModelIo { + detail: error.to_string(), + }) + } +} + +fn validate_audio_window(audio_window: &[f32]) -> Result<(), SpeakerOnnxError> { + let expected_samples = PYANNOTE_WINDOW_S as usize * PYANNOTE_SAMPLE_RATE_HZ as usize; + if audio_window.len() != expected_samples { + return Err(SpeakerOnnxError::InvalidAudioWindow { + expected_samples, + actual_samples: audio_window.len(), + }); + } + Ok(()) +} + +fn validate_session_io(session: &Session) -> Result<(), SpeakerOnnxError> { + let inputs = session.inputs(); + let outputs = session.outputs(); + if inputs.len() != 1 { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("expected one input, got {}", inputs.len()), + }); + } + if outputs.len() != 1 { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("expected one output, got {}", outputs.len()), + }); + } + expect_tensor( + "input", + inputs[0].name(), + inputs[0].dtype(), + INPUT_NAME, + &[ExpectedDim::Any, ExpectedDim::Any, ExpectedDim::Any], + )?; + expect_tensor( + "output", + outputs[0].name(), + outputs[0].dtype(), + OUTPUT_NAME, + &[ + ExpectedDim::Any, + ExpectedDim::Any, + ExpectedDim::Exact(PYANNOTE_CLASS_COUNT as i64), + ], + )?; + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::repo_root; + + #[test] + fn committed_pyannote_model_accepts_zero_window_and_returns_589x7_finite_values() { + let model_path = repo_root().join( + "packages/solstone-journal-models/solstone_journal_models/assets/pyannote-segmentation-3.0.onnx", + ); + let mut segmenter = PyannoteSegmenter::open(&model_path, &[SpeakerExecutionProvider::Cpu]) + .expect("segmenter"); + let audio = vec![0.0; PYANNOTE_WINDOW_S as usize * PYANNOTE_SAMPLE_RATE_HZ as usize]; + + let log_probs = segmenter.infer_window(&audio).expect("log probs"); + + assert_eq!(log_probs.frames(), PYANNOTE_FRAMES_PER_WINDOW); + assert_eq!(log_probs.bins(), PYANNOTE_CLASS_COUNT); + assert!(log_probs.data().iter().all(|value| value.is_finite())); + } + + #[test] + fn pyannote_rejects_wrong_length_window() { + let error = validate_audio_window(&[0.0; 42]).unwrap_err(); + + assert_eq!( + error, + SpeakerOnnxError::InvalidAudioWindow { + expected_samples: PYANNOTE_WINDOW_S as usize * PYANNOTE_SAMPLE_RATE_HZ as usize, + actual_samples: 42, + } + ); + } +} diff --git a/core/crates/solstone-core-speakers-onnx/src/session.rs b/core/crates/solstone-core-speakers-onnx/src/session.rs new file mode 100644 index 000000000..b810f90a7 --- /dev/null +++ b/core/crates/solstone-core-speakers-onnx/src/session.rs @@ -0,0 +1,107 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::path::Path; + +use ort::ep::{CPU, CoreML, ExecutionProviderDispatch}; +use ort::session::Session; +use ort::value::{TensorElementType, ValueType}; + +use crate::{SpeakerExecutionProvider, SpeakerOnnxError}; + +pub(crate) fn open_session( + model_path: &Path, + providers: &[SpeakerExecutionProvider], +) -> Result { + if providers.is_empty() { + return Err(SpeakerOnnxError::EmptyProviderPlan); + } + let dispatches = provider_dispatches(providers)?; + Ok(Session::builder()? + .with_execution_providers(dispatches)? + .commit_from_file(model_path)?) +} + +fn provider_dispatches( + providers: &[SpeakerExecutionProvider], +) -> Result, SpeakerOnnxError> { + let mut dispatches = Vec::with_capacity(providers.len()); + for provider in providers { + match provider { + SpeakerExecutionProvider::CoreMl => { + if !cfg!(target_vendor = "apple") { + return Err(SpeakerOnnxError::ProviderUnavailable { provider: "coreml" }); + } + dispatches.push(CoreML::default().build()); + } + SpeakerExecutionProvider::Cpu => { + dispatches.push(CPU::default().build()); + } + } + } + Ok(dispatches) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum ExpectedDim { + Any, + Exact(i64), +} + +pub(crate) fn expect_tensor( + label: &str, + name: &str, + value_type: &ValueType, + expected_name: &str, + expected_shape: &[ExpectedDim], +) -> Result<(), SpeakerOnnxError> { + if name != expected_name { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("{label} name {name:?} is not {expected_name:?}"), + }); + } + let ValueType::Tensor { ty, shape, .. } = value_type else { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("{label} {name:?} is not a tensor"), + }); + }; + if *ty != TensorElementType::Float32 { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("{label} {name:?} is {ty}, not float32"), + }); + } + if shape.len() != expected_shape.len() { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("{label} {name:?} shape {shape} has wrong rank"), + }); + } + for (index, (actual, expected)) in shape.iter().zip(expected_shape).enumerate() { + match expected { + ExpectedDim::Any => {} + ExpectedDim::Exact(value) if actual == value => {} + ExpectedDim::Exact(value) => { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("{label} {name:?} dim {index} is {actual}, not {value}"), + }); + } + } + } + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn coreml_provider_open_is_rejected_on_non_apple_builds() { + if cfg!(target_vendor = "apple") { + return; + } + let error = provider_dispatches(&[SpeakerExecutionProvider::CoreMl]).unwrap_err(); + assert_eq!( + error, + SpeakerOnnxError::ProviderUnavailable { provider: "coreml" } + ); + } +} diff --git a/core/crates/solstone-core-speakers-onnx/src/wespeaker.rs b/core/crates/solstone-core-speakers-onnx/src/wespeaker.rs new file mode 100644 index 000000000..21cbfda6f --- /dev/null +++ b/core/crates/solstone-core-speakers-onnx/src/wespeaker.rs @@ -0,0 +1,143 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::path::Path; + +use ort::session::Session; +use ort::value::Tensor; +use solstone_core_speakers::{FeatureMatrix, WESPEAKER_EMBEDDING_SIZE, WESPEAKER_MEL_BINS}; + +use crate::session::{ExpectedDim, expect_tensor, open_session}; +use crate::{SpeakerExecutionProvider, SpeakerOnnxError}; + +const INPUT_NAME: &str = "feats"; +const OUTPUT_NAME: &str = "embs"; + +#[derive(Debug, Clone, PartialEq)] +pub struct SpeakerEmbedding { + values: [f32; WESPEAKER_EMBEDDING_SIZE], +} + +impl SpeakerEmbedding { + pub fn values(&self) -> &[f32; WESPEAKER_EMBEDDING_SIZE] { + &self.values + } +} + +#[derive(Debug)] +pub struct WespeakerEmbedder { + session: Session, + input_name: String, + output_name: String, +} + +impl WespeakerEmbedder { + pub fn open( + model_path: &Path, + providers: &[SpeakerExecutionProvider], + ) -> Result { + let session = open_session(model_path, providers)?; + validate_session_io(&session)?; + Ok(Self { + session, + input_name: INPUT_NAME.to_string(), + output_name: OUTPUT_NAME.to_string(), + }) + } + + pub fn embed( + &mut self, + features: &FeatureMatrix, + ) -> Result { + if features.frames() == 0 || features.bins() != WESPEAKER_MEL_BINS { + return Err(SpeakerOnnxError::InvalidFeatureMatrix { + frames: features.frames(), + bins: features.bins(), + }); + } + let input = Tensor::from_array(( + [1_usize, features.frames(), WESPEAKER_MEL_BINS], + features.data().to_vec().into_boxed_slice(), + ))?; + let mut outputs = self + .session + .run(ort::inputs![self.input_name.as_str() => input])?; + let output = + outputs + .remove(&self.output_name) + .ok_or_else(|| SpeakerOnnxError::MissingOutput { + name: self.output_name.clone(), + })?; + let (shape, values) = output.try_extract_tensor::()?; + if shape[..] != [1, WESPEAKER_EMBEDDING_SIZE as i64] { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("output shape {shape} is not [1, {WESPEAKER_EMBEDDING_SIZE}]"), + }); + } + let mut embedding = [0.0; WESPEAKER_EMBEDDING_SIZE]; + embedding.copy_from_slice(values); + Ok(SpeakerEmbedding { values: embedding }) + } +} + +fn validate_session_io(session: &Session) -> Result<(), SpeakerOnnxError> { + let inputs = session.inputs(); + let outputs = session.outputs(); + if inputs.len() != 1 { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("expected one input, got {}", inputs.len()), + }); + } + if outputs.len() != 1 { + return Err(SpeakerOnnxError::InvalidModelIo { + detail: format!("expected one output, got {}", outputs.len()), + }); + } + expect_tensor( + "input", + inputs[0].name(), + inputs[0].dtype(), + INPUT_NAME, + &[ + ExpectedDim::Any, + ExpectedDim::Any, + ExpectedDim::Exact(WESPEAKER_MEL_BINS as i64), + ], + )?; + expect_tensor( + "output", + outputs[0].name(), + outputs[0].dtype(), + OUTPUT_NAME, + &[ + ExpectedDim::Any, + ExpectedDim::Exact(WESPEAKER_EMBEDDING_SIZE as i64), + ], + )?; + Ok(()) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::{decode_waveform, fixture, repo_root}; + use solstone_core_speakers::{WESPEAKER_SAMPLE_RATE_HZ, compute_wespeaker_filterbank_cmn}; + + #[test] + fn committed_wespeaker_model_accepts_fixture_features_and_returns_256_floats() { + let fixture = fixture(); + let audio = decode_waveform(&fixture); + let features = + compute_wespeaker_filterbank_cmn(&audio, WESPEAKER_SAMPLE_RATE_HZ).expect("features"); + let model_path = repo_root().join( + "packages/solstone-journal-models/solstone_journal_models/assets/wespeaker-resnet34-256.onnx", + ); + let mut embedder = WespeakerEmbedder::open(&model_path, &[SpeakerExecutionProvider::Cpu]) + .expect("embedder"); + + let embedding = embedder.embed(&features).expect("embedding"); + + assert_eq!(embedding.values().len(), WESPEAKER_EMBEDDING_SIZE); + assert!(embedding.values().iter().all(|value| value.is_finite())); + } +} diff --git a/core/crates/solstone-core-speakers/src/filterbank.rs b/core/crates/solstone-core-speakers/src/filterbank.rs new file mode 100644 index 000000000..6d37cd543 --- /dev/null +++ b/core/crates/solstone-core-speakers/src/filterbank.rs @@ -0,0 +1,380 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::f32::consts::PI; + +use crate::{FeatureMatrix, SpeakerFeatureError}; + +pub const WESPEAKER_SAMPLE_RATE_HZ: u32 = 16_000; +pub const WESPEAKER_MEL_BINS: usize = 80; +pub const WESPEAKER_FRAME_LENGTH_SAMPLES: usize = 400; +pub const WESPEAKER_FRAME_SHIFT_SAMPLES: usize = 160; +pub const WESPEAKER_FFT_SIZE: usize = 512; +pub const WESPEAKER_EMBEDDING_SIZE: usize = 256; + +const LOW_FREQ_HZ: f32 = 20.0; +const HIGH_FREQ_HZ: f32 = WESPEAKER_SAMPLE_RATE_HZ as f32 / 2.0; +const PREEMPH_COEFF: f32 = 0.97; +const AUDIO_SCALE: f32 = 32768.0; +const ROW_NORM_EPSILON: f32 = 1e-9; + +pub fn compute_wespeaker_filterbank_cmn( + audio: &[f32], + sample_rate_hz: u32, +) -> Result { + if sample_rate_hz != WESPEAKER_SAMPLE_RATE_HZ { + return Err(SpeakerFeatureError::UnsupportedSampleRate { + expected: WESPEAKER_SAMPLE_RATE_HZ, + actual: sample_rate_hz, + }); + } + if let Some((index, _sample)) = audio + .iter() + .enumerate() + .find(|(_index, sample)| !sample.is_finite()) + { + return Err(SpeakerFeatureError::NonFiniteAudioSample { index }); + } + + let frames = num_snipped_frames(audio.len()); + let mut data = vec![0.0; frames * WESPEAKER_MEL_BINS]; + if frames == 0 { + return FeatureMatrix::from_row_major(0, WESPEAKER_MEL_BINS, data); + } + + let window = povey_window(); + let mel_bank = MelBank::production(); + for frame in 0..frames { + let mut padded = padded_frame(audio, frame, &window); + let power = power_spectrum_512(&mut padded); + let output_start = frame * WESPEAKER_MEL_BINS; + mel_bank.compute_log_energies( + &power, + &mut data[output_start..output_start + WESPEAKER_MEL_BINS], + ); + } + subtract_column_mean(frames, WESPEAKER_MEL_BINS, &mut data); + FeatureMatrix::from_row_major(frames, WESPEAKER_MEL_BINS, data) +} + +pub fn row_l2_normalize(features: &FeatureMatrix) -> FeatureMatrix { + let mut data = features.data().to_vec(); + for row in data.chunks_mut(features.bins()) { + let norm = row.iter().map(|value| value * value).sum::().sqrt(); + let denom = if norm > ROW_NORM_EPSILON { norm } else { 1.0 }; + for value in row { + *value /= denom; + } + } + FeatureMatrix::from_row_major(features.frames(), features.bins(), data) + .expect("row-l2 normalization preserves feature shape") +} + +fn num_snipped_frames(samples: usize) -> usize { + if samples < WESPEAKER_FRAME_LENGTH_SAMPLES { + 0 + } else { + 1 + (samples - WESPEAKER_FRAME_LENGTH_SAMPLES) / WESPEAKER_FRAME_SHIFT_SAMPLES + } +} + +fn povey_window() -> [f32; WESPEAKER_FRAME_LENGTH_SAMPLES] { + let mut window = [0.0; WESPEAKER_FRAME_LENGTH_SAMPLES]; + let scale = 2.0_f64 * std::f64::consts::PI / (WESPEAKER_FRAME_LENGTH_SAMPLES as f64 - 1.0); + for (index, value) in window.iter_mut().enumerate() { + *value = (0.5_f64 - 0.5_f64 * (scale * index as f64).cos()).powf(0.85) as f32; + } + window +} + +fn padded_frame( + audio: &[f32], + frame: usize, + window: &[f32; WESPEAKER_FRAME_LENGTH_SAMPLES], +) -> [f32; WESPEAKER_FFT_SIZE] { + let mut padded = [0.0; WESPEAKER_FFT_SIZE]; + let start = frame * WESPEAKER_FRAME_SHIFT_SAMPLES; + for index in 0..WESPEAKER_FRAME_LENGTH_SAMPLES { + padded[index] = audio[start + index] * AUDIO_SCALE; + } + remove_dc_offset(&mut padded[..WESPEAKER_FRAME_LENGTH_SAMPLES]); + preemphasize(&mut padded[..WESPEAKER_FRAME_LENGTH_SAMPLES]); + apply_window(&mut padded[..WESPEAKER_FRAME_LENGTH_SAMPLES], window); + padded +} + +fn remove_dc_offset(frame: &mut [f32]) { + let mean = frame.iter().sum::() / frame.len() as f32; + for sample in frame { + *sample -= mean; + } +} + +fn preemphasize(frame: &mut [f32]) { + for index in (1..frame.len()).rev() { + frame[index] -= PREEMPH_COEFF * frame[index - 1]; + } + frame[0] -= PREEMPH_COEFF * frame[0]; +} + +fn apply_window(frame: &mut [f32], window: &[f32; WESPEAKER_FRAME_LENGTH_SAMPLES]) { + for (sample, weight) in frame.iter_mut().zip(window) { + *sample *= *weight; + } +} + +fn power_spectrum_512(input: &mut [f32; WESPEAKER_FFT_SIZE]) -> [f32; WESPEAKER_FFT_SIZE / 2 + 1] { + let mut real = *input; + let mut imag = [0.0; WESPEAKER_FFT_SIZE]; + fft_512(&mut real, &mut imag); + let mut power = [0.0; WESPEAKER_FFT_SIZE / 2 + 1]; + for index in 0..power.len() { + power[index] = real[index] * real[index] + imag[index] * imag[index]; + } + power +} + +fn fft_512(real: &mut [f32; WESPEAKER_FFT_SIZE], imag: &mut [f32; WESPEAKER_FFT_SIZE]) { + let mut j = 0; + for i in 1..WESPEAKER_FFT_SIZE { + let mut bit = WESPEAKER_FFT_SIZE >> 1; + while j & bit != 0 { + j ^= bit; + bit >>= 1; + } + j ^= bit; + if i < j { + real.swap(i, j); + imag.swap(i, j); + } + } + + let mut len = 2; + while len <= WESPEAKER_FFT_SIZE { + let angle = -2.0 * PI / len as f32; + let wlen_real = angle.cos(); + let wlen_imag = angle.sin(); + for start in (0..WESPEAKER_FFT_SIZE).step_by(len) { + let mut w_real = 1.0; + let mut w_imag = 0.0; + for offset in 0..(len / 2) { + let even = start + offset; + let odd = even + len / 2; + let odd_real = real[odd] * w_real - imag[odd] * w_imag; + let odd_imag = real[odd] * w_imag + imag[odd] * w_real; + real[odd] = real[even] - odd_real; + imag[odd] = imag[even] - odd_imag; + real[even] += odd_real; + imag[even] += odd_imag; + let next_real = w_real * wlen_real - w_imag * wlen_imag; + w_imag = w_real * wlen_imag + w_imag * wlen_real; + w_real = next_real; + } + } + len *= 2; + } +} + +fn subtract_column_mean(frames: usize, bins: usize, data: &mut [f32]) { + for bin in 0..bins { + let mut sum = 0.0; + for frame in 0..frames { + sum += data[frame * bins + bin]; + } + let mean = sum / frames as f32; + for frame in 0..frames { + data[frame * bins + bin] -= mean; + } + } +} + +#[derive(Debug, Clone)] +struct MelBin { + offset: usize, + weights: Vec, +} + +#[derive(Debug, Clone)] +struct MelBank { + bins: Vec, +} + +impl MelBank { + fn production() -> Self { + let mel_low = mel_scale(LOW_FREQ_HZ); + let mel_high = mel_scale(HIGH_FREQ_HZ); + let delta = (mel_high - mel_low) / (WESPEAKER_MEL_BINS as f32 + 1.0); + let fft_bin_width = WESPEAKER_SAMPLE_RATE_HZ as f32 / WESPEAKER_FFT_SIZE as f32; + let mut bins = Vec::with_capacity(WESPEAKER_MEL_BINS); + for bin in 0..WESPEAKER_MEL_BINS { + let left = mel_low + bin as f32 * delta; + let center = mel_low + (bin as f32 + 1.0) * delta; + let right = mel_low + (bin as f32 + 2.0) * delta; + let mut dense = [0.0; WESPEAKER_FFT_SIZE / 2]; + let mut first = None; + let mut last = 0; + for (fft_bin, weight) in dense.iter_mut().enumerate() { + let freq = fft_bin_width * fft_bin as f32; + let mel = mel_scale(freq); + if mel > left && mel < right { + *weight = if mel <= center { + (mel - left) / (center - left) + } else { + (right - mel) / (right - center) + }; + first.get_or_insert(fft_bin); + last = fft_bin; + } + } + let first = first.expect("production mel bin should have weights"); + bins.push(MelBin { + offset: first, + weights: dense[first..=last].to_vec(), + }); + } + MelBank { bins } + } + + fn compute_log_energies(&self, power: &[f32; WESPEAKER_FFT_SIZE / 2 + 1], output: &mut [f32]) { + for (bin, out) in self.bins.iter().zip(output.iter_mut()) { + let mut energy = 0.0; + for (index, weight) in bin.weights.iter().enumerate() { + energy += weight * power[bin.offset + index]; + } + *out = energy.max(f32::EPSILON).ln(); + } + } + + #[cfg(test)] + fn weight_columns(&self) -> usize { + WESPEAKER_FFT_SIZE / 2 + } +} + +fn mel_scale(freq: f32) -> f32 { + 1127.0 * (1.0 + freq / 700.0).ln() +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::{ + assert_matrix_within, assert_region_within, decode_waveform, filterbank_fixture, + fixture_matrix, fixture_range, matrix_comparison_error, max_abs_diff, + source_tree_needle_count, + }; + + const TOLERANCE: f32 = 1e-2; + + #[test] + fn nyquist_power_bin_is_computed_but_not_consumed_by_mel_bank() { + let bank = MelBank::production(); + assert_eq!(bank.weight_columns(), WESPEAKER_FFT_SIZE / 2); + + let mut power = [0.0; WESPEAKER_FFT_SIZE / 2 + 1]; + let mut baseline = [0.0; WESPEAKER_MEL_BINS]; + bank.compute_log_energies(&power, &mut baseline); + + power[WESPEAKER_FFT_SIZE / 2] = 1.0e12; + let mut perturbed = [0.0; WESPEAKER_MEL_BINS]; + bank.compute_log_energies(&power, &mut perturbed); + + assert_eq!(baseline, perturbed); + } + + #[test] + fn povey_window_spans_400_samples_and_padded_tail_stays_zero() { + let window = povey_window(); + assert_eq!(window.len(), WESPEAKER_FRAME_LENGTH_SAMPLES); + + let audio = vec![0.25; WESPEAKER_FRAME_LENGTH_SAMPLES]; + let padded = padded_frame(&audio, 0, &window); + assert!( + padded[WESPEAKER_FRAME_LENGTH_SAMPLES..] + .iter() + .all(|sample| *sample == 0.0) + ); + } + + #[test] + fn filterbank_cmn_and_row_l2_match_fixture_regions_separately() { + let fixture = filterbank_fixture(); + let audio = decode_waveform(&fixture); + let cmn = + compute_wespeaker_filterbank_cmn(&audio, WESPEAKER_SAMPLE_RATE_HZ).expect("features"); + let row_l2 = row_l2_normalize(&cmn); + let expected_cmn = fixture_matrix(&fixture, "filterbank_cmn"); + let expected_row_l2 = fixture_matrix(&fixture, "row_l2_normalized"); + let near_rows = fixture_range(&fixture, "near_silent_rows"); + let broad_rows = fixture_range(&fixture, "broadband_rows"); + eprintln!( + "filterbank_cmn_max_abs_diff={}", + max_abs_diff(cmn.data(), &expected_cmn) + ); + + assert_matrix_within("filterbank_cmn", cmn.data(), &expected_cmn, TOLERANCE); + assert_matrix_within( + "row_l2_normalized", + row_l2.data(), + &expected_row_l2, + TOLERANCE, + ); + assert_region_within( + "filterbank_cmn near_silent", + cmn.data(), + &expected_cmn, + cmn.bins(), + near_rows.clone(), + TOLERANCE, + ); + assert_region_within( + "filterbank_cmn broadband", + cmn.data(), + &expected_cmn, + cmn.bins(), + broad_rows.clone(), + TOLERANCE, + ); + assert_region_within( + "row_l2_normalized near_silent", + row_l2.data(), + &expected_row_l2, + row_l2.bins(), + near_rows, + TOLERANCE, + ); + assert_region_within( + "row_l2_normalized broadband", + row_l2.data(), + &expected_row_l2, + row_l2.bins(), + broad_rows, + TOLERANCE, + ); + } + + #[test] + fn fixture_comparison_fails_when_value_exceeds_tolerance() { + let fixture = filterbank_fixture(); + let audio = decode_waveform(&fixture); + let cmn = + compute_wespeaker_filterbank_cmn(&audio, WESPEAKER_SAMPLE_RATE_HZ).expect("features"); + let mut expected = fixture_matrix(&fixture, "filterbank_cmn"); + expected[0] = cmn.data()[0] + TOLERANCE * 2.0; + + let result = matrix_comparison_error("filterbank_cmn", cmn.data(), &expected, TOLERANCE); + assert!(result.is_some()); + } + + #[test] + fn mel_bank_construction_has_single_production_site_scans_src_tree() { + let needle = concat!("MelBank", " { bins"); + let (visited_files, count) = source_tree_needle_count(needle); + + assert!( + visited_files >= 3, + "source walk visited too few Rust files: {visited_files}" + ); + assert_eq!(count, 1); + } +} diff --git a/core/crates/solstone-core-speakers/src/lib.rs b/core/crates/solstone-core-speakers/src/lib.rs index e17187303..587250cc0 100644 --- a/core/crates/solstone-core-speakers/src/lib.rs +++ b/core/crates/solstone-core-speakers/src/lib.rs @@ -1,24 +1,28 @@ // SPDX-License-Identifier: AGPL-3.0-only // Copyright (c) 2026 sol pbc +mod filterbank; +mod segmentation; + use std::error::Error; -use std::f32::consts::PI; use std::fmt; pub mod diarization; -pub const WESPEAKER_SAMPLE_RATE_HZ: u32 = 16_000; -pub const WESPEAKER_MEL_BINS: usize = 80; -pub const WESPEAKER_FRAME_LENGTH_SAMPLES: usize = 400; -pub const WESPEAKER_FRAME_SHIFT_SAMPLES: usize = 160; -pub const WESPEAKER_FFT_SIZE: usize = 512; -pub const WESPEAKER_EMBEDDING_SIZE: usize = 256; - -const LOW_FREQ_HZ: f32 = 20.0; -const HIGH_FREQ_HZ: f32 = WESPEAKER_SAMPLE_RATE_HZ as f32 / 2.0; -const PREEMPH_COEFF: f32 = 0.97; -const AUDIO_SCALE: f32 = 32768.0; -const ROW_NORM_EPSILON: f32 = 1e-9; +pub use filterbank::{ + WESPEAKER_EMBEDDING_SIZE, WESPEAKER_FFT_SIZE, WESPEAKER_FRAME_LENGTH_SAMPLES, + WESPEAKER_FRAME_SHIFT_SAMPLES, WESPEAKER_MEL_BINS, WESPEAKER_SAMPLE_RATE_HZ, + compute_wespeaker_filterbank_cmn, row_l2_normalize, +}; +pub use segmentation::{ + DIARIZE_MIN_OVERLAP, PYANNOTE_CLASS_COUNT, PYANNOTE_DIARIZE_STRIDE_S, + PYANNOTE_FRAMES_PER_WINDOW, PYANNOTE_OVERLAP_CLASSES, PYANNOTE_OVERLAP_STRIDE_S, + PYANNOTE_SAMPLE_RATE_HZ, PYANNOTE_SINGLE_SPEAKER_CLASSES, PYANNOTE_WINDOW_S, + PyannoteSegmentationPassResult, SLOT_ACTIVE_MIN_SHARE, SPEAKER_EVIDENCE_MULTI_MIN, + SPEAKER_EVIDENCE_SINGLE_MAX, SpeakerEvidence, SpeakerEvidenceDecision, + SpeakerSegmentationError, SpeakerWindowStats, compute_speaker_window_stats, + decide_speaker_evidence, run_pyannote_segmentation_pass, +}; #[derive(Debug, Clone, PartialEq)] pub struct FeatureMatrix { @@ -113,366 +117,25 @@ impl fmt::Display for SpeakerFeatureError { impl Error for SpeakerFeatureError {} -pub fn compute_wespeaker_filterbank_cmn( - audio: &[f32], - sample_rate_hz: u32, -) -> Result { - if sample_rate_hz != WESPEAKER_SAMPLE_RATE_HZ { - return Err(SpeakerFeatureError::UnsupportedSampleRate { - expected: WESPEAKER_SAMPLE_RATE_HZ, - actual: sample_rate_hz, - }); - } - if let Some((index, _sample)) = audio - .iter() - .enumerate() - .find(|(_index, sample)| !sample.is_finite()) - { - return Err(SpeakerFeatureError::NonFiniteAudioSample { index }); - } - - let frames = num_snipped_frames(audio.len()); - let mut data = vec![0.0; frames * WESPEAKER_MEL_BINS]; - if frames == 0 { - return FeatureMatrix::from_row_major(0, WESPEAKER_MEL_BINS, data); - } - - let window = povey_window(); - let mel_bank = MelBank::production(); - for frame in 0..frames { - let mut padded = padded_frame(audio, frame, &window); - let power = power_spectrum_512(&mut padded); - let output_start = frame * WESPEAKER_MEL_BINS; - mel_bank.compute_log_energies( - &power, - &mut data[output_start..output_start + WESPEAKER_MEL_BINS], - ); - } - subtract_column_mean(frames, WESPEAKER_MEL_BINS, &mut data); - FeatureMatrix::from_row_major(frames, WESPEAKER_MEL_BINS, data) -} - -pub fn row_l2_normalize(features: &FeatureMatrix) -> FeatureMatrix { - let mut data = features.data.clone(); - for row in data.chunks_mut(features.bins) { - let norm = row.iter().map(|value| value * value).sum::().sqrt(); - let denom = if norm > ROW_NORM_EPSILON { norm } else { 1.0 }; - for value in row { - *value /= denom; - } - } - FeatureMatrix { - frames: features.frames, - bins: features.bins, - data, - } -} - -fn num_snipped_frames(samples: usize) -> usize { - if samples < WESPEAKER_FRAME_LENGTH_SAMPLES { - 0 - } else { - 1 + (samples - WESPEAKER_FRAME_LENGTH_SAMPLES) / WESPEAKER_FRAME_SHIFT_SAMPLES - } -} - -fn povey_window() -> [f32; WESPEAKER_FRAME_LENGTH_SAMPLES] { - let mut window = [0.0; WESPEAKER_FRAME_LENGTH_SAMPLES]; - let scale = 2.0_f64 * std::f64::consts::PI / (WESPEAKER_FRAME_LENGTH_SAMPLES as f64 - 1.0); - for (index, value) in window.iter_mut().enumerate() { - *value = (0.5_f64 - 0.5_f64 * (scale * index as f64).cos()).powf(0.85) as f32; - } - window -} - -fn padded_frame( - audio: &[f32], - frame: usize, - window: &[f32; WESPEAKER_FRAME_LENGTH_SAMPLES], -) -> [f32; WESPEAKER_FFT_SIZE] { - let mut padded = [0.0; WESPEAKER_FFT_SIZE]; - let start = frame * WESPEAKER_FRAME_SHIFT_SAMPLES; - for index in 0..WESPEAKER_FRAME_LENGTH_SAMPLES { - padded[index] = audio[start + index] * AUDIO_SCALE; - } - remove_dc_offset(&mut padded[..WESPEAKER_FRAME_LENGTH_SAMPLES]); - preemphasize(&mut padded[..WESPEAKER_FRAME_LENGTH_SAMPLES]); - apply_window(&mut padded[..WESPEAKER_FRAME_LENGTH_SAMPLES], window); - padded -} - -fn remove_dc_offset(frame: &mut [f32]) { - let mean = frame.iter().sum::() / frame.len() as f32; - for sample in frame { - *sample -= mean; - } -} - -fn preemphasize(frame: &mut [f32]) { - for index in (1..frame.len()).rev() { - frame[index] -= PREEMPH_COEFF * frame[index - 1]; - } - frame[0] -= PREEMPH_COEFF * frame[0]; -} - -fn apply_window(frame: &mut [f32], window: &[f32; WESPEAKER_FRAME_LENGTH_SAMPLES]) { - for (sample, weight) in frame.iter_mut().zip(window) { - *sample *= *weight; - } -} - -fn power_spectrum_512(input: &mut [f32; WESPEAKER_FFT_SIZE]) -> [f32; WESPEAKER_FFT_SIZE / 2 + 1] { - let mut real = *input; - let mut imag = [0.0; WESPEAKER_FFT_SIZE]; - fft_512(&mut real, &mut imag); - let mut power = [0.0; WESPEAKER_FFT_SIZE / 2 + 1]; - for index in 0..power.len() { - power[index] = real[index] * real[index] + imag[index] * imag[index]; - } - power -} - -fn fft_512(real: &mut [f32; WESPEAKER_FFT_SIZE], imag: &mut [f32; WESPEAKER_FFT_SIZE]) { - let mut j = 0; - for i in 1..WESPEAKER_FFT_SIZE { - let mut bit = WESPEAKER_FFT_SIZE >> 1; - while j & bit != 0 { - j ^= bit; - bit >>= 1; - } - j ^= bit; - if i < j { - real.swap(i, j); - imag.swap(i, j); - } - } - - let mut len = 2; - while len <= WESPEAKER_FFT_SIZE { - let angle = -2.0 * PI / len as f32; - let wlen_real = angle.cos(); - let wlen_imag = angle.sin(); - for start in (0..WESPEAKER_FFT_SIZE).step_by(len) { - let mut w_real = 1.0; - let mut w_imag = 0.0; - for offset in 0..(len / 2) { - let even = start + offset; - let odd = even + len / 2; - let odd_real = real[odd] * w_real - imag[odd] * w_imag; - let odd_imag = real[odd] * w_imag + imag[odd] * w_real; - real[odd] = real[even] - odd_real; - imag[odd] = imag[even] - odd_imag; - real[even] += odd_real; - imag[even] += odd_imag; - let next_real = w_real * wlen_real - w_imag * wlen_imag; - w_imag = w_real * wlen_imag + w_imag * wlen_real; - w_real = next_real; - } - } - len *= 2; - } -} - -fn subtract_column_mean(frames: usize, bins: usize, data: &mut [f32]) { - for bin in 0..bins { - let mut sum = 0.0; - for frame in 0..frames { - sum += data[frame * bins + bin]; - } - let mean = sum / frames as f32; - for frame in 0..frames { - data[frame * bins + bin] -= mean; - } - } -} - -#[derive(Debug, Clone)] -struct MelBin { - offset: usize, - weights: Vec, -} - -#[derive(Debug, Clone)] -struct MelBank { - bins: Vec, -} - -impl MelBank { - fn production() -> Self { - let mel_low = mel_scale(LOW_FREQ_HZ); - let mel_high = mel_scale(HIGH_FREQ_HZ); - let delta = (mel_high - mel_low) / (WESPEAKER_MEL_BINS as f32 + 1.0); - let fft_bin_width = WESPEAKER_SAMPLE_RATE_HZ as f32 / WESPEAKER_FFT_SIZE as f32; - let mut bins = Vec::with_capacity(WESPEAKER_MEL_BINS); - for bin in 0..WESPEAKER_MEL_BINS { - let left = mel_low + bin as f32 * delta; - let center = mel_low + (bin as f32 + 1.0) * delta; - let right = mel_low + (bin as f32 + 2.0) * delta; - let mut dense = [0.0; WESPEAKER_FFT_SIZE / 2]; - let mut first = None; - let mut last = 0; - for (fft_bin, weight) in dense.iter_mut().enumerate() { - let freq = fft_bin_width * fft_bin as f32; - let mel = mel_scale(freq); - if mel > left && mel < right { - *weight = if mel <= center { - (mel - left) / (center - left) - } else { - (right - mel) / (right - center) - }; - first.get_or_insert(fft_bin); - last = fft_bin; - } - } - let first = first.expect("production mel bin should have weights"); - bins.push(MelBin { - offset: first, - weights: dense[first..=last].to_vec(), - }); - } - MelBank { bins } - } - - fn compute_log_energies(&self, power: &[f32; WESPEAKER_FFT_SIZE / 2 + 1], output: &mut [f32]) { - for (bin, out) in self.bins.iter().zip(output.iter_mut()) { - let mut energy = 0.0; - for (index, weight) in bin.weights.iter().enumerate() { - energy += weight * power[bin.offset + index]; - } - *out = energy.max(f32::EPSILON).ln(); - } - } - - #[cfg(test)] - fn weight_columns(&self) -> usize { - WESPEAKER_FFT_SIZE / 2 - } -} - -fn mel_scale(freq: f32) -> f32 { - 1127.0 * (1.0 + freq / 700.0).ln() -} - #[cfg(test)] -mod tests { - use super::*; +pub(crate) mod test_support { use serde_json::Value; + use std::path::{Path, PathBuf}; - const FIXTURE: &str = include_str!("../../../fixtures/speaker_filterbank.json"); - const TOLERANCE: f32 = 1e-2; - - #[test] - fn nyquist_power_bin_is_computed_but_not_consumed_by_mel_bank() { - let bank = MelBank::production(); - assert_eq!(bank.weight_columns(), WESPEAKER_FFT_SIZE / 2); + pub(crate) const FILTERBANK_FIXTURE: &str = + include_str!("../../../fixtures/speaker_filterbank.json"); + pub(crate) const STAGE_FIXTURE: &str = + include_str!("../../../fixtures/speaker_stage_boundaries.json"); - let mut power = [0.0; WESPEAKER_FFT_SIZE / 2 + 1]; - let mut baseline = [0.0; WESPEAKER_MEL_BINS]; - bank.compute_log_energies(&power, &mut baseline); - - power[WESPEAKER_FFT_SIZE / 2] = 1.0e12; - let mut perturbed = [0.0; WESPEAKER_MEL_BINS]; - bank.compute_log_energies(&power, &mut perturbed); - - assert_eq!(baseline, perturbed); - } - - #[test] - fn povey_window_spans_400_samples_and_padded_tail_stays_zero() { - let window = povey_window(); - assert_eq!(window.len(), WESPEAKER_FRAME_LENGTH_SAMPLES); - - let audio = vec![0.25; WESPEAKER_FRAME_LENGTH_SAMPLES]; - let padded = padded_frame(&audio, 0, &window); - assert!( - padded[WESPEAKER_FRAME_LENGTH_SAMPLES..] - .iter() - .all(|sample| *sample == 0.0) - ); - } - - #[test] - fn filterbank_cmn_and_row_l2_match_fixture_regions_separately() { - let fixture = fixture(); - let audio = decode_waveform(&fixture); - let cmn = - compute_wespeaker_filterbank_cmn(&audio, WESPEAKER_SAMPLE_RATE_HZ).expect("features"); - let row_l2 = row_l2_normalize(&cmn); - let expected_cmn = fixture_matrix(&fixture, "filterbank_cmn"); - let expected_row_l2 = fixture_matrix(&fixture, "row_l2_normalized"); - let near_rows = fixture_range(&fixture, "near_silent_rows"); - let broad_rows = fixture_range(&fixture, "broadband_rows"); - eprintln!( - "filterbank_cmn_max_abs_diff={}", - max_abs_diff(cmn.data(), &expected_cmn) - ); - - assert_matrix_within("filterbank_cmn", cmn.data(), &expected_cmn, TOLERANCE); - assert_matrix_within( - "row_l2_normalized", - row_l2.data(), - &expected_row_l2, - TOLERANCE, - ); - assert_region_within( - "filterbank_cmn near_silent", - cmn.data(), - &expected_cmn, - cmn.bins(), - near_rows.clone(), - TOLERANCE, - ); - assert_region_within( - "filterbank_cmn broadband", - cmn.data(), - &expected_cmn, - cmn.bins(), - broad_rows.clone(), - TOLERANCE, - ); - assert_region_within( - "row_l2_normalized near_silent", - row_l2.data(), - &expected_row_l2, - row_l2.bins(), - near_rows, - TOLERANCE, - ); - assert_region_within( - "row_l2_normalized broadband", - row_l2.data(), - &expected_row_l2, - row_l2.bins(), - broad_rows, - TOLERANCE, - ); - } - - #[test] - fn fixture_comparison_fails_when_value_exceeds_tolerance() { - let fixture = fixture(); - let audio = decode_waveform(&fixture); - let cmn = - compute_wespeaker_filterbank_cmn(&audio, WESPEAKER_SAMPLE_RATE_HZ).expect("features"); - let mut expected = fixture_matrix(&fixture, "filterbank_cmn"); - expected[0] = cmn.data()[0] + TOLERANCE * 2.0; - - let result = matrix_comparison_error("filterbank_cmn", cmn.data(), &expected, TOLERANCE); - assert!(result.is_some()); + pub(crate) fn filterbank_fixture() -> Value { + serde_json::from_str(FILTERBANK_FIXTURE).expect("filterbank fixture JSON") } - #[test] - fn mel_bank_construction_has_single_production_site() { - let source = include_str!("lib.rs"); - let needle = concat!("MelBank", " { bins"); - assert_eq!(source.matches(needle).count(), 1); + pub(crate) fn stage_fixture() -> Value { + serde_json::from_str(STAGE_FIXTURE).expect("stage fixture JSON") } - fn fixture() -> Value { - serde_json::from_str(FIXTURE).expect("fixture JSON") - } - - fn decode_waveform(fixture: &Value) -> Vec { + pub(crate) fn decode_waveform(fixture: &Value) -> Vec { let encoded = fixture["waveform"]["samples_f32_le_base64"] .as_str() .expect("waveform base64"); @@ -483,7 +146,7 @@ mod tests { .collect() } - fn fixture_matrix(fixture: &Value, name: &str) -> Vec { + pub(crate) fn fixture_matrix(fixture: &Value, name: &str) -> Vec { let rows = fixture["matrices"][name]["rows"] .as_array() .expect("matrix rows"); @@ -497,20 +160,25 @@ mod tests { .collect() } - fn fixture_range(fixture: &Value, name: &str) -> std::ops::Range { + pub(crate) fn fixture_range(fixture: &Value, name: &str) -> std::ops::Range { let values = fixture["waveform"][name].as_array().expect("row range"); let start = values[0].as_u64().expect("range start") as usize; let end = values[1].as_u64().expect("range end") as usize; start..end } - fn assert_matrix_within(name: &str, actual: &[f32], expected: &[f32], tolerance: f32) { + pub(crate) fn assert_matrix_within( + name: &str, + actual: &[f32], + expected: &[f32], + tolerance: f32, + ) { if let Some(error) = matrix_comparison_error(name, actual, expected, tolerance) { panic!("{error}"); } } - fn assert_region_within( + pub(crate) fn assert_region_within( name: &str, actual: &[f32], expected: &[f32], @@ -523,7 +191,7 @@ mod tests { assert_matrix_within(name, &actual[start..end], &expected[start..end], tolerance); } - fn matrix_comparison_error( + pub(crate) fn matrix_comparison_error( name: &str, actual: &[f32], expected: &[f32], @@ -546,7 +214,7 @@ mod tests { } } - fn max_abs_diff(actual: &[f32], expected: &[f32]) -> f32 { + pub(crate) fn max_abs_diff(actual: &[f32], expected: &[f32]) -> f32 { max_abs_diff_with_index(actual, expected).0 } @@ -563,7 +231,7 @@ mod tests { (max_abs, max_index) } - fn decode_base64(input: &str) -> Vec { + pub(crate) fn decode_base64(input: &str) -> Vec { let mut out = Vec::with_capacity(input.len() / 4 * 3); let mut quartet = [0_u8; 4]; let mut len = 0; @@ -592,4 +260,34 @@ mod tests { assert_eq!(len, 0); out } + + pub(crate) fn source_tree_needle_count(needle: &str) -> (usize, usize) { + let src = Path::new(env!("CARGO_MANIFEST_DIR")).join("src"); + let mut files = Vec::new(); + collect_rust_files(&src, &mut files); + let count = files + .iter() + .map(|path| { + std::fs::read_to_string(path) + .unwrap_or_else(|error| panic!("failed to read {}: {error}", path.display())) + .matches(needle) + .count() + }) + .sum(); + (files.len(), count) + } + + fn collect_rust_files(path: &Path, files: &mut Vec) { + for entry in std::fs::read_dir(path) + .unwrap_or_else(|error| panic!("failed to read {}: {error}", path.display())) + { + let entry = entry.expect("directory entry"); + let path = entry.path(); + if path.is_dir() { + collect_rust_files(&path, files); + } else if path.extension().and_then(|value| value.to_str()) == Some("rs") { + files.push(path); + } + } + } } diff --git a/core/crates/solstone-core-speakers/src/segmentation.rs b/core/crates/solstone-core-speakers/src/segmentation.rs new file mode 100644 index 000000000..91eda1b02 --- /dev/null +++ b/core/crates/solstone-core-speakers/src/segmentation.rs @@ -0,0 +1,958 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::borrow::Cow; +use std::error::Error; +use std::fmt; + +use crate::{FeatureMatrix, SpeakerFeatureError}; + +pub const PYANNOTE_SAMPLE_RATE_HZ: u32 = 16_000; +pub const PYANNOTE_WINDOW_S: u32 = 10; +pub const PYANNOTE_OVERLAP_STRIDE_S: u32 = 5; +pub const PYANNOTE_DIARIZE_STRIDE_S: u32 = 2; +pub const PYANNOTE_FRAMES_PER_WINDOW: usize = 589; +pub const PYANNOTE_CLASS_COUNT: usize = 7; +pub const PYANNOTE_OVERLAP_CLASSES: [usize; 3] = [4, 5, 6]; +pub const PYANNOTE_SINGLE_SPEAKER_CLASSES: [usize; 3] = [1, 2, 3]; +pub const SLOT_ACTIVE_MIN_SHARE: f64 = 0.10; +pub const SPEAKER_EVIDENCE_MULTI_MIN: f64 = 0.05; +pub const SPEAKER_EVIDENCE_SINGLE_MAX: f64 = 0.05; +pub const DIARIZE_MIN_OVERLAP: f64 = 0.05; + +#[derive(Debug, Clone, PartialEq)] +pub struct PyannoteSegmentationPassResult { + pub overlap_fraction: f64, + pub avg_log_probs: FeatureMatrix, + pub window_stats: Vec, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct SpeakerWindowStats { + pub speech_frames: usize, + pub active_slot_count: usize, + pub overlap_frames: usize, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum SpeakerEvidence { + NoSpeech, + Single, + Multi, +} + +impl SpeakerEvidence { + pub fn as_str(&self) -> &'static str { + match self { + Self::NoSpeech => "none", + Self::Single => "single", + Self::Multi => "multi", + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq)] +pub struct SpeakerEvidenceDecision { + pub speaker_evidence: SpeakerEvidence, + pub multi_window_fraction: f64, + pub mean_window_overlap_share: f64, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum SpeakerSegmentationError { + UnsupportedSampleRate { + expected: u32, + actual: u32, + }, + InvalidStride { + stride_s: u32, + }, + NonFiniteAudioSample { + index: usize, + }, + ShapeOverflow { + frames: usize, + classes: usize, + }, + WindowLogProbShapeMismatch { + window_index: usize, + expected_frames: usize, + expected_classes: usize, + actual_frames: usize, + actual_classes: usize, + actual_len: usize, + }, + Inference { + window_index: usize, + source: E, + }, +} + +impl fmt::Display for SpeakerSegmentationError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::UnsupportedSampleRate { expected, actual } => write!( + formatter, + "pyannote segmentation requires sample rate {expected}, got {actual}" + ), + Self::InvalidStride { stride_s } => { + write!( + formatter, + "pyannote segmentation stride must be positive, got {stride_s}" + ) + } + Self::NonFiniteAudioSample { index } => { + write!(formatter, "audio sample at index {index} is not finite") + } + Self::ShapeOverflow { frames, classes } => write!( + formatter, + "pyannote segmentation matrix shape overflow: frames={frames} classes={classes}" + ), + Self::WindowLogProbShapeMismatch { + window_index, + expected_frames, + expected_classes, + actual_frames, + actual_classes, + actual_len, + } => write!( + formatter, + "pyannote window {window_index} log-prob shape mismatch: expected frames={expected_frames} classes={expected_classes}, got frames={actual_frames} classes={actual_classes} len={actual_len}" + ), + Self::Inference { + window_index, + source, + } => write!( + formatter, + "pyannote window {window_index} inference failed: {source}" + ), + } + } +} + +impl Error for SpeakerSegmentationError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + Self::Inference { source, .. } => Some(source), + _ => None, + } + } +} + +pub fn run_pyannote_segmentation_pass( + audio: &[f32], + sample_rate_hz: u32, + stride_s: u32, + mut infer_window: F, +) -> Result> +where + F: FnMut(usize, &[f32]) -> Result, +{ + if sample_rate_hz != PYANNOTE_SAMPLE_RATE_HZ { + return Err(SpeakerSegmentationError::UnsupportedSampleRate { + expected: PYANNOTE_SAMPLE_RATE_HZ, + actual: sample_rate_hz, + }); + } + if stride_s == 0 { + return Err(SpeakerSegmentationError::InvalidStride { stride_s }); + } + if let Some((index, _sample)) = audio + .iter() + .enumerate() + .find(|(_index, sample)| !sample.is_finite()) + { + return Err(SpeakerSegmentationError::NonFiniteAudioSample { index }); + } + + let window_samples = PYANNOTE_WINDOW_S as usize * sample_rate_hz as usize; + let stride_samples = stride_s as usize * sample_rate_hz as usize; + let audio_padded: Cow<'_, [f32]> = if audio.len() < window_samples { + let mut padded = Vec::with_capacity(window_samples); + padded.extend_from_slice(audio); + padded.resize(window_samples, 0.0); + Cow::Owned(padded) + } else { + Cow::Borrowed(audio) + }; + let starts = window_starts(audio_padded.len(), window_samples, stride_samples); + let samples_per_frame = window_samples as f64 / PYANNOTE_FRAMES_PER_WINDOW as f64; + let num_frames = (audio_padded.len() as f64 / samples_per_frame).ceil() as usize; + let len = num_frames.checked_mul(PYANNOTE_CLASS_COUNT).ok_or( + SpeakerSegmentationError::ShapeOverflow { + frames: num_frames, + classes: PYANNOTE_CLASS_COUNT, + }, + )?; + let mut accum = vec![0.0_f64; len]; + let mut counts = vec![0_usize; num_frames]; + let mut window_stats = Vec::with_capacity(starts.len()); + + for (window_index, start_sample) in starts.iter().copied().enumerate() { + let chunk = &audio_padded[start_sample..start_sample + window_samples]; + let log_probs = infer_window(window_index, chunk).map_err(|source| { + SpeakerSegmentationError::Inference { + window_index, + source, + } + })?; + validate_window_log_probs(window_index, &log_probs)?; + let frame_start = frame_start_for_sample(start_sample, samples_per_frame); + let requested_frame_end = frame_start + log_probs.frames(); + let frame_end = requested_frame_end.min(num_frames); + let used_frames = frame_end.saturating_sub(frame_start); + let stats_matrix = truncated_log_probs(&log_probs, used_frames) + .expect("validated pyannote log-probs can be truncated"); + window_stats.push( + compute_speaker_window_stats(&stats_matrix) + .expect("validated pyannote log-probs have the expected class count"), + ); + + for local_frame in 0..used_frames { + let global_frame = frame_start + local_frame; + let source_start = local_frame * PYANNOTE_CLASS_COUNT; + let target_start = global_frame * PYANNOTE_CLASS_COUNT; + for class in 0..PYANNOTE_CLASS_COUNT { + accum[target_start + class] += log_probs.data()[source_start + class] as f64; + } + counts[global_frame] += 1; + } + } + + let avg_log_probs = + average_accumulated_log_probs(&accum, &counts, num_frames, PYANNOTE_CLASS_COUNT) + .expect("accumulator shape is constructed from checked dimensions"); + let overlap_fraction = compute_overlap_fraction(&avg_log_probs); + Ok(PyannoteSegmentationPassResult { + overlap_fraction, + avg_log_probs, + window_stats, + }) +} + +pub fn compute_speaker_window_stats( + log_probs: &FeatureMatrix, +) -> Result { + if log_probs.bins() != PYANNOTE_CLASS_COUNT { + return Err(SpeakerFeatureError::ShapeMismatch { + frames: log_probs.frames(), + bins: PYANNOTE_CLASS_COUNT, + len: log_probs.data().len(), + }); + } + + let mut counts = [0_usize; PYANNOTE_CLASS_COUNT]; + for frame in 0..log_probs.frames() { + let row = log_probs.row(frame).expect("frame index is in bounds"); + counts[argmax(row)] += 1; + } + let speech_frames = counts[1..].iter().sum(); + if speech_frames == 0 { + return Ok(SpeakerWindowStats { + speech_frames: 0, + active_slot_count: 0, + overlap_frames: 0, + }); + } + let active_slot_count = counts[1..4] + .iter() + .filter(|count| (**count as f64 / speech_frames as f64) >= SLOT_ACTIVE_MIN_SHARE) + .count(); + let overlap_frames = PYANNOTE_OVERLAP_CLASSES + .iter() + .map(|class| counts[*class]) + .sum(); + Ok(SpeakerWindowStats { + speech_frames, + active_slot_count, + overlap_frames, + }) +} + +pub fn decide_speaker_evidence( + overlap_fraction: f64, + window_stats: &[SpeakerWindowStats], +) -> SpeakerEvidenceDecision { + let speech_windows: Vec<&SpeakerWindowStats> = window_stats + .iter() + .filter(|row| row.speech_frames > 0) + .collect(); + if speech_windows.is_empty() { + return SpeakerEvidenceDecision { + speaker_evidence: SpeakerEvidence::NoSpeech, + multi_window_fraction: 0.0, + mean_window_overlap_share: 0.0, + }; + } + + let multi_window_count = speech_windows + .iter() + .filter(|row| row.active_slot_count > 1) + .count(); + let multi_window_fraction = multi_window_count as f64 / speech_windows.len() as f64; + let mean_window_overlap_share = speech_windows + .iter() + .map(|row| row.overlap_frames as f64 / row.speech_frames as f64) + .sum::() + / speech_windows.len() as f64; + + let speaker_evidence = if multi_window_fraction >= SPEAKER_EVIDENCE_MULTI_MIN + || overlap_fraction >= DIARIZE_MIN_OVERLAP + { + SpeakerEvidence::Multi + } else if multi_window_fraction < SPEAKER_EVIDENCE_SINGLE_MAX + && mean_window_overlap_share < DIARIZE_MIN_OVERLAP + { + SpeakerEvidence::Single + } else { + SpeakerEvidence::Multi + }; + + SpeakerEvidenceDecision { + speaker_evidence, + multi_window_fraction, + mean_window_overlap_share, + } +} + +fn window_starts(len_padded: usize, window_samples: usize, stride_samples: usize) -> Vec { + let mut starts = Vec::new(); + let mut start = 0_usize; + while start + window_samples <= len_padded { + starts.push(start); + start = match start.checked_add(stride_samples) { + Some(next) => next, + None => break, + }; + } + let final_start = len_padded.saturating_sub(window_samples); + if starts.last().copied() != Some(final_start) { + starts.push(final_start); + } + starts +} + +fn frame_start_for_sample(start_sample: usize, samples_per_frame: f64) -> usize { + (start_sample as f64 / samples_per_frame).round_ties_even() as usize +} + +fn validate_window_log_probs( + window_index: usize, + log_probs: &FeatureMatrix, +) -> Result<(), SpeakerSegmentationError> { + if log_probs.frames() != PYANNOTE_FRAMES_PER_WINDOW || log_probs.bins() != PYANNOTE_CLASS_COUNT + { + return Err(SpeakerSegmentationError::WindowLogProbShapeMismatch { + window_index, + expected_frames: PYANNOTE_FRAMES_PER_WINDOW, + expected_classes: PYANNOTE_CLASS_COUNT, + actual_frames: log_probs.frames(), + actual_classes: log_probs.bins(), + actual_len: log_probs.data().len(), + }); + } + Ok(()) +} + +fn truncated_log_probs( + log_probs: &FeatureMatrix, + used_frames: usize, +) -> Result { + if used_frames == log_probs.frames() { + return Ok(log_probs.clone()); + } + let len = + used_frames + .checked_mul(log_probs.bins()) + .ok_or(SpeakerFeatureError::ShapeOverflow { + frames: used_frames, + bins: log_probs.bins(), + })?; + FeatureMatrix::from_row_major( + used_frames, + log_probs.bins(), + log_probs.data()[..len].to_vec(), + ) +} + +fn average_accumulated_log_probs( + accum: &[f64], + counts: &[usize], + frames: usize, + classes: usize, +) -> Result { + let expected = frames + .checked_mul(classes) + .ok_or(SpeakerFeatureError::ShapeOverflow { + frames, + bins: classes, + })?; + if accum.len() != expected { + return Err(SpeakerFeatureError::ShapeMismatch { + frames, + bins: classes, + len: accum.len(), + }); + } + let mut data = Vec::with_capacity(expected); + for frame in 0..frames { + let count = counts.get(frame).copied().unwrap_or(0).max(1) as f64; + for class in 0..classes { + data.push((accum[frame * classes + class] / count) as f32); + } + } + FeatureMatrix::from_row_major(frames, classes, data) +} + +fn compute_overlap_fraction(avg_log_probs: &FeatureMatrix) -> f64 { + let mut speech_count = 0_usize; + let mut overlap_count = 0_usize; + for frame in 0..avg_log_probs.frames() { + let row = avg_log_probs.row(frame).expect("frame index is in bounds"); + let class = argmax(row); + if class >= 1 { + speech_count += 1; + if PYANNOTE_OVERLAP_CLASSES.contains(&class) { + overlap_count += 1; + } + } + } + if speech_count == 0 { + 0.0 + } else { + overlap_count as f64 / speech_count as f64 + } +} + +fn argmax(row: &[f32]) -> usize { + let mut max_index = 0; + let mut max_value = row[0]; + for (index, value) in row.iter().enumerate().skip(1) { + if *value > max_value { + max_index = index; + max_value = *value; + } + } + max_index +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::test_support::{matrix_comparison_error, stage_fixture}; + use serde_json::Value; + use std::convert::Infallible; + + #[derive(Debug, Clone, PartialEq, Eq)] + struct StubError; + + impl fmt::Display for StubError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("stub error") + } + } + + impl Error for StubError {} + + #[test] + fn segmentation_constants_match_fixture_identity() { + let fixture = stage_fixture(); + let constants = &fixture["identity"]["source_constants"]; + + assert_eq!( + PYANNOTE_SAMPLE_RATE_HZ as u64, + constants["diarize"]["SAMPLE_RATE"] + .as_u64() + .expect("sample rate") + ); + assert_eq!( + PYANNOTE_WINDOW_S as u64, + constants["overlap"]["WINDOW_S"] + .as_u64() + .expect("window seconds") + ); + assert_eq!( + PYANNOTE_OVERLAP_STRIDE_S as u64, + constants["overlap"]["STRIDE_S"] + .as_u64() + .expect("overlap stride") + ); + assert_eq!( + PYANNOTE_DIARIZE_STRIDE_S as u64, + constants["overlap"]["_DIARIZE_STRIDE_S"] + .as_u64() + .expect("diarize stride") + ); + assert_eq!( + PYANNOTE_FRAMES_PER_WINDOW as u64, + constants["overlap"]["FRAMES_PER_WINDOW"] + .as_u64() + .expect("frames per window") + ); + assert_eq!( + usize_array(&constants["overlap"]["OVERLAP_CLASSES"]), + PYANNOTE_OVERLAP_CLASSES + ); + assert_eq!( + usize_array(&constants["diarize"]["SINGLE_SPEAKER_CLASSES"]), + PYANNOTE_SINGLE_SPEAKER_CLASSES + ); + assert_eq!( + SLOT_ACTIVE_MIN_SHARE, + constants["encoder_config"]["SLOT_ACTIVE_MIN_SHARE"] + .as_f64() + .expect("slot active min share") + ); + assert_eq!( + SPEAKER_EVIDENCE_MULTI_MIN, + constants["encoder_config"]["SPEAKER_EVIDENCE_MULTI_MIN"] + .as_f64() + .expect("speaker evidence multi min") + ); + assert_eq!( + SPEAKER_EVIDENCE_SINGLE_MAX, + constants["encoder_config"]["SPEAKER_EVIDENCE_SINGLE_MAX"] + .as_f64() + .expect("speaker evidence single max") + ); + assert_eq!( + DIARIZE_MIN_OVERLAP, + constants["encoder_config"]["DIARIZE_MIN_OVERLAP"] + .as_f64() + .expect("diarize min overlap") + ); + } + + #[test] + fn speaker_evidence_wire_strings_match_fixture_decisions() { + let fixture = stage_fixture(); + let cases = &fixture["speaker_evidence"]; + + assert_eq!( + SpeakerEvidence::NoSpeech.as_str(), + cases["none"]["decision"]["speaker_evidence"] + .as_str() + .expect("none decision") + ); + assert_eq!( + SpeakerEvidence::Single.as_str(), + cases["single"]["decision"]["speaker_evidence"] + .as_str() + .expect("single decision") + ); + assert_eq!( + SpeakerEvidence::Multi.as_str(), + cases["multi_by_active_slots"]["decision"]["speaker_evidence"] + .as_str() + .expect("multi decision") + ); + } + + #[test] + fn start_sequence_stride_aligned_buffer_length() { + let audio = indexed_audio(20 * PYANNOTE_SAMPLE_RATE_HZ as usize); + let mut starts = Vec::new(); + + run_pyannote_segmentation_pass( + &audio, + PYANNOTE_SAMPLE_RATE_HZ, + PYANNOTE_OVERLAP_STRIDE_S, + |_, window| { + starts.push(window[0] as usize); + Ok::(class_matrix(&vec![0; PYANNOTE_FRAMES_PER_WINDOW])) + }, + ) + .expect("segmentation pass"); + + assert_eq!(starts, vec![0, 80_000, 160_000]); + } + + #[test] + fn start_sequence_non_stride_aligned_buffer_appends_final_window() { + let audio = indexed_audio(13 * PYANNOTE_SAMPLE_RATE_HZ as usize); + let mut starts = Vec::new(); + + run_pyannote_segmentation_pass( + &audio, + PYANNOTE_SAMPLE_RATE_HZ, + PYANNOTE_OVERLAP_STRIDE_S, + |_, window| { + starts.push(window[0] as usize); + Ok::(class_matrix(&vec![0; PYANNOTE_FRAMES_PER_WINDOW])) + }, + ) + .expect("segmentation pass"); + + assert_eq!(starts, vec![0, 48_000]); + } + + #[test] + fn buffer_shorter_than_one_window_is_zero_padded() { + let audio = vec![0.25; 3 * PYANNOTE_SAMPLE_RATE_HZ as usize]; + let mut inspected = false; + + run_pyannote_segmentation_pass( + &audio, + PYANNOTE_SAMPLE_RATE_HZ, + PYANNOTE_OVERLAP_STRIDE_S, + |_, window| { + inspected = true; + assert_eq!(window.len(), 10 * PYANNOTE_SAMPLE_RATE_HZ as usize); + assert!(window[..audio.len()].iter().all(|sample| *sample == 0.25)); + assert!(window[audio.len()..].iter().all(|sample| *sample == 0.0)); + Ok::(class_matrix(&vec![0; PYANNOTE_FRAMES_PER_WINDOW])) + }, + ) + .expect("segmentation pass"); + + assert!(inspected); + } + + #[test] + fn frame_start_uses_round_ties_even_for_stride5_294_5_tie() { + let window_samples = PYANNOTE_WINDOW_S as usize * PYANNOTE_SAMPLE_RATE_HZ as usize; + let samples_per_frame = window_samples as f64 / PYANNOTE_FRAMES_PER_WINDOW as f64; + let start_sample = PYANNOTE_OVERLAP_STRIDE_S as usize * PYANNOTE_SAMPLE_RATE_HZ as usize; + let ratio = start_sample as f64 / samples_per_frame; + + assert_eq!(start_sample, 80_000); + assert_eq!(ratio, 294.5); + // Python round() uses banker's rounding, so this must be ties-to-even. + // Rust f64::round() is half-away-from-zero and would shift this window to 295. + assert_eq!(ratio.round() as usize, 295); + assert_eq!(frame_start_for_sample(start_sample, samples_per_frame), 294); + } + + #[test] + fn segmentation_accumulates_f64_floors_counts_then_narrows_to_f32() { + let audio = vec![0.0; 14 * PYANNOTE_SAMPLE_RATE_HZ as usize]; + let result = run_pyannote_segmentation_pass( + &audio, + PYANNOTE_SAMPLE_RATE_HZ, + PYANNOTE_DIARIZE_STRIDE_S, + |window_index, _window| { + let mut matrix = zero_matrix(); + match window_index { + 0 => set_log_prob(&mut matrix, 236, 1, 100_000_000.0), + 1 => set_log_prob(&mut matrix, 118, 1, 1.0), + 2 => set_log_prob(&mut matrix, 0, 1, -100_000_000.0), + _ => {} + } + Ok::(matrix) + }, + ) + .expect("segmentation pass"); + let expected = (1.0_f64 / 3.0) as f32; + let actual = result.avg_log_probs.row(236).expect("frame 236")[1]; + + assert_eq!(actual, expected); + + let floored = average_accumulated_log_probs(&[5.0, 7.0], &[0], 1, 2).expect("average"); + assert_eq!(floored.data(), &[5.0, 7.0]); + + let expected_matrix = vec![expected]; + let actual_matrix = vec![actual]; + assert!( + matrix_comparison_error("accumulated_mean", &actual_matrix, &expected_matrix, 0.0) + .is_none() + ); + } + + #[test] + fn overlap_fraction_uses_averaged_f32_argmax_not_per_window_argmax() { + let audio = vec![0.0; 14 * PYANNOTE_SAMPLE_RATE_HZ as usize]; + let result = run_pyannote_segmentation_pass( + &audio, + PYANNOTE_SAMPLE_RATE_HZ, + PYANNOTE_DIARIZE_STRIDE_S, + |window_index, _window| { + let mut matrix = zero_matrix(); + match window_index { + 0 => set_log_prob(&mut matrix, 236, 4, 10.0), + 1 => set_log_prob(&mut matrix, 118, 1, 10.0), + 2 => set_log_prob(&mut matrix, 0, 1, 10.0), + _ => {} + } + Ok::(matrix) + }, + ) + .expect("segmentation pass"); + let row = result.avg_log_probs.row(236).expect("frame 236"); + + assert!(row[1] > row[4]); + assert_eq!(result.overlap_fraction, 0.0); + } + + #[test] + fn no_speech_buffer_returns_zero_overlap_fraction() { + let audio = vec![0.0; 10 * PYANNOTE_SAMPLE_RATE_HZ as usize]; + + let result = run_pyannote_segmentation_pass( + &audio, + PYANNOTE_SAMPLE_RATE_HZ, + PYANNOTE_OVERLAP_STRIDE_S, + |_, _window| { + Ok::(class_matrix(&vec![0; PYANNOTE_FRAMES_PER_WINDOW])) + }, + ) + .expect("segmentation pass"); + + assert_eq!(result.overlap_fraction, 0.0); + } + + #[test] + fn window_stats_use_post_tail_truncation_argmax() { + let mut classes = vec![4; PYANNOTE_FRAMES_PER_WINDOW]; + classes[0] = 1; + classes[1] = 1; + let full = class_matrix(&classes); + let truncated = truncated_log_probs(&full, 2).expect("truncated matrix"); + + assert_eq!( + compute_speaker_window_stats(&truncated).expect("stats"), + SpeakerWindowStats { + speech_frames: 2, + active_slot_count: 1, + overlap_frames: 0, + } + ); + assert_ne!( + compute_speaker_window_stats(&full).expect("stats"), + compute_speaker_window_stats(&truncated).expect("stats") + ); + } + + #[test] + fn window_stats_for_all_fixture_cases_match_committed_windows() { + let fixture = stage_fixture(); + let evidence = &fixture["speaker_evidence"]; + let ambiguity_overlap_frames = + (DIARIZE_MIN_OVERLAP * PYANNOTE_FRAMES_PER_WINDOW as f64).ceil() as usize; + + let cases = [ + ("none", vec![0; PYANNOTE_FRAMES_PER_WINDOW]), + ("single", vec![1; PYANNOTE_FRAMES_PER_WINDOW]), + ( + "multi_by_active_slots", + [vec![1; PYANNOTE_FRAMES_PER_WINDOW / 2], { + vec![2; PYANNOTE_FRAMES_PER_WINDOW - PYANNOTE_FRAMES_PER_WINDOW / 2] + }] + .concat(), + ), + ( + "else_branch_overlap_ambiguity", + [ + vec![1; PYANNOTE_FRAMES_PER_WINDOW - ambiguity_overlap_frames], + vec![4; ambiguity_overlap_frames], + ] + .concat(), + ), + ]; + + for (name, classes) in cases { + let stats = compute_speaker_window_stats(&class_matrix(&classes)).expect("stats"); + let expected = &evidence[name]["windows"][0]; + + assert_eq!( + stats.speech_frames as u64, + expected["speech_frames"].as_u64().expect("speech frames"), + "{name}" + ); + assert_eq!( + stats.active_slot_count as u64, + expected["active_slot_count"] + .as_u64() + .expect("active slot count"), + "{name}" + ); + assert_eq!( + stats.overlap_frames as u64, + expected["overlap_frames"].as_u64().expect("overlap frames"), + "{name}" + ); + } + } + + #[test] + fn zero_speech_window_yields_zero_statistics() { + assert_eq!( + compute_speaker_window_stats(&class_matrix(&vec![0; PYANNOTE_FRAMES_PER_WINDOW])) + .expect("stats"), + SpeakerWindowStats { + speech_frames: 0, + active_slot_count: 0, + overlap_frames: 0, + } + ); + } + + #[test] + fn decisions_for_all_fixture_cases_match_committed_values() { + let fixture = stage_fixture(); + let tolerance = fixture["comparison"]["cluster_score_abs_tolerance"] + .as_f64() + .expect("cluster score tolerance"); + let evidence = &fixture["speaker_evidence"]; + + for name in [ + "none", + "single", + "multi_by_active_slots", + "else_branch_overlap_ambiguity", + ] { + let case = &evidence[name]; + let decision = decide_speaker_evidence( + case["overlap_fraction"].as_f64().expect("overlap fraction"), + &[fixture_window_stats(case)], + ); + let expected = &case["decision"]; + + assert_eq!( + decision.speaker_evidence.as_str(), + expected["speaker_evidence"].as_str().expect("evidence"), + "{name}" + ); + assert_within( + decision.multi_window_fraction, + expected["multi_window_fraction"] + .as_f64() + .expect("multi fraction"), + tolerance, + name, + ); + assert_within( + decision.mean_window_overlap_share, + expected["mean_window_overlap_share"] + .as_f64() + .expect("mean overlap"), + tolerance, + name, + ); + } + } + + #[test] + fn mixed_zero_speech_and_speech_windows_ignore_zero_speech_denominators() { + let decision = decide_speaker_evidence( + 0.0, + &[ + SpeakerWindowStats { + speech_frames: 0, + active_slot_count: 0, + overlap_frames: 0, + }, + SpeakerWindowStats { + speech_frames: 100, + active_slot_count: 1, + overlap_frames: 4, + }, + SpeakerWindowStats { + speech_frames: 0, + active_slot_count: 0, + overlap_frames: 0, + }, + ], + ); + + assert_eq!(decision.speaker_evidence, SpeakerEvidence::Single); + assert_eq!(decision.multi_window_fraction, 0.0); + assert_eq!(decision.mean_window_overlap_share, 0.04); + } + + #[test] + fn malformed_window_logprob_shape_reports_specific_error_variant() { + let audio = vec![0.0; 10 * PYANNOTE_SAMPLE_RATE_HZ as usize]; + let error = run_pyannote_segmentation_pass( + &audio, + PYANNOTE_SAMPLE_RATE_HZ, + PYANNOTE_OVERLAP_STRIDE_S, + |_, _window| { + Ok::( + FeatureMatrix::from_row_major( + PYANNOTE_FRAMES_PER_WINDOW - 1, + PYANNOTE_CLASS_COUNT, + vec![0.0; (PYANNOTE_FRAMES_PER_WINDOW - 1) * PYANNOTE_CLASS_COUNT], + ) + .expect("stub matrix"), + ) + }, + ) + .unwrap_err(); + + assert_eq!( + error, + SpeakerSegmentationError::WindowLogProbShapeMismatch { + window_index: 0, + expected_frames: PYANNOTE_FRAMES_PER_WINDOW, + expected_classes: PYANNOTE_CLASS_COUNT, + actual_frames: PYANNOTE_FRAMES_PER_WINDOW - 1, + actual_classes: PYANNOTE_CLASS_COUNT, + actual_len: (PYANNOTE_FRAMES_PER_WINDOW - 1) * PYANNOTE_CLASS_COUNT, + } + ); + } + + #[test] + fn segmentation_comparison_fails_when_value_exceeds_tolerance() { + let actual = vec![0.25_f32, 0.5]; + let mut expected = actual.clone(); + let tolerance = 1e-6_f32; + expected[1] += tolerance * 2.0; + + assert!(matrix_comparison_error("segmentation", &actual, &expected, tolerance).is_some()); + } + + fn indexed_audio(samples: usize) -> Vec { + (0..samples).map(|index| index as f32).collect() + } + + fn zero_matrix() -> FeatureMatrix { + FeatureMatrix::from_row_major( + PYANNOTE_FRAMES_PER_WINDOW, + PYANNOTE_CLASS_COUNT, + vec![0.0; PYANNOTE_FRAMES_PER_WINDOW * PYANNOTE_CLASS_COUNT], + ) + .expect("zero matrix") + } + + fn class_matrix(classes: &[usize]) -> FeatureMatrix { + let mut data = vec![-10.0; classes.len() * PYANNOTE_CLASS_COUNT]; + for (frame, class) in classes.iter().enumerate() { + data[frame * PYANNOTE_CLASS_COUNT + *class] = 10.0; + } + FeatureMatrix::from_row_major(classes.len(), PYANNOTE_CLASS_COUNT, data) + .expect("class matrix") + } + + fn set_log_prob(matrix: &mut FeatureMatrix, frame: usize, class: usize, value: f32) { + let mut data = matrix.data().to_vec(); + data[frame * matrix.bins() + class] = value; + *matrix = FeatureMatrix::from_row_major(matrix.frames(), matrix.bins(), data) + .expect("mutated matrix"); + } + + fn usize_array(value: &Value) -> [usize; N] { + let array = value.as_array().expect("array"); + assert_eq!(array.len(), N); + std::array::from_fn(|index| array[index].as_u64().expect("usize value") as usize) + } + + fn fixture_window_stats(case: &Value) -> SpeakerWindowStats { + let window = &case["windows"][0]; + SpeakerWindowStats { + speech_frames: window["speech_frames"].as_u64().expect("speech frames") as usize, + active_slot_count: window["active_slot_count"] + .as_u64() + .expect("active slot count") as usize, + overlap_frames: window["overlap_frames"].as_u64().expect("overlap frames") as usize, + } + } + + fn assert_within(actual: f64, expected: f64, tolerance: f64, label: &str) { + let diff = (actual - expected).abs(); + assert!( + diff <= tolerance, + "{label}: actual={actual} expected={expected} tolerance={tolerance}" + ); + } +} -- 2.51.2