#![forbid(unsafe_code)]
mod constant_folding;
mod dead_node;
mod error;
mod fusion;
mod pass;
pub use constant_folding::ConstantFolding;
pub use dead_node::DeadNodeElimination;
pub use error::{OptimizerError, Result};
pub use fusion::{CONTRIB_DOMAIN, FusionPattern, OpFusion, PatternMatch, default_fusion_patterns};
pub use pass::{InitializerResolver, OptimizationPass, PassContext, run_passes};
pub fn default_passes() -> Vec<Box<dyn OptimizationPass>> {
vec![
Box::new(ConstantFolding),
Box::new(DeadNodeElimination),
Box::new(OpFusion::new()),
]
}
#[cfg(test)]
mod tests {
use super::*;
use onnx_runtime_ir::{DataType, Graph, Node, NodeId, static_shape};
#[test]
fn default_passes_lists_three() {
let passes = default_passes();
assert_eq!(passes.len(), 3);
assert_eq!(passes[0].name(), "ConstantFolding");
assert_eq!(passes[1].name(), "DeadNodeElimination");
assert_eq!(passes[2].name(), "OpFusion");
}
#[test]
fn run_passes_pipeline_on_matmul_add_with_dead_branch() {
let mut g = Graph::new();
g.opset_imports.insert(String::new(), 17);
let mk =
|g: &mut Graph, n: &str| g.create_named_value(n, DataType::Float32, static_shape([4]));
let a = mk(&mut g, "a");
let w = mk(&mut g, "w");
let bias = mk(&mut g, "bias");
g.add_input(a);
g.add_input(w);
g.add_input(bias);
let m = mk(&mut g, "m");
g.insert_node(Node::new(
NodeId(0),
"MatMul",
vec![Some(a), Some(w)],
vec![m],
));
let out = mk(&mut g, "out");
g.insert_node(Node::new(
NodeId(0),
"Add",
vec![Some(m), Some(bias)],
vec![out],
));
g.add_output(out);
let dead = mk(&mut g, "dead");
g.insert_node(Node::new(NodeId(0), "Neg", vec![Some(a)], vec![dead]));
run_passes(&mut g, &default_passes(), &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 1);
assert_eq!(g.nodes.values().next().unwrap().op_type, "FusedMatMulBias");
assert!(g.validate().is_ok());
}
#[test]
fn run_passes_is_ok_on_empty_graph() {
let mut g = Graph::new();
run_passes(&mut g, &default_passes(), &PassContext::new()).unwrap();
assert_eq!(g.num_nodes(), 0);
}
}