use std::collections::HashMap;
use onnx_runtime_ir::{Node, normalize_domain};
use crate::context::{InferenceContext, MergePolicy, NodeIo, SymbolInterner};
use crate::error::ShapeInferError;
pub type InferenceFn = fn(&mut InferenceContext) -> Result<(), ShapeInferError>;
#[derive(Default)]
pub struct InferenceRegistry {
handlers: HashMap<(String, String), Vec<(u64, InferenceFn)>>,
}
impl InferenceRegistry {
pub fn empty() -> Self {
Self::default()
}
pub fn default_registry() -> Self {
let mut reg = Self::empty();
crate::handlers::register_all(&mut reg);
reg
}
pub fn register(&mut self, domain: &str, op: &str, min_opset: u64, rule: InferenceFn) {
let key = (normalize_domain(domain).to_string(), op.to_string());
let entry = self.handlers.entry(key).or_default();
match entry.binary_search_by_key(&min_opset, |(v, _)| *v) {
Ok(idx) => entry[idx] = (min_opset, rule), Err(idx) => entry.insert(idx, (min_opset, rule)),
}
}
pub fn get(&self, domain: &str, op: &str, version: u64) -> Option<InferenceFn> {
let key = (normalize_domain(domain).to_string(), op.to_string());
let entry = self.handlers.get(&key)?;
let mut chosen = None;
for &(min_opset, rule) in entry {
if min_opset <= version {
chosen = Some(rule);
} else {
break;
}
}
chosen
}
pub fn infer_node(
&self,
node: &Node,
opset_imports: &HashMap<String, u64>,
inputs: Vec<NodeIo>,
policy: MergePolicy,
interner: &mut SymbolInterner,
) -> Result<Vec<NodeIo>, ShapeInferError> {
let version = {
if node.is_default_domain() {
opset_imports.get("").copied().unwrap_or(1)
} else {
opset_imports.get(&node.domain).copied().unwrap_or(1)
}
};
let Some(rule) = self.get(&node.domain, &node.op_type, version) else {
return Ok(vec![NodeIo::default(); node.outputs.len()]);
};
let mut ctx = InferenceContext::new(node, inputs, opset_imports, policy, interner);
rule(&mut ctx)?;
Ok(ctx.into_outputs())
}
}