Files
rustytorch/crates/training/rtx-auto/benches/auto_bench.rs
T
2026-03-04 00:08:42 +00:00

312 lines
11 KiB
Rust

use criterion::{BenchmarkId, Criterion, black_box, criterion_group, criterion_main};
use rtx_auto::{
agents::{
DataEngineeringAgent, KernelSynthesizerAgent, ParallelPlannerAgent, QuantGuardianAgent,
},
initialize_autonomous_system,
proposal::{Proposal, ProposalType, ProposalValidator},
rollback::RollbackManager,
};
use rtx_graph::ComputeGraph;
use rtx_runtime::Runtime;
use rtx_tensor::{DataType, Shape, Tensor};
use std::collections::HashMap;
use tokio::runtime::Runtime as TokioRuntime;
fn bench_data_agent_proposals(c: &mut Criterion) {
let rt = TokioRuntime::new().unwrap();
c.bench_with_input(
BenchmarkId::new("data_agent_layout_proposals", "1024x1024"),
&1024,
|b, &size| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let agent = DataEngineeringAgent::new(&runtime).unwrap();
let tensor =
Tensor::zeros(&Shape::new(vec![size, size]), DataType::F32).unwrap();
let proposals = agent.generate_layout_proposals(&tensor).await.unwrap();
black_box(proposals)
})
})
},
);
let sizes = vec![256, 512, 1024, 2048];
for size in sizes {
c.bench_with_input(
BenchmarkId::new("data_agent_coalescing", size),
&size,
|b, &size| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let agent = DataEngineeringAgent::new(&runtime).unwrap();
let tensors: Vec<Tensor> = (0..4)
.map(|_| {
Tensor::zeros(&Shape::new(vec![size, size]), DataType::F32).unwrap()
})
.collect();
let proposals =
agent.generate_coalescing_proposals(&tensors).await.unwrap();
black_box(proposals)
})
})
},
);
}
}
fn bench_parallel_planner_proposals(c: &mut Criterion) {
let rt = TokioRuntime::new().unwrap();
c.bench_with_input(
BenchmarkId::new("parallel_planner_data_parallel", "8_nodes"),
&8,
|b, &node_count| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let agent = ParallelPlannerAgent::new(&runtime).unwrap();
let mut graph = ComputeGraph::new();
// Add nodes to the graph
let mut prev_node = None;
for i in 0..node_count {
let deps = if let Some(prev) = prev_node {
vec![prev]
} else {
vec![]
};
let node = graph
.add_node(&format!("node_{}", i), deps, vec![])
.unwrap();
prev_node = Some(node);
}
let proposals = agent
.generate_data_parallel_proposals(&graph)
.await
.unwrap();
black_box(proposals)
})
})
},
);
}
fn bench_quant_guardian_accuracy(c: &mut Criterion) {
let rt = TokioRuntime::new().unwrap();
let sizes = vec![64, 128, 256, 512];
for size in sizes {
c.bench_with_input(
BenchmarkId::new("quant_guardian_accuracy_monitoring", size),
&size,
|b, &size| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let agent = QuantGuardianAgent::new(&runtime).unwrap();
let reference =
Tensor::randn(&Shape::new(vec![size, size]), DataType::F32).unwrap();
let quantized = reference.clone(); // In practice would be actually quantized
let metrics = agent
.monitor_accuracy(&reference, &quantized)
.await
.unwrap();
black_box(metrics)
})
})
},
);
}
c.bench_with_input(
BenchmarkId::new("quant_guardian_proposals", "256x256"),
&256,
|b, &size| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let agent = QuantGuardianAgent::new(&runtime).unwrap();
let tensor =
Tensor::randn(&Shape::new(vec![size, size]), DataType::F32).unwrap();
let proposals = agent
.generate_quantization_proposals(&tensor, 0.95)
.await
.unwrap();
black_box(proposals)
})
})
},
);
}
fn bench_kernel_synthesizer_fusion(c: &mut Criterion) {
let rt = TokioRuntime::new().unwrap();
c.bench_with_input(
BenchmarkId::new("kernel_synthesizer_fusion", "4_kernels"),
&4,
|b, &kernel_count| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let agent = KernelSynthesizerAgent::new(&runtime).unwrap();
let kernels: Vec<_> = (0..kernel_count)
.map(|i| rtx_kernel::KernelSpec::new(&format!("kernel_{}", i)).build())
.collect();
let proposals = agent.generate_fusion_proposals(&kernels).await.unwrap();
black_box(proposals)
})
})
},
);
}
fn bench_proposal_validation(c: &mut Criterion) {
let rt = TokioRuntime::new().unwrap();
let proposal_counts = vec![5, 10, 20, 50];
for count in proposal_counts {
c.bench_with_input(
BenchmarkId::new("proposal_validation_ranking", count),
&count,
|b, &count| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let validator = ProposalValidator::new(&runtime).unwrap();
let proposals: Vec<Proposal> = (0..count)
.map(|i| {
Proposal::new(
ProposalType::DataLayout,
format!("Test proposal {}", i),
1.0 + (i as f32) * 0.1,
)
})
.collect();
let ranked = validator.rank_proposals(&proposals).await.unwrap();
black_box(ranked)
})
})
},
);
}
c.bench_function("proposal_conflict_detection", |b| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let validator = ProposalValidator::new(&runtime).unwrap();
let proposals = vec![
Proposal::new(ProposalType::DataLayout, "Layout 1".to_string(), 1.5),
Proposal::new(ProposalType::DataLayout, "Layout 2".to_string(), 1.3),
Proposal::new(ProposalType::KernelFusion, "Fusion".to_string(), 2.0),
];
let conflicts = validator.detect_conflicts(&proposals).await.unwrap();
black_box(conflicts)
})
})
});
}
fn bench_rollback_operations(c: &mut Criterion) {
let rt = TokioRuntime::new().unwrap();
c.bench_function("checkpoint_creation", |b| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let manager = RollbackManager::new(&runtime).unwrap();
let mut state = HashMap::new();
state.insert("accuracy".to_string(), vec![0.95, 0.92, 0.88]);
state.insert("loss".to_string(), vec![0.1, 0.15, 0.12]);
let checkpoint = manager
.create_checkpoint(
rtx_auto::rollback::CheckpointType::Manual,
state,
"Benchmark checkpoint".to_string(),
)
.await
.unwrap();
black_box(checkpoint)
})
})
});
c.bench_function("rollback_threshold_check", |b| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let manager = RollbackManager::new(&runtime).unwrap();
let mut current_metrics = HashMap::new();
current_metrics.insert("accuracy".to_string(), vec![0.85]);
current_metrics.insert("latency".to_string(), vec![150.0]);
let should_rollback = manager
.should_trigger_automatic_rollback(&current_metrics)
.await
.unwrap();
black_box(should_rollback)
})
})
});
}
fn bench_autonomous_system_initialization(c: &mut Criterion) {
let rt = TokioRuntime::new().unwrap();
c.bench_function("autonomous_system_init", |b| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let optimizer = initialize_autonomous_system(&runtime).await.unwrap();
black_box(optimizer)
})
})
});
c.bench_function("agent_health_check", |b| {
b.iter(|| {
rt.block_on(async {
let runtime = Runtime::new().await.unwrap();
let optimizer = initialize_autonomous_system(&runtime).await.unwrap();
let health = optimizer.check_agent_health();
black_box(health)
})
})
});
}
criterion_group!(
benches,
bench_data_agent_proposals,
bench_parallel_planner_proposals,
bench_quant_guardian_accuracy,
bench_kernel_synthesizer_fusion,
bench_proposal_validation,
bench_rollback_operations,
bench_autonomous_system_initialization
);
criterion_main!(benches);