Initial commit
This commit is contained in:
@@ -0,0 +1,345 @@
|
||||
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()
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user