1use std::{
6 fmt,
7 path::{Path, PathBuf},
8};
9
10use crate::environment;
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
13pub enum ExecutionProvider {
14 #[default]
15 Cpu,
16 DirectMl,
17 CoreMl,
18}
19
20impl ExecutionProvider {
21 pub fn supported_on_this_os(self) -> bool {
22 match self {
23 Self::Cpu => true,
24 Self::DirectMl => cfg!(target_os = "windows"),
25 Self::CoreMl => cfg!(target_vendor = "apple"),
26 }
27 }
28}
29
30#[derive(Debug)]
31pub enum SessionError {
32 MissingArtifact(PathBuf),
33 ProviderUnsupported(ExecutionProvider),
35 Runtime(String),
36}
37
38impl fmt::Display for SessionError {
39 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
40 match self {
41 Self::MissingArtifact(p) => write!(f, "model file missing: {}", p.display()),
42 Self::ProviderUnsupported(p) => {
43 write!(f, "execution provider {p:?} unsupported on this OS")
44 }
45 Self::Runtime(m) => write!(f, "{m}"),
46 }
47 }
48}
49
50impl std::error::Error for SessionError {}
51
52#[derive(Debug, Clone, Default)]
53pub struct SessionOptions {
54 pub provider: ExecutionProvider,
55 pub device_id: Option<i32>,
57 pub threads: Option<usize>,
59 pub directml_static_shape_workaround: bool,
64}
65
66impl SessionOptions {
67 pub fn new(provider: ExecutionProvider) -> Self {
68 Self {
69 provider,
70 ..Self::default()
71 }
72 }
73
74 pub fn with_device_id(mut self, device_id: Option<i32>) -> Self {
75 self.device_id = device_id;
76 self
77 }
78
79 pub fn with_threads(mut self, threads: Option<usize>) -> Self {
80 self.threads = threads;
81 self
82 }
83
84 pub fn cpu_for_stage(&self) -> Self {
86 Self {
87 provider: ExecutionProvider::Cpu,
88 ..self.clone()
89 }
90 }
91
92 pub fn build_from_file(&self, path: &Path) -> Result<ort::session::Session, SessionError> {
93 if !path.is_file() {
94 return Err(SessionError::MissingArtifact(path.to_path_buf()));
95 }
96 self.build(|b| b.commit_from_file(path), &path.display().to_string())
97 }
98
99 pub fn build_from_memory(&self, bytes: &[u8]) -> Result<ort::session::Session, SessionError> {
100 self.build(|b| b.commit_from_memory(bytes), "<memory>")
101 }
102
103 fn build(
104 &self,
105 commit: impl FnOnce(
106 &mut ort::session::builder::SessionBuilder,
107 ) -> ort::Result<ort::session::Session>,
108 what: &str,
109 ) -> Result<ort::session::Session, SessionError> {
110 if !self.provider.supported_on_this_os() {
111 return Err(SessionError::ProviderUnsupported(self.provider));
112 }
113 let mut builder = ort::session::Session::builder().map_err(rt("session builder"))?;
114 let shared_pool = environment::shared_pool_active();
115
116 let directml = self.provider == ExecutionProvider::DirectMl;
117 if directml {
118 builder = builder
119 .with_independent_thread_pool()
120 .map_err(|e| map_ort_error("independent DirectML pool", e))?;
121 }
122 if self.provider == ExecutionProvider::Cpu && self.threads.is_some() && shared_pool {
124 builder = builder
125 .with_independent_thread_pool()
126 .map_err(|e| map_ort_error("independent CPU stage", e))?;
127 }
128 if let Some(threads) = self.threads {
129 builder = builder
130 .with_intra_threads(threads)
131 .map_err(|e| map_ort_error("intra threads", e))?;
132 }
133 if self.directml_static_shape_workaround && directml {
134 builder = builder
135 .with_config_entry("ep.dml.disable_graph_fusion", "1")
136 .map_err(|e| map_ort_error("disable_graph_fusion", e))?
137 .with_optimization_level(ort::session::builder::GraphOptimizationLevel::Level1)
138 .map_err(|e| map_ort_error("optimization level", e))?;
139 }
140 builder = match self.provider {
141 ExecutionProvider::Cpu => builder,
142 ExecutionProvider::DirectMl => {
143 let mut ep = ort::ep::DirectML::default();
144 if let Some(id) = self.device_id {
145 ep = ep.with_device_id(id);
146 }
147 builder
148 .with_execution_providers([ep.build().error_on_failure()])
149 .map_err(|e| map_ort_error("directml", e))?
150 }
151 ExecutionProvider::CoreMl => builder
152 .with_execution_providers([ort::ep::CoreML::default().build().error_on_failure()])
153 .map_err(|e| map_ort_error("coreml", e))?,
154 };
155 commit(&mut builder).map_err(|e| SessionError::Runtime(format!("load {what}: {e}")))
156 }
157}
158
159fn rt(label: &str) -> impl FnOnce(ort::Error) -> SessionError {
160 let label = label.to_owned();
161 move |error| SessionError::Runtime(format!("{label}: {error}"))
162}
163
164fn map_ort_error<R>(label: &str, error: ort::Error<R>) -> SessionError {
165 SessionError::Runtime(format!("{label}: {error}"))
166}