use rtx_polygraph::{ error::PolygraphError, fusion::{FusionAnalyzer, FusionOpportunity, FusionType}, ir::{DataType, IRNode, IRNodeType, NodeId, Shape}, }; #[test] fn test_dense_dense_fusion_opportunity() { let mut analyzer = FusionAnalyzer::new(); // Create two consecutive MatMul operations 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])], ); analyzer.add_node(node1).unwrap(); analyzer.add_node(node2).unwrap(); let opportunities = analyzer.find_fusion_opportunities(); assert!(!opportunities.is_empty()); let fusion_op = &opportunities[0]; assert!(matches!(fusion_op.fusion_type(), FusionType::DenseDense)); assert_eq!(fusion_op.participating_nodes().len(), 2); assert!(fusion_op.estimated_speedup() > 1.0); } #[test] fn test_graph_dense_fusion_opportunity() { let mut analyzer = FusionAnalyzer::new(); // Graph convolution followed by dense matmul let graph_node = IRNode::new( NodeId(1), IRNodeType::GraphConvolution { aggregation: rtx_polygraph::ir::AggregationType::Sum, }, vec![NodeId(0)], vec![DataType::F32], vec![Shape::new(vec![100, 64])], ); let dense_node = IRNode::new( NodeId(2), IRNodeType::MatMul { transpose_a: false, transpose_b: false, }, vec![NodeId(1)], vec![DataType::F32], vec![Shape::new(vec![100, 128])], ); analyzer.add_node(graph_node).unwrap(); analyzer.add_node(dense_node).unwrap(); let opportunities = analyzer.find_fusion_opportunities(); assert!(!opportunities.is_empty()); let fusion_op = &opportunities[0]; assert!(matches!(fusion_op.fusion_type(), FusionType::GraphDense)); assert!(fusion_op.estimated_speedup() > 1.2); // Graph-dense fusion should have higher speedup } #[test] fn test_sparse_dense_fusion_opportunity() { let mut analyzer = FusionAnalyzer::new(); let sparse_node = IRNode::new( NodeId(1), IRNodeType::SparseMatMul { format: rtx_polygraph::ir::SparseFormat::CSR, }, vec![NodeId(0)], vec![DataType::F32], vec![Shape::new(vec![256, 128])], ); let dense_node = IRNode::new( NodeId(2), IRNodeType::MatMul { transpose_a: false, transpose_b: false, }, vec![NodeId(1)], vec![DataType::F32], vec![Shape::new(vec![256, 64])], ); analyzer.add_node(sparse_node).unwrap(); analyzer.add_node(dense_node).unwrap(); let opportunities = analyzer.find_fusion_opportunities(); assert!(!opportunities.is_empty()); let fusion_op = &opportunities[0]; assert!(matches!(fusion_op.fusion_type(), FusionType::SparseDense)); } #[test] fn test_no_fusion_opportunity_for_incompatible_shapes() { let mut analyzer = FusionAnalyzer::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::MatMul { transpose_a: false, transpose_b: false, }, vec![NodeId(3)], // Different input - not consecutive vec![DataType::F32], vec![Shape::new(vec![32, 64])], // Incompatible shape ); analyzer.add_node(node1).unwrap(); analyzer.add_node(node2).unwrap(); let opportunities = analyzer.find_fusion_opportunities(); assert!(opportunities.is_empty()); } #[test] fn test_fft_elementwise_fusion() { let mut analyzer = FusionAnalyzer::new(); let fft_node = IRNode::new( NodeId(1), IRNodeType::FFT { inverse: false, axes: vec![0], }, vec![NodeId(0)], vec![DataType::C64], vec![Shape::new(vec![1024])], ); let elementwise_node = IRNode::new( NodeId(2), IRNodeType::ElementwiseMul, vec![NodeId(1)], vec![DataType::C64], vec![Shape::new(vec![1024])], ); analyzer.add_node(fft_node).unwrap(); analyzer.add_node(elementwise_node).unwrap(); let opportunities = analyzer.find_fusion_opportunities(); assert!(!opportunities.is_empty()); let fusion_op = &opportunities[0]; assert!(matches!( fusion_op.fusion_type(), FusionType::FFTElementwise )); } #[test] fn test_fusion_legality_checking() { let mut analyzer = FusionAnalyzer::new(); // Control flow should prevent fusion let conditional_node = IRNode::new( NodeId(1), IRNodeType::Conditional { condition_id: NodeId(0), }, vec![NodeId(2), NodeId(3)], vec![DataType::F32], vec![Shape::new(vec![64, 64])], ); let matmul_node = IRNode::new( NodeId(4), IRNodeType::MatMul { transpose_a: false, transpose_b: false, }, vec![NodeId(1)], vec![DataType::F32], vec![Shape::new(vec![64, 128])], ); analyzer.add_node(conditional_node).unwrap(); analyzer.add_node(matmul_node).unwrap(); let opportunities = analyzer.find_fusion_opportunities(); // Should not find fusion across control flow boundary let valid_opportunities: Vec<_> = opportunities.iter().filter(|op| op.is_legal()).collect(); assert!(valid_opportunities.is_empty()); } #[test] fn test_fusion_analyzer_memory_constraint_checking() { let mut analyzer = FusionAnalyzer::with_memory_limit(1024 * 1024); // 1MB limit // Create large operation that would exceed memory let large_node = IRNode::new( NodeId(1), IRNodeType::MatMul { transpose_a: false, transpose_b: false, }, vec![NodeId(0)], vec![DataType::F32], vec![Shape::new(vec![1024, 1024])], // 4MB matrix ); let result = analyzer.add_node(large_node); assert!(result.is_err()); assert!(matches!( result.unwrap_err(), PolygraphError::InsufficientMemory { .. } )); }