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    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    /// Requested provider cannot exist on this OS. Never silently downgraded.
34    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    /// DirectML adapter index; `None` lets the runtime pick.
56    pub device_id: Option<i32>,
57    /// Explicit intra-op threads. Opts the session out of the shared pool.
58    pub threads: Option<usize>,
59    /// Set for graphs with a fixed time dimension run on DirectML: applies the
60    /// HeardRight-measured fusion workaround (ORT-DirectML 1.24.4 silently
61    /// returns near-zero output for static-shape conformer encoders). Remove
62    /// once a fixed onnxruntime-directml release passes the corpus.
63    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    /// Same options forced to CPU, for dispatch-bound stages (mel, decoder/joint).
85    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        // Explicit CPU threads need a private pool; ORT ignores them otherwise.
123        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}