use super::ExecutionProvider;
use super::OnnxError;
use super::ort_session::OrtSession;
use super::runtime::{InferenceError, InferenceRuntime, InferenceTensor, NamedTensor};
use std::cell::Cell;
use std::path::Path;
const BACKEND_AUTO: u8 = 0;
const BACKEND_ORT: u8 = 1;
#[cfg(feature = "backend-tract")]
const BACKEND_TRACT: u8 = 2;
thread_local! {
static BACKEND_FORCE: Cell<u8> = const { Cell::new(BACKEND_AUTO) };
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum InferenceBackend {
Ort,
#[cfg(feature = "backend-tract")]
Tract,
}
impl InferenceBackend {
pub fn resolve() -> Self {
let forced = BACKEND_FORCE.with(Cell::get);
match forced {
BACKEND_ORT => return Self::Ort,
#[cfg(feature = "backend-tract")]
BACKEND_TRACT => return Self::Tract,
_ => {}
}
match std::env::var("POLYVOICE_INFERENCE_BACKEND") {
Ok(v) => {
let lower = v.to_ascii_lowercase();
match lower.as_str() {
"ort" | "onnxruntime" | "onnx-runtime" => Self::Ort,
"tract" => {
#[cfg(feature = "backend-tract")]
{
Self::Tract
}
#[cfg(not(feature = "backend-tract"))]
{
tracing::warn!(
"POLYVOICE_INFERENCE_BACKEND=tract but the `backend-tract` \
feature is not enabled — falling back to ort"
);
Self::Ort
}
}
other => {
tracing::warn!("unknown POLYVOICE_INFERENCE_BACKEND={other:?}; using ort");
Self::Ort
}
}
}
Err(_) => Self::Ort,
}
}
pub fn force(backend: Option<Self>) {
let code = match backend {
None => BACKEND_AUTO,
Some(Self::Ort) => BACKEND_ORT,
#[cfg(feature = "backend-tract")]
Some(Self::Tract) => BACKEND_TRACT,
};
BACKEND_FORCE.with(|c| c.set(code));
}
}
#[derive(Debug)]
pub enum RuntimeSession {
Ort(OrtSession),
#[cfg(feature = "backend-tract")]
Tract(super::tract_session::TractSession),
}
impl RuntimeSession {
pub fn from_path(
model_path: &Path,
ep: ExecutionProvider,
intra_threads: Option<usize>,
) -> Result<Self, OnnxError> {
match InferenceBackend::resolve() {
InferenceBackend::Ort => Ok(Self::Ort(OrtSession::from_path(
model_path,
ep,
intra_threads,
)?)),
#[cfg(feature = "backend-tract")]
InferenceBackend::Tract => Ok(Self::Tract(
super::tract_session::TractSession::from_path(model_path, intra_threads)?,
)),
}
}
pub fn backend(&self) -> InferenceBackend {
match self {
Self::Ort(_) => InferenceBackend::Ort,
#[cfg(feature = "backend-tract")]
Self::Tract(_) => InferenceBackend::Tract,
}
}
}
impl InferenceRuntime for RuntimeSession {
fn input_names(&self) -> &[String] {
match self {
Self::Ort(s) => s.input_names(),
#[cfg(feature = "backend-tract")]
Self::Tract(s) => s.input_names(),
}
}
fn run(&mut self, inputs: &[NamedTensor<'_>]) -> Result<Vec<InferenceTensor>, InferenceError> {
match self {
Self::Ort(s) => s.run(inputs),
#[cfg(feature = "backend-tract")]
Self::Tract(s) => s.run(inputs),
}
}
fn run_ordered(
&mut self,
inputs: &[&InferenceTensor],
) -> Result<Vec<InferenceTensor>, InferenceError> {
match self {
Self::Ort(s) => s.run_ordered(inputs),
#[cfg(feature = "backend-tract")]
Self::Tract(s) => s.run_ordered(inputs),
}
}
}