Files
rustytorch/crates/specialized/rtx-polygraph/tests/fusion_analysis_tests.rs
T
2026-03-04 00:08:42 +00:00

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 { .. }
));
}