From 24d639224ca11112025289065d7538851606b56e Mon Sep 17 00:00:00 2001 From: hachem Date: Mon, 24 Aug 2026 14:48:36 +0200 Subject: [chore]: unwrap project --- examples/tester/common.rs | 215 ++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 215 insertions(+) create mode 100644 examples/tester/common.rs (limited to 'examples/tester/common.rs') diff --git a/examples/tester/common.rs b/examples/tester/common.rs new file mode 100644 index 0000000..9e576da --- /dev/null +++ b/examples/tester/common.rs @@ -0,0 +1,215 @@ +use psi::{QuantumCircuit, QuantumState, Runtime, Vector}; +use psi::{HorizontalRenderer, VerticalRenderer}; +use std::time::{Duration, Instant}; + +/// A named list of circuit builders used by the benchmark/test suites. +pub type CircuitCases = Vec<(&'static str, Box QuantumCircuit>)>; + +pub struct BenchmarkResult { + pub name: String, + pub basic_time: Duration, + pub mt_time: Duration, + pub results_match: bool, +} + +pub fn benchmark_circuit(name: &str, circuit_builder: F) -> BenchmarkResult +where + F: Fn() -> QuantumCircuit, +{ + let mut circuit_st = circuit_builder(); + let mut circuit_mt = circuit_builder(); + + let start_st = Instant::now(); + circuit_st.compute_with(Runtime::BasicRT); + let basic_time = start_st.elapsed(); + + let start_mt = Instant::now(); + circuit_mt.compute_with(Runtime::BasicRTMT); + let mt_time = start_mt.elapsed(); + + let state_st = circuit_st.state(); + let state_mt = circuit_mt.state(); + + let results_match = states_equal(state_st, state_mt); + + BenchmarkResult { + name: name.to_string(), + basic_time, + mt_time, + results_match, + } +} + +pub fn states_equal(a: &QuantumState, b: &QuantumState) -> bool { + if a.size() != b.size() { + return false; + } + for i in 0..a.size() { + let amp_a = a.get(i); + let amp_b = b.get(i); + let diff_real = (amp_a.real - amp_b.real).abs(); + let diff_imag = (amp_a.imaginary - amp_b.imaginary).abs(); + if diff_real > 1e-10 || diff_imag > 1e-10 { + return false; + } + } + true +} + +pub fn format_duration(d: Duration) -> String { + if d.as_secs() > 0 { + format!("{:.3}s", d.as_secs_f64()) + } else if d.as_millis() > 0 { + format!("{:.3}ms", d.as_secs_f64() * 1000.0) + } else { + format!("{:.3}μs", d.as_secs_f64() * 1_000_000.0) + } +} + +pub fn print_section(title: &str) { + let width = 61; + let padding = width - title.len() - 2; + println!("┌{}┐", "─".repeat(width)); + println!("│ {}{} │", title, " ".repeat(padding)); + println!("└{}┘\n", "─".repeat(width)); +} + +pub fn print_circuit(circuit: &QuantumCircuit) { + println!("Horizontal:\n{}", HorizontalRenderer::new(circuit)); + println!("Vertical:\n{}", VerticalRenderer::new(circuit)); +} + +pub fn print_benchmark_table(results: &[BenchmarkResult]) { + 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) + ); + let title_sep = format!( + "╠{}╤{}╤{}╤{}╤{}╣", + "═".repeat(c1 + 2), + "═".repeat(c2 + 2), + "═".repeat(c3 + 2), + "═".repeat(c4 + 2), + "═".repeat(c5 + 2) + ); + let header_sep = format!( + "╠{}╪{}╪{}╪{}╪{}╣", + "═".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) + ); + + let total_width = c1 + c2 + c3 + c4 + c5 + 14; + + println!("\n{}", top); + println!( + "║{:^width$}║", + "RUNTIME BENCHMARK RESULTS", + width = total_width + ); + println!("{}", title_sep); + println!( + "║ {:c2$} │ {:>c3$} │ {:>c4$} │ {:^c5$} ║", + name, basic, mt, speedup, matched, + ); + } + + println!("{}", bottom); +} + +pub fn print_summary(results: &[BenchmarkResult]) { + let all_match = results.iter().all(|r| r.results_match); + println!("\n"); + if all_match { + println!("✓ All circuits produced identical results with both runtimes!"); + } else { + println!("✗ WARNING: Some circuits produced different results!"); + } + + let total_basic: Duration = results.iter().map(|r| r.basic_time).sum(); + let total_mt: Duration = results.iter().map(|r| r.mt_time).sum(); + let overall_speedup = total_basic.as_secs_f64() / total_mt.as_secs_f64(); + + println!( + "\nTotal time - BasicRT: {} | BasicRTMT: {} | Overall speedup: {:.2}x", + format_duration(total_basic), + format_duration(total_mt), + overall_speedup + ); +} -- cgit v1.3