314 lines
7.9 KiB
Rust
314 lines
7.9 KiB
Rust
use rtx_polygraph::{
|
|
error::PolygraphError,
|
|
ir::{DataType, IRGraph, IRNode, IRNodeType, NodeId, Shape},
|
|
optimizer::{DeadCodeEliminationPass, FusionPass, MemoryOptimizationPass, OptimizationPass},
|
|
};
|
|
|
|
#[test]
|
|
fn test_fusion_pass_application() {
|
|
let mut graph = IRGraph::new();
|
|
|
|
// Create fusible sequence: MatMul -> MatMul
|
|
let node1 = IRNode::new(
|
|
NodeId(1),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(0)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 128])],
|
|
);
|
|
|
|
let node2 = IRNode::new(
|
|
NodeId(2),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(1)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 256])],
|
|
);
|
|
|
|
graph.add_node(node1);
|
|
graph.add_node(node2);
|
|
|
|
let original_node_count = graph.node_count();
|
|
assert_eq!(original_node_count, 2);
|
|
|
|
let mut pass = FusionPass::new();
|
|
let result = pass.apply(&mut graph);
|
|
|
|
assert!(result.is_ok());
|
|
let optimized_node_count = graph.node_count();
|
|
|
|
// After fusion, should have fewer nodes
|
|
assert!(optimized_node_count < original_node_count);
|
|
}
|
|
|
|
#[test]
|
|
fn test_dead_code_elimination_pass() {
|
|
let mut graph = IRGraph::new();
|
|
|
|
// Create a node with no users
|
|
let unused_node = IRNode::new(
|
|
NodeId(1),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(0)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 128])],
|
|
);
|
|
|
|
// Create a used node
|
|
let used_node = IRNode::new(
|
|
NodeId(2),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(0)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 256])],
|
|
);
|
|
|
|
// Mark used_node as output
|
|
graph.add_node(unused_node);
|
|
graph.add_node(used_node);
|
|
graph.mark_output(NodeId(2));
|
|
|
|
let original_count = graph.node_count();
|
|
assert_eq!(original_count, 2);
|
|
|
|
let mut pass = DeadCodeEliminationPass::new();
|
|
let result = pass.apply(&mut graph);
|
|
|
|
assert!(result.is_ok());
|
|
|
|
// Dead node should be eliminated
|
|
let final_count = graph.node_count();
|
|
assert_eq!(final_count, 1);
|
|
assert!(graph.contains_node(NodeId(2))); // Used node remains
|
|
assert!(!graph.contains_node(NodeId(1))); // Unused node removed
|
|
}
|
|
|
|
#[test]
|
|
fn test_memory_optimization_pass() {
|
|
let mut graph = IRGraph::new();
|
|
|
|
// Create nodes with overlapping lifetimes that can share memory
|
|
let node1 = IRNode::new(
|
|
NodeId(1),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(0)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![1024, 1024])], // Large tensor
|
|
);
|
|
|
|
let node2 = IRNode::new(
|
|
NodeId(2),
|
|
IRNodeType::Add,
|
|
vec![NodeId(1)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![1024, 1024])],
|
|
);
|
|
|
|
let node3 = IRNode::new(
|
|
NodeId(3),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(2)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![1024, 512])],
|
|
);
|
|
|
|
graph.add_node(node1);
|
|
graph.add_node(node2);
|
|
graph.add_node(node3);
|
|
graph.mark_output(NodeId(3));
|
|
|
|
let original_memory = graph.estimated_memory_usage();
|
|
|
|
let mut pass = MemoryOptimizationPass::new();
|
|
let result = pass.apply(&mut graph);
|
|
|
|
assert!(result.is_ok());
|
|
|
|
let optimized_memory = graph.estimated_memory_usage();
|
|
assert!(optimized_memory < original_memory);
|
|
}
|
|
|
|
#[test]
|
|
fn test_optimization_pass_chain() {
|
|
let mut graph = IRGraph::new();
|
|
|
|
// Create complex graph with fusion opportunities and dead code
|
|
let nodes = vec![
|
|
IRNode::new(
|
|
NodeId(1),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(0)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 128])],
|
|
),
|
|
IRNode::new(
|
|
NodeId(2),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(1)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 256])],
|
|
),
|
|
IRNode::new(
|
|
NodeId(3),
|
|
IRNodeType::Add, // Unused node
|
|
vec![NodeId(0)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 64])],
|
|
),
|
|
IRNode::new(
|
|
NodeId(4),
|
|
IRNodeType::ReLU,
|
|
vec![NodeId(2)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 256])],
|
|
),
|
|
];
|
|
|
|
for node in nodes {
|
|
graph.add_node(node);
|
|
}
|
|
graph.mark_output(NodeId(4));
|
|
|
|
let original_count = graph.node_count();
|
|
assert_eq!(original_count, 4);
|
|
|
|
// Apply optimization passes in sequence
|
|
let mut fusion_pass = FusionPass::new();
|
|
let mut dce_pass = DeadCodeEliminationPass::new();
|
|
let mut memory_pass = MemoryOptimizationPass::new();
|
|
|
|
fusion_pass.apply(&mut graph).expect("Fusion pass failed");
|
|
dce_pass.apply(&mut graph).expect("DCE pass failed");
|
|
memory_pass.apply(&mut graph).expect("Memory pass failed");
|
|
|
|
let final_count = graph.node_count();
|
|
assert!(final_count < original_count);
|
|
}
|
|
|
|
#[test]
|
|
fn test_optimization_pass_preserves_semantics() {
|
|
let mut graph = IRGraph::new();
|
|
|
|
let node1 = IRNode::new(
|
|
NodeId(1),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(0)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 128])],
|
|
);
|
|
|
|
let node2 = IRNode::new(
|
|
NodeId(2),
|
|
IRNodeType::Add,
|
|
vec![NodeId(1)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 128])],
|
|
);
|
|
|
|
graph.add_node(node1);
|
|
graph.add_node(node2);
|
|
graph.mark_output(NodeId(2));
|
|
|
|
// Capture original semantics
|
|
let original_outputs = graph.output_shapes();
|
|
let original_dtypes = graph.output_dtypes();
|
|
|
|
let mut pass = FusionPass::new();
|
|
pass.apply(&mut graph).expect("Pass failed");
|
|
|
|
// Verify semantics preserved
|
|
let final_outputs = graph.output_shapes();
|
|
let final_dtypes = graph.output_dtypes();
|
|
|
|
assert_eq!(original_outputs, final_outputs);
|
|
assert_eq!(original_dtypes, final_dtypes);
|
|
}
|
|
|
|
#[test]
|
|
fn test_optimization_pass_error_handling() {
|
|
let mut graph = IRGraph::new();
|
|
|
|
// Create invalid graph structure
|
|
let node = IRNode::new(
|
|
NodeId(1),
|
|
IRNodeType::MatMul {
|
|
transpose_a: false,
|
|
transpose_b: false,
|
|
},
|
|
vec![NodeId(999)], // Non-existent input
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 128])],
|
|
);
|
|
|
|
graph.add_node(node);
|
|
graph.mark_output(NodeId(1));
|
|
|
|
let mut pass = FusionPass::new();
|
|
let result = pass.apply(&mut graph);
|
|
|
|
assert!(result.is_err());
|
|
assert!(matches!(
|
|
result.unwrap_err(),
|
|
PolygraphError::InvalidGraphStructure { .. }
|
|
));
|
|
}
|
|
|
|
#[test]
|
|
fn test_custom_optimization_pass() {
|
|
struct NoOpPass;
|
|
|
|
impl OptimizationPass for NoOpPass {
|
|
fn name(&self) -> &str {
|
|
"noop"
|
|
}
|
|
|
|
fn apply(&mut self, graph: &mut IRGraph) -> Result<(), PolygraphError> {
|
|
// Do nothing - just validate we can implement custom passes
|
|
Ok(())
|
|
}
|
|
}
|
|
|
|
let mut graph = IRGraph::new();
|
|
let node = IRNode::new(
|
|
NodeId(1),
|
|
IRNodeType::Add,
|
|
vec![NodeId(0)],
|
|
vec![DataType::F32],
|
|
vec![Shape::new(vec![64, 128])],
|
|
);
|
|
graph.add_node(node);
|
|
|
|
let mut pass = NoOpPass;
|
|
let result = pass.apply(&mut graph);
|
|
|
|
assert!(result.is_ok());
|
|
assert_eq!(pass.name(), "noop");
|
|
}
|