Skip to main content

onnx_runtime_optimizer/
lib.rs

1//! # `onnx-runtime-optimizer`
2//!
3//! Device-independent graph→graph optimization passes for the ORT 2.0 runtime
4//! (see `docs/ORT2.md` §18 "Optimization Passes"). This is the first Phase-2
5//! crate: pure, safe Rust graph rewriting over [`onnx_runtime_ir`] — **no**
6//! CUDA, no ORT C library, no FFI.
7//!
8//! ## What lives here
9//!
10//! | Concept | Type |
11//! |---------|------|
12//! | Pass contract | [`OptimizationPass`], [`PassContext`], [`run_passes`] |
13//! | Dead-code removal | [`DeadNodeElimination`] |
14//! | Bounded constant folding | [`ConstantFolding`] |
15//! | Operator fusion | [`OpFusion`], [`FusionPattern`], [`PatternMatch`] |
16//! | Errors | [`OptimizerError`], [`Result`] |
17//!
18//! ## Pipeline
19//!
20//! [`default_passes`] returns only the device-independent passes implemented
21//! here, in pipeline order: `ConstantFolding → DeadNodeElimination → OpFusion`.
22//!
23//! ### Deferred (Phase 2b / Phase 3)
24//!
25//! The full pipeline in `docs/ORT2.md` §18.1 also lists passes that depend on
26//! crates or analyses not yet built. They are intentionally **not** implemented
27//! here and are listed in [`default_passes`]'s source in their eventual
28//! pipeline position: `ShapeInference` (the loader owns inference for now),
29//! `AttentionFusionPass`, `LayoutPropagation`, `PlacementOptimizer`,
30//! `TransferInsertion`, `InPlaceDetection`, `MemoryPlanning`,
31//! `CudaGraphRegionDetection`, and `OverlapScheduling`.
32
33#![forbid(unsafe_code)]
34
35mod constant_folding;
36mod dead_node;
37mod error;
38mod fusion;
39mod pass;
40
41pub use constant_folding::ConstantFolding;
42pub use dead_node::DeadNodeElimination;
43pub use error::{OptimizerError, Result};
44pub use fusion::{CONTRIB_DOMAIN, FusionPattern, OpFusion, PatternMatch, default_fusion_patterns};
45pub use pass::{InitializerResolver, OptimizationPass, PassContext, run_passes};
46
47/// The device-independent Phase-1 pass pipeline, in run order.
48///
49/// ```text
50/// ConstantFolding  →  DeadNodeElimination  →  OpFusion
51/// ```
52///
53/// Constant folding runs first so it can materialize shape-computation
54/// constants, then dead-node elimination prunes any node left unreachable, and
55/// finally op fusion collapses recognized op sequences.
56///
57/// **Deferred passes** (each in its eventual pipeline slot; see the crate-level
58/// docs for why): after `ConstantFolding` would come `ShapeInference`; after
59/// `OpFusion` would come `AttentionFusionPass`, then `LayoutPropagation`,
60/// `PlacementOptimizer`, `TransferInsertion`, `InPlaceDetection`,
61/// `MemoryPlanning`, `CudaGraphRegionDetection`, and `OverlapScheduling`.
62pub fn default_passes() -> Vec<Box<dyn OptimizationPass>> {
63    vec![
64        Box::new(ConstantFolding),
65        // ShapeInference — deferred (Phase 2b): the loader owns inference.
66        Box::new(DeadNodeElimination),
67        Box::new(OpFusion::new()),
68        // AttentionFusionPass, LayoutPropagation, PlacementOptimizer,
69        // TransferInsertion, InPlaceDetection, MemoryPlanning,
70        // CudaGraphRegionDetection, OverlapScheduling — deferred (Phase 2b/3).
71    ]
72}
73
74#[cfg(test)]
75mod tests {
76    use super::*;
77    use onnx_runtime_ir::{DataType, Graph, Node, NodeId, static_shape};
78
79    #[test]
80    fn default_passes_lists_three() {
81        let passes = default_passes();
82        assert_eq!(passes.len(), 3);
83        assert_eq!(passes[0].name(), "ConstantFolding");
84        assert_eq!(passes[1].name(), "DeadNodeElimination");
85        assert_eq!(passes[2].name(), "OpFusion");
86    }
87
88    #[test]
89    fn run_passes_pipeline_on_matmul_add_with_dead_branch() {
90        // MatMul+Add feeding an output, plus a dead Neg branch off `a`.
91        let mut g = Graph::new();
92        g.opset_imports.insert(String::new(), 17);
93        let mk =
94            |g: &mut Graph, n: &str| g.create_named_value(n, DataType::Float32, static_shape([4]));
95        let a = mk(&mut g, "a");
96        let w = mk(&mut g, "w");
97        let bias = mk(&mut g, "bias");
98        g.add_input(a);
99        g.add_input(w);
100        g.add_input(bias);
101        let m = mk(&mut g, "m");
102        g.insert_node(Node::new(
103            NodeId(0),
104            "MatMul",
105            vec![Some(a), Some(w)],
106            vec![m],
107        ));
108        let out = mk(&mut g, "out");
109        g.insert_node(Node::new(
110            NodeId(0),
111            "Add",
112            vec![Some(m), Some(bias)],
113            vec![out],
114        ));
115        g.add_output(out);
116        // Dead branch.
117        let dead = mk(&mut g, "dead");
118        g.insert_node(Node::new(NodeId(0), "Neg", vec![Some(a)], vec![dead]));
119
120        run_passes(&mut g, &default_passes(), &PassContext::new()).unwrap();
121
122        // Dead Neg removed by DCE; MatMul+Add fused by OpFusion.
123        assert_eq!(g.num_nodes(), 1);
124        assert_eq!(g.nodes.values().next().unwrap().op_type, "FusedMatMulBias");
125        assert!(g.validate().is_ok());
126    }
127
128    #[test]
129    fn run_passes_is_ok_on_empty_graph() {
130        let mut g = Graph::new();
131        run_passes(&mut g, &default_passes(), &PassContext::new()).unwrap();
132        assert_eq!(g.num_nodes(), 0);
133    }
134}