use std::collections::HashMap;
use onnx_runtime_ir::{DataType, Dim, Node, SymbolId, ValueId};
use crate::dim_expr::DimExpr;
use crate::error::ShapeInferError;
use crate::shape_data::ShapeData;
pub type TypedShape = Vec<DimExpr>;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TypeInfo {
pub dtype: DataType,
pub shape: TypedShape,
}
impl TypeInfo {
pub fn new(dtype: DataType, shape: TypedShape) -> Self {
Self { dtype, shape }
}
pub fn rank(&self) -> usize {
self.shape.len()
}
}
#[derive(Clone, Debug, Default)]
pub struct NodeIo {
pub type_info: Option<TypeInfo>,
pub shape_data: Option<ShapeData>,
}
impl NodeIo {
pub fn typed(type_info: TypeInfo) -> Self {
Self {
type_info: Some(type_info),
shape_data: None,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum MergePolicy {
#[default]
Permissive,
Strict,
}
#[derive(Debug)]
pub struct SymbolInterner {
next: u32,
cache: HashMap<DimExpr, SymbolId>,
fresh: Vec<SymbolId>,
}
impl SymbolInterner {
pub fn new(next: u32) -> Self {
Self {
next,
cache: HashMap::new(),
fresh: Vec::new(),
}
}
pub fn fresh_symbol(&mut self) -> SymbolId {
let id = SymbolId(self.next);
self.next = self.next.saturating_add(1);
self.fresh.push(id);
id
}
pub fn fresh_dim(&mut self) -> DimExpr {
DimExpr::symbol(self.fresh_symbol())
}
pub fn lower(&mut self, expr: &DimExpr) -> Dim {
if expr.is_overflow() {
return Dim::Symbolic(self.fresh_symbol());
}
if let Some(n) = expr.as_const() {
if n >= 0 {
return Dim::Static(n as usize);
}
return Dim::Symbolic(self.fresh_symbol());
}
if let Some(s) = expr.as_symbol() {
return Dim::Symbolic(s);
}
if let Some(&id) = self.cache.get(expr) {
return Dim::Symbolic(id);
}
let id = self.fresh_symbol();
self.cache.insert(expr.clone(), id);
Dim::Symbolic(id)
}
pub fn fresh_symbols(&self) -> &[SymbolId] {
&self.fresh
}
}
pub struct InferenceContext<'a> {
pub node: &'a Node,
opset_imports: &'a HashMap<String, u64>,
policy: MergePolicy,
inputs: Vec<NodeIo>,
outputs: Vec<NodeIo>,
interner: &'a mut SymbolInterner,
}
impl<'a> InferenceContext<'a> {
pub fn new(
node: &'a Node,
inputs: Vec<NodeIo>,
opset_imports: &'a HashMap<String, u64>,
policy: MergePolicy,
interner: &'a mut SymbolInterner,
) -> Self {
let outputs = vec![NodeIo::default(); node.outputs.len()];
Self {
node,
opset_imports,
policy,
inputs,
outputs,
interner,
}
}
pub fn op(&self) -> &str {
&self.node.op_type
}
pub fn num_inputs(&self) -> usize {
self.inputs.len()
}
pub fn num_outputs(&self) -> usize {
self.outputs.len()
}
pub fn has_input(&self, i: usize) -> bool {
self.node
.inputs
.get(i)
.map(Option::is_some)
.unwrap_or(false)
}
pub fn input_type(&self, i: usize) -> Option<&TypeInfo> {
self.inputs.get(i)?.type_info.as_ref()
}
pub fn input_shape(&self, i: usize) -> Option<&[DimExpr]> {
self.input_type(i).map(|t| t.shape.as_slice())
}
pub fn input_dtype(&self, i: usize) -> Option<DataType> {
self.input_type(i).map(|t| t.dtype)
}
pub fn input_rank(&self, i: usize) -> Option<usize> {
self.input_type(i).map(TypeInfo::rank)
}
pub fn input_shape_data(&self, i: usize) -> Option<&ShapeData> {
self.inputs.get(i)?.shape_data.as_ref()
}
pub fn set_output_type(&mut self, i: usize, type_info: TypeInfo) {
if let Some(slot) = self.outputs.get_mut(i) {
slot.type_info = Some(type_info);
}
}
pub fn set_output(&mut self, i: usize, dtype: DataType, shape: TypedShape) {
self.set_output_type(i, TypeInfo::new(dtype, shape));
}
pub fn set_output_shape_data(&mut self, i: usize, data: ShapeData) {
if let Some(slot) = self.outputs.get_mut(i) {
slot.shape_data = Some(data);
}
}
pub fn into_outputs(self) -> Vec<NodeIo> {
self.outputs
}
pub fn policy(&self) -> MergePolicy {
self.policy
}
pub fn opset(&self, domain: &str) -> u64 {
if domain.is_empty() || domain == "ai.onnx" {
self.opset_imports
.get("")
.or_else(|| self.opset_imports.get("ai.onnx"))
.copied()
.unwrap_or(1)
} else {
self.opset_imports.get(domain).copied().unwrap_or(1)
}
}
pub fn fresh_dim(&mut self) -> DimExpr {
self.interner.fresh_dim()
}
pub fn broadcast(
&mut self,
a: &[DimExpr],
b: &[DimExpr],
) -> Result<TypedShape, ShapeInferError> {
let rank = a.len().max(b.len());
let mut out = Vec::with_capacity(rank);
for axis in 0..rank {
let da = dim_from_right(a, rank, axis);
let db = dim_from_right(b, rank, axis);
out.push(self.broadcast_dim(&da, &db)?);
}
Ok(out)
}
pub fn broadcast_dim(&mut self, a: &DimExpr, b: &DimExpr) -> Result<DimExpr, ShapeInferError> {
let ac = a.as_const();
let bc = b.as_const();
if ac == Some(1) {
return Ok(b.clone());
}
if bc == Some(1) {
return Ok(a.clone());
}
if a == b {
return Ok(a.clone());
}
match (ac, bc) {
(Some(x), Some(y)) => {
if x == y {
Ok(a.clone())
} else if self.policy == MergePolicy::Strict {
Err(ShapeInferError::Invalid {
op: self.node.op_type.clone(),
detail: format!("incompatible broadcast dims {x} and {y}"),
})
} else {
Ok(self.fresh_dim())
}
}
(Some(_), None) => Ok(a.clone()),
(None, Some(_)) => Ok(b.clone()),
(None, None) => match (a.as_symbol(), b.as_symbol()) {
(Some(sa), Some(sb)) => {
Ok(if sa.0 <= sb.0 { a.clone() } else { b.clone() })
}
_ => Ok(self.fresh_dim()),
},
}
}
}
fn dim_from_right(shape: &[DimExpr], rank: usize, axis: usize) -> DimExpr {
let offset = rank - shape.len();
if axis < offset {
DimExpr::constant(1)
} else {
shape[axis - offset].clone()
}
}
pub fn merge_shapes(
value: ValueId,
inferred: &[DimExpr],
declared: &[Dim],
policy: MergePolicy,
) -> Result<Vec<DimExpr>, ShapeInferError> {
if inferred.len() != declared.len() {
if policy == MergePolicy::Strict {
return Err(ShapeInferError::RankConflict {
value,
inferred: inferred.len(),
declared: declared.len(),
});
}
return Ok(inferred.to_vec());
}
let mut out = Vec::with_capacity(inferred.len());
for (axis, (inf, dec)) in inferred.iter().zip(declared.iter()).enumerate() {
let dec_expr: DimExpr = (*dec).into();
let merged = match (inf.as_const(), dec_expr.as_const()) {
(Some(a), Some(b)) if a != b => {
if policy == MergePolicy::Strict {
return Err(ShapeInferError::ShapeConflict {
value,
axis,
inferred: a,
declared: b,
});
}
inf.clone()
}
(Some(_), _) => inf.clone(),
(None, Some(_)) => dec_expr,
(None, None) => inf.clone(),
};
out.push(merged);
}
Ok(out)
}