use ort::ep::ExecutionProviderDispatch;
use ort::session::builder::SessionBuilder;
use crate::config::Accelerator;
use crate::error::{OcrError, Result};
struct Provider {
name: &'static str,
available: bool,
dispatch: ExecutionProviderDispatch,
}
pub(crate) fn apply_accelerator(
builder: SessionBuilder,
accelerator: Accelerator,
) -> Result<(SessionBuilder, Accelerator)> {
match accelerator {
Accelerator::Cpu => Ok((builder, Accelerator::Cpu)),
Accelerator::Auto => Ok(register_preferred(builder)),
explicit => register_explicit(builder, explicit),
}
}
fn register_preferred(builder: SessionBuilder) -> (SessionBuilder, Accelerator) {
let mut builder = builder;
for &candidate in preferred_accelerators() {
let Some(provider) = provider(candidate) else {
continue;
};
if !provider.available {
tracing::debug!(
accelerator = candidate.as_str(),
provider = provider.name,
"skipping an execution provider the linked ONNX Runtime was not built with"
);
continue;
}
match builder.with_execution_providers([provider.dispatch.error_on_failure()]) {
Ok(configured) => {
tracing::info!(
accelerator = candidate.as_str(),
provider = provider.name,
"registered execution provider"
);
return (configured, candidate);
}
Err(error) => {
tracing::warn!(
accelerator = candidate.as_str(),
provider = provider.name,
%error,
"execution provider could not be registered; trying the next candidate"
);
builder = error.recover();
}
}
}
tracing::info!(
accelerator = "cpu",
"no accelerator registered; using the CPU execution provider"
);
(builder, Accelerator::Cpu)
}
fn register_explicit(builder: SessionBuilder, accelerator: Accelerator) -> Result<(SessionBuilder, Accelerator)> {
let Some(feature) = cargo_feature(accelerator) else {
return Err(OcrError::config(format!(
"the `ort` backend has no execution provider for accelerator `{}`",
accelerator.as_str()
)));
};
let Some(provider) = provider(accelerator) else {
return Err(OcrError::config(format!(
"accelerator `{}` is not compiled into this build of sceptre; rebuild with the `{feature}` cargo feature",
accelerator.as_str(),
)));
};
if !provider.available {
return Err(OcrError::config(format!(
"accelerator `{}` is unavailable: the linked ONNX Runtime was built without `{}`",
accelerator.as_str(),
provider.name
)));
}
match builder.with_execution_providers([provider.dispatch.error_on_failure()]) {
Ok(configured) => {
tracing::info!(
accelerator = accelerator.as_str(),
provider = provider.name,
"registered execution provider"
);
Ok((configured, accelerator))
}
Err(error) => Err(OcrError::inference(format!(
"ONNX Runtime backend failed to register the `{}` execution provider for accelerator `{}`: {error}",
provider.name,
accelerator.as_str()
))),
}
}
fn preferred_accelerators() -> &'static [Accelerator] {
#[cfg(any(target_os = "macos", target_os = "ios"))]
{
&[Accelerator::CoreMl]
}
#[cfg(target_os = "windows")]
{
&[Accelerator::DirectMl, Accelerator::Cuda]
}
#[cfg(not(any(target_os = "macos", target_os = "ios", target_os = "windows")))]
{
&[Accelerator::Cuda]
}
}
fn cargo_feature(accelerator: Accelerator) -> Option<&'static str> {
match accelerator {
Accelerator::Cpu | Accelerator::Auto | Accelerator::Metal => None,
Accelerator::CoreMl => Some("ort-coreml"),
Accelerator::DirectMl => Some("ort-directml"),
Accelerator::Cuda => Some("ort-cuda"),
}
}
fn provider(accelerator: Accelerator) -> Option<Provider> {
match accelerator {
Accelerator::Cpu | Accelerator::Auto | Accelerator::Metal => None,
Accelerator::CoreMl => coreml_provider(),
Accelerator::DirectMl => directml_provider(),
Accelerator::Cuda => cuda_provider(),
}
}
#[cfg(feature = "ort-coreml")]
fn coreml_provider() -> Option<Provider> {
use ort::ep::ExecutionProvider;
let provider = ort::ep::CoreML::default();
Some(Provider {
name: provider.name(),
available: provider.is_available().unwrap_or(false),
dispatch: provider.build(),
})
}
#[cfg(not(feature = "ort-coreml"))]
fn coreml_provider() -> Option<Provider> {
None
}
#[cfg(feature = "ort-directml")]
fn directml_provider() -> Option<Provider> {
use ort::ep::ExecutionProvider;
let provider = ort::ep::DirectML::default();
Some(Provider {
name: provider.name(),
available: provider.is_available().unwrap_or(false),
dispatch: provider.build(),
})
}
#[cfg(not(feature = "ort-directml"))]
fn directml_provider() -> Option<Provider> {
None
}
#[cfg(feature = "ort-cuda")]
fn cuda_provider() -> Option<Provider> {
use ort::ep::ExecutionProvider;
let provider = ort::ep::CUDA::default();
Some(Provider {
name: provider.name(),
available: provider.is_available().unwrap_or(false),
dispatch: provider.build(),
})
}
#[cfg(not(feature = "ort-cuda"))]
fn cuda_provider() -> Option<Provider> {
None
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cpu_and_auto_have_no_execution_provider_of_their_own() {
assert!(provider(Accelerator::Cpu).is_none());
assert!(provider(Accelerator::Auto).is_none());
}
#[test]
fn every_execution_provider_names_the_cargo_feature_that_enables_it() {
assert_eq!(cargo_feature(Accelerator::CoreMl), Some("ort-coreml"));
assert_eq!(cargo_feature(Accelerator::DirectMl), Some("ort-directml"));
assert_eq!(cargo_feature(Accelerator::Cuda), Some("ort-cuda"));
}
#[test]
fn metal_has_no_ort_execution_provider() {
assert_eq!(cargo_feature(Accelerator::Metal), None);
assert!(provider(Accelerator::Metal).is_none());
}
#[test]
fn preferred_accelerators_never_list_cpu_or_auto() {
let preferred = preferred_accelerators();
assert!(!preferred.is_empty(), "every platform needs at least one candidate");
assert!(
preferred.iter().all(|candidate| !candidate.is_cpu_only()),
"the CPU fallback is implicit, not a candidate: {preferred:?}"
);
}
#[cfg(target_os = "macos")]
#[test]
fn macos_prefers_coreml() {
assert_eq!(preferred_accelerators(), &[Accelerator::CoreMl]);
}
#[cfg(all(feature = "ort-bundled", not(feature = "ort-cuda")))]
#[test]
fn should_reject_an_accelerator_that_is_not_compiled_in() {
let builder = ort::session::Session::builder().expect("build session options");
let Err(error) = apply_accelerator(builder, Accelerator::Cuda) else {
panic!("cuda is not compiled in and must be rejected");
};
let message = error.to_string();
assert!(message.contains("cuda"), "message must name the accelerator: {message}");
assert!(message.contains("ort-cuda"), "message must name the feature: {message}");
}
#[cfg(all(feature = "ort-bundled", feature = "ort-coreml", target_os = "macos"))]
#[test]
fn should_register_coreml_on_macos_when_compiled_in() {
let builder = ort::session::Session::builder().expect("build session options");
let (_builder, registered) = apply_accelerator(builder, Accelerator::CoreMl).expect("coreml must register");
assert_eq!(registered, Accelerator::CoreMl);
}
#[cfg(feature = "ort-bundled")]
#[test]
fn should_leave_the_cpu_selection_on_the_implicit_provider() {
let builder = ort::session::Session::builder().expect("build session options");
let (_builder, registered) = apply_accelerator(builder, Accelerator::Cpu).expect("cpu never fails");
assert_eq!(registered, Accelerator::Cpu);
}
}