use std::collections::HashMap;
use onnx_runtime_ir::{Graph, ValueId};
use crate::error::ValidateError;
use crate::liveness::compute_liveness;
use crate::options::PlanOptions;
use crate::oracle::static_size_oracle;
use crate::plan::{ActivationPlan, SlotId};
use crate::view_map::ViewMap;
pub fn validate<F>(
plan: &ActivationPlan,
graph: &Graph,
view_map: &ViewMap,
size_oracle: F,
options: &PlanOptions,
) -> Result<(), ValidateError>
where
F: Fn(ValueId) -> Option<usize>,
{
let live = compute_liveness(graph, view_map, options).map_err(|_| ValidateError::Cycle)?;
let slot_cap: HashMap<SlotId, usize> = plan
.slots
.iter()
.map(|s| (s.id, s.capacity_bytes))
.collect();
let mut by_slot: HashMap<SlotId, Vec<ValueId>> = HashMap::new();
for &owner in live.intervals.keys() {
let Some(&sid) = plan.assignments.get(&owner) else {
return Err(ValidateError::MissingAssignment { value: owner });
};
let Some(&cap) = slot_cap.get(&sid) else {
return Err(ValidateError::UnknownSlot {
value: owner,
slot: sid,
});
};
if let Some(need) = size_oracle(owner)
&& need > cap
{
return Err(ValidateError::UndersizedSlot {
value: owner,
slot: sid,
needed: need,
capacity: cap,
});
}
by_slot.entry(sid).or_default().push(owner);
}
for (&sid, members) in &by_slot {
for i in 0..members.len() {
for j in (i + 1)..members.len() {
let a = members[i];
let b = members[j];
if live.intervals[&a].overlaps(&live.intervals[&b]) {
return Err(ValidateError::SlotConflict { a, b, slot: sid });
}
}
}
}
for vid in graph.values.keys() {
if !view_map.is_view(vid) {
continue;
}
let root = view_map.root(vid);
let Some(root_interval) = live.intervals.get(&root) else {
continue; };
let view_use = graph
.consumers(vid)
.into_iter()
.filter_map(|consumer| live.order_index.get(&consumer).copied())
.max();
let view_use = match view_use {
Some(u) => u.max(view_output_end(graph, &live, vid)),
None => view_output_end(graph, &live, vid),
};
if view_use > root_interval.use_end {
return Err(ValidateError::ViewOutlivesSource {
view: vid,
source_owner: root,
view_use,
source_end: root_interval.use_end,
});
}
}
Ok(())
}
fn view_output_end(graph: &Graph, live: &crate::liveness::Liveness, vid: ValueId) -> usize {
if graph.outputs.contains(&vid) {
live.last_index
} else {
0
}
}
pub fn validate_static(
plan: &ActivationPlan,
graph: &Graph,
view_map: &ViewMap,
options: &PlanOptions,
) -> Result<(), ValidateError> {
let oracle = static_size_oracle(graph);
validate(plan, graph, view_map, oracle, options)
}