Skip to main content

remove_node_bench/
remove_node_bench.rs

1use std::hint::black_box;
2use std::time::{Duration, Instant};
3
4use onnx_runtime_ir::{DataType, Graph, Node, NodeId, static_shape};
5
6fn hub_graph(node_count: usize) -> (Graph, Vec<NodeId>) {
7    let mut graph = Graph::new();
8    let hub = graph.create_value(DataType::Float32, static_shape([1]));
9    graph.add_input(hub);
10    let mut nodes = Vec::with_capacity(node_count);
11    for _ in 0..node_count {
12        let output = graph.create_value(DataType::Float32, static_shape([1]));
13        nodes.push(graph.insert_node(Node::new(
14            NodeId(0),
15            "Identity",
16            vec![Some(hub)],
17            vec![output],
18        )));
19    }
20    (graph, nodes)
21}
22
23fn sequential_remove(node_count: usize) -> Duration {
24    let (mut graph, nodes) = hub_graph(node_count);
25    let start = Instant::now();
26    for node in nodes {
27        graph.remove_node(node);
28    }
29    let elapsed = start.elapsed();
30    black_box(graph);
31    elapsed
32}
33
34fn single_hub_disconnect(node_count: usize, repeats: usize) -> Duration {
35    let (base, nodes) = hub_graph(node_count);
36    let target = nodes[0];
37    let mut elapsed = Duration::ZERO;
38    for _ in 0..repeats {
39        let mut graph = base.clone();
40        let start = Instant::now();
41        graph.remove_node(target);
42        elapsed += start.elapsed();
43        black_box(graph);
44    }
45    elapsed / repeats as u32
46}
47
48fn median(mut samples: Vec<Duration>) -> Duration {
49    samples.sort_unstable();
50    samples[samples.len() / 2]
51}
52
53fn main() {
54    let sizes: Vec<usize> = std::env::args()
55        .skip(1)
56        .map(|arg| arg.parse().expect("node counts must be positive integers"))
57        .collect();
58    let sizes = if sizes.is_empty() {
59        vec![10_000, 20_000]
60    } else {
61        sizes
62    };
63
64    println!("nodes,sequential_remove_ms,single_hub_disconnect_us");
65    for node_count in sizes {
66        let remove = median(
67            (0..3)
68                .map(|_| sequential_remove(node_count))
69                .collect::<Vec<_>>(),
70        );
71        let disconnect = single_hub_disconnect(node_count, 100);
72        println!(
73            "{node_count},{:.3},{:.3}",
74            remove.as_secs_f64() * 1_000.0,
75            disconnect.as_secs_f64() * 1_000_000.0
76        );
77    }
78}