1#[cfg(any(target_vendor = "apple", target_os = "windows"))]
4use ort::ep::ExecutionProvider as _;
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum ProviderKind {
8 DirectMl,
9 CoreMl,
10 Cpu,
11}
12
13#[derive(Debug, Clone, PartialEq, Eq)]
14pub struct ProviderDiagnostic {
15 pub provider: ProviderKind,
16 pub available: bool,
17 pub detail: String,
18}
19
20pub fn probe_providers() -> Vec<ProviderDiagnostic> {
29 vec![
30 ProviderDiagnostic {
31 provider: ProviderKind::Cpu,
32 available: true,
33 detail: "always compiled in".into(),
34 },
35 diagnose(ProviderKind::CoreMl, coreml_probe()),
36 diagnose(ProviderKind::DirectMl, directml_probe()),
37 ]
38}
39
40type Probe = Option<Result<bool, String>>;
43
44fn diagnose(provider: ProviderKind, probe: Probe) -> ProviderDiagnostic {
45 let (available, detail) = match probe {
46 None => (false, "not supported on this OS".to_string()),
47 Some(Ok(true)) => (true, "compiled into the loaded runtime".to_string()),
48 Some(Ok(false)) => (false, "not compiled into the loaded runtime".to_string()),
49 Some(Err(e)) => (false, format!("probe failed: {e}")),
50 };
51 ProviderDiagnostic {
52 provider,
53 available,
54 detail,
55 }
56}
57
58#[cfg(target_vendor = "apple")]
59fn coreml_probe() -> Probe {
60 Some(guarded(|| ort::ep::CoreML::default().is_available()))
61}
62
63#[cfg(not(target_vendor = "apple"))]
64fn coreml_probe() -> Probe {
65 None
66}
67
68#[cfg(target_os = "windows")]
69fn directml_probe() -> Probe {
70 Some(guarded(|| ort::ep::DirectML::default().is_available()))
71}
72
73#[cfg(not(target_os = "windows"))]
74fn directml_probe() -> Probe {
75 None
76}
77
78#[cfg(any(target_vendor = "apple", target_os = "windows"))]
79fn guarded(f: impl FnOnce() -> ort::Result<bool> + std::panic::UnwindSafe) -> Result<bool, String> {
80 match std::panic::catch_unwind(f) {
81 Ok(result) => result.map_err(|e| e.to_string()),
82 Err(_) => Err("ONNX Runtime library could not be loaded".into()),
83 }
84}