Skip to main content

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;
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/// Normalise the default ONNX domain: the empty string and `"ai.onnx"` are the
22/// same domain.
23fn norm_domain(domain: &str) -> &str {
24    if domain == "ai.onnx" { "" } else { domain }
25}
26
27/// A registry mapping `(domain, op_type, opset)` to an [`InferenceFn`].
28#[derive(Default)]
29pub struct InferenceRegistry {
30    /// `(domain, op)` → ascending list of `(min_opset, rule)`.
31    handlers: HashMap<(String, String), Vec<(u64, InferenceFn)>>,
32}
33
34impl InferenceRegistry {
35    /// An empty registry (no rules).
36    pub fn empty() -> Self {
37        Self::default()
38    }
39
40    /// A registry populated with every built-in rule.
41    pub fn default_registry() -> Self {
42        let mut reg = Self::empty();
43        crate::handlers::register_all(&mut reg);
44        reg
45    }
46
47    /// Register `rule` for `(domain, op)` applying from opset `min_opset`
48    /// upward. A later registration at a higher `min_opset` supersedes this one
49    /// for those versions.
50    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), // replace same-version rule
55            Err(idx) => entry.insert(idx, (min_opset, rule)),
56        }
57    }
58
59    /// Look up the rule for `(domain, op)` effective at opset `version`: the
60    /// registration with the greatest `min_opset <= version`.
61    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    /// Infer a single node's outputs.
76    ///
77    /// Returns one [`NodeIo`] per output slot. An unregistered op (or one whose
78    /// rule declines to resolve an output) yields empty [`NodeIo`]s — the
79    /// permissive "leave it unknown" outcome, never an error.
80    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}