use criterion::{BenchmarkId, Criterion, black_box, criterion_group, criterion_main}; use rtx_polygraph::{ cache::KernelCache, fusion::FusionAnalyzer, ir::{DataType, IRNode, IRNodeType, NodeId, Shape}, }; fn bench_fusion_analysis(c: &mut Criterion) { let mut group = c.benchmark_group("fusion_analysis"); for size in [64, 128, 256, 512].iter() { group.bench_with_input(BenchmarkId::new("dense_dense", size), size, |b, &size| { b.iter(|| { let mut analyzer = FusionAnalyzer::new(); for i in 0..10 { let node = IRNode::new( NodeId(i as u32), IRNodeType::MatMul { transpose_a: false, transpose_b: false, }, if i == 0 { vec![NodeId(100)] } else { vec![NodeId((i - 1) as u32)] }, vec![DataType::F32], vec![Shape::new(vec![size, size])], ); analyzer.add_node(node); } black_box(analyzer.find_fusion_opportunities()) }) }); group.bench_with_input(BenchmarkId::new("graph_dense", size), size, |b, &size| { b.iter(|| { let mut analyzer = FusionAnalyzer::new(); for i in 0..5 { let graph_node = IRNode::new( NodeId((i * 2) as u32), IRNodeType::GraphConvolution { aggregation: rtx_polygraph::ir::AggregationType::Sum, }, vec![NodeId(if i == 0 { 100 } else { (i * 2 - 1) as u32 })], vec![DataType::F32], vec![Shape::new(vec![size, size / 2])], ); let dense_node = IRNode::new( NodeId((i * 2 + 1) as u32), IRNodeType::MatMul { transpose_a: false, transpose_b: false, }, vec![NodeId((i * 2) as u32)], vec![DataType::F32], vec![Shape::new(vec![size, size])], ); analyzer.add_node(graph_node); analyzer.add_node(dense_node); } black_box(analyzer.find_fusion_opportunities()) }) }); } group.finish(); } fn bench_cache_operations(c: &mut Criterion) { let mut group = c.benchmark_group("cache_operations"); for cache_size in [100, 1000, 10000].iter() { group.bench_with_input( BenchmarkId::new("cache_lookup", cache_size), cache_size, |b, &cache_size| { let mut cache = KernelCache::with_capacity(cache_size); // Pre-populate cache for i in 0..cache_size { let node = IRNode::new( NodeId(i as u32), IRNodeType::MatMul { transpose_a: i % 2 == 0, transpose_b: i % 3 == 0, }, vec![NodeId((i + 1000) as u32)], vec![DataType::F32], vec![Shape::new(vec![64 + (i % 10), 128 + (i % 10)])], ); let key = rtx_polygraph::cache::CacheKey::from_node(&node); let kernel = rtx_polygraph::cache::CachedKernel::new( format!("kernel_{}", i), vec![i as u8; 1024], std::time::Duration::from_millis(100), ); cache.insert(key, kernel); } b.iter(|| { for i in 0..100 { let lookup_idx = i % cache_size; let node = IRNode::new( NodeId(lookup_idx as u32), IRNodeType::MatMul { transpose_a: lookup_idx % 2 == 0, transpose_b: lookup_idx % 3 == 0, }, vec![NodeId((lookup_idx + 1000) as u32)], vec![DataType::F32], vec![Shape::new(vec![ 64 + (lookup_idx % 10), 128 + (lookup_idx % 10), ])], ); let key = rtx_polygraph::cache::CacheKey::from_node(&node); black_box(cache.get(&key)); } }) }, ); } group.finish(); } fn bench_ir_node_creation(c: &mut Criterion) { let mut group = c.benchmark_group("ir_node_creation"); group.bench_function("matmul_nodes", |b| { b.iter(|| { for i in 0..1000 { let node = IRNode::new( NodeId(i), IRNodeType::MatMul { transpose_a: i % 2 == 0, transpose_b: i % 3 == 0, }, vec![NodeId(i + 1)], vec![DataType::F32], vec![Shape::new(vec![64, 128])], ); black_box(node); } }) }); group.bench_function("sparse_nodes", |b| { b.iter(|| { for i in 0..1000 { let node = IRNode::new( NodeId(i), IRNodeType::SparseMatMul { format: rtx_polygraph::ir::SparseFormat::CSR, }, vec![NodeId(i + 1)], vec![DataType::F32], vec![Shape::new(vec![256, 512])], ); black_box(node); } }) }); group.bench_function("graph_nodes", |b| { b.iter(|| { for i in 0..1000 { let node = IRNode::new( NodeId(i), IRNodeType::GraphConvolution { aggregation: rtx_polygraph::ir::AggregationType::Sum, }, vec![NodeId(i + 1)], vec![DataType::F32], vec![Shape::new(vec![100, 64])], ); black_box(node); } }) }); group.finish(); } fn bench_signature_generation(c: &mut Criterion) { let mut group = c.benchmark_group("signature_generation"); let test_nodes = vec![ IRNode::new( NodeId(1), IRNodeType::MatMul { transpose_a: false, transpose_b: true, }, vec![NodeId(0)], vec![DataType::F32], vec![Shape::new(vec![1024, 2048])], ), IRNode::new( NodeId(2), IRNodeType::SparseMatMul { format: rtx_polygraph::ir::SparseFormat::COO, }, vec![NodeId(1)], vec![DataType::F32], vec![Shape::new(vec![2048, 1024])], ), IRNode::new( NodeId(3), IRNodeType::GraphConvolution { aggregation: rtx_polygraph::ir::AggregationType::Mean, }, vec![NodeId(2)], vec![DataType::F32], vec![Shape::new(vec![500, 128])], ), ]; group.bench_function("node_signatures", |b| { b.iter(|| { for node in &test_nodes { black_box(node.signature()); } }) }); group.finish(); } criterion_group!( benches, bench_fusion_analysis, bench_cache_operations, bench_ir_node_creation, bench_signature_generation ); criterion_main!(benches);