onnx_runtime_shape_inference/registry.rs
1//! The extensible, opset-aware operator registry.
2//!
3//! Inference rules are keyed by `(domain, op_type)` and, within a key, by the
4//! opset version at which they were introduced. Registration is *range-based*:
5//! a rule registered at version `N` applies to every opset `>= N` until a later
6//! registration supersedes it — mirroring how ONNX operator schemas evolve.
7//! Unregistered operators are not an error; their outputs are simply left
8//! unresolved (permissive behaviour).
9
10use std::collections::HashMap;
11
12use onnx_runtime_ir::{Node, normalize_domain};
13
14use crate::context::{InferenceContext, MergePolicy, NodeIo, SymbolInterner};
15use crate::error::ShapeInferError;
16
17/// An operator inference rule: reads inputs from the [`InferenceContext`] and
18/// sets its outputs' types (and, where applicable, shape-data).
19pub type InferenceFn = fn(&mut InferenceContext) -> Result<(), ShapeInferError>;
20
21/// A registry mapping `(domain, op_type, opset)` to an [`InferenceFn`].
22#[derive(Default)]
23pub struct InferenceRegistry {
24 /// `(domain, op)` → ascending list of `(min_opset, rule)`.
25 handlers: HashMap<(String, String), Vec<(u64, InferenceFn)>>,
26}
27
28impl InferenceRegistry {
29 /// An empty registry (no rules).
30 pub fn empty() -> Self {
31 Self::default()
32 }
33
34 /// A registry populated with every built-in rule.
35 pub fn default_registry() -> Self {
36 let mut reg = Self::empty();
37 crate::handlers::register_all(&mut reg);
38 reg
39 }
40
41 /// Register `rule` for `(domain, op)` applying from opset `min_opset`
42 /// upward. A later registration at a higher `min_opset` supersedes this one
43 /// for those versions.
44 pub fn register(&mut self, domain: &str, op: &str, min_opset: u64, rule: InferenceFn) {
45 let key = (normalize_domain(domain).to_string(), op.to_string());
46 let entry = self.handlers.entry(key).or_default();
47 match entry.binary_search_by_key(&min_opset, |(v, _)| *v) {
48 Ok(idx) => entry[idx] = (min_opset, rule), // replace same-version rule
49 Err(idx) => entry.insert(idx, (min_opset, rule)),
50 }
51 }
52
53 /// Look up the rule for `(domain, op)` effective at opset `version`: the
54 /// registration with the greatest `min_opset <= version`.
55 pub fn get(&self, domain: &str, op: &str, version: u64) -> Option<InferenceFn> {
56 let key = (normalize_domain(domain).to_string(), op.to_string());
57 let entry = self.handlers.get(&key)?;
58 let mut chosen = None;
59 for &(min_opset, rule) in entry {
60 if min_opset <= version {
61 chosen = Some(rule);
62 } else {
63 break;
64 }
65 }
66 chosen
67 }
68
69 /// Infer a single node's outputs.
70 ///
71 /// Returns one [`NodeIo`] per output slot. An unregistered op (or one whose
72 /// rule declines to resolve an output) yields empty [`NodeIo`]s — the
73 /// permissive "leave it unknown" outcome, never an error.
74 pub fn infer_node(
75 &self,
76 node: &Node,
77 opset_imports: &HashMap<String, u64>,
78 inputs: Vec<NodeIo>,
79 policy: MergePolicy,
80 interner: &mut SymbolInterner,
81 ) -> Result<Vec<NodeIo>, ShapeInferError> {
82 let version = {
83 // Loaded IR is canonical (`normalize_domain` applied at load), so the
84 // default domain is `""` for both node domains and opset-import keys.
85 if node.is_default_domain() {
86 opset_imports.get("").copied().unwrap_or(1)
87 } else {
88 opset_imports.get(&node.domain).copied().unwrap_or(1)
89 }
90 };
91 let Some(rule) = self.get(&node.domain, &node.op_type, version) else {
92 return Ok(vec![NodeIo::default(); node.outputs.len()]);
93 };
94 let mut ctx = InferenceContext::new(node, inputs, opset_imports, policy, interner);
95 rule(&mut ctx)?;
96 Ok(ctx.into_outputs())
97 }
98}