Skip to main content

rightkit_ort/
session.rs

1//! Session construction and execution-provider selection (merged from
2//! HeardRight `heardright-onnx-asr/session.rs`; `HR_ONNX_*_EXPERIMENT`
3//! environment switches dropped as product policy).
4
5use 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    /// The accelerator this OS can host: CoreML on Apple targets, DirectML on
22    /// Windows, CPU elsewhere. Runtime selection for callers that want "the
23    /// platform EP" without naming it.
24    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    /// Requested provider cannot exist on this OS. Never silently downgraded.
47    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    /// DirectML adapter index; `None` lets the runtime pick.
69    pub device_id: Option<i32>,
70    /// Explicit intra-op threads. Opts the session out of the shared pool.
71    pub threads: Option<usize>,
72    /// Set for graphs with a fixed time dimension run on DirectML: applies the
73    /// HeardRight-measured fusion workaround (ORT-DirectML 1.24.4 silently
74    /// returns near-zero output for static-shape conformer encoders). Remove
75    /// once a fixed onnxruntime-directml release passes the corpus.
76    pub directml_static_shape_workaround: bool,
77    /// When the requested accelerator is unsupported on this OS, fails to
78    /// register, or fails to load the model, build a CPU session instead.
79    /// Off by default: failures are errors. The fallback is never silent;
80    /// [`BuiltSession::fallback`] records why it happened.
81    pub cpu_fallback: bool,
82}
83
84/// A committed session plus the provider that actually serves it.
85#[derive(Debug)]
86pub struct BuiltSession {
87    pub session: ort::session::Session,
88    pub provider: ExecutionProvider,
89    /// `Some(reason)` when [`SessionOptions::cpu_fallback`] replaced the
90    /// requested accelerator with CPU.
91    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    /// Same options forced to CPU, for dispatch-bound stages (mel, decoder/joint).
118    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    /// Like [`Self::build_from_file`], also reporting the effective provider.
134    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    /// Like [`Self::build_from_memory`], also reporting the effective provider.
142    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        // Explicit CPU threads need a private pool; ORT ignores them otherwise.
192        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    /// Register the requested accelerator. ort rc.13 compiles each EP only
214    /// behind its Cargo feature, which rightkit-ort enables per OS (CoreML on
215    /// Apple, DirectML on Windows); `supported_on_this_os` already rejected
216    /// the other combinations, so the non-hosting arms are unreachable.
217    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}