use crate::error::{PolygraphError, Result}; use crate::ir::{DataType, IRNode, IRNodeType, NodeId}; use serde::{Deserialize, Serialize}; use std::collections::HashMap; #[derive(Debug, Clone, PartialEq, Hash, Serialize, Deserialize)] pub enum FusionType { DenseDense, GraphDense, SparseDense, FFTElementwise, ElementwiseChain, } #[derive(Debug, Clone)] pub struct FusionOpportunity { nodes: Vec, fusion_type: FusionType, speedup_estimate: f32, is_legal: bool, } impl FusionOpportunity { pub fn new(nodes: Vec, fusion_type: FusionType, speedup_estimate: f32) -> Self { let is_legal = Self::check_legality(&nodes, &fusion_type); Self { nodes, fusion_type, speedup_estimate, is_legal, } } pub fn participating_nodes(&self) -> &[IRNode] { &self.nodes } pub fn fusion_type(&self) -> &FusionType { &self.fusion_type } pub fn estimated_speedup(&self) -> f32 { self.speedup_estimate } pub fn is_legal(&self) -> bool { self.is_legal } fn check_legality(nodes: &[IRNode], fusion_type: &FusionType) -> bool { if nodes.len() < 2 { return false; } // Check for control flow boundaries for node in nodes { if matches!(node.node_type(), IRNodeType::Conditional { .. }) { return false; // Can't fuse across control flow } } // Check data dependency chain for i in 1..nodes.len() { let prev_id = nodes[i - 1].id(); if !nodes[i].inputs().contains(&prev_id) { return false; // Not a valid chain } } // Type-specific legality checks match fusion_type { FusionType::DenseDense => nodes .iter() .all(|n| matches!(n.node_type(), IRNodeType::MatMul { .. })), FusionType::GraphDense => { if nodes.len() != 2 { return false; } matches!(nodes[0].node_type(), IRNodeType::GraphConvolution { .. }) && matches!(nodes[1].node_type(), IRNodeType::MatMul { .. }) } FusionType::SparseDense => { if nodes.len() != 2 { return false; } matches!(nodes[0].node_type(), IRNodeType::SparseMatMul { .. }) && matches!(nodes[1].node_type(), IRNodeType::MatMul { .. }) } FusionType::FFTElementwise => { if nodes.len() != 2 { return false; } matches!(nodes[0].node_type(), IRNodeType::FFT { .. }) && matches!(nodes[1].node_type(), IRNodeType::ElementwiseMul) } FusionType::ElementwiseChain => nodes.iter().all(|n| { matches!( n.node_type(), IRNodeType::Add | IRNodeType::ElementwiseMul | IRNodeType::ReLU ) }), } } } pub struct FusionAnalyzer { nodes: HashMap, memory_limit: Option, } impl FusionAnalyzer { pub fn new() -> Self { Self { nodes: HashMap::new(), memory_limit: None, } } pub fn with_memory_limit(limit_bytes: usize) -> Self { Self { nodes: HashMap::new(), memory_limit: Some(limit_bytes), } } pub fn add_node(&mut self, node: IRNode) -> Result<()> { // Check memory constraints if set if let Some(limit) = self.memory_limit { let node_memory = self.estimate_node_memory(&node); if node_memory > limit { return Err(PolygraphError::InsufficientMemory { required_mb: (node_memory / (1024 * 1024)) as u64, available_mb: (limit / (1024 * 1024)) as u64, }); } } self.nodes.insert(node.id(), node); Ok(()) } fn estimate_node_memory(&self, node: &IRNode) -> usize { node.output_shapes() .iter() .zip(node.output_dtypes()) .map(|(shape, dtype)| { let element_size = match dtype { DataType::F32 | DataType::I32 => 4, DataType::F64 | DataType::I64 | DataType::C64 => 8, DataType::C128 => 16, }; shape.size() * element_size }) .sum() } pub fn find_fusion_opportunities(&self) -> Vec { let mut opportunities = Vec::new(); // Find dense-dense fusion opportunities opportunities.extend(self.find_dense_dense_fusion()); // Find graph-dense fusion opportunities opportunities.extend(self.find_graph_dense_fusion()); // Find sparse-dense fusion opportunities opportunities.extend(self.find_sparse_dense_fusion()); // Find FFT-elementwise fusion opportunities opportunities.extend(self.find_fft_elementwise_fusion()); // Filter out illegal fusions opportunities .into_iter() .filter(|op| op.is_legal()) .collect() } fn find_dense_dense_fusion(&self) -> Vec { let mut opportunities = Vec::new(); let matmul_nodes: Vec<_> = self .nodes .values() .filter(|n| matches!(n.node_type(), IRNodeType::MatMul { .. })) .collect(); // Look for consecutive MatMul operations for i in 0..matmul_nodes.len() { for j in 0..matmul_nodes.len() { if i == j { continue; } let node1 = matmul_nodes[i]; let node2 = matmul_nodes[j]; // Check if node2 uses output of node1 if node2.inputs().contains(&node1.id()) { // Check shape compatibility if self.check_shape_compatibility(node1, node2) { let speedup = self.estimate_dense_dense_speedup(node1, node2); let opportunity = FusionOpportunity::new( vec![node1.clone(), node2.clone()], FusionType::DenseDense, speedup, ); opportunities.push(opportunity); } } } } opportunities } fn find_graph_dense_fusion(&self) -> Vec { let mut opportunities = Vec::new(); let graph_nodes: Vec<_> = self .nodes .values() .filter(|n| matches!(n.node_type(), IRNodeType::GraphConvolution { .. })) .collect(); let dense_nodes: Vec<_> = self .nodes .values() .filter(|n| matches!(n.node_type(), IRNodeType::MatMul { .. })) .collect(); for graph_node in graph_nodes { for dense_node in &dense_nodes { if dense_node.inputs().contains(&graph_node.id()) && self.check_shape_compatibility(graph_node, dense_node) { let speedup = 1.5; // Graph-dense fusion has higher speedup potential let opportunity = FusionOpportunity::new( vec![graph_node.clone(), (*dense_node).clone()], FusionType::GraphDense, speedup, ); opportunities.push(opportunity); } } } opportunities } fn find_sparse_dense_fusion(&self) -> Vec { let mut opportunities = Vec::new(); let sparse_nodes: Vec<_> = self .nodes .values() .filter(|n| matches!(n.node_type(), IRNodeType::SparseMatMul { .. })) .collect(); let dense_nodes: Vec<_> = self .nodes .values() .filter(|n| matches!(n.node_type(), IRNodeType::MatMul { .. })) .collect(); for sparse_node in sparse_nodes { for dense_node in &dense_nodes { if dense_node.inputs().contains(&sparse_node.id()) && self.check_shape_compatibility(sparse_node, dense_node) { let speedup = 1.3; // Sparse-dense fusion speedup let opportunity = FusionOpportunity::new( vec![sparse_node.clone(), (*dense_node).clone()], FusionType::SparseDense, speedup, ); opportunities.push(opportunity); } } } opportunities } fn find_fft_elementwise_fusion(&self) -> Vec { let mut opportunities = Vec::new(); let fft_nodes: Vec<_> = self .nodes .values() .filter(|n| matches!(n.node_type(), IRNodeType::FFT { .. })) .collect(); let elementwise_nodes: Vec<_> = self .nodes .values() .filter(|n| matches!(n.node_type(), IRNodeType::ElementwiseMul)) .collect(); for fft_node in fft_nodes { for elem_node in &elementwise_nodes { if elem_node.inputs().contains(&fft_node.id()) && self.check_shape_compatibility(fft_node, elem_node) { let speedup = 1.8; // FFT-elementwise fusion can be very effective let opportunity = FusionOpportunity::new( vec![fft_node.clone(), (*elem_node).clone()], FusionType::FFTElementwise, speedup, ); opportunities.push(opportunity); } } } opportunities } fn check_shape_compatibility(&self, node1: &IRNode, node2: &IRNode) -> bool { if node1.output_shapes().is_empty() || node2.output_shapes().is_empty() { return false; } let output_shape1 = &node1.output_shapes()[0]; // For simplicity, assume compatible if output of first matches expected input of second // In a real implementation, this would be more sophisticated !output_shape1.dims().is_empty() } fn estimate_dense_dense_speedup(&self, node1: &IRNode, node2: &IRNode) -> f32 { // Simple heuristic: larger operations benefit more from fusion let size1 = node1.output_shapes()[0].size(); let size2 = node2.output_shapes()[0].size(); let avg_size = (size1 + size2) as f32 / 2.0; // Base speedup + size factor 1.2 + (avg_size.log2() / 20.0).min(0.8) } } impl Default for FusionAnalyzer { fn default() -> Self { Self::new() } }