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;
pub struct FbankOnnxExtractor {
pool: crate::utils::ObjectPool<RuntimeSession>,
embedding_dim: usize,
fbank: FbankExtractor,
}
#[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,
},
}
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 pool_size = crate::onnx::resolve_session_pool_size(pool_size);
let mut sessions = Vec::with_capacity(pool_size);
let intra = crate::onnx::resolve_intra_threads(pool_size);
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),
embedding_dim,
fbank: FbankExtractor::new(crate::features::FbankConfig::default()),
})
}
#[cfg_attr(not(any(test, feature = "embedder")), allow(dead_code))]
pub(crate) fn pool_size(&self) -> usize {
self.pool.capacity()
}
}
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)
}
}
#[allow(clippy::unwrap_used)]
#[cfg(test)]
mod tests {
use super::*;
use crate::onnx::{ExecutionProvider, InferenceBackend};
use std::path::PathBuf;
const RESNET34: &str = "models/wespeaker_resnet34.onnx";
const RESNET34_DIM: usize = 256;
fn resnet34_path() -> Option<PathBuf> {
let p = Path::new(RESNET34);
if p.is_file() {
Some(p.to_path_buf())
} else {
None
}
}
fn sine_pcm(secs: f32, sr: u32) -> Vec<f32> {
let n = (secs * sr as f32) as usize;
(0..n)
.map(|i| {
let t = i as f32 / sr as f32;
0.3 * (2.0 * std::f32::consts::PI * 300.0 * t).sin()
})
.collect()
}
fn build_err(r: Result<FbankOnnxExtractor, FbankExtractorError>) -> FbankExtractorError {
match r {
Err(e) => e,
Ok(_) => panic!("expected construction to fail"),
}
}
#[test]
fn new_rejects_zero_pool_size() {
let err = build_err(FbankOnnxExtractor::new(
Path::new("models/__missing__.onnx"),
RESNET34_DIM,
0,
ExecutionProvider::Cpu,
));
assert!(matches!(err, FbankExtractorError::EmptyPool));
assert_eq!(err.to_string(), "pool_size must be > 0");
}
#[test]
fn new_reports_session_build_error_for_missing_model() {
let err = build_err(FbankOnnxExtractor::new(
Path::new("models/__definitely_missing__.onnx"),
RESNET34_DIM,
2,
ExecutionProvider::Cpu,
));
match err {
FbankExtractorError::SessionBuild { index, source } => {
assert_eq!(index, 0);
let msg = FbankExtractorError::SessionBuild { index, source }.to_string();
assert!(msg.starts_with("session 0:"), "unexpected: {msg}");
}
other => panic!("expected SessionBuild, got {other:?}"),
}
}
#[test]
#[cfg_attr(miri, ignore)]
fn new_builds_pool_and_reports_size() {
let Some(path) = resnet34_path() else {
eprintln!("skip: {RESNET34} missing");
return;
};
#[cfg(feature = "onnx")]
InferenceBackend::force(Some(InferenceBackend::Ort));
let ext = FbankOnnxExtractor::new(&path, RESNET34_DIM, 2, ExecutionProvider::Cpu).unwrap();
assert_eq!(ext.pool_size(), 2);
assert_eq!(ext.dim(), RESNET34_DIM);
InferenceBackend::force(None);
}
#[test]
#[cfg_attr(miri, ignore)]
fn embed_returns_unit_norm_embedding() {
let Some(path) = resnet34_path() else {
eprintln!("skip: {RESNET34} missing");
return;
};
#[cfg(feature = "onnx")]
InferenceBackend::force(Some(InferenceBackend::Ort));
let ext = FbankOnnxExtractor::new(&path, RESNET34_DIM, 1, ExecutionProvider::Cpu).unwrap();
let pcm = sine_pcm(1.0, 16_000);
let emb = ext.embed(&pcm).unwrap();
assert_eq!(emb.len(), RESNET34_DIM);
assert!(emb.iter().all(|v| v.is_finite()));
let norm: f32 = emb.iter().map(|v| v * v).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-4, "expected unit norm, got {norm}");
let emb2 = ext.embed(&pcm).unwrap();
assert_eq!(emb, emb2);
InferenceBackend::force(None);
}
#[test]
#[cfg_attr(miri, ignore)]
fn embed_zero_pads_short_input() {
let Some(path) = resnet34_path() else {
eprintln!("skip: {RESNET34} missing");
return;
};
#[cfg(feature = "onnx")]
InferenceBackend::force(Some(InferenceBackend::Ort));
let ext = FbankOnnxExtractor::new(&path, RESNET34_DIM, 1, ExecutionProvider::Cpu).unwrap();
let pcm = sine_pcm(0.005, 16_000);
assert!(pcm.len() < 400);
let emb = ext.embed(&pcm).unwrap();
assert_eq!(emb.len(), RESNET34_DIM);
assert!(emb.iter().all(|v| v.is_finite()));
InferenceBackend::force(None);
}
#[test]
#[cfg_attr(miri, ignore)]
fn embed_detects_dim_mismatch() {
let Some(path) = resnet34_path() else {
eprintln!("skip: {RESNET34} missing");
return;
};
#[cfg(feature = "onnx")]
InferenceBackend::force(Some(InferenceBackend::Ort));
let ext = FbankOnnxExtractor::new(&path, 192, 1, ExecutionProvider::Cpu).unwrap();
let err = ext.embed(&sine_pcm(1.0, 16_000)).unwrap_err();
match err {
EmbedderError::DimMismatch { expected, actual } => {
assert_eq!(expected, 192);
assert_eq!(actual, RESNET34_DIM);
}
other => panic!("expected DimMismatch, got {other:?}"),
}
InferenceBackend::force(None);
}
#[cfg(all(feature = "backend-tract", feature = "onnx"))]
fn cosine(a: &[f32], b: &[f32]) -> f64 {
let mut dot = 0.0f64;
let mut na = 0.0f64;
let mut nb = 0.0f64;
for (&x, &y) in a.iter().zip(b.iter()) {
dot += f64::from(x) * f64::from(y);
na += f64::from(x) * f64::from(x);
nb += f64::from(y) * f64::from(y);
}
dot / (na.sqrt() * nb.sqrt()).max(1e-12)
}
#[cfg(all(feature = "backend-tract", feature = "onnx"))]
#[test]
#[cfg_attr(miri, ignore)]
fn embed_ort_vs_tract_real_segments() {
let fp32 = resnet34_path();
let int8 = {
let p = Path::new("models/int8/resnet34_int8.onnx");
p.is_file().then(|| p.to_path_buf())
};
let wav = Path::new(
"benchmarks/results/powerset-tract-rtf-der-2026-08-12/smoke-vox3/audio/fuzfh.wav",
);
let wav = if wav.is_file() {
wav
} else {
Path::new("data/voxconverse-test/audio/fuzfh.wav")
};
if !wav.is_file() {
eprintln!("skip embed_ort_vs_tract: fuzfh missing");
return;
}
let (audio, sr) = crate::wav::read_wav(wav).expect("wav");
assert_eq!(sr, 16_000);
let ranges = [(0.0f32, 12.78f32), (13.01, 15.53), (15.69, 26.01)];
let clips: Vec<Vec<f32>> = ranges
.iter()
.map(|&(a, b)| {
let s = (a * 16_000.0) as usize;
let e = ((b * 16_000.0) as usize).min(audio.len());
audio[s..e].to_vec()
})
.collect();
for (label, model) in [("fp32", fp32), ("int8", int8)] {
let Some(path) = model else {
eprintln!("skip {label}: model missing");
continue;
};
InferenceBackend::force(Some(InferenceBackend::Ort));
let ort_ext =
FbankOnnxExtractor::new(&path, RESNET34_DIM, 1, ExecutionProvider::Cpu).unwrap();
let ort_embs: Vec<_> = clips.iter().map(|c| ort_ext.embed(c).unwrap()).collect();
InferenceBackend::force(None);
InferenceBackend::force(Some(InferenceBackend::Tract));
let tract_ext =
FbankOnnxExtractor::new(&path, RESNET34_DIM, 1, ExecutionProvider::Cpu).unwrap();
let tract_embs: Vec<_> = clips.iter().map(|c| tract_ext.embed(c).unwrap()).collect();
InferenceBackend::force(None);
for i in 0..clips.len() {
let c = cosine(&ort_embs[i], &tract_embs[i]);
eprintln!(
"embed {label} seg{i}: ort↔tract cosine={c:.6} len={}",
clips[i].len()
);
if label == "fp32" {
assert!(
c > 0.99,
"FP32 ResNet ort↔tract cosine must be ~1, got {c} on seg{i}"
);
}
}
let o12 = cosine(&ort_embs[0], &ort_embs[1]);
let t12 = cosine(&tract_embs[0], &tract_embs[1]);
let o02 = cosine(&ort_embs[0], &ort_embs[2]);
let t02 = cosine(&tract_embs[0], &tract_embs[2]);
eprintln!(
"embed {label} pairwise: ort 0↔1={o12:.4} 0↔2={o02:.4} | tract 0↔1={t12:.4} 0↔2={t02:.4}"
);
if label == "int8" {
assert!(
t12 > 0.8 && t02 > 0.8,
"expected INT8 tract pairwise collapse (got {t12}, {t02}); \
if this fails, INT8 tract may have become safe"
);
}
}
}
}