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}" + ); + } +}