use crate::embedder::{Embedder, EmbedderError};
use crate::features::{FbankExtractor, apply_cmvn};
use crate::onnx::{InferenceRuntime, InferenceTensor, RuntimeSession};
use crate::utils::l2_normalize;
use std::path::Path;
#[cfg(feature = "onnx")]
pub struct FbankOnnxExtractor {
pool: crate::utils::ObjectPool<RuntimeSession>,
#[allow(dead_code)]
pool_size: usize,
embedding_dim: usize,
fbank: FbankExtractor,
}
#[cfg(feature = "onnx")]
#[derive(Clone, thiserror::Error, Debug)]
pub enum FbankExtractorError {
#[error("pool_size must be > 0")]
EmptyPool,
#[error("session {index}: {source}")]
SessionBuild {
index: usize,
#[source]
source: crate::onnx::OnnxError,
},
}
#[cfg(feature = "onnx")]
impl FbankOnnxExtractor {
pub fn new(
model_path: &Path,
embedding_dim: usize,
pool_size: usize,
ep: crate::onnx::ExecutionProvider,
) -> Result<Self, FbankExtractorError> {
if pool_size == 0 {
return Err(FbankExtractorError::EmptyPool);
}
let mut sessions = Vec::with_capacity(pool_size);
let intra = std::thread::available_parallelism()
.map(|n| (n.get() / pool_size).max(1))
.unwrap_or(1);
for i in 0..pool_size {
let session =
crate::onnx::build_session_with_ep(model_path, ep, Some(intra)).map_err(|e| {
FbankExtractorError::SessionBuild {
index: i,
source: e,
}
})?;
sessions.push(session);
}
Ok(Self {
pool: crate::utils::ObjectPool::new(sessions),
pool_size,
embedding_dim,
fbank: FbankExtractor::new(crate::features::FbankConfig::default()),
})
}
#[allow(dead_code)]
pub(crate) fn pool_size(&self) -> usize {
self.pool_size
}
}
#[cfg(feature = "onnx")]
impl Embedder for FbankOnnxExtractor {
fn dim(&self) -> usize {
self.embedding_dim
}
fn embed(&self, samples: &[f32]) -> Result<Vec<f32>, EmbedderError> {
let mut session = self.pool.checkout();
let min_samples = self.fbank.config.win_length;
let padded: Vec<f32>;
let samples = if samples.len() < min_samples {
padded = {
let mut v = vec![0.0_f32; min_samples];
v[..samples.len()].copy_from_slice(samples);
v
};
&padded
} else {
samples
};
let fbank = self
.fbank
.extract(samples)
.map_err(|e| EmbedderError::InferenceFailed {
detail: e.to_string(),
})?;
if fbank.is_empty() {
let sample_rate = self.fbank.config.sample_rate as f32;
return Err(EmbedderError::AudioTooShort {
actual_secs: samples.len() as f32 / sample_rate,
min_secs: min_samples as f32 / sample_rate,
});
}
let fbank = apply_cmvn(&fbank);
let n_frames = fbank.len();
let n_mels = fbank[0].len();
let flat: Vec<f32> = fbank.into_iter().flatten().collect();
let input = InferenceTensor::f32(vec![1, n_frames, n_mels], flat);
let outputs =
session
.run_ordered(&[&input])
.map_err(|e| EmbedderError::InferenceFailed {
detail: e.to_string(),
})?;
let first = outputs
.into_iter()
.next()
.ok_or_else(|| EmbedderError::InferenceFailed {
detail: "ONNX model produced no outputs".to_string(),
})?;
let data = first
.into_f32()
.map_err(|e| EmbedderError::InferenceFailed {
detail: e.to_string(),
})?;
let data_len = data.len();
if data_len != self.embedding_dim {
return Err(EmbedderError::DimMismatch {
expected: self.embedding_dim,
actual: data_len,
});
}
let mut embedding = data;
l2_normalize(&mut embedding);
Ok(embedding)
}
}