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, 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    /// Number of distinct `(domain, operator)` keys in the registry.
70    pub fn operator_count(&self) -> usize {
71        self.handlers.len()
72    }
73
74    /// Number of opset-versioned inference rule entries.
75    pub fn entry_count(&self) -> usize {
76        self.handlers.values().map(Vec::len).sum()
77    }
78
79    /// Every registered rule as `(domain, operator, min_opset)`, sorted.
80    ///
81    /// This is the full identity of the catalog — its keys and opset floors,
82    /// not its handler bodies — and the thing to pin. A swap of the
83    /// [`InferenceFn`] behind an unchanged triple is deliberately out of scope
84    /// here; that is what the behavioural rule tests cover. Neither count above
85    /// can see a change that preserves them:
86    ///
87    /// - a **rename** drops one key and adds another, so `operator_count` and
88    ///   `entry_count` both hold;
89    /// - an **opset move** rewrites an existing entry's `min_opset` in place,
90    ///   so `entry_count` holds too.
91    ///
92    /// Both are silent in production rather than loud: [`Self::get`] returns
93    /// `None` for a key it does not know *and* for a version below every
94    /// registration, and [`Self::infer_node`] treats `None` permissively —
95    /// outputs are left unknown and the model still runs.
96    pub fn operator_versions(&self) -> Vec<(&str, &str, u64)> {
97        let mut rules: Vec<(&str, &str, u64)> = self
98            .handlers
99            .iter()
100            .flat_map(|((domain, op), entries)| {
101                entries
102                    .iter()
103                    .map(move |(min_opset, _)| (domain.as_str(), op.as_str(), *min_opset))
104            })
105            .collect();
106        rules.sort_unstable();
107        rules
108    }
109
110    /// Infer a single node's outputs.
111    ///
112    /// Returns one [`NodeIo`] per output slot. An unregistered op (or one whose
113    /// rule declines to resolve an output) yields empty [`NodeIo`]s — the
114    /// permissive "leave it unknown" outcome, never an error.
115    pub fn infer_node(
116        &self,
117        node: &Node,
118        opset_imports: &HashMap<String, u64>,
119        inputs: Vec<NodeIo>,
120        policy: MergePolicy,
121        interner: &mut SymbolInterner,
122    ) -> Result<Vec<NodeIo>, ShapeInferError> {
123        let version = if let Some(version) = node.local_opset() {
124            version
125        } else {
126            // Loaded IR is canonical (`normalize_domain` applied at load), so the
127            // default domain is `""` for both node domains and opset-import keys.
128            if node.is_default_domain() {
129                opset_imports.get("").copied().unwrap_or(1)
130            } else {
131                opset_imports.get(&node.domain).copied().unwrap_or(1)
132            }
133        };
134        let Some(rule) = self.get(&node.domain, &node.op_type, version) else {
135            return Ok(vec![NodeIo::default(); node.outputs.len()]);
136        };
137        let mut ctx = InferenceContext::new(node, inputs, opset_imports, policy, interner);
138        rule(&mut ctx)?;
139        Ok(ctx.into_outputs())
140    }
141}