247 lines
6.5 KiB
Rust
247 lines
6.5 KiB
Rust
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 { .. }
|
|
));
|
|
}
|