346 lines
11 KiB
Rust
346 lines
11 KiB
Rust
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<IRNode>,
|
|
fusion_type: FusionType,
|
|
speedup_estimate: f32,
|
|
is_legal: bool,
|
|
}
|
|
|
|
impl FusionOpportunity {
|
|
pub fn new(nodes: Vec<IRNode>, 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<NodeId, IRNode>,
|
|
memory_limit: Option<usize>,
|
|
}
|
|
|
|
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<FusionOpportunity> {
|
|
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<FusionOpportunity> {
|
|
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<FusionOpportunity> {
|
|
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<FusionOpportunity> {
|
|
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<FusionOpportunity> {
|
|
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()
|
|
}
|
|
}
|