1use ort::ep::ExecutionProvider as _;
4
5#[derive(Debug, Clone, Copy, PartialEq, Eq)]
6pub enum ProviderKind {
7 DirectMl,
8 CoreMl,
9 Cpu,
10}
11
12#[derive(Debug, Clone, PartialEq, Eq)]
13pub struct ProviderDiagnostic {
14 pub provider: ProviderKind,
15 pub available: bool,
16 pub detail: String,
17}
18
19pub fn probe_providers() -> Vec<ProviderDiagnostic> {
28 let mut out = vec![ProviderDiagnostic {
29 provider: ProviderKind::Cpu,
30 available: true,
31 detail: "always compiled in".into(),
32 }];
33 let coreml = ort::ep::CoreML::default();
34 let directml = ort::ep::DirectML::default();
35 for (kind, supported, available) in [
36 (
37 ProviderKind::CoreMl,
38 coreml.supported_by_platform(),
39 guarded(|| coreml.is_available()),
40 ),
41 (
42 ProviderKind::DirectMl,
43 directml.supported_by_platform(),
44 guarded(|| directml.is_available()),
45 ),
46 ] {
47 let (available, detail) = match (supported, available) {
48 (false, _) => (false, "not supported on this OS".to_string()),
49 (true, Ok(true)) => (true, "compiled into the loaded runtime".to_string()),
50 (true, Ok(false)) => (false, "not compiled into the loaded runtime".to_string()),
51 (true, Err(e)) => (false, format!("probe failed: {e}")),
52 };
53 out.push(ProviderDiagnostic {
54 provider: kind,
55 available,
56 detail,
57 });
58 }
59 out
60}
61
62fn guarded(f: impl FnOnce() -> ort::Result<bool> + std::panic::UnwindSafe) -> Result<bool, String> {
63 match std::panic::catch_unwind(f) {
64 Ok(result) => result.map_err(|e| e.to_string()),
65 Err(_) => Err("ONNX Runtime library could not be loaded".into()),
66 }
67}