aboutsummaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorhachem <im@hachem.wtf>2025-12-10 06:59:01 +0100
committerhachem <im@hachem.wtf>2025-12-10 06:59:01 +0100
commita01eb2353803c1a14337ccf689f94b5b9c35f5a4 (patch)
treed56da2e675897a7ba1ba5735eccf15394cb1c613
parent140daa490485812ef50756796435538b9f6d9428 (diff)
[add]: Kernel+Kernel Batching
-rw-r--r--libpsi-core/src/core/kernel.rs222
-rw-r--r--libpsi-core/src/core/mod.rs2
-rw-r--r--libpsi-core/src/core/runtime.rs74
-rw-r--r--libpsi-core/src/lib.rs1
-rw-r--r--tester/src/common.rs129
-rw-r--r--tester/src/kernels.rs218
-rw-r--r--tester/src/main.rs8
7 files changed, 608 insertions, 46 deletions
diff --git a/libpsi-core/src/core/kernel.rs b/libpsi-core/src/core/kernel.rs
new file mode 100644
index 0000000..3d75425
--- /dev/null
+++ b/libpsi-core/src/core/kernel.rs
@@ -0,0 +1,222 @@
+use crate::{complex, Complex, Matrix};
+use rayon::prelude::*;
+
+#[derive(Clone)]
+pub struct Kernel {
+ pub matrix: Matrix<Complex<f64>>,
+ pub targets: Vec<usize>,
+ pub name: String,
+}
+
+impl Kernel {
+ pub fn new(name: &str, matrix: Matrix<Complex<f64>>, targets: Vec<usize>) -> Self {
+ Self {
+ matrix,
+ targets,
+ name: name.to_string(),
+ }
+ }
+
+ pub fn num_qubits(&self) -> usize {
+ self.targets.len()
+ }
+
+ pub fn can_fuse_with(&self, other: &Kernel) -> bool {
+ if self.targets.len() != 1 || other.targets.len() != 1 {
+ return false;
+ }
+ self.targets[0] == other.targets[0]
+ }
+
+ pub fn fuse(&self, other: &Kernel) -> Option<Kernel> {
+ if !self.can_fuse_with(other) {
+ return None;
+ }
+ let fused_matrix = other.matrix.dot(&self.matrix)?;
+ Some(Kernel {
+ matrix: fused_matrix,
+ targets: self.targets.clone(),
+ name: format!("{}+{}", self.name, other.name),
+ })
+ }
+}
+
+pub struct KernelBatch {
+ kernels: Vec<Kernel>,
+ num_qubits: usize,
+}
+
+impl KernelBatch {
+ pub fn new(num_qubits: usize) -> Self {
+ Self {
+ kernels: Vec::new(),
+ num_qubits,
+ }
+ }
+
+ pub fn add(&mut self, kernel: Kernel) {
+ self.kernels.push(kernel);
+ }
+
+ pub fn len(&self) -> usize {
+ self.kernels.len()
+ }
+
+ pub fn is_empty(&self) -> bool {
+ self.kernels.is_empty()
+ }
+
+ pub fn kernels(&self) -> &[Kernel] {
+ &self.kernels
+ }
+
+ pub fn optimize(&mut self) {
+ if self.kernels.len() < 2 {
+ return;
+ }
+
+ let mut optimized: Vec<Kernel> = Vec::with_capacity(self.kernels.len());
+ let mut i = 0;
+
+ while i < self.kernels.len() {
+ let current = &self.kernels[i];
+
+ if i + 1 < self.kernels.len() {
+ let next = &self.kernels[i + 1];
+ if let Some(fused) = current.fuse(next) {
+ optimized.push(fused);
+ i += 2;
+ continue;
+ }
+ }
+
+ optimized.push(current.clone());
+ i += 1;
+ }
+
+ self.kernels = optimized;
+ }
+
+ pub fn execute(&self, state: &mut Vec<Complex<f64>>) {
+ for kernel in &self.kernels {
+ *state = apply_kernel(state, kernel, self.num_qubits);
+ }
+ }
+
+ pub fn execute_parallel(&self, state: &mut Vec<Complex<f64>>) {
+ for kernel in &self.kernels {
+ *state = apply_kernel_parallel(state, kernel, self.num_qubits);
+ }
+ }
+}
+
+fn apply_kernel(state: &[Complex<f64>], kernel: &Kernel, num_qubits: usize) -> Vec<Complex<f64>> {
+ let dim = 1 << num_qubits;
+ let g = kernel.targets.len();
+ let gate_dim = 1 << g;
+
+ let target_bits: Vec<usize> = kernel.targets.iter().map(|&t| num_qubits - 1 - t).collect();
+
+ let mut non_target_mask: usize = (1 << num_qubits) - 1;
+ for &pos in &target_bits {
+ non_target_mask &= !(1 << pos);
+ }
+
+ let mut new_state = vec![complex!(0.0, 0.0); dim];
+
+ for i in 0..dim {
+ let mut target_idx = 0usize;
+ for (k, &pos) in target_bits.iter().enumerate() {
+ if (i >> pos) & 1 == 1 {
+ target_idx |= 1 << (g - 1 - k);
+ }
+ }
+
+ let mut sum = complex!(0.0, 0.0);
+
+ for j in 0..gate_dim {
+ let gate_elem = kernel.matrix.data[target_idx * gate_dim + j];
+
+ if gate_elem.real.abs() < 1e-15 && gate_elem.imaginary.abs() < 1e-15 {
+ continue;
+ }
+
+ let mut source_idx = i & non_target_mask;
+ for (k, &pos) in target_bits.iter().enumerate() {
+ if (j >> (g - 1 - k)) & 1 == 1 {
+ source_idx |= 1 << pos;
+ }
+ }
+
+ sum = sum + gate_elem * state[source_idx];
+ }
+
+ new_state[i] = sum;
+ }
+
+ new_state
+}
+
+fn apply_kernel_parallel(
+ state: &[Complex<f64>],
+ kernel: &Kernel,
+ num_qubits: usize,
+) -> Vec<Complex<f64>> {
+ let dim = 1 << num_qubits;
+ let g = kernel.targets.len();
+ let gate_dim = 1 << g;
+
+ let target_bits: Vec<usize> = kernel.targets.iter().map(|&t| num_qubits - 1 - t).collect();
+
+ let mut non_target_mask: usize = (1 << num_qubits) - 1;
+ for &pos in &target_bits {
+ non_target_mask &= !(1 << pos);
+ }
+
+ (0..dim)
+ .into_par_iter()
+ .map(|i| {
+ let mut target_idx = 0usize;
+ for (k, &pos) in target_bits.iter().enumerate() {
+ if (i >> pos) & 1 == 1 {
+ target_idx |= 1 << (g - 1 - k);
+ }
+ }
+
+ let mut sum = complex!(0.0, 0.0);
+
+ for j in 0..gate_dim {
+ let gate_elem = kernel.matrix.data[target_idx * gate_dim + j];
+
+ if gate_elem.real.abs() < 1e-15 && gate_elem.imaginary.abs() < 1e-15 {
+ continue;
+ }
+
+ let mut source_idx = i & non_target_mask;
+ for (k, &pos) in target_bits.iter().enumerate() {
+ if (j >> (g - 1 - k)) & 1 == 1 {
+ source_idx |= 1 << pos;
+ }
+ }
+
+ sum = sum + gate_elem * state[source_idx];
+ }
+
+ sum
+ })
+ .collect()
+}
+
+pub struct KernelBuilder {
+ num_qubits: usize,
+}
+
+impl KernelBuilder {
+ pub fn new(num_qubits: usize) -> Self {
+ Self { num_qubits }
+ }
+
+ pub fn num_qubits(&self) -> usize {
+ self.num_qubits
+ }
+}
diff --git a/libpsi-core/src/core/mod.rs b/libpsi-core/src/core/mod.rs
index 697b0d7..343e319 100644
--- a/libpsi-core/src/core/mod.rs
+++ b/libpsi-core/src/core/mod.rs
@@ -2,6 +2,7 @@ pub mod circuit;
pub mod classical_components;
pub mod custom_gate;
pub mod gates;
+pub mod kernel;
pub mod quantum_components;
pub mod runtime;
@@ -9,5 +10,6 @@ pub use circuit::*;
pub use classical_components::*;
pub use custom_gate::*;
pub use gates::*;
+pub use kernel::*;
pub use quantum_components::*;
pub use runtime::*;
diff --git a/libpsi-core/src/core/runtime.rs b/libpsi-core/src/core/runtime.rs
index 0a78f0f..69b9d01 100644
--- a/libpsi-core/src/core/runtime.rs
+++ b/libpsi-core/src/core/runtime.rs
@@ -1,4 +1,4 @@
-use super::{GateOp, QuantumGate, QuantumRegister, QuantumState};
+use super::{GateOp, Kernel, KernelBatch, QuantumGate, QuantumRegister, QuantumState};
use crate::gates::{
cp_matrix, crx_matrix, cry_matrix, crz_matrix, p_matrix, rx_matrix, ry_matrix,
rz_matrix, u1_matrix, u2_matrix, u3_matrix, CNOT, CZ, FREDKIN, HADAMARD,
@@ -9,7 +9,6 @@ use crate::maths::vector::Vector;
use crate::{complex, Complex, Matrix};
use rayon::prelude::*;
-/// Minimum number of qubits to enable parallelism (2^8 = 256 state vector elements)
const PARALLEL_THRESHOLD: usize = 8;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
@@ -17,6 +16,8 @@ pub enum Runtime {
#[default]
BasicRT,
BasicRTMT,
+ BatchedRT,
+ BatchedRTMT,
WFEvolution,
WFEvolutionMT,
GPUAccelerated,
@@ -27,6 +28,8 @@ impl Runtime {
match self {
Runtime::BasicRT => Self::compute_basic(num_qubits, operations),
Runtime::BasicRTMT => Self::compute_basic_mt(num_qubits, operations),
+ Runtime::BatchedRT => Self::compute_batched(num_qubits, operations, false),
+ Runtime::BatchedRTMT => Self::compute_batched(num_qubits, operations, true),
Runtime::WFEvolution => {
unimplemented!("WFEvolution (Schrödinger equation) runtime not yet implemented")
}
@@ -41,6 +44,73 @@ impl Runtime {
}
}
+ pub fn build_kernel_batch(num_qubits: usize, operations: &[GateOp]) -> KernelBatch {
+ let mut batch = KernelBatch::new(num_qubits);
+
+ for op in operations {
+ if let Some(kernel) = Self::op_to_kernel(op) {
+ batch.add(kernel);
+ }
+ }
+
+ batch
+ }
+
+ fn op_to_kernel(op: &GateOp) -> Option<Kernel> {
+ let (matrix, targets, name): (Matrix<Complex<f64>>, Vec<usize>, &str) = match op {
+ GateOp::H(t) => (HADAMARD.matrix.clone(), vec![*t], "H"),
+ GateOp::X(t) => (PAULI_X.matrix.clone(), vec![*t], "X"),
+ GateOp::Y(t) => (PAULI_Y.matrix.clone(), vec![*t], "Y"),
+ GateOp::Z(t) => (PAULI_Z.matrix.clone(), vec![*t], "Z"),
+ GateOp::S(t) => (S_GATE.matrix.clone(), vec![*t], "S"),
+ GateOp::T(t) => (T_GATE.matrix.clone(), vec![*t], "T"),
+ GateOp::Sdg(t) => (SDG_GATE.matrix.clone(), vec![*t], "Sdg"),
+ GateOp::Tdg(t) => (TDG_GATE.matrix.clone(), vec![*t], "Tdg"),
+ GateOp::Sx(t) => (SX_GATE.matrix.clone(), vec![*t], "Sx"),
+ GateOp::Sxdg(t) => (SXDG_GATE.matrix.clone(), vec![*t], "Sxdg"),
+ GateOp::Rx(t, theta) => (rx_matrix(*theta), vec![*t], "Rx"),
+ GateOp::Ry(t, theta) => (ry_matrix(*theta), vec![*t], "Ry"),
+ GateOp::Rz(t, theta) => (rz_matrix(*theta), vec![*t], "Rz"),
+ GateOp::P(t, theta) => (p_matrix(*theta), vec![*t], "P"),
+ GateOp::U1(t, lambda) => (u1_matrix(*lambda), vec![*t], "U1"),
+ GateOp::U2(t, phi, lambda) => (u2_matrix(*phi, *lambda), vec![*t], "U2"),
+ GateOp::U3(t, theta, phi, lambda) => (u3_matrix(*theta, *phi, *lambda), vec![*t], "U3"),
+ GateOp::CNOT(c, t) => (CNOT.matrix.clone(), vec![*c, *t], "CNOT"),
+ GateOp::CZ(c, t) => (CZ.matrix.clone(), vec![*c, *t], "CZ"),
+ GateOp::SWAP(a, b) => (SWAP.matrix.clone(), vec![*a, *b], "SWAP"),
+ GateOp::CRx(c, t, theta) => (crx_matrix(*theta), vec![*c, *t], "CRx"),
+ GateOp::CRy(c, t, theta) => (cry_matrix(*theta), vec![*c, *t], "CRy"),
+ GateOp::CRz(c, t, theta) => (crz_matrix(*theta), vec![*c, *t], "CRz"),
+ GateOp::CP(c, t, theta) => (cp_matrix(*theta), vec![*c, *t], "CP"),
+ GateOp::CCNOT(c1, c2, t) => (TOFFOLI.matrix.clone(), vec![*c1, *c2, *t], "CCNOT"),
+ GateOp::CSWAP(c, t1, t2) => (FREDKIN.matrix.clone(), vec![*c, *t1, *t2], "CSWAP"),
+ GateOp::Measure(_, _) => return None,
+ GateOp::Custom(gate, tgts) => {
+ let qg = gate.to_quantum_gate();
+ (qg.matrix, tgts.clone(), "Custom")
+ }
+ };
+
+ Some(Kernel::new(name, matrix, targets))
+ }
+
+ fn compute_batched(num_qubits: usize, operations: &[GateOp], parallel: bool) -> QuantumState {
+ let dim = 1 << num_qubits;
+ let mut state: Vec<Complex<f64>> = vec![complex!(0.0, 0.0); dim];
+ state[0] = complex!(1.0, 0.0);
+
+ let mut batch = Self::build_kernel_batch(num_qubits, operations);
+ batch.optimize();
+
+ if parallel && num_qubits >= PARALLEL_THRESHOLD {
+ batch.execute_parallel(&mut state);
+ } else {
+ batch.execute(&mut state);
+ }
+
+ QuantumState::new(state)
+ }
+
fn compute_basic(num_qubits: usize, operations: &[GateOp]) -> QuantumState {
let names: Vec<String> = (0..num_qubits).map(|i| format!("q{}", i)).collect();
let leaked_names: &'static [String] = Box::leak(names.into_boxed_slice());
diff --git a/libpsi-core/src/lib.rs b/libpsi-core/src/lib.rs
index 7b99170..ff3253c 100644
--- a/libpsi-core/src/lib.rs
+++ b/libpsi-core/src/lib.rs
@@ -11,5 +11,6 @@ pub use core::circuit::*;
pub use core::classical_components::*;
pub use core::custom_gate::*;
pub use core::gates;
+pub use core::kernel::*;
pub use core::quantum_components::*;
pub use core::runtime::*;
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);
}