use std::path::PathBuf;
use crate::error::DecibriError;
pub(crate) struct OnnxTensorView<'a> {
pub shape: &'a [i64],
pub data: OnnxTensorData<'a>,
}
pub(crate) enum OnnxTensorData<'a> {
F32(&'a [f32]),
I64(&'a [i64]),
}
pub(crate) struct OnnxInputs<'a> {
pub items: &'a [(&'a str, OnnxTensorView<'a>)],
}
pub(crate) struct OnnxOutputs {
pub tensors: Vec<(String, OnnxOutputTensor)>,
}
pub(crate) struct OnnxOutputTensor {
#[allow(dead_code)]
pub shape: Vec<i64>,
pub data: OnnxTensorOwned,
}
pub(crate) enum OnnxTensorOwned {
F32(Vec<f32>),
#[allow(dead_code)]
I64(Vec<i64>),
}
impl OnnxOutputs {
pub(crate) fn get(&self, name: &str) -> Option<&OnnxOutputTensor> {
self.tensors.iter().find(|(n, _)| n == name).map(|(_, t)| t)
}
}
pub(crate) trait OnnxSession: Send + Sync {
fn run(&mut self, inputs: OnnxInputs<'_>) -> Result<OnnxOutputs, DecibriError>;
}
pub(crate) struct OnnxSessionBuilder {
model_path: PathBuf,
intra_threads: usize,
}
impl OnnxSessionBuilder {
pub(crate) fn from_file(path: impl Into<PathBuf>) -> Self {
Self {
model_path: path.into(),
intra_threads: 1,
}
}
pub(crate) fn with_intra_threads(mut self, n: usize) -> Self {
self.intra_threads = n;
self
}
pub(crate) fn build(self) -> Result<Box<dyn OnnxSession>, DecibriError> {
let session = ort_impl::OrtSession::open(&self.model_path, self.intra_threads)?;
Ok(Box::new(session))
}
}
mod ort_impl {
use std::borrow::Cow;
use std::path::Path;
use ort::session::{Session, SessionInputValue};
use ort::value::{Tensor, TensorElementType};
use super::{
DecibriError, OnnxInputs, OnnxOutputTensor, OnnxOutputs, OnnxSession, OnnxTensorData,
OnnxTensorOwned,
};
pub(super) struct OrtSession {
inner: Session,
}
impl OrtSession {
pub(super) fn open(path: &Path, intra: usize) -> Result<Self, DecibriError> {
let session = Session::builder()
.map_err(DecibriError::OrtSessionBuildFailed)?
.with_intra_threads(intra)
.map_err(|e| DecibriError::OrtThreadsConfigFailed(e.into()))?
.commit_from_file(path)
.map_err(|e| DecibriError::VadModelLoadFailed {
path: path.to_path_buf(),
source: e,
})?;
Ok(Self { inner: session })
}
}
fn known_kind(name: &str) -> &'static str {
match name {
"input" => "input",
"state" => "state",
"sr" => "sr",
"output" => "output",
"stateN" => "state",
_ => "trait_tensor",
}
}
impl OnnxSession for OrtSession {
fn run(&mut self, inputs: OnnxInputs<'_>) -> Result<OnnxOutputs, DecibriError> {
let mut input_pairs: Vec<(Cow<'_, str>, SessionInputValue<'_>)> =
Vec::with_capacity(inputs.items.len());
for (name, view) in inputs.items.iter() {
let kind = known_kind(name);
let shape: Vec<i64> = view.shape.to_vec();
let value: SessionInputValue<'_> = match &view.data {
OnnxTensorData::F32(slice) => {
let tensor = Tensor::from_array((shape, slice.to_vec()))
.map_err(|e| DecibriError::OrtTensorCreateFailed { kind, source: e })?;
tensor.into()
}
OnnxTensorData::I64(slice) => {
let tensor = Tensor::from_array((shape, slice.to_vec()))
.map_err(|e| DecibriError::OrtTensorCreateFailed { kind, source: e })?;
tensor.into()
}
};
input_pairs.push((Cow::Owned((*name).to_string()), value));
}
let outputs = self
.inner
.run(input_pairs)
.map_err(DecibriError::OrtInferenceFailed)?;
let names: Vec<String> = outputs.keys().map(|k| k.to_string()).collect();
let mut tensors: Vec<(String, OnnxOutputTensor)> = Vec::with_capacity(names.len());
for name in names {
let kind = known_kind(&name);
let value = &outputs[name.as_str()];
let dtype = value.dtype().tensor_type();
let owned = match dtype {
Some(TensorElementType::Float32) => {
let (shape, data) = value.try_extract_tensor::<f32>().map_err(|e| {
DecibriError::OrtTensorExtractFailed { kind, source: e }
})?;
OnnxOutputTensor {
shape: shape.to_vec(),
data: OnnxTensorOwned::F32(data.to_vec()),
}
}
Some(TensorElementType::Int64) => {
let (shape, data) = value.try_extract_tensor::<i64>().map_err(|e| {
DecibriError::OrtTensorExtractFailed { kind, source: e }
})?;
OnnxOutputTensor {
shape: shape.to_vec(),
data: OnnxTensorOwned::I64(data.to_vec()),
}
}
_ => {
return Err(DecibriError::OnnxBackendFailed {
backend: "ort",
source: format!(
"unsupported output tensor element type for {name}: {dtype:?}"
)
.into(),
});
}
};
tensors.push((name, owned));
}
Ok(OnnxOutputs { tensors })
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[allow(dead_code)]
fn assert_send_sync<T: Send + Sync>() {}
#[test]
fn box_dyn_session_is_send_sync() {
assert_send_sync::<Box<dyn OnnxSession>>();
}
struct MockSession {
expected_input_count: usize,
canned_outputs: Vec<(String, OnnxOutputTensor)>,
}
impl OnnxSession for MockSession {
fn run(&mut self, inputs: OnnxInputs<'_>) -> Result<OnnxOutputs, DecibriError> {
assert_eq!(
inputs.items.len(),
self.expected_input_count,
"MockSession got unexpected number of inputs"
);
let mut tensors = Vec::with_capacity(self.canned_outputs.len());
for (name, t) in self.canned_outputs.iter() {
let cloned = OnnxOutputTensor {
shape: t.shape.clone(),
data: match &t.data {
OnnxTensorOwned::F32(v) => OnnxTensorOwned::F32(v.clone()),
OnnxTensorOwned::I64(v) => OnnxTensorOwned::I64(v.clone()),
},
};
tensors.push((name.clone(), cloned));
}
Ok(OnnxOutputs { tensors })
}
}
#[test]
fn mock_session_round_trip() {
let canned = vec![(
"output".to_string(),
OnnxOutputTensor {
shape: vec![1, 1],
data: OnnxTensorOwned::F32(vec![0.42]),
},
)];
let mut session: Box<dyn OnnxSession> = Box::new(MockSession {
expected_input_count: 3,
canned_outputs: canned,
});
let audio = vec![0.0f32; 512];
let state = vec![0.0f32; 256];
let sr = vec![16000i64];
let inputs = OnnxInputs {
items: &[
(
"input",
OnnxTensorView {
shape: &[1, 512],
data: OnnxTensorData::F32(&audio),
},
),
(
"state",
OnnxTensorView {
shape: &[2, 1, 128],
data: OnnxTensorData::F32(&state),
},
),
(
"sr",
OnnxTensorView {
shape: &[1],
data: OnnxTensorData::I64(&sr),
},
),
],
};
let outputs = session.run(inputs).expect("MockSession should succeed");
assert_eq!(outputs.tensors.len(), 1);
let probe = outputs.get("output").expect("output present");
assert_eq!(probe.shape, vec![1, 1]);
match &probe.data {
OnnxTensorOwned::F32(v) => assert_eq!(v.as_slice(), &[0.42f32]),
OnnxTensorOwned::I64(_) => panic!("expected F32"),
}
}
#[test]
fn output_lookup_by_name_returns_none_for_missing() {
let outputs = OnnxOutputs {
tensors: vec![(
"output".to_string(),
OnnxOutputTensor {
shape: vec![1],
data: OnnxTensorOwned::F32(vec![1.0]),
},
)],
};
assert!(outputs.get("output").is_some());
assert!(outputs.get("not_a_name").is_none());
}
#[test]
fn input_slice_preserves_order_and_names() {
let a = vec![1.0f32, 2.0, 3.0];
let b = vec![10i64, 20];
let inputs = OnnxInputs {
items: &[
(
"first",
OnnxTensorView {
shape: &[3],
data: OnnxTensorData::F32(&a),
},
),
(
"second",
OnnxTensorView {
shape: &[2],
data: OnnxTensorData::I64(&b),
},
),
],
};
assert_eq!(inputs.items.len(), 2);
assert_eq!(inputs.items[0].0, "first");
assert_eq!(inputs.items[1].0, "second");
match inputs.items[0].1.data {
OnnxTensorData::F32(s) => assert_eq!(s, a.as_slice()),
OnnxTensorData::I64(_) => panic!("expected F32"),
}
match inputs.items[1].1.data {
OnnxTensorData::I64(s) => assert_eq!(s, b.as_slice()),
OnnxTensorData::F32(_) => panic!("expected I64"),
}
}
#[test]
fn ort_backed_session_runs_silero_inference() {
use std::path::Path;
let manifest_dir = env!("CARGO_MANIFEST_DIR");
let model = Path::new(manifest_dir)
.join("..")
.join("..")
.join("models")
.join("silero_vad.onnx");
if !model.is_file() {
eprintln!(
"skipping ort_backed_session_runs_silero_inference: model not found at {}",
model.display()
);
return;
}
crate::vad::init_ort_once(None).expect("ORT init should succeed");
let mut session = OnnxSessionBuilder::from_file(&model)
.with_intra_threads(1)
.build()
.expect("ORT session build should succeed for bundled model");
let audio = vec![0.0f32; 512];
let state = vec![0.0f32; 256];
let sr = vec![16000i64];
let outputs = session
.run(OnnxInputs {
items: &[
(
"input",
OnnxTensorView {
shape: &[1, 512],
data: OnnxTensorData::F32(&audio),
},
),
(
"state",
OnnxTensorView {
shape: &[2, 1, 128],
data: OnnxTensorData::F32(&state),
},
),
(
"sr",
OnnxTensorView {
shape: &[1],
data: OnnxTensorData::I64(&sr),
},
),
],
})
.expect("ORT session run should succeed");
let probe = outputs
.get("output")
.expect("Silero emits an `output` tensor");
match &probe.data {
OnnxTensorOwned::F32(v) => assert_eq!(v.len(), 1, "Silero `output` is one f32"),
OnnxTensorOwned::I64(_) => panic!("Silero `output` is f32, not i64"),
}
let state_n = outputs
.get("stateN")
.expect("Silero emits a `stateN` tensor");
match &state_n.data {
OnnxTensorOwned::F32(v) => assert_eq!(v.len(), 256, "Silero `stateN` is 256 f32s"),
OnnxTensorOwned::I64(_) => panic!("Silero `stateN` is f32, not i64"),
}
}
}