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)]
#[path = "overlap_mask_tests.rs"]
mod overlap_mask_tests;
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[path = "trait_tests.rs"]
mod trait_tests;
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[path = "pool_tests.rs"]
mod pool_tests;
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[path = "error_display_tests.rs"]
mod error_display_tests;
#[allow(clippy::unwrap_used)]
#[cfg(test)]
#[path = "dummy_extractor_tests.rs"]
mod dummy_extractor_tests;
#[allow(clippy::unwrap_used)]
#[cfg(all(test, feature = "onnx", feature = "embedder"))]
#[path = "onnx_adapter_tests.rs"]
mod onnx_adapter_tests;