#[cfg(any(target_vendor = "apple", target_os = "windows"))]
use ort::ep::ExecutionProvider as _;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderKind {
DirectMl,
CoreMl,
Cpu,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ProviderDiagnostic {
pub provider: ProviderKind,
pub available: bool,
pub detail: String,
}
pub fn probe_providers() -> Vec<ProviderDiagnostic> {
vec![
ProviderDiagnostic {
provider: ProviderKind::Cpu,
available: true,
detail: "always compiled in".into(),
},
diagnose(ProviderKind::CoreMl, coreml_probe()),
diagnose(ProviderKind::DirectMl, directml_probe()),
]
}
type Probe = Option<Result<bool, String>>;
fn diagnose(provider: ProviderKind, probe: Probe) -> ProviderDiagnostic {
let (available, detail) = match probe {
None => (false, "not supported on this OS".to_string()),
Some(Ok(true)) => (true, "compiled into the loaded runtime".to_string()),
Some(Ok(false)) => (false, "not compiled into the loaded runtime".to_string()),
Some(Err(e)) => (false, format!("probe failed: {e}")),
};
ProviderDiagnostic {
provider,
available,
detail,
}
}
#[cfg(target_vendor = "apple")]
fn coreml_probe() -> Probe {
Some(guarded(|| ort::ep::CoreML::default().is_available()))
}
#[cfg(not(target_vendor = "apple"))]
fn coreml_probe() -> Probe {
None
}
#[cfg(target_os = "windows")]
fn directml_probe() -> Probe {
Some(guarded(|| ort::ep::DirectML::default().is_available()))
}
#[cfg(not(target_os = "windows"))]
fn directml_probe() -> Probe {
None
}
#[cfg(any(target_vendor = "apple", target_os = "windows"))]
fn guarded(f: impl FnOnce() -> ort::Result<bool> + std::panic::UnwindSafe) -> Result<bool, String> {
match std::panic::catch_unwind(f) {
Ok(result) => result.map_err(|e| e.to_string()),
Err(_) => Err("ONNX Runtime library could not be loaded".into()),
}
}