diff options
Diffstat (limited to 'libpsi-core/src/core/runtime.rs')
| -rw-r--r-- | libpsi-core/src/core/runtime.rs | 258 |
1 files changed, 224 insertions, 34 deletions
diff --git a/libpsi-core/src/core/runtime.rs b/libpsi-core/src/core/runtime.rs index a75435e..382fa2f 100644 --- a/libpsi-core/src/core/runtime.rs +++ b/libpsi-core/src/core/runtime.rs @@ -1,9 +1,13 @@ -use super::{GateOp, Kernel, KernelBatch, QuantumGate, QuantumRegister, QuantumState}; +use super::{ + GateOp, Kernel, KernelBatch, QuantumGate, QuantumRegister, QuantumState, + StructureAwareKernelBatch, +}; 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, PAULI_X, PAULI_Y, PAULI_Z, SDG_GATE, SWAP, SXDG_GATE, SX_GATE, S_GATE, TDG_GATE, TOFFOLI, T_GATE, }; +use crate::maths::simd::{apply_single_qubit_gate_simd, apply_single_qubit_gate_simd_parallel}; use crate::maths::vector::Vector; use crate::{complex, Complex, Matrix}; use rayon::prelude::*; @@ -11,6 +15,129 @@ use rayon::prelude::*; const PARALLEL_THRESHOLD: usize = 8; #[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub struct RuntimeConfig { + pub parallel: bool, + pub simd: bool, + pub batched: bool, + pub structure_aware: bool, + pub parallel_threshold: usize, +} + +impl RuntimeConfig { + pub fn new() -> Self { + Self { + parallel: false, + simd: false, + batched: false, + structure_aware: false, + parallel_threshold: PARALLEL_THRESHOLD, + } + } + + pub fn parallel(mut self) -> Self { + self.parallel = true; + self + } + + pub fn simd(mut self) -> Self { + self.simd = true; + self + } + + pub fn batched(mut self) -> Self { + self.batched = true; + self + } + + pub fn structure_aware(mut self) -> Self { + self.structure_aware = true; + self + } + + pub fn with_threshold(mut self, threshold: usize) -> Self { + self.parallel_threshold = threshold; + self + } + + pub fn optimal() -> Self { + Self::new().structure_aware().simd().parallel() + } + + pub fn compute(&self, num_qubits: usize, operations: &[GateOp]) -> 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 use_parallel = self.parallel && num_qubits >= self.parallel_threshold; + + if self.structure_aware { + let mut batch = Runtime::build_structure_aware_batch(num_qubits, operations); + batch.optimise(); + self.execute_kernels(&mut state, batch.kernels(), num_qubits, use_parallel); + } else if self.batched { + let mut batch = Runtime::build_kernel_batch(num_qubits, operations); + batch.optimize(); + self.execute_kernels(&mut state, batch.kernels(), num_qubits, use_parallel); + } else { + let batch = Runtime::build_kernel_batch(num_qubits, operations); + self.execute_kernels(&mut state, batch.kernels(), num_qubits, use_parallel); + } + + QuantumState::new(state) + } + + fn execute_kernels( + &self, + state: &mut Vec<Complex<f64>>, + kernels: &[Kernel], + num_qubits: usize, + use_parallel: bool, + ) { + for kernel in kernels { + if self.simd && kernel.targets.len() == 1 { + let gate = matrix_to_2x2(&kernel.matrix); + if use_parallel { + apply_single_qubit_gate_simd_parallel( + state, + &gate, + kernel.targets[0], + num_qubits, + ); + } else { + apply_single_qubit_gate_simd(state, &gate, kernel.targets[0], num_qubits); + } + } else if use_parallel { + *state = apply_gate_parallel(state, &kernel.matrix, &kernel.targets, num_qubits); + } else { + *state = apply_kernel_direct(state, kernel, num_qubits); + } + } + } +} + +impl std::fmt::Display for RuntimeConfig { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let mut features = Vec::new(); + if self.structure_aware { + features.push("structure-aware"); + } + if self.batched && !self.structure_aware { + features.push("batched"); + } + if self.simd { + features.push("SIMD"); + } + if self.parallel { + features.push("parallel"); + } + if features.is_empty() { + features.push("basic"); + } + write!(f, "Runtime[{}]", features.join("+")) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] pub enum Runtime { #[default] BasicRT, @@ -19,20 +146,43 @@ pub enum Runtime { BatchedRTMT, SimdRT, SimdRTMT, + StructureAwareRT, + StructureAwareMT, WFEvolution, WFEvolutionMT, GPUAccelerated, + Custom(RuntimeConfig), } impl Runtime { + pub fn custom() -> RuntimeConfig { + RuntimeConfig::new() + } + + pub fn optimal() -> RuntimeConfig { + RuntimeConfig::optimal() + } + + pub fn to_config(&self) -> RuntimeConfig { + match self { + Runtime::BasicRT => RuntimeConfig::new(), + Runtime::BasicRTMT => RuntimeConfig::new().parallel(), + Runtime::BatchedRT => RuntimeConfig::new().batched(), + Runtime::BatchedRTMT => RuntimeConfig::new().batched().parallel(), + Runtime::SimdRT => RuntimeConfig::new().batched().simd(), + Runtime::SimdRTMT => RuntimeConfig::new().batched().simd().parallel(), + Runtime::StructureAwareRT => RuntimeConfig::new().structure_aware().simd(), + Runtime::StructureAwareMT => RuntimeConfig::new().structure_aware().simd().parallel(), + Runtime::Custom(config) => *config, + _ => RuntimeConfig::new(), + } + } + pub fn compute(&self, num_qubits: usize, operations: &[GateOp]) -> QuantumState { 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::SimdRT => Self::compute_simd(num_qubits, operations, false), - Runtime::SimdRTMT => Self::compute_simd(num_qubits, operations, true), + Runtime::Custom(config) => config.compute(num_qubits, operations), Runtime::WFEvolution => { unimplemented!("WFEvolution (Schrödinger equation) runtime not yet implemented") } @@ -44,6 +194,7 @@ impl Runtime { Runtime::GPUAccelerated => { unimplemented!("GPUAccelerated runtime not yet implemented") } + _ => self.to_config().compute(num_qubits, operations), } } @@ -97,38 +248,19 @@ impl Runtime { 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_simd(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(); + pub fn build_structure_aware_batch( + num_qubits: usize, + operations: &[GateOp], + ) -> StructureAwareKernelBatch { + let mut batch = StructureAwareKernelBatch::new(num_qubits); - if parallel && num_qubits >= PARALLEL_THRESHOLD { - batch.execute_simd_parallel(&mut state); - } else { - batch.execute_simd(&mut state); + for op in operations { + if let Some(kernel) = Self::op_to_kernel(op) { + batch.add(kernel); + } } - QuantumState::new(state) + batch } fn compute_basic(num_qubits: usize, operations: &[GateOp]) -> QuantumState { @@ -393,3 +525,61 @@ fn apply_gate_parallel( new_state } + +fn matrix_to_2x2(matrix: &Matrix<Complex<f64>>) -> [[Complex<f64>; 2]; 2] { + [ + [matrix.data[0], matrix.data[1]], + [matrix.data[2], matrix.data[3]], + ] +} + +fn apply_kernel_direct( + 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 +} |
