use std::collections::HashMap;
use onnx_runtime_ir::{DataType, Dim, Node, SymbolId, ValueId, normalize_domain};
use crate::dim_expr::DimExpr;
use crate::error::ShapeInferError;
use crate::infer::ANON_SYMBOL_FLOOR;
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, PartialEq, Eq)]
pub struct TensorType {
pub dtype: DataType,
pub shape: Option<TypedShape>,
}
impl TensorType {
pub fn new(dtype: DataType, shape: TypedShape) -> Self {
Self {
dtype,
shape: Some(shape),
}
}
pub fn dtype_only(dtype: DataType) -> Self {
Self { dtype, shape: None }
}
pub fn to_type_info(&self) -> Option<TypeInfo> {
self.shape
.as_ref()
.map(|shape| TypeInfo::new(self.dtype, shape.clone()))
}
}
impl From<TypeInfo> for TensorType {
fn from(type_info: TypeInfo) -> Self {
Self {
dtype: type_info.dtype,
shape: Some(type_info.shape),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ValueType {
Tensor(TensorType),
Sequence(Box<ValueType>),
Optional(Box<ValueType>),
Map(DataType, Box<ValueType>),
}
impl ValueType {
pub fn tensor(dtype: DataType, shape: TypedShape) -> Self {
Self::Tensor(TensorType::new(dtype, shape))
}
pub fn sequence(element: ValueType) -> Self {
Self::Sequence(Box::new(element))
}
pub fn as_tensor(&self) -> Option<&TensorType> {
match self {
Self::Tensor(tensor) => Some(tensor),
_ => None,
}
}
pub fn as_sequence_element(&self) -> Option<&ValueType> {
match self {
Self::Sequence(element) => Some(element),
_ => None,
}
}
}
#[derive(Clone, Debug, Default)]
pub struct NodeIo {
pub type_info: Option<TypeInfo>,
pub shape_data: Option<ShapeData>,
pub value_type: Option<ValueType>,
}
impl NodeIo {
pub fn typed(type_info: TypeInfo) -> Self {
Self {
type_info: Some(type_info),
shape_data: None,
value_type: None,
}
}
pub fn container(value_type: ValueType) -> Self {
Self {
type_info: None,
shape_data: None,
value_type: Some(value_type),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum MergePolicy {
#[default]
Permissive,
Strict,
}
#[derive(Debug)]
pub struct SymbolInterner {
next: u32,
initial_floor: u32,
cache: HashMap<DimExpr, SymbolId>,
fresh: Vec<SymbolId>,
unifications: Vec<(SymbolId, SymbolId)>,
derivations: Vec<(SymbolId, SymbolId)>,
opaque: Vec<SymbolId>,
}
impl SymbolInterner {
pub fn new(next: u32) -> Self {
Self {
next,
initial_floor: next,
cache: HashMap::new(),
fresh: Vec::new(),
unifications: Vec::new(),
derivations: Vec::new(),
opaque: Vec::new(),
}
}
pub fn initial_floor(&self) -> u32 {
self.initial_floor
}
fn record_unification(&mut self, a: SymbolId, b: SymbolId) {
self.unifications.push((a, b));
}
fn record_derivation(&mut self, derived: SymbolId, source: SymbolId) {
self.derivations.push((derived, source));
}
fn record_opaque(&mut self, sym: SymbolId) {
self.opaque.push(sym);
}
pub fn unifications(&self) -> &[(SymbolId, SymbolId)] {
&self.unifications
}
pub fn derivations(&self) -> &[(SymbolId, SymbolId)] {
&self.derivations
}
pub fn opaque(&self) -> &[SymbolId] {
&self.opaque
}
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() {
let id = self.fresh_symbol();
self.record_opaque(id);
return Dim::Symbolic(id);
}
if let Some(n) = expr.as_const() {
if n >= 0 {
return Dim::Static(n as usize);
}
let id = self.fresh_symbol();
self.record_opaque(id);
return Dim::Symbolic(id);
}
if let Some(s) = expr.as_symbol() {
return Dim::Symbolic(s);
}
if let Some(&id) = self.cache.get(expr) {
self.record_expr_derivation(id, expr);
return Dim::Symbolic(id);
}
let id = self.fresh_symbol();
self.cache.insert(expr.clone(), id);
self.record_expr_derivation(id, expr);
Dim::Symbolic(id)
}
fn record_expr_derivation(&mut self, derived: SymbolId, expr: &DimExpr) {
let mut seen = std::collections::HashSet::new();
for source in expr.symbol_ids() {
if source != derived && seen.insert(source) {
self.record_derivation(derived, source);
}
}
}
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_value_type(&self, i: usize) -> Option<&ValueType> {
self.inputs.get(i)?.value_type.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_value_type(&mut self, i: usize, value_type: ValueType) {
if let Some(slot) = self.outputs.get_mut(i) {
slot.value_type = Some(value_type);
}
}
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 {
let domain = normalize_domain(domain);
if domain == self.node.domain
&& let Some(version) = self.node.local_opset()
{
return version;
}
self.opset_imports.get(domain).copied().unwrap_or(1)
}
pub fn fresh_dim(&mut self) -> DimExpr {
self.interner.fresh_dim()
}
pub(crate) fn interner_mut(&mut self) -> &mut SymbolInterner {
self.interner
}
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)) if sa.0 >= ANON_SYMBOL_FLOOR || sb.0 >= ANON_SYMBOL_FLOOR => {
self.interner.record_unification(sa, sb);
Ok(if sa.0 <= sb.0 { a.clone() } else { b.clone() })
}
_ => Ok(self.fresh_broadcast_dim(a, b)),
},
}
}
fn fresh_broadcast_dim(&mut self, a: &DimExpr, b: &DimExpr) -> DimExpr {
let fresh = self.interner.fresh_symbol();
let mut seen = std::collections::HashSet::new();
for source in a.symbol_ids().chain(b.symbol_ids()) {
if source != fresh && seen.insert(source) {
self.interner.record_derivation(fresh, source);
}
}
DimExpr::symbol(fresh)
}
}
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)
}
pub(crate) fn merge_element_shape(
interner: &mut SymbolInterner,
a: &[DimExpr],
b: &[DimExpr],
) -> Option<TypedShape> {
if a.len() != b.len() {
return None;
}
let merged = a
.iter()
.zip(b.iter())
.map(|(da, db)| {
if da == db {
da.clone()
} else {
interner.fresh_dim()
}
})
.collect();
Some(merged)
}
pub(crate) fn unify_tensor_type(
interner: &mut SymbolInterner,
op: &str,
acc: TensorType,
other: TensorType,
) -> Result<TensorType, ShapeInferError> {
if acc.dtype != other.dtype {
return Err(ShapeInferError::Invalid {
op: op.into(),
detail: format!(
"sequence elements must share a dtype, found {:?} and {:?}",
acc.dtype, other.dtype
),
});
}
let shape = match (acc.shape, other.shape) {
(Some(acc_shape), Some(other_shape)) => {
merge_element_shape(interner, &acc_shape, &other_shape)
}
_ => None,
};
Ok(TensorType {
dtype: acc.dtype,
shape,
})
}
pub(crate) fn unify_value_type(
interner: &mut SymbolInterner,
op: &str,
a: &ValueType,
b: &ValueType,
) -> Result<ValueType, ShapeInferError> {
match (a, b) {
(ValueType::Tensor(a), ValueType::Tensor(b)) => Ok(ValueType::Tensor(unify_tensor_type(
interner,
op,
a.clone(),
b.clone(),
)?)),
(ValueType::Sequence(a), ValueType::Sequence(b)) => {
Ok(ValueType::sequence(unify_value_type(interner, op, a, b)?))
}
(ValueType::Optional(a), ValueType::Optional(b)) => Ok(ValueType::Optional(Box::new(
unify_value_type(interner, op, a, b)?,
))),
(ValueType::Map(ak, av), ValueType::Map(bk, bv)) if ak == bk => Ok(ValueType::Map(
*ak,
Box::new(unify_value_type(interner, op, av, bv)?),
)),
_ => Err(ShapeInferError::Invalid {
op: op.into(),
detail: format!("container types disagree: {a:?} vs {b:?}"),
}),
}
}
#[cfg(test)]
mod opset_resolution_tests {
use super::*;
use onnx_runtime_ir::NodeId;
fn context_for(version: Option<i64>, imports: &HashMap<String, u64>) -> u64 {
let mut node = Node::new(NodeId(0), "Swish", vec![], vec![]);
node.version = version;
let mut interner = SymbolInterner::new(0);
let context = InferenceContext::new(
&node,
Vec::new(),
imports,
MergePolicy::default(),
&mut interner,
);
context.opset("")
}
#[test]
fn a_node_version_overrides_the_graph_import() {
let imports = HashMap::from([(String::new(), 13)]);
assert_eq!(context_for(Some(24), &imports), 24);
}
#[test]
fn implausible_versions_defer_to_the_graph() {
let imports = HashMap::from([(String::new(), 13)]);
for version in [-1, 0, i64::MAX, i64::from(i32::MAX) + 1] {
assert_eq!(
context_for(Some(version), &imports),
13,
"version {version} is not usable and must not override the graph"
);
}
}
}