use std::collections::HashMap;
use onnx_runtime_ir::{Dim, Graph, Node, SymbolConstraints, SymbolId, ValueId, WeightRef};
use crate::context::{MergePolicy, NodeIo, SymbolInterner, TypeInfo, TypedShape, merge_shapes};
use crate::dim_expr::DimExpr;
use crate::error::ShapeInferError;
use crate::registry::InferenceRegistry;
use crate::report::InferenceReport;
use crate::shape_data::ShapeData;
const ANON_SYMBOL_FLOOR: u32 = 0x8000_0000;
impl InferenceRegistry {
pub fn infer_graph(
&self,
graph: &mut Graph,
opset_imports: &HashMap<String, u64>,
policy: MergePolicy,
) -> Result<InferenceReport, ShapeInferError> {
let order = graph
.topological_order()
.map_err(|_| ShapeInferError::CycleDetected)?;
let mut interner = SymbolInterner::new(seed_next_symbol(graph));
let mut types: HashMap<ValueId, TypeInfo> = HashMap::new();
let mut shape_data: HashMap<ValueId, ShapeData> = HashMap::new();
seed_sources(graph, &mut types, &mut shape_data);
let declared_out: HashMap<ValueId, Vec<Dim>> = graph
.outputs
.iter()
.filter_map(|&vid| graph.try_value(vid).map(|v| (vid, v.shape.clone())))
.collect();
for nid in order {
let node = graph.node(nid).clone();
let inputs = gather_inputs(&node, &types, &shape_data);
let outputs = self.infer_node(&node, opset_imports, inputs, policy, &mut interner)?;
for (slot, io) in node.outputs.iter().zip(outputs) {
if let Some(ti) = io.type_info {
types.insert(*slot, ti);
}
if let Some(sd) = io.shape_data
&& sd.within_bounds()
{
shape_data.insert(*slot, sd);
}
}
}
for (&vid, declared) in &declared_out {
if let Some(ti) = types.get(&vid) {
let merged = merge_shapes(vid, &ti.shape, declared, policy)?;
let dtype = ti.dtype;
types.insert(vid, TypeInfo::new(dtype, merged));
}
}
let mut resolved = Vec::new();
for (&vid, ti) in &types {
if graph.try_value(vid).is_none() {
continue;
}
let dims: Vec<Dim> = ti.shape.iter().map(|d| interner.lower(d)).collect();
let value = graph.value_mut(vid);
value.shape = dims;
value.dtype = ti.dtype;
resolved.push(vid);
}
for &sym in interner.fresh_symbols() {
graph
.symbol_constraints
.entry(sym)
.or_insert_with(|| SymbolConstraints::new(sym, None));
}
let unresolved: Vec<ValueId> = graph
.values
.keys()
.filter(|vid| !types.contains_key(vid))
.collect();
Ok(InferenceReport {
total_values: graph.num_values(),
fresh_symbols: interner.fresh_symbols().len(),
resolved,
unresolved,
})
}
}
fn seed_sources(
graph: &Graph,
types: &mut HashMap<ValueId, TypeInfo>,
shape_data: &mut HashMap<ValueId, ShapeData>,
) {
for (vid, value) in graph.values.iter() {
if value.producer.is_some() {
continue;
}
let shape: TypedShape = value.shape.iter().map(|&d| DimExpr::from(d)).collect();
types.insert(vid, TypeInfo::new(value.dtype, shape));
}
for (&vid, weight) in &graph.initializers {
if let WeightRef::Inline(t) = weight
&& let Some(sd) = ShapeData::from_tensor(t.dtype, &t.dims, &t.data)
{
shape_data.insert(vid, sd);
}
}
}
fn gather_inputs(
node: &Node,
types: &HashMap<ValueId, TypeInfo>,
shape_data: &HashMap<ValueId, ShapeData>,
) -> Vec<NodeIo> {
node.inputs
.iter()
.map(|slot| match slot {
Some(vid) => NodeIo {
type_info: types.get(vid).cloned(),
shape_data: shape_data.get(vid).cloned(),
},
None => NodeIo::default(),
})
.collect()
}
fn seed_next_symbol(graph: &Graph) -> u32 {
let mut max = ANON_SYMBOL_FLOOR.saturating_sub(1);
for &SymbolId(id) in graph.symbol_constraints.keys() {
max = max.max(id);
}
for value in graph.values.values() {
for dim in &value.shape {
if let Dim::Symbolic(SymbolId(id)) = dim {
max = max.max(*id);
}
}
}
max.saturating_add(1)
}