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

247 lines
8.0 KiB
Rust

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);