Skip to main content

smriti/ml/
runtime.rs

1//! ONNX Runtime initialization and management
2
3use std::path::{Path, PathBuf};
4use std::sync::atomic::{AtomicBool, Ordering};
5use std::sync::{Mutex, OnceLock};
6
7use ort::execution_providers::ExecutionProvider;
8use ort::session::builder::GraphOptimizationLevel;
9use ort::session::Session;
10
11/// Latch: only log the provider-probe result the first time a session is
12/// built. Subsequent sessions reuse whichever provider won without spamming.
13static EP_LOGGED: AtomicBool = AtomicBool::new(false);
14
15/// Best-known label for the actual execution provider in use, populated
16/// the first time a session is built. Read by Settings to surface
17/// "Face inference: GPU (DirectML)" or "Face inference: CPU".
18///
19/// Best-effort: ORT silently falls back across providers per-node, so
20/// this reflects the *best* provider that probed available, not a
21/// proof that every op runs on it.
22static ACTIVE_PROVIDER: OnceLock<&'static str> = OnceLock::new();
23static ORT_INITIALIZED: OnceLock<()> = OnceLock::new();
24static ORT_INIT_LOCK: Mutex<()> = Mutex::new(());
25
26/// User-facing label for the active execution provider, e.g. "DirectML",
27/// "CUDA", "CoreML", "CPU". Returns "CPU" before any session is built.
28pub fn active_execution_provider() -> &'static str {
29    ACTIVE_PROVIDER.get().copied().unwrap_or("CPU")
30}
31
32/// Platform-specific ONNX Runtime library name
33#[cfg(target_os = "windows")]
34const ORT_LIB_NAME: &str = "onnxruntime.dll";
35#[cfg(target_os = "macos")]
36const ORT_LIB_NAME: &str = "libonnxruntime.dylib";
37#[cfg(not(any(target_os = "windows", target_os = "macos")))]
38const ORT_LIB_NAME: &str = "libonnxruntime.so";
39
40/// ONNX Runtime environment wrapper
41///
42/// Manages the global ONNX Runtime environment and provides
43/// helper methods for loading models.
44///
45/// Uses `load-dynamic` feature: the runtime library is loaded at runtime.
46/// The library is resolved in this order:
47/// 1. `ORT_DYLIB_PATH` environment variable (if set)
48/// 2. `libs/onnxruntime/<LIB>` relative to the executable
49/// 3. `libs/onnxruntime/<LIB>` relative to the current working directory
50pub struct OnnxRuntime;
51
52impl OnnxRuntime {
53    fn usable_runtime_file(path: &Path) -> bool {
54        std::fs::metadata(path)
55            .is_ok_and(|metadata| metadata.is_file() && metadata.len() >= 1024 * 1024)
56    }
57
58    fn runtime_file_name_matches(name: &str) -> bool {
59        #[cfg(target_os = "macos")]
60        {
61            name.starts_with("libonnxruntime") && name.ends_with(".dylib")
62        }
63        #[cfg(not(target_os = "macos"))]
64        {
65            name.starts_with(ORT_LIB_NAME)
66        }
67    }
68
69    fn find_runtime_in_dir(dir: &Path) -> Option<PathBuf> {
70        let direct = dir.join(ORT_LIB_NAME);
71        if Self::usable_runtime_file(&direct) {
72            return Some(direct);
73        }
74
75        // Also check for versioned variants (e.g. libonnxruntime.so.1.23.0)
76        let entries = std::fs::read_dir(dir).ok()?;
77        for entry in entries.flatten() {
78            let path = entry.path();
79            if let Some(name) = path.file_name().and_then(|n| n.to_str()) {
80                if Self::runtime_file_name_matches(name) && Self::usable_runtime_file(&path) {
81                    return Some(path);
82                }
83            }
84        }
85
86        None
87    }
88
89    /// Resolve the path to the ONNX Runtime shared library.
90    ///
91    /// Checks ORT_DYLIB_PATH env var first, then looks in libs/onnxruntime/
92    /// relative to the executable and current working directory.
93    fn resolve_dylib_path() -> Option<PathBuf> {
94        // 1. Check ORT_DYLIB_PATH environment variable
95        if let Ok(path) = std::env::var("ORT_DYLIB_PATH") {
96            let p = PathBuf::from(&path);
97            if p.exists() {
98                tracing::info!("Using ONNX Runtime from ORT_DYLIB_PATH: {}", p.display());
99                return Some(p);
100            }
101            tracing::warn!("ORT_DYLIB_PATH set but file not found: {}", path);
102        }
103
104        // 1b. Check installed optional asset-pack roots. The asset pack
105        // stores runtimes under platform subdirectories; bootstrap owns
106        // that layout so health checks and dynamic loading agree.
107        if let Some(candidate) = crate::bootstrap::onnx_runtime_path() {
108            tracing::info!(
109                "Using ONNX Runtime from optional asset-pack path: {}",
110                candidate.display()
111            );
112            return Some(candidate);
113        }
114
115        let rel_dir = Path::new("libs").join("onnxruntime");
116
117        // 2. Relative to the executable. Also tries two levels up
118        // because `cargo tauri dev` runs the binary from
119        // `target/debug/`, and the dev-tree `libs/onnxruntime/` lives
120        // at the workspace root (target/debug/../.. = workspace).
121        if let Ok(exe) = std::env::current_exe() {
122            if let Some(exe_dir) = exe.parent() {
123                for base in [exe_dir.to_path_buf(), exe_dir.join("..").join("..")] {
124                    let candidate_dir = base.join(&rel_dir);
125                    if let Some(candidate) = Self::find_runtime_in_dir(&candidate_dir) {
126                        tracing::info!(
127                            "Using ONNX Runtime from exe-relative path: {}",
128                            candidate.display()
129                        );
130                        return Some(candidate);
131                    }
132                }
133            }
134        }
135
136        // 3. Relative to the current working directory. Also walks one
137        // level up because `cargo tauri dev` sets CWD to src-tauri/,
138        // not the workspace root.
139        if let Ok(cwd) = std::env::current_dir() {
140            let mut bases: Vec<PathBuf> = vec![cwd.clone()];
141            if let Some(parent) = cwd.parent() {
142                bases.push(parent.to_path_buf());
143            }
144            for base in bases {
145                let candidate_dir = base.join(&rel_dir);
146                if let Some(candidate) = Self::find_runtime_in_dir(&candidate_dir) {
147                    tracing::info!(
148                        "Using ONNX Runtime from cwd-relative path: {}",
149                        candidate.display()
150                    );
151                    return Some(candidate);
152                }
153            }
154        }
155
156        None
157    }
158
159    /// Initialize the ONNX Runtime global environment.
160    ///
161    /// Dynamically loads libonnxruntime.so at runtime. The library is searched
162    /// in the order described on [`OnnxRuntime`].
163    ///
164    /// This should be called once at application startup, before creating any sessions.
165    pub fn init() -> ort::Result<Self> {
166        if ORT_INITIALIZED.get().is_some() {
167            return Ok(Self);
168        }
169
170        let _guard = ORT_INIT_LOCK.lock().map_err(|_| {
171            ort::Error::new("ONNX Runtime initialization lock is poisoned".to_string())
172        })?;
173        if ORT_INITIALIZED.get().is_some() {
174            return Ok(Self);
175        }
176
177        if let Some(dylib_path) = Self::resolve_dylib_path() {
178            ort::init_from(&dylib_path)?.commit();
179            let _ = ORT_INITIALIZED.set(());
180            tracing::info!(
181                "ONNX Runtime initialized (dynamic) from: {}",
182                dylib_path.display()
183            );
184        } else {
185            return Err(ort::Error::new(
186                format!(
187                    "ONNX Runtime library not found. Set ORT_DYLIB_PATH or place {} (1.23.x) in libs/onnxruntime/",
188                    ORT_LIB_NAME
189                ),
190            ));
191        }
192        Ok(Self)
193    }
194
195    /// Load an ONNX model with a specific number of intra-op threads.
196    ///
197    /// Attempts to register GPU execution providers in platform-specific
198    /// priority order, then falls back to CPU if none initialize. The CPU
199    /// provider is always appended last so session creation never fails due
200    /// to missing GPU drivers or runtime libraries.
201    ///
202    /// Platform priority:
203    ///   Windows: DirectML  (covers NVIDIA/AMD/Intel/Qualcomm via D3D12)
204    ///   Linux:   CUDA      (NVIDIA). ROCm could be added via a cargo flag later.
205    ///   macOS:   CoreML    (Apple Silicon + AMD on Intel Macs)
206    ///
207    /// `intra_threads` applies to the CPU provider; GPU providers ignore it.
208    pub fn load_model_with_threads<P: AsRef<Path>>(
209        &self,
210        path: P,
211        intra_threads: usize,
212    ) -> ort::Result<Session> {
213        // Build the priority list. Every entry is a Dispatch value that ORT
214        // will try to initialize; if the underlying native library or driver
215        // is unavailable, ORT silently skips it and continues down the list.
216        let mut providers: Vec<ort::execution_providers::ExecutionProviderDispatch> = Vec::new();
217
218        // 1. Platform-native GPU EPs (existing priority, unchanged).
219        #[cfg(target_os = "windows")]
220        {
221            providers.push(ort::execution_providers::DirectMLExecutionProvider::default().build());
222        }
223
224        #[cfg(target_os = "linux")]
225        {
226            providers.push(ort::execution_providers::CUDAExecutionProvider::default().build());
227
228            // OpenVINO EP: only engages if the host has a custom-built libonnxruntime.so
229            // with --use_openvino AND OpenVINO toolkit installed. On a stock install
230            // this dispatch silently fails-over to the next EP, which is exactly what
231            // we want — no crash, no warning spam.
232            providers.push(
233                ort::execution_providers::OpenVINO::default()
234                    .with_device_type("GPU")
235                    .build(),
236            );
237        }
238
239        #[cfg(target_os = "macos")]
240        {
241            providers.push(ort::execution_providers::CoreMLExecutionProvider::default().build());
242        }
243
244        // 2. CPU-side accelerators — included in the standard ORT binary.
245        //    Both are silently skipped if their kernels don't apply.
246        providers.push(
247            ort::execution_providers::OneDNN::default()
248                .with_arena_allocator(true)
249                .build(),
250        );
251        providers.push(ort::execution_providers::XNNPACK::default().build());
252
253        // 3. Vanilla CPU last.
254        providers.push(ort::execution_providers::CPUExecutionProvider::default().build());
255
256        if !EP_LOGGED.swap(true, Ordering::Relaxed) {
257            Self::probe_and_log_providers();
258        }
259
260        let session = Session::builder()?
261            .with_optimization_level(GraphOptimizationLevel::Level3)?
262            .with_intra_threads(intra_threads)?
263            .with_execution_providers(providers)?
264            .commit_from_file(path)?;
265
266        Ok(session)
267    }
268
269    /// One-shot probe: log which execution providers appear usable on this
270    /// machine. Best-effort — actual provider binding happens per-session.
271    fn probe_and_log_providers() {
272        let mut available: Vec<&'static str> = Vec::new();
273
274        #[cfg(target_os = "windows")]
275        {
276            let ep = ort::execution_providers::DirectMLExecutionProvider::default();
277            if ep.is_available().unwrap_or(false) {
278                available.push("DirectML");
279            }
280        }
281
282        #[cfg(target_os = "linux")]
283        {
284            let ep = ort::execution_providers::CUDAExecutionProvider::default();
285            if ep.is_available().unwrap_or(false) {
286                available.push("CUDA");
287            }
288
289            let ep = ort::execution_providers::OpenVINO::default();
290            if ep.is_available().unwrap_or(false) {
291                available.push("OpenVINO");
292            }
293        }
294
295        #[cfg(target_os = "macos")]
296        {
297            let ep = ort::execution_providers::CoreMLExecutionProvider::default();
298            if ep.is_available().unwrap_or(false) {
299                available.push("CoreML");
300            }
301        }
302
303        {
304            let ep = ort::execution_providers::OneDNN::default();
305            if ep.is_available().unwrap_or(false) {
306                available.push("OneDNN");
307            }
308        }
309
310        {
311            let ep = ort::execution_providers::XNNPACK::default();
312            if ep.is_available().unwrap_or(false) {
313                available.push("XNNPACK");
314            }
315        }
316
317        available.push("CPU");
318
319        // Capture the best (front-of-list) provider so Settings can
320        // show what's actually engaging.
321        let chosen = *available.first().unwrap_or(&"CPU");
322        let _ = ACTIVE_PROVIDER.set(chosen);
323
324        if available.len() > 1 {
325            tracing::info!(
326                "Face inference will try execution providers in order: {}",
327                available.join(" -> ")
328            );
329        } else {
330            tracing::info!("Face inference will run on CPU (no GPU providers available)");
331        }
332    }
333}