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"); }