aboutsummaryrefslogtreecommitdiff
path: root/tester
diff options
context:
space:
mode:
Diffstat (limited to 'tester')
-rw-r--r--tester/src/common.rs129
-rw-r--r--tester/src/kernels.rs218
-rw-r--r--tester/src/main.rs8
3 files changed, 311 insertions, 44 deletions
diff --git a/tester/src/common.rs b/tester/src/common.rs
index 90eabd8..63a7c4c 100644
--- a/tester/src/common.rs
+++ b/tester/src/common.rs
@@ -59,7 +59,7 @@ pub fn format_duration(d: Duration) -> String {
} else if d.as_millis() > 0 {
format!("{:.3}ms", d.as_secs_f64() * 1000.0)
} else {
- format!("{:.3}us", d.as_secs_f64() * 1_000_000.0)
+ format!("{:.3}μs", d.as_secs_f64() * 1_000_000.0)
}
}
@@ -77,46 +77,95 @@ pub fn print_circuit(circuit: &QuantumCircuit) {
}
pub fn print_benchmark_table(results: &[BenchmarkResult]) {
- const C1: usize = 30;
- const C2: usize = 12;
- const C3: usize = 12;
- const C4: usize = 10;
- const C5: usize = 5;
+ if results.is_empty() {
+ return;
+ }
+
+ let headers = ["Circuit", "BasicRT", "BasicRTMT", "Speedup", "Match"];
+
+ let formatted: Vec<(String, String, String, String, String)> = results
+ .iter()
+ .map(|r| {
+ let speedup = r.basic_time.as_secs_f64() / r.mt_time.as_secs_f64();
+ (
+ r.name.clone(),
+ format_duration(r.basic_time),
+ format_duration(r.mt_time),
+ if speedup.is_finite() {
+ format!("{:.2}x", speedup)
+ } else {
+ "N/A".to_string()
+ },
+ if r.results_match { "✓" } else { "✗" }.to_string(),
+ )
+ })
+ .collect();
+
+ let c1 = formatted
+ .iter()
+ .map(|r| r.0.len())
+ .max()
+ .unwrap()
+ .max(headers[0].len());
+ let c2 = formatted
+ .iter()
+ .map(|r| r.1.len())
+ .max()
+ .unwrap()
+ .max(headers[1].len());
+ let c3 = formatted
+ .iter()
+ .map(|r| r.2.len())
+ .max()
+ .unwrap()
+ .max(headers[2].len());
+ let c4 = formatted
+ .iter()
+ .map(|r| r.3.len())
+ .max()
+ .unwrap()
+ .max(headers[3].len());
+ let c5 = formatted
+ .iter()
+ .map(|r| r.4.chars().count())
+ .max()
+ .unwrap()
+ .max(headers[4].len());
let top = format!(
"╔{}═{}═{}═{}═{}╗",
- "═".repeat(C1 + 2),
- "═".repeat(C2 + 2),
- "═".repeat(C3 + 2),
- "═".repeat(C4 + 2),
- "═".repeat(C5 + 2)
+ "═".repeat(c1 + 2),
+ "═".repeat(c2 + 2),
+ "═".repeat(c3 + 2),
+ "═".repeat(c4 + 2),
+ "═".repeat(c5 + 2)
);
- let title = format!(
+ let title_sep = format!(
"╠{}╤{}╤{}╤{}╤{}╣",
- "═".repeat(C1 + 2),
- "═".repeat(C2 + 2),
- "═".repeat(C3 + 2),
- "═".repeat(C4 + 2),
- "═".repeat(C5 + 2)
+ "═".repeat(c1 + 2),
+ "═".repeat(c2 + 2),
+ "═".repeat(c3 + 2),
+ "═".repeat(c4 + 2),
+ "═".repeat(c5 + 2)
);
- let header = format!(
+ let header_sep = format!(
"╠{}╪{}╪{}╪{}╪{}╣",
- "═".repeat(C1 + 2),
- "═".repeat(C2 + 2),
- "═".repeat(C3 + 2),
- "═".repeat(C4 + 2),
- "═".repeat(C5 + 2)
+ "═".repeat(c1 + 2),
+ "═".repeat(c2 + 2),
+ "═".repeat(c3 + 2),
+ "═".repeat(c4 + 2),
+ "═".repeat(c5 + 2)
);
let bottom = format!(
"╚{}╧{}╧{}╧{}╧{}╝",
- "═".repeat(C1 + 2),
- "═".repeat(C2 + 2),
- "═".repeat(C3 + 2),
- "═".repeat(C4 + 2),
- "═".repeat(C5 + 2)
+ "═".repeat(c1 + 2),
+ "═".repeat(c2 + 2),
+ "═".repeat(c3 + 2),
+ "═".repeat(c4 + 2),
+ "═".repeat(c5 + 2)
);
- let total_width = C1 + C2 + C3 + C4 + C5 + 14;
+ let total_width = c1 + c2 + c3 + c4 + c5 + 14;
println!("\n{}", top);
println!(
@@ -124,25 +173,17 @@ pub fn print_benchmark_table(results: &[BenchmarkResult]) {
"RUNTIME BENCHMARK RESULTS",
width = total_width
);
- println!("{}", title);
+ println!("{}", title_sep);
println!(
- "║ {:<C1$} │ {:^C2$} │ {:^C3$} │ {:^C4$} │ {:^C5$} ║",
- "Circuit", "BasicRT", "BasicRTMT", "Speedup", "Match",
+ "║ {:<c1$} │ {:^c2$} │ {:^c3$} │ {:^c4$} │ {:^c5$} ║",
+ headers[0], headers[1], headers[2], headers[3], headers[4],
);
- println!("{}", header);
-
- for r in results {
- let speedup = r.basic_time.as_secs_f64() / r.mt_time.as_secs_f64();
- let speedup_str = format!("{:.2}x", speedup);
- let match_str = if r.results_match { "✓" } else { "✗" };
+ println!("{}", header_sep);
+ for (name, basic, mt, speedup, matched) in &formatted {
println!(
- "║ {:<C1$} │ {:>C2$} │ {:>C3$} │ {:>C4$} │ {:^C5$} ║",
- r.name,
- format_duration(r.basic_time),
- format_duration(r.mt_time),
- speedup_str,
- match_str,
+ "║ {:<c1$} │ {:>c2$} │ {:>c3$} │ {:>c4$} │ {:^c5$} ║",
+ name, basic, mt, speedup, matched,
);
}
diff --git a/tester/src/kernels.rs b/tester/src/kernels.rs
new file mode 100644
index 0000000..fabe394
--- /dev/null
+++ b/tester/src/kernels.rs
@@ -0,0 +1,218 @@
+use crate::common::{print_section, states_equal, BenchmarkResult};
+use libpsi_core::{QuantumCircuit, Runtime};
+use std::f64::consts::PI;
+use std::time::Instant;
+
+pub fn run_all(results: &mut Vec<BenchmarkResult>) {
+ println!("═══════════════════════════════════════════════════════════════");
+ println!(" KERNEL BATCHING TESTS");
+ println!("═══════════════════════════════════════════════════════════════\n");
+
+ test_kernel_fusion(results);
+ test_batched_vs_basic(results);
+ test_batched_large_circuits(results);
+}
+
+pub fn test_kernel_fusion(results: &mut Vec<BenchmarkResult>) {
+ print_section("Kernel Fusion Test");
+
+ let builder = || {
+ let mut circuit = QuantumCircuit::new(2);
+ circuit.h(0).t(0).s(0).x(0).h(1).z(1);
+ circuit
+ };
+
+ let circuit = builder();
+ let batch = Runtime::build_kernel_batch(2, circuit.operations());
+ let original_count = batch.len();
+ println!("Original kernels: {}", original_count);
+ for (i, k) in batch.kernels().iter().enumerate() {
+ println!(" {}: {} on {:?}", i, k.name, k.targets);
+ }
+
+ let mut optimized_batch = Runtime::build_kernel_batch(2, circuit.operations());
+ optimized_batch.optimize();
+ let optimized_count = optimized_batch.len();
+ println!("\nOptimized kernels: {}", optimized_count);
+ for (i, k) in optimized_batch.kernels().iter().enumerate() {
+ println!(" {}: {} on {:?}", i, k.name, k.targets);
+ }
+
+ let reduction = ((original_count - optimized_count) as f64 / original_count as f64) * 100.0;
+ println!(
+ "\nKernel reduction: {} → {} ({:.0}% fewer)",
+ original_count, optimized_count, reduction
+ );
+
+ let mut basic = builder();
+ let start = Instant::now();
+ basic.compute_with(Runtime::BasicRT);
+ let basic_time = start.elapsed();
+
+ let mut batched = builder();
+ let start = Instant::now();
+ batched.compute_with(Runtime::BatchedRT);
+ let batched_time = start.elapsed();
+
+ let match_result = states_equal(basic.state(), batched.state());
+ println!("Results match: {}\n", if match_result { "✓" } else { "✗" });
+
+ results.push(BenchmarkResult {
+ name: format!("Fusion ({}→{} kernels)", original_count, optimized_count),
+ basic_time,
+ mt_time: batched_time,
+ results_match: match_result,
+ });
+
+ let fusion_heavy = || {
+ let mut circuit = QuantumCircuit::new(1);
+ circuit.h(0).t(0).s(0).x(0).y(0).z(0).h(0).t(0);
+ circuit
+ };
+
+ let circuit2 = fusion_heavy();
+ let batch2 = Runtime::build_kernel_batch(1, circuit2.operations());
+ let orig2 = batch2.len();
+ let mut opt_batch2 = Runtime::build_kernel_batch(1, circuit2.operations());
+ opt_batch2.optimize();
+ let opt2 = opt_batch2.len();
+
+ let mut basic2 = fusion_heavy();
+ let start = Instant::now();
+ basic2.compute_with(Runtime::BasicRT);
+ let basic_time2 = start.elapsed();
+
+ let mut batched2 = fusion_heavy();
+ let start = Instant::now();
+ batched2.compute_with(Runtime::BatchedRT);
+ let batched_time2 = start.elapsed();
+
+ let match2 = states_equal(basic2.state(), batched2.state());
+
+ results.push(BenchmarkResult {
+ name: format!("Heavy fusion ({}→{} kernels)", orig2, opt2),
+ basic_time: basic_time2,
+ mt_time: batched_time2,
+ results_match: match2,
+ });
+}
+
+pub fn test_batched_vs_basic(results: &mut Vec<BenchmarkResult>) {
+ print_section("Batched vs Basic Runtime Comparison");
+
+ let test_cases: Vec<(&str, Box<dyn Fn() -> QuantumCircuit>)> = vec![
+ (
+ "Bell State",
+ Box::new(|| {
+ let mut c = QuantumCircuit::new(2);
+ c.h(0).cnot(0, 1);
+ c
+ }),
+ ),
+ (
+ "GHZ State",
+ Box::new(|| {
+ let mut c = QuantumCircuit::new(3);
+ c.h(0).cnot(0, 1).cnot(0, 2);
+ c
+ }),
+ ),
+ (
+ "Rotation Chain",
+ Box::new(|| {
+ let mut c = QuantumCircuit::new(3);
+ c.rx(0, PI / 4.0)
+ .ry(0, PI / 4.0)
+ .rz(0, PI / 4.0)
+ .rx(1, PI / 3.0)
+ .ry(1, PI / 3.0);
+ c
+ }),
+ ),
+ (
+ "Mixed Gates",
+ Box::new(|| {
+ let mut c = QuantumCircuit::new(4);
+ c.h(0).h(1).h(2).h(3).cnot(0, 1).cnot(2, 3).cz(1, 2);
+ c
+ }),
+ ),
+ ];
+
+ for (name, builder) in test_cases {
+ let mut basic = builder();
+ let start = Instant::now();
+ basic.compute_with(Runtime::BasicRT);
+ let basic_time = start.elapsed();
+
+ let mut batched = builder();
+ let start = Instant::now();
+ batched.compute_with(Runtime::BatchedRT);
+ let batched_time = start.elapsed();
+
+ let match_result = states_equal(basic.state(), batched.state());
+
+ println!(
+ "{}: Basic={:.2}μs, Batched={:.2}μs, Match={}",
+ name,
+ basic_time.as_secs_f64() * 1_000_000.0,
+ batched_time.as_secs_f64() * 1_000_000.0,
+ if match_result { "✓" } else { "✗" }
+ );
+
+ results.push(BenchmarkResult {
+ name: format!("Batched: {}", name),
+ basic_time,
+ mt_time: batched_time,
+ results_match: match_result,
+ });
+ }
+ println!();
+}
+
+pub fn test_batched_large_circuits(results: &mut Vec<BenchmarkResult>) {
+ print_section("Batched Runtime on Large Circuits");
+
+ let sizes = [8, 10, 12];
+
+ for &n in &sizes {
+ let builder = || {
+ let mut circuit = QuantumCircuit::new(n);
+ for i in 0..n {
+ circuit.h(i);
+ }
+ for i in 0..(n - 1) {
+ circuit.cnot(i, i + 1);
+ }
+ circuit
+ };
+
+ let mut basic_mt = builder();
+ let start = Instant::now();
+ basic_mt.compute_with(Runtime::BasicRTMT);
+ let basic_mt_time = start.elapsed();
+
+ let mut batched_mt = builder();
+ let start = Instant::now();
+ batched_mt.compute_with(Runtime::BatchedRTMT);
+ let batched_mt_time = start.elapsed();
+
+ let match_result = states_equal(basic_mt.state(), batched_mt.state());
+
+ println!(
+ "{}-qubit: BasicRTMT={:.3}ms, BatchedRTMT={:.3}ms, Match={}",
+ n,
+ basic_mt_time.as_secs_f64() * 1000.0,
+ batched_mt_time.as_secs_f64() * 1000.0,
+ if match_result { "✓" } else { "✗" }
+ );
+
+ results.push(BenchmarkResult {
+ name: format!("{}-qubit batched", n),
+ basic_time: basic_mt_time,
+ mt_time: batched_mt_time,
+ results_match: match_result,
+ });
+ }
+ println!();
+}
diff --git a/tester/src/main.rs b/tester/src/main.rs
index ac0a61f..1fb4d32 100644
--- a/tester/src/main.rs
+++ b/tester/src/main.rs
@@ -2,6 +2,7 @@ mod benchmarks;
mod clifford;
mod common;
mod custom_gates;
+mod kernels;
mod non_clifford;
use common::{print_benchmark_table, print_summary, BenchmarkResult};
@@ -21,6 +22,7 @@ fn print_usage() {
println!(" clifford Run Clifford gate tests only");
println!(" non-clifford Run non-Clifford gate tests only");
println!(" custom Run custom gate tests only");
+ println!(" kernels Run kernel batching tests only");
println!(" bench Run benchmark tests only");
println!(" help Show this help message");
println!();
@@ -28,6 +30,7 @@ fn print_usage() {
println!(" tester # Run all tests");
println!(" tester clifford # Run only Clifford gate tests");
println!(" tester non-clifford # Run only rotation/parametric gate tests");
+ println!(" tester kernels # Run only kernel batching tests");
println!(" tester custom bench # Run custom gates and benchmarks");
}
@@ -50,6 +53,7 @@ fn main() {
let run_clifford = run_all || args.iter().any(|a| a == "clifford");
let run_non_clifford = run_all || args.iter().any(|a| a == "non-clifford");
let run_custom = run_all || args.iter().any(|a| a == "custom");
+ let run_kernels = run_all || args.iter().any(|a| a == "kernels");
let run_bench = run_all || args.iter().any(|a| a == "bench");
if run_clifford {
@@ -64,6 +68,10 @@ fn main() {
custom_gates::run_all(&mut results);
}
+ if run_kernels {
+ kernels::run_all(&mut results);
+ }
+
if run_bench {
benchmarks::run_all(&mut results);
}