Initial commit

This commit is contained in:
redclawsystems
2026-03-04 00:08:42 +00:00
commit 4d88dc0584
4449 changed files with 1556714 additions and 0 deletions
@@ -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()
}
}