1use std::collections::HashMap;
4use std::path::Path;
5
6use onnx_runtime_ir::{DataType, Node, Shape, TensorLayout};
7
8use crate::error::Result;
9use crate::kernel::{Kernel, KernelMatch};
10use crate::provider::{EpConfig, EpId, ExecutionProvider};
11
12#[derive(Clone, PartialEq, Eq, Hash, Debug)]
14pub struct OpKey {
15 pub op_type: String,
16 pub domain: String,
17 pub since_version: u64,
18}
19
20impl OpKey {
21 pub fn new(op_type: impl Into<String>, domain: impl Into<String>, since_version: u64) -> Self {
22 Self {
23 op_type: op_type.into(),
24 domain: domain.into(),
25 since_version,
26 }
27 }
28}
29
30fn norm_domain(domain: &str) -> &str {
34 if domain == "ai.onnx" { "" } else { domain }
35}
36
37pub trait KernelFactory: Send + Sync {
39 fn create(&self, node: &Node, input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>>;
40}
41
42#[derive(Default)]
44pub struct OpRegistry {
45 entries: HashMap<OpKey, Box<dyn KernelFactory>>,
46 by_op: HashMap<String, HashMap<String, Vec<u64>>>,
48}
49
50impl OpRegistry {
51 pub fn new() -> Self {
52 Self::default()
53 }
54
55 pub fn register(&mut self, mut key: OpKey, factory: Box<dyn KernelFactory>) {
57 key.domain = norm_domain(&key.domain).to_owned();
58 let versions = self
59 .by_op
60 .entry(key.domain.clone())
61 .or_default()
62 .entry(key.op_type.clone())
63 .or_default();
64 if let Err(index) = versions.binary_search(&key.since_version) {
65 versions.insert(index, key.since_version);
66 }
67 self.entries.insert(key, factory);
68 }
69
70 pub fn lookup(&self, op_type: &str, domain: &str, opset: u64) -> Option<&dyn KernelFactory> {
73 let domain = norm_domain(domain);
74 let versions = self.by_op.get(domain)?.get(op_type)?;
75 let index = versions.partition_point(|&version| version <= opset);
76 let since_version = *versions.get(index.checked_sub(1)?)?;
77 self.entries
78 .get(&OpKey::new(op_type, domain, since_version))
79 .map(Box::as_ref)
80 }
81
82 pub fn supports(&self, op_type: &str, domain: &str, opset: u64) -> bool {
85 let domain = norm_domain(domain);
86 self.by_op
87 .get(domain)
88 .and_then(|ops| ops.get(op_type))
89 .and_then(|versions| versions.first())
90 .is_some_and(|&since_version| since_version <= opset)
91 }
92
93 pub fn earliest_since_version(&self, op_type: &str, domain: &str) -> Option<u64> {
96 let domain = norm_domain(domain);
97 self.by_op.get(domain)?.get(op_type)?.first().copied()
98 }
99
100 pub fn len(&self) -> usize {
102 self.entries.len()
103 }
104
105 pub fn is_empty(&self) -> bool {
107 self.entries.is_empty()
108 }
109}
110
111#[cfg(test)]
112mod tests {
113 use super::*;
114
115 struct DummyFactory(u64);
116
117 impl KernelFactory for DummyFactory {
118 fn create(&self, _node: &Node, _input_shapes: &[Vec<usize>]) -> Result<Box<dyn Kernel>> {
119 let _ = self.0;
120 unreachable!("registry tests do not create kernels")
121 }
122 }
123
124 #[test]
125 fn indexed_queries_match_linear_reference() {
126 let mut registry = OpRegistry::new();
127 let mut state = 0x9e37_79b9_u64;
128 let ops = ["Add", "Mul", "Gemm", "Attention"];
129 let domains = ["", "ai.onnx", "com.microsoft", "pkg.nxrt"];
130
131 for factory_id in 0..256 {
132 state = state
133 .wrapping_mul(6_364_136_223_846_793_005)
134 .wrapping_add(1);
135 let op_type = ops[(state as usize) % ops.len()];
136 state = state
137 .wrapping_mul(6_364_136_223_846_793_005)
138 .wrapping_add(1);
139 let domain = domains[(state as usize) % domains.len()];
140 state = state
141 .wrapping_mul(6_364_136_223_846_793_005)
142 .wrapping_add(1);
143 let since_version = state % 25;
144 registry.register(
145 OpKey::new(op_type, domain, since_version),
146 Box::new(DummyFactory(factory_id)),
147 );
148 }
149
150 for _ in 0..512 {
151 state = state
152 .wrapping_mul(6_364_136_223_846_793_005)
153 .wrapping_add(1);
154 let op_type = ops[(state as usize) % ops.len()];
155 state = state
156 .wrapping_mul(6_364_136_223_846_793_005)
157 .wrapping_add(1);
158 let domain = domains[(state as usize) % domains.len()];
159 state = state
160 .wrapping_mul(6_364_136_223_846_793_005)
161 .wrapping_add(1);
162 let opset = state % 30;
163 let domain = norm_domain(domain);
164
165 let linear_lookup = registry
166 .entries
167 .iter()
168 .filter(|(key, _)| {
169 key.op_type == op_type && key.domain == domain && key.since_version <= opset
170 })
171 .max_by_key(|(key, _)| key.since_version)
172 .map(|(_, factory)| factory.as_ref());
173 match (registry.lookup(op_type, domain, opset), linear_lookup) {
174 (Some(indexed), Some(linear)) => assert!(std::ptr::eq(indexed, linear)),
175 (None, None) => {}
176 _ => panic!("indexed lookup differed from linear reference"),
177 }
178
179 let linear_supports = registry.entries.keys().any(|key| {
180 key.op_type == op_type && key.domain == domain && key.since_version <= opset
181 });
182 assert_eq!(registry.supports(op_type, domain, opset), linear_supports);
183
184 let linear_earliest = registry
185 .entries
186 .keys()
187 .filter(|key| key.op_type == op_type && key.domain == domain)
188 .map(|key| key.since_version)
189 .min();
190 assert_eq!(
191 registry.earliest_since_version(op_type, domain),
192 linear_earliest
193 );
194 }
195 }
196}
197
198#[derive(Default)]
200pub struct EpRegistry {
201 eps: Vec<Box<dyn ExecutionProvider>>,
202 priority: Vec<EpId>,
204}
205
206impl EpRegistry {
207 pub fn new() -> Self {
208 Self::default()
209 }
210
211 pub fn register(&mut self, ep: Box<dyn ExecutionProvider>) -> EpId {
213 let id = EpId(self.eps.len() as u32);
214 self.eps.push(ep);
215 self.priority.push(id);
216 id
217 }
218
219 pub fn load_legacy(&mut self, path: &Path, config: &EpConfig) -> Result<EpId> {
221 let _ = (path, config);
222 todo!("ort2-ep-api Phase 2: dlopen legacy ORT plugin EP and adapt its vtable")
223 }
224
225 pub fn set_priority(&mut self, order: Vec<EpId>) {
227 self.priority = order;
228 }
229
230 pub fn get(&self, id: EpId) -> Option<&dyn ExecutionProvider> {
232 self.eps.get(id.0 as usize).map(|b| b.as_ref())
233 }
234
235 pub fn priority(&self) -> &[EpId] {
237 &self.priority
238 }
239
240 pub fn candidates_for_op(
242 &self,
243 op: &Node,
244 opset: u64,
245 shapes: &[Shape],
246 input_dtypes: &[DataType],
247 layouts: &[TensorLayout],
248 ) -> Vec<(EpId, KernelMatch)> {
249 let mut out = Vec::new();
250 for &id in &self.priority {
251 if let Some(ep) = self.get(id) {
252 let m = ep.supports_op(op, opset, shapes, input_dtypes, layouts);
253 if m.is_supported() {
254 out.push((id, m));
255 }
256 }
257 }
258 out
259 }
260}