use onnx_runtime_ir::{DataType, Graph, Node, NodeId, Shape, ValueId, static_shape};
use onnx_runtime_memory::{
PlanOptions, PlanStatus, SlotId, ValidateError, ViewMap, peak_activation_bytes_at_bounds,
plan_activations, plan_activations_static, validate_static,
};
use std::collections::HashMap;
const F32: DataType = DataType::Float32;
fn op(g: &mut Graph, op_type: &str, inputs: &[ValueId], out_shape: Shape) -> ValueId {
let out = g.create_value(F32, out_shape);
let ins = inputs.iter().map(|v| Some(*v)).collect();
g.insert_node(Node::new(NodeId(0), op_type, ins, vec![out]));
out
}
fn plan_of(g: &Graph, vm: &ViewMap) -> onnx_runtime_memory::ActivationPlan {
let status = plan_activations_static(g, vm, &PlanOptions::default()).unwrap();
match status {
PlanStatus::Complete(p) => {
validate_static(&p, g, vm, &PlanOptions::default()).expect("plan must validate");
p
}
PlanStatus::Deferred { unknown_sizes } => {
panic!("unexpected deferred plan: {unknown_sizes:?}")
}
}
}
#[test]
fn linear_chain_double_buffers() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([100]));
g.add_input(inp);
let a = op(&mut g, "Relu", &[inp], static_shape([100]));
let b = op(&mut g, "Relu", &[a], static_shape([100]));
let c = op(&mut g, "Relu", &[b], static_shape([100]));
let out = op(&mut g, "Relu", &[c], static_shape([100]));
g.add_output(out);
let vm = ViewMap::new();
let plan = plan_of(&g, &vm);
assert_eq!(plan.naive_bytes, 4 * 400);
assert_eq!(plan.num_slots, 2, "linear chain should reuse into 2 slots");
assert_eq!(plan.peak_bytes, 2 * 400);
assert!(
(plan.savings_ratio - 0.5).abs() < 1e-9,
"expected 50% savings"
);
assert!(!plan.assignments.contains_key(&inp));
}
#[test]
fn diamond_needs_two_concurrent() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([10]));
g.add_input(inp);
let a = op(&mut g, "Relu", &[inp], static_shape([10]));
let b = op(&mut g, "Relu", &[a], static_shape([10]));
let c = op(&mut g, "Relu", &[a], static_shape([10]));
let d = op(&mut g, "Add", &[b, c], static_shape([10]));
g.add_output(d);
let vm = ViewMap::new();
let plan = plan_of(&g, &vm);
assert_ne!(plan.assignments[&b], plan.assignments[&c]);
assert!(plan.num_slots >= 2);
assert!(plan.peak_bytes < plan.naive_bytes);
}
#[test]
fn residual_skip_pins_source_slot() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([8]));
g.add_input(inp);
let x = op(&mut g, "Relu", &[inp], static_shape([8]));
let h1 = op(&mut g, "Relu", &[x], static_shape([8]));
let h2 = op(&mut g, "Relu", &[h1], static_shape([8]));
let out = op(&mut g, "Add", &[h2, x], static_shape([8]));
g.add_output(out);
let vm = ViewMap::new();
let plan = plan_of(&g, &vm);
assert_ne!(plan.assignments[&x], plan.assignments[&h1]);
assert_ne!(plan.assignments[&x], plan.assignments[&h2]);
}
#[test]
fn disjoint_chains_share_slots() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([16]));
g.add_input(inp);
let a1 = op(&mut g, "Relu", &[inp], static_shape([16]));
let a2 = op(&mut g, "Relu", &[a1], static_shape([16]));
let b1 = op(&mut g, "Relu", &[a2], static_shape([16]));
let b2 = op(&mut g, "Relu", &[b1], static_shape([16]));
g.add_output(b2);
let vm = ViewMap::new();
let plan = plan_of(&g, &vm);
assert_eq!(
plan.assignments[&a1], plan.assignments[&b1],
"disjoint lifetimes should share a slot"
);
assert_eq!(plan.num_slots, 2);
}
#[test]
fn view_folding_extends_source_and_owns_no_slot() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([32]));
g.add_input(inp);
let src = op(&mut g, "Relu", &[inp], static_shape([32]));
let view = op(&mut g, "Slice", &[src], static_shape([16]));
let other = op(&mut g, "Relu", &[src], static_shape([32]));
let out = op(&mut g, "Add", &[other, view], static_shape([32]));
g.add_output(out);
let vm = ViewMap::from_pairs([(view, src)]);
let plan = plan_of(&g, &vm);
assert!(!plan.assignments.contains_key(&view));
assert_ne!(plan.assignments[&src], plan.assignments[&other]);
}
#[test]
fn view_of_view_folds_to_root() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([64]));
g.add_input(inp);
let src = op(&mut g, "Relu", &[inp], static_shape([64]));
let v1 = op(&mut g, "Reshape", &[src], static_shape([8, 8]));
let v2 = op(&mut g, "Transpose", &[v1], static_shape([8, 8]));
let other = op(&mut g, "Relu", &[src], static_shape([64]));
let out = op(&mut g, "Flatten", &[v2], static_shape([64]));
g.add_output(out);
let out2 = op(&mut g, "Add", &[out, other], static_shape([64]));
g.add_output(out2);
let vm = ViewMap::from_pairs([(v1, src), (v2, v1)]);
assert_eq!(vm.root(v2), src);
let plan = plan_of(&g, &vm);
assert!(!plan.assignments.contains_key(&v1));
assert!(!plan.assignments.contains_key(&v2));
assert_ne!(plan.assignments[&src], plan.assignments[&other]);
}
#[test]
fn graph_output_slot_never_reused() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([4]));
g.add_input(inp);
let a = op(&mut g, "Relu", &[inp], static_shape([4]));
let out = op(&mut g, "Relu", &[a], static_shape([4]));
g.add_output(out);
let vm = ViewMap::new();
let plan = plan_of(&g, &vm);
let out_slot = plan.assignments[&out];
let shared = plan
.assignments
.iter()
.filter(|(v, s)| **s == out_slot && **v != out)
.count();
assert_eq!(shared, 0, "output slot must be exclusive");
}
#[test]
fn symbolic_shape_defers_plan() {
let mut g = Graph::new();
let seq = g.intern_symbol("seq");
let inp = g.create_named_value("in", F32, static_shape([8]));
g.add_input(inp);
let dyn_shape: Shape = vec![seq.into(), 8usize.into()];
let a = op(&mut g, "Relu", &[inp], dyn_shape);
let out = op(&mut g, "Relu", &[a], static_shape([8]));
g.add_output(out);
let vm = ViewMap::new();
let status = plan_activations_static(&g, &vm, &PlanOptions::default()).unwrap();
match status {
PlanStatus::Deferred { unknown_sizes } => {
assert!(
unknown_sizes.contains(&a),
"symbolic value must be deferred"
);
}
PlanStatus::Complete(_) => panic!("expected a deferred plan for a symbolic shape"),
}
}
#[test]
fn custom_oracle_enables_runtime_planning() {
let mut g = Graph::new();
let seq = g.intern_symbol("seq");
let inp = g.create_named_value("in", F32, vec![seq.into()]);
g.add_input(inp);
let a = op(&mut g, "Relu", &[inp], vec![seq.into()]);
let out = op(&mut g, "Relu", &[a], vec![seq.into()]);
g.add_output(out);
let vm = ViewMap::new();
let oracle = |_v: ValueId| Some(128 * 4);
let status = plan_activations(&g, &vm, oracle, &PlanOptions::default()).unwrap();
let plan = status.unwrap_complete();
assert_eq!(plan.num_slots, 2);
validate_static(&plan, &g, &vm, &PlanOptions::default())
.expect("static validate should skip symbolic sizes");
}
#[test]
fn validate_catches_overlap_conflict() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([4]));
g.add_input(inp);
let a = op(&mut g, "Relu", &[inp], static_shape([4]));
let b = op(&mut g, "Relu", &[a], static_shape([4]));
let out = op(&mut g, "Add", &[a, b], static_shape([4]));
g.add_output(out);
let vm = ViewMap::new();
let mut plan = plan_of(&g, &vm);
let bad = SlotId(0);
plan.assignments.insert(a, bad);
plan.assignments.insert(b, bad);
let err = validate_static(&plan, &g, &vm, &PlanOptions::default()).unwrap_err();
assert!(
matches!(err, ValidateError::SlotConflict { slot, .. } if slot == bad),
"expected a SlotConflict, got {err:?}"
);
}
#[test]
fn long_chain_savings_ratio() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([1000]));
g.add_input(inp);
let mut cur = op(&mut g, "Relu", &[inp], static_shape([1000]));
let mut count = 1;
for _ in 0..9 {
cur = op(&mut g, "Relu", &[cur], static_shape([1000]));
count += 1;
}
g.add_output(cur);
let vm = ViewMap::new();
let plan = plan_of(&g, &vm);
assert_eq!(count, 10);
assert_eq!(plan.naive_bytes, 10 * 1000 * 4);
assert_eq!(plan.num_slots, 2, "chain of any length reuses into 2 slots");
assert_eq!(plan.peak_bytes, 2 * 1000 * 4);
let expected = 1.0 - (2.0 / 10.0);
assert!(
(plan.savings_ratio - expected).abs() < 1e-9,
"savings_ratio {} != {}",
plan.savings_ratio,
expected
);
println!(
"10-node chain: naive={}B peak={}B slots={} savings_ratio={:.3}",
plan.naive_bytes, plan.peak_bytes, plan.num_slots, plan.savings_ratio
);
}
#[test]
fn include_graph_inputs_option() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([50]));
g.add_input(inp);
let a = op(&mut g, "Relu", &[inp], static_shape([50]));
let out = op(&mut g, "Relu", &[a], static_shape([50]));
g.add_output(out);
let vm = ViewMap::new();
let opts = PlanOptions::default().with_graph_inputs(true);
let status = plan_activations_static(&g, &vm, &opts).unwrap();
let plan = status.unwrap_complete();
validate_static(&plan, &g, &vm, &opts).unwrap();
assert!(plan.assignments.contains_key(&inp));
assert_eq!(plan.naive_bytes, 3 * 50 * 4);
}
#[test]
fn dead_on_arrival_value_is_graceful() {
let mut g = Graph::new();
let inp = g.create_named_value("in", F32, static_shape([4]));
g.add_input(inp);
let a = op(&mut g, "Relu", &[inp], static_shape([4]));
let _dead = op(&mut g, "Relu", &[a], static_shape([4]));
let out = op(&mut g, "Relu", &[a], static_shape([4]));
g.add_output(out);
let vm = ViewMap::new();
let plan = plan_of(&g, &vm);
assert!(plan.num_slots >= 1);
}
fn dynamic_chain() -> (Graph, ViewMap, onnx_runtime_ir::SymbolId) {
let mut g = Graph::new();
let seq = g.intern_symbol("S");
let shape: Shape = vec![onnx_runtime_ir::Dim::Static(4), seq.into()];
let inp = g.create_named_value("in", F32, shape.clone());
g.add_input(inp);
let a = op(&mut g, "Relu", &[inp], shape.clone());
let b = op(&mut g, "Relu", &[a], shape.clone());
let out = op(&mut g, "Relu", &[b], shape);
g.add_output(out);
(g, ViewMap::new(), seq)
}
#[test]
fn static_planning_defers_on_a_dynamic_graph() {
let (g, vm, _seq) = dynamic_chain();
let status = plan_activations_static(&g, &vm, &PlanOptions::default()).unwrap();
assert!(
status.is_deferred(),
"a graph with a symbolic dimension cannot be sized statically"
);
}
#[test]
fn bounds_turn_a_deferred_plan_into_a_reservation() {
let (g, vm, seq) = dynamic_chain();
let bounds = HashMap::from([(seq, 128usize)]);
let peak = peak_activation_bytes_at_bounds(&g, &vm, &bounds, &PlanOptions::default())
.unwrap()
.expect("every symbol is bound, so the plan must complete");
assert_eq!(
peak,
2 * 2048,
"the reservation must be the concurrent peak"
);
}
#[test]
fn the_reservation_scales_with_the_bound() {
let (g, vm, seq) = dynamic_chain();
let small = peak_activation_bytes_at_bounds(
&g,
&vm,
&HashMap::from([(seq, 16usize)]),
&PlanOptions::default(),
)
.unwrap()
.unwrap();
let large = peak_activation_bytes_at_bounds(
&g,
&vm,
&HashMap::from([(seq, 160usize)]),
&PlanOptions::default(),
)
.unwrap()
.unwrap();
assert_eq!(
large,
small * 10,
"ten times the admitted sequence length must reserve ten times the bytes"
);
}
#[test]
fn an_unbound_symbol_defers_rather_than_reserving_nothing() {
let mut g = Graph::new();
let seq = g.intern_symbol("S");
let other = g.intern_symbol("T");
let shape: Shape = vec![seq.into(), other.into()];
let inp = g.create_named_value("in", F32, shape.clone());
g.add_input(inp);
let out = op(&mut g, "Relu", &[inp], shape);
g.add_output(out);
let vm = ViewMap::new();
let bounds = HashMap::from([(seq, 8usize)]);
let peak = peak_activation_bytes_at_bounds(&g, &vm, &bounds, &PlanOptions::default()).unwrap();
assert!(
peak.is_none(),
"a partially bound graph must defer, not reserve a number derived from guessing"
);
}