pub trait Embedder: Send + Sync {
fn dim(&self) -> usize;
fn embed(&self, audio: &[f32]) -> Result<Vec<f32>, EmbedderError>;
fn embed_batch(&self, audios: &[&[f32]]) -> Result<Vec<Vec<f32>>, EmbedderError> {
audios.iter().map(|a| self.embed(a)).collect()
}
}
#[non_exhaustive]
#[derive(Debug, Clone, thiserror::Error)]
pub enum EmbedderError {
#[error("audio too short for this embedder: {actual_secs:.3}s < {min_secs:.3}s")]
AudioTooShort { actual_secs: f32, min_secs: f32 },
#[error("ONNX inference failed: {detail}")]
InferenceFailed { detail: String },
#[error("resource exhausted: {detail}")]
ResourceExhausted { detail: String },
#[error("expected embedding dim {expected}, got {actual}")]
DimMismatch { expected: usize, actual: usize },
#[error("model file io error on {path}: {detail}")]
ModelIo {
path: std::path::PathBuf,
detail: String,
},
#[cfg(feature = "onnx")]
#[error("failed to build embedder for {path}: {source}")]
SessionBuild {
path: std::path::PathBuf,
#[source]
source: crate::fbank_onnx::FbankExtractorError,
},
#[error("legacy adapter error: {0}")]
Legacy(String),
}
impl EmbedderError {
pub fn is_resource_exhausted(&self) -> bool {
match self {
Self::ResourceExhausted { .. } => true,
Self::InferenceFailed { detail } | Self::Legacy(detail) => {
detail_looks_exhausted(detail)
}
_ => false,
}
}
}
fn detail_looks_exhausted(detail: &str) -> bool {
detail.contains("pool exhausted")
}
pub struct DummyExtractor {
dim: usize,
seed: std::sync::atomic::AtomicU64,
}
impl DummyExtractor {
pub fn new(dim: usize) -> Self {
Self {
dim,
seed: std::sync::atomic::AtomicU64::new(1),
}
}
fn next_unit_vector(&self) -> Vec<f32> {
let mut seed = self.seed.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let mut vec = vec![0.0f32; self.dim];
for v in &mut vec {
seed = seed.wrapping_mul(1103515245).wrapping_add(12345);
*v = ((seed % 1000) as f32 / 1000.0) - 0.5;
}
crate::utils::l2_normalize(&mut vec);
vec
}
}
impl Embedder for DummyExtractor {
fn dim(&self) -> usize {
self.dim
}
fn embed(&self, _audio: &[f32]) -> Result<Vec<f32>, EmbedderError> {
Ok(self.next_unit_vector())
}
}
pub fn apply_overlap_mask(
audio: &[f32],
overlap_regions: &[(f32, f32)],
sample_rate: u32,
) -> Vec<f32> {
let mut out = audio.to_vec();
if out.is_empty() {
return out;
}
let sr = sample_rate as f32;
for &(start_s, end_s) in overlap_regions {
if !end_s.is_finite() || !start_s.is_finite() || end_s <= start_s {
continue;
}
let start = (start_s * sr).max(0.0).floor() as usize;
let end = (end_s * sr).max(0.0).ceil() as usize;
let end = end.min(out.len());
if start >= end || start >= out.len() {
continue;
}
for v in &mut out[start..end] {
*v = 0.0;
}
}
out
}
#[cfg(test)]
pub(crate) struct EmbedderPool<E: Embedder> {
pool: crate::utils::ObjectPool<E>,
dim: usize,
capacity: usize,
}
#[cfg(test)]
impl<E: Embedder> EmbedderPool<E> {
pub fn new(embedders: Vec<E>) -> Result<Self, EmbedderError> {
let dim = embedders.first().map(|e| e.dim()).unwrap_or(0);
for e in embedders.iter().skip(1) {
let actual = e.dim();
if actual != dim {
return Err(EmbedderError::DimMismatch {
expected: dim,
actual,
});
}
}
let capacity = embedders.len();
Ok(Self {
pool: crate::utils::ObjectPool::new(embedders),
dim,
capacity,
})
}
pub fn dim(&self) -> usize {
self.dim
}
pub fn is_empty(&self) -> bool {
self.capacity == 0
}
pub fn embed(&self, audio: &[f32]) -> Result<Vec<f32>, EmbedderError> {
if self.is_empty() {
return Err(EmbedderError::ResourceExhausted {
detail: "empty embedder pool".to_owned(),
});
}
let embedder = self.pool.checkout();
embedder.embed(audio)
}
}
#[cfg(all(feature = "onnx", feature = "embedder"))]
fn parallel_embed_batch<E: Embedder>(
embedder: &E,
audios: &[&[f32]],
max_threads: usize,
) -> Result<Vec<Vec<f32>>, EmbedderError> {
let n = audios.len();
if n == 0 {
return Ok(Vec::new());
}
let num_threads = std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(4)
.min(max_threads.max(1))
.min(n);
let chunk_size = n.div_ceil(num_threads);
let chunks: Vec<&[&[f32]]> = audios.chunks(chunk_size).collect();
std::thread::scope(|s| {
let handles: Vec<_> = chunks
.into_iter()
.map(|chunk| {
s.spawn(move || {
chunk
.iter()
.map(|audio| embedder.embed(audio))
.collect::<Vec<_>>()
})
})
.collect();
let mut all_results = Vec::with_capacity(n);
for h in handles {
let chunk_results = h
.join()
.map_err(|_| EmbedderError::Legacy("embed_batch thread panicked".to_string()))?;
all_results.extend(chunk_results);
}
all_results.into_iter().collect::<Result<Vec<_>, _>>()
})
}
#[cfg(all(feature = "onnx", feature = "embedder"))]
mod onnx_adapters {
use super::*;
use crate::fbank_onnx::FbankOnnxExtractor;
use std::path::Path;
struct FbankAdapter {
inner: FbankOnnxExtractor,
dim: usize,
}
impl FbankAdapter {
fn new(
path: impl AsRef<Path>,
dim: usize,
pool_size: usize,
ep: crate::onnx::ExecutionProvider,
) -> Result<Self, EmbedderError> {
let inner =
FbankOnnxExtractor::new(path.as_ref(), dim, pool_size, ep).map_err(|e| {
EmbedderError::SessionBuild {
path: path.as_ref().to_path_buf(),
source: e,
}
})?;
Ok(Self { inner, dim })
}
}
impl Embedder for FbankAdapter {
fn dim(&self) -> usize {
self.dim
}
fn embed(&self, audio: &[f32]) -> Result<Vec<f32>, EmbedderError> {
self.inner.embed(audio)
}
fn embed_batch(&self, audios: &[&[f32]]) -> Result<Vec<Vec<f32>>, EmbedderError> {
parallel_embed_batch(self, audios, self.inner.pool_size())
}
}
macro_rules! named_fbank_adapter {
($(#[$meta:meta])* $name:ident) => {
$(#[$meta])*
pub struct $name(FbankAdapter);
impl Embedder for $name {
fn dim(&self) -> usize {
self.0.dim()
}
fn embed(&self, audio: &[f32]) -> Result<Vec<f32>, EmbedderError> {
self.0.embed(audio)
}
fn embed_batch(
&self,
audios: &[&[f32]],
) -> Result<Vec<Vec<f32>>, EmbedderError> {
self.0.embed_batch(audios)
}
}
};
}
named_fbank_adapter! {
ResNet34Adapter
}
impl ResNet34Adapter {
pub fn new(
path: impl AsRef<Path>,
pool_size: usize,
ep: crate::onnx::ExecutionProvider,
) -> Result<Self, EmbedderError> {
FbankAdapter::new(path, 256, pool_size, ep).map(Self)
}
}
named_fbank_adapter! {
CamPlusPlusExtractor
}
impl CamPlusPlusExtractor {
pub fn new(
path: impl AsRef<Path>,
dim: usize,
pool_size: usize,
ep: crate::onnx::ExecutionProvider,
) -> Result<Self, EmbedderError> {
FbankAdapter::new(path, dim, pool_size, ep).map(Self)
}
}
named_fbank_adapter! {
ERes2NetV2Extractor
}
impl ERes2NetV2Extractor {
pub const DIM: usize = 192;
pub fn new(
path: impl AsRef<Path>,
pool_size: usize,
ep: crate::onnx::ExecutionProvider,
) -> Result<Self, EmbedderError> {
Self::with_dim(path, Self::DIM, pool_size, ep)
}
pub fn with_dim(
path: impl AsRef<Path>,
dim: usize,
pool_size: usize,
ep: crate::onnx::ExecutionProvider,
) -> Result<Self, EmbedderError> {
FbankAdapter::new(path, dim, pool_size, ep).map(Self)
}
}
}
#[cfg(all(feature = "onnx", feature = "embedder"))]
pub use onnx_adapters::{CamPlusPlusExtractor, ERes2NetV2Extractor, ResNet34Adapter};
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod overlap_mask_tests {
use super::*;
#[test]
fn no_overlap_regions_pass_through() {
let audio = vec![1.0_f32; 16_000];
let masked = apply_overlap_mask(&audio, &[], 16_000);
assert_eq!(masked, audio);
}
#[test]
fn single_overlap_region_is_zeroed() {
let audio = vec![1.0_f32; 16_000];
let masked = apply_overlap_mask(&audio, &[(0.5, 0.7)], 16_000);
for (i, &v) in masked.iter().enumerate() {
if (8000..11200).contains(&i) {
assert_eq!(v, 0.0, "sample {i} should be zeroed");
} else {
assert_eq!(v, 1.0, "sample {i} should pass through");
}
}
}
#[test]
fn empty_input_returns_empty() {
let masked = apply_overlap_mask(&[], &[(0.0, 1.0)], 16_000);
assert!(masked.is_empty());
}
#[test]
fn out_of_bounds_overlap_is_clamped() {
let audio = vec![1.0_f32; 100];
let masked = apply_overlap_mask(&audio, &[(0.5, 1.0)], 16_000);
assert_eq!(masked, audio, "out-of-bounds overlap is a no-op");
}
#[test]
fn negative_overlap_start_is_clamped_to_zero() {
let audio = vec![1.0_f32; 16_000];
let masked = apply_overlap_mask(&audio, &[(-1.0, 0.5)], 16_000);
for &v in masked.iter().take(8000) {
assert_eq!(v, 0.0);
}
for &v in masked.iter().skip(8000) {
assert_eq!(v, 1.0);
}
}
#[test]
fn multiple_overlap_regions_all_zeroed() {
let audio = vec![1.0_f32; 16_000];
let masked = apply_overlap_mask(&audio, &[(0.1, 0.2), (0.5, 0.6), (0.9, 1.0)], 16_000);
let zero_ranges = [(1600..3200), (8000..9600), (14_400..16_000)];
for (i, &v) in masked.iter().enumerate() {
let in_zero = zero_ranges.iter().any(|r| r.contains(&i));
if in_zero {
assert_eq!(v, 0.0, "sample {i} should be zeroed");
} else {
assert_eq!(v, 1.0, "sample {i} should pass through");
}
}
}
#[test]
fn invalid_overlap_with_end_before_start_is_no_op() {
let audio = vec![1.0_f32; 16_000];
let masked = apply_overlap_mask(&audio, &[(0.7, 0.5)], 16_000);
assert_eq!(masked, audio, "end<start is silently skipped");
}
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod trait_tests {
use super::*;
struct ConstantEmbedder {
values: Vec<f32>,
}
impl Embedder for ConstantEmbedder {
fn dim(&self) -> usize {
self.values.len()
}
fn embed(&self, _audio: &[f32]) -> Result<Vec<f32>, EmbedderError> {
Ok(self.values.clone())
}
}
#[test]
fn embedder_trait_object_is_dyn_compatible() {
let e = ConstantEmbedder {
values: vec![0.1, 0.2, 0.3],
};
let _b: Box<dyn Embedder> = Box::new(e);
}
#[test]
fn embedder_default_batch_is_serial() {
let e = ConstantEmbedder {
values: vec![0.5; 4],
};
let inputs: Vec<&[f32]> = vec![&[][..], &[][..], &[][..]];
let out = e.embed_batch(&inputs).unwrap();
assert_eq!(out.len(), 3);
assert!(out.iter().all(|v| v.len() == 4 && v[0] == 0.5));
}
#[test]
fn embedder_dim_matches_output() {
let e = ConstantEmbedder {
values: vec![1.0; 192],
};
assert_eq!(e.dim(), 192);
assert_eq!(e.embed(&[]).unwrap().len(), 192);
}
#[test]
fn embedder_error_audio_too_short_displays() {
let err = EmbedderError::AudioTooShort {
actual_secs: 0.05,
min_secs: 0.25,
};
let msg = format!("{err}");
assert!(msg.contains("0.05"));
assert!(msg.contains("0.25"));
}
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod pool_tests {
use super::*;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
struct CountingEmbedder {
counter: Arc<AtomicUsize>,
dim: usize,
}
impl Embedder for CountingEmbedder {
fn dim(&self) -> usize {
self.dim
}
fn embed(&self, _audio: &[f32]) -> Result<Vec<f32>, EmbedderError> {
self.counter.fetch_add(1, Ordering::SeqCst);
Ok(vec![0.0; self.dim])
}
}
fn make_pool(n: usize) -> (EmbedderPool<CountingEmbedder>, Arc<AtomicUsize>) {
let counter = Arc::new(AtomicUsize::new(0));
let mut embedders = Vec::with_capacity(n);
for _ in 0..n {
embedders.push(CountingEmbedder {
counter: counter.clone(),
dim: 192,
});
}
let pool = EmbedderPool::new(embedders).unwrap();
(pool, counter)
}
#[test]
fn pool_with_single_embedder_round_trip() {
let (pool, counter) = make_pool(1);
let result = pool.embed(&[0.0_f32; 100]).unwrap();
assert_eq!(result.len(), 192);
assert_eq!(counter.load(Ordering::SeqCst), 1);
}
#[test]
fn pool_dim_is_consistent() {
let (pool, _) = make_pool(4);
assert_eq!(pool.dim(), 192);
}
#[test]
fn pool_serial_embed_increments_counter_per_call() {
let (pool, counter) = make_pool(2);
for _ in 0..5 {
pool.embed(&[0.0_f32; 100]).unwrap();
}
assert_eq!(counter.load(Ordering::SeqCst), 5);
}
#[test]
fn pool_with_zero_embedders_errors() {
let pool: EmbedderPool<CountingEmbedder> = EmbedderPool::new(Vec::new()).unwrap();
let err = pool
.embed(&[0.0_f32; 100])
.expect_err("empty pool must fail");
assert!(
matches!(err, EmbedderError::ResourceExhausted { .. }),
"empty pool is resource exhaustion, got {err}"
);
assert!(err.is_resource_exhausted());
}
#[test]
fn pool_rejects_mismatched_embedder_dims() {
let counter = Arc::new(AtomicUsize::new(0));
let embedders = vec![
CountingEmbedder {
counter: counter.clone(),
dim: 192,
},
CountingEmbedder {
counter: counter.clone(),
dim: 256,
},
];
let err = match EmbedderPool::new(embedders) {
Err(e) => e,
Ok(_) => panic!("mismatched dims must fail"),
};
assert!(
matches!(
err,
EmbedderError::DimMismatch {
expected: 192,
actual: 256
}
),
"expected DimMismatch(192, 256), got {err}"
);
}
#[test]
fn resource_exhausted_classifier() {
let typed = EmbedderError::ResourceExhausted {
detail: "speaker sessions busy".into(),
};
assert!(typed.is_resource_exhausted());
let legacy_string = EmbedderError::InferenceFailed {
detail: "onnx session pool exhausted".into(),
};
assert!(legacy_string.is_resource_exhausted());
let other = EmbedderError::DimMismatch {
expected: 1,
actual: 2,
};
assert!(!other.is_resource_exhausted());
}
#[test]
fn pipeline_and_streaming_error_helpers() {
use crate::pipeline::LegacyPipelineError;
use crate::streaming::StreamingError;
let emb = EmbedderError::ResourceExhausted {
detail: "busy".into(),
};
let pe = LegacyPipelineError::Embedding(emb.clone());
let se = StreamingError::Embedding(emb);
assert!(pe.is_resource_exhausted());
assert!(se.is_resource_exhausted());
let non_embedding = LegacyPipelineError::AudioTooLong {
actual_secs: 2.0,
max_secs: 1.0,
};
assert!(!non_embedding.is_resource_exhausted());
}
}