#![allow(deprecated)]
use crate::embedding::{EmbeddingError, EmbeddingExtractor};
use crate::features::{FbankExtractor, apply_cmvn};
use crate::onnx::{InferenceRuntime, InferenceTensor, RuntimeSession};
use crate::types::DiarizationConfig;
use crate::utils::l2_normalize;
use std::path::Path;
#[cfg(feature = "onnx")]
#[deprecated(
since = "0.7.0",
note = "use the v1.0 Embedder trait in polyvoice::embedder"
)]
pub struct FbankOnnxExtractor {
pool: crate::utils::ObjectPool<RuntimeSession>,
embedding_dim: usize,
fbank: FbankExtractor,
}
#[cfg(feature = "onnx")]
impl FbankOnnxExtractor {
pub fn new(
model_path: &Path,
embedding_dim: usize,
pool_size: usize,
ep: crate::onnx::ExecutionProvider,
) -> anyhow::Result<Self> {
if pool_size == 0 {
anyhow::bail!("pool_size must be > 0");
}
let mut sessions = Vec::with_capacity(pool_size);
for i in 0..pool_size {
let session = crate::onnx::build_session_with_ep(model_path, ep, Some(1))
.map_err(|e| EmbeddingError::InferenceFailed(format!("session {i}: {e}")))?;
sessions.push(session);
}
Ok(Self {
pool: crate::utils::ObjectPool::new(sessions),
embedding_dim,
fbank: FbankExtractor::new(crate::features::FbankConfig::default()),
})
}
}
#[cfg(feature = "onnx")]
impl EmbeddingExtractor for FbankOnnxExtractor {
fn extract(
&self,
samples: &[f32],
_config: &DiarizationConfig,
) -> Result<Vec<f32>, EmbeddingError> {
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| EmbeddingError::InferenceFailed(e.to_string()))?;
if fbank.is_empty() {
return Err(EmbeddingError::InvalidInput {
expected: self.fbank.config.win_length,
got: samples.len(),
});
}
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| EmbeddingError::InferenceFailed(e.to_string()))?;
let first = outputs.into_iter().next().ok_or_else(|| {
EmbeddingError::InferenceFailed("ONNX model produced no outputs".to_string())
})?;
let data = first
.into_f32()
.map_err(|e| EmbeddingError::InferenceFailed(e.to_string()))?;
let data_len = data.len();
if data_len != self.embedding_dim {
return Err(EmbeddingError::InferenceFailed(format!(
"expected embedding dim {}, got {}",
self.embedding_dim, data_len
)));
}
let mut embedding = data;
l2_normalize(&mut embedding);
Ok(embedding)
}
fn embedding_dim(&self) -> usize {
self.embedding_dim
}
}
#[cfg(not(feature = "onnx"))]
#[derive(Debug)]
#[deprecated(
since = "0.7.0",
note = "use the v1.0 Embedder trait in polyvoice::embedder"
)]
pub struct FbankOnnxExtractor;
#[cfg(not(feature = "onnx"))]
impl FbankOnnxExtractor {
pub fn new(
_model_path: &Path,
_embedding_dim: usize,
_pool_size: usize,
) -> anyhow::Result<Self> {
anyhow::bail!("the `onnx` feature is not enabled")
}
}
#[cfg(all(test, not(feature = "onnx")))]
mod tests {
use super::*;
use std::path::Path;
#[test]
fn fbank_onnx_extractor_new_without_onnx_fails() {
let result = FbankOnnxExtractor::new(Path::new("dummy.onnx"), 256, 1);
assert!(result.is_err());
let err = match result {
Err(e) => e.to_string(),
Ok(_) => panic!("expected error"),
};
assert!(
err.contains("onnx") || err.contains("not enabled"),
"expected onnx-related error, got: {err}"
);
}
#[test]
fn fbank_onnx_extractor_stub_exists() {
let _ = FbankOnnxExtractor;
}
}