onnx_runtime_shape_inference/
registry.rs1use std::collections::HashMap;
11
12use onnx_runtime_ir::Node;
13
14use crate::context::{InferenceContext, MergePolicy, NodeIo, SymbolInterner};
15use crate::error::ShapeInferError;
16
17pub type InferenceFn = fn(&mut InferenceContext) -> Result<(), ShapeInferError>;
20
21fn norm_domain(domain: &str) -> &str {
24 if domain == "ai.onnx" { "" } else { domain }
25}
26
27#[derive(Default)]
29pub struct InferenceRegistry {
30 handlers: HashMap<(String, String), Vec<(u64, InferenceFn)>>,
32}
33
34impl InferenceRegistry {
35 pub fn empty() -> Self {
37 Self::default()
38 }
39
40 pub fn default_registry() -> Self {
42 let mut reg = Self::empty();
43 crate::handlers::register_all(&mut reg);
44 reg
45 }
46
47 pub fn register(&mut self, domain: &str, op: &str, min_opset: u64, rule: InferenceFn) {
51 let key = (norm_domain(domain).to_string(), op.to_string());
52 let entry = self.handlers.entry(key).or_default();
53 match entry.binary_search_by_key(&min_opset, |(v, _)| *v) {
54 Ok(idx) => entry[idx] = (min_opset, rule), Err(idx) => entry.insert(idx, (min_opset, rule)),
56 }
57 }
58
59 pub fn get(&self, domain: &str, op: &str, version: u64) -> Option<InferenceFn> {
62 let key = (norm_domain(domain).to_string(), op.to_string());
63 let entry = self.handlers.get(&key)?;
64 let mut chosen = None;
65 for &(min_opset, rule) in entry {
66 if min_opset <= version {
67 chosen = Some(rule);
68 } else {
69 break;
70 }
71 }
72 chosen
73 }
74
75 pub fn infer_node(
81 &self,
82 node: &Node,
83 opset_imports: &HashMap<String, u64>,
84 inputs: Vec<NodeIo>,
85 policy: MergePolicy,
86 interner: &mut SymbolInterner,
87 ) -> Result<Vec<NodeIo>, ShapeInferError> {
88 let version = {
89 let domain = norm_domain(&node.domain);
90 if domain.is_empty() {
91 opset_imports
92 .get("")
93 .or_else(|| opset_imports.get("ai.onnx"))
94 .copied()
95 .unwrap_or(1)
96 } else {
97 opset_imports.get(domain).copied().unwrap_or(1)
98 }
99 };
100 let Some(rule) = self.get(&node.domain, &node.op_type, version) else {
101 return Ok(vec![NodeIo::default(); node.outputs.len()]);
102 };
103 let mut ctx = InferenceContext::new(node, inputs, opset_imports, policy, interner);
104 rule(&mut ctx)?;
105 Ok(ctx.into_outputs())
106 }
107}