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::{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 = |g: &mut Graph, n: &str| {
94            g.create_named_value(n, DataType::Float32, static_shape([4]))
95        };
96        let a = mk(&mut g, "a");
97        let w = mk(&mut g, "w");
98        let bias = mk(&mut g, "bias");
99        g.add_input(a);
100        g.add_input(w);
101        g.add_input(bias);
102        let m = mk(&mut g, "m");
103        g.insert_node(Node::new(NodeId(0), "MatMul", vec![Some(a), Some(w)], vec![m]));
104        let out = mk(&mut g, "out");
105        g.insert_node(Node::new(NodeId(0), "Add", vec![Some(m), Some(bias)], vec![out]));
106        g.add_output(out);
107        // Dead branch.
108        let dead = mk(&mut g, "dead");
109        g.insert_node(Node::new(NodeId(0), "Neg", vec![Some(a)], vec![dead]));
110
111        run_passes(&mut g, &default_passes(), &PassContext::new()).unwrap();
112
113        // Dead Neg removed by DCE; MatMul+Add fused by OpFusion.
114        assert_eq!(g.num_nodes(), 1);
115        assert_eq!(g.nodes.values().next().unwrap().op_type, "FusedMatMulBias");
116        assert!(g.validate().is_ok());
117    }
118
119    #[test]
120    fn run_passes_is_ok_on_empty_graph() {
121        let mut g = Graph::new();
122        run_passes(&mut g, &default_passes(), &PassContext::new()).unwrap();
123        assert_eq!(g.num_nodes(), 0);
124    }
125}