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 platform_accelerator() -> Self {
25 if cfg!(target_vendor = "apple") {
26 Self::CoreMl
27 } else if cfg!(target_os = "windows") {
28 Self::DirectMl
29 } else {
30 Self::Cpu
31 }
32 }
33
34 pub fn supported_on_this_os(self) -> bool {
35 match self {
36 Self::Cpu => true,
37 Self::DirectMl => cfg!(target_os = "windows"),
38 Self::CoreMl => cfg!(target_vendor = "apple"),
39 }
40 }
41}
42
43#[derive(Debug)]
44pub enum SessionError {
45 MissingArtifact(PathBuf),
46 ProviderUnsupported(ExecutionProvider),
48 Runtime(String),
49}
50
51impl fmt::Display for SessionError {
52 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
53 match self {
54 Self::MissingArtifact(p) => write!(f, "model file missing: {}", p.display()),
55 Self::ProviderUnsupported(p) => {
56 write!(f, "execution provider {p:?} unsupported on this OS")
57 }
58 Self::Runtime(m) => write!(f, "{m}"),
59 }
60 }
61}
62
63impl std::error::Error for SessionError {}
64
65#[derive(Debug, Clone, Default)]
66pub struct SessionOptions {
67 pub provider: ExecutionProvider,
68 pub device_id: Option<i32>,
70 pub threads: Option<usize>,
72 pub directml_static_shape_workaround: bool,
77 pub cpu_fallback: bool,
82}
83
84#[derive(Debug)]
86pub struct BuiltSession {
87 pub session: ort::session::Session,
88 pub provider: ExecutionProvider,
89 pub fallback: Option<String>,
92}
93
94impl SessionOptions {
95 pub fn new(provider: ExecutionProvider) -> Self {
96 Self {
97 provider,
98 ..Self::default()
99 }
100 }
101
102 pub fn with_device_id(mut self, device_id: Option<i32>) -> Self {
103 self.device_id = device_id;
104 self
105 }
106
107 pub fn with_threads(mut self, threads: Option<usize>) -> Self {
108 self.threads = threads;
109 self
110 }
111
112 pub fn with_cpu_fallback(mut self, cpu_fallback: bool) -> Self {
113 self.cpu_fallback = cpu_fallback;
114 self
115 }
116
117 pub fn cpu_for_stage(&self) -> Self {
119 Self {
120 provider: ExecutionProvider::Cpu,
121 ..self.clone()
122 }
123 }
124
125 pub fn build_from_file(&self, path: &Path) -> Result<ort::session::Session, SessionError> {
126 self.build_from_file_reported(path).map(|b| b.session)
127 }
128
129 pub fn build_from_memory(&self, bytes: &[u8]) -> Result<ort::session::Session, SessionError> {
130 self.build_from_memory_reported(bytes).map(|b| b.session)
131 }
132
133 pub fn build_from_file_reported(&self, path: &Path) -> Result<BuiltSession, SessionError> {
135 if !path.is_file() {
136 return Err(SessionError::MissingArtifact(path.to_path_buf()));
137 }
138 self.build_with_fallback(|b| b.commit_from_file(path), &path.display().to_string())
139 }
140
141 pub fn build_from_memory_reported(&self, bytes: &[u8]) -> Result<BuiltSession, SessionError> {
143 self.build_with_fallback(|b| b.commit_from_memory(bytes), "<memory>")
144 }
145
146 fn build_with_fallback(
147 &self,
148 commit: impl Fn(
149 &mut ort::session::builder::SessionBuilder,
150 ) -> ort::Result<ort::session::Session>,
151 what: &str,
152 ) -> Result<BuiltSession, SessionError> {
153 match self.build(&commit, what) {
154 Ok(session) => Ok(BuiltSession {
155 session,
156 provider: self.provider,
157 fallback: None,
158 }),
159 Err(SessionError::MissingArtifact(p)) => Err(SessionError::MissingArtifact(p)),
160 Err(error) if self.cpu_fallback && self.provider != ExecutionProvider::Cpu => {
161 let session = self.cpu_for_stage().build(&commit, what)?;
162 Ok(BuiltSession {
163 session,
164 provider: ExecutionProvider::Cpu,
165 fallback: Some(format!("{:?} -> Cpu: {error}", self.provider)),
166 })
167 }
168 Err(error) => Err(error),
169 }
170 }
171
172 fn build(
173 &self,
174 commit: &impl Fn(
175 &mut ort::session::builder::SessionBuilder,
176 ) -> ort::Result<ort::session::Session>,
177 what: &str,
178 ) -> Result<ort::session::Session, SessionError> {
179 if !self.provider.supported_on_this_os() {
180 return Err(SessionError::ProviderUnsupported(self.provider));
181 }
182 let mut builder = ort::session::Session::builder().map_err(rt("session builder"))?;
183 let shared_pool = environment::shared_pool_active();
184
185 let directml = self.provider == ExecutionProvider::DirectMl;
186 if directml {
187 builder = builder
188 .with_independent_thread_pool()
189 .map_err(|e| map_ort_error("independent DirectML pool", e))?;
190 }
191 if self.provider == ExecutionProvider::Cpu && self.threads.is_some() && shared_pool {
193 builder = builder
194 .with_independent_thread_pool()
195 .map_err(|e| map_ort_error("independent CPU stage", e))?;
196 }
197 if let Some(threads) = self.threads {
198 builder = builder
199 .with_intra_threads(threads)
200 .map_err(|e| map_ort_error("intra threads", e))?;
201 }
202 if self.directml_static_shape_workaround && directml {
203 builder = builder
204 .with_config_entry("ep.dml.disable_graph_fusion", "1")
205 .map_err(|e| map_ort_error("disable_graph_fusion", e))?
206 .with_optimization_level(ort::session::builder::GraphOptimizationLevel::Level1)
207 .map_err(|e| map_ort_error("optimization level", e))?;
208 }
209 builder = self.register_provider(builder)?;
210 commit(&mut builder).map_err(|e| SessionError::Runtime(format!("load {what}: {e}")))
211 }
212
213 fn register_provider(
218 &self,
219 builder: ort::session::builder::SessionBuilder,
220 ) -> Result<ort::session::builder::SessionBuilder, SessionError> {
221 match self.provider {
222 ExecutionProvider::Cpu => Ok(builder),
223 #[cfg(target_os = "windows")]
224 ExecutionProvider::DirectMl => {
225 let mut ep = ort::ep::DirectML::default();
226 if let Some(id) = self.device_id {
227 ep = ep.with_device_id(id);
228 }
229 builder
230 .with_execution_providers([ep.build().error_on_failure()])
231 .map_err(|e| map_ort_error("directml", e))
232 }
233 #[cfg(target_vendor = "apple")]
234 ExecutionProvider::CoreMl => builder
235 .with_execution_providers([ort::ep::CoreML::default().build().error_on_failure()])
236 .map_err(|e| map_ort_error("coreml", e)),
237 #[allow(unreachable_patterns)]
238 other => Err(SessionError::ProviderUnsupported(other)),
239 }
240 }
241}
242
243fn rt(label: &str) -> impl FnOnce(ort::Error) -> SessionError {
244 let label = label.to_owned();
245 move |error| SessionError::Runtime(format!("{label}: {error}"))
246}
247
248fn map_ort_error<R>(label: &str, error: ort::Error<R>) -> SessionError {
249 SessionError::Runtime(format!("{label}: {error}"))
250}