use std::collections::HashMap;
use scirs2_core::Complex64;
use super::variational_algorithms::{ParameterizedQuantumCircuit, QuantumGate};
use crate::{CircuitResult, DeviceError, DeviceResult};
const MAX_SIMULATED_QUBITS: usize = 26;
pub fn simulate_statevector(circuit: &ParameterizedQuantumCircuit) -> DeviceResult<Vec<Complex64>> {
let num_qubits = circuit.num_qubits();
if num_qubits > MAX_SIMULATED_QUBITS {
return Err(DeviceError::InvalidInput(format!(
"Local state-vector simulation supports at most {MAX_SIMULATED_QUBITS} qubits, \
but circuit has {num_qubits}"
)));
}
let dim = 1usize << num_qubits;
let mut state = vec![Complex64::new(0.0, 0.0); dim];
state[0] = Complex64::new(1.0, 0.0);
for gate in circuit.gates() {
apply_gate(&mut state, num_qubits, gate)?;
}
Ok(state)
}
fn apply_gate(state: &mut [Complex64], num_qubits: usize, gate: &QuantumGate) -> DeviceResult<()> {
match *gate {
QuantumGate::H(q) => {
let s = std::f64::consts::FRAC_1_SQRT_2;
apply_single_qubit(
state,
num_qubits,
q,
[
Complex64::new(s, 0.0),
Complex64::new(s, 0.0),
Complex64::new(s, 0.0),
Complex64::new(-s, 0.0),
],
)
}
QuantumGate::X(q) => apply_single_qubit(
state,
num_qubits,
q,
[
Complex64::new(0.0, 0.0),
Complex64::new(1.0, 0.0),
Complex64::new(1.0, 0.0),
Complex64::new(0.0, 0.0),
],
),
QuantumGate::Y(q) => apply_single_qubit(
state,
num_qubits,
q,
[
Complex64::new(0.0, 0.0),
Complex64::new(0.0, -1.0),
Complex64::new(0.0, 1.0),
Complex64::new(0.0, 0.0),
],
),
QuantumGate::Z(q) => apply_single_qubit(
state,
num_qubits,
q,
[
Complex64::new(1.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(-1.0, 0.0),
],
),
QuantumGate::SDagger(q) => apply_single_qubit(
state,
num_qubits,
q,
[
Complex64::new(1.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(0.0, 0.0),
Complex64::new(0.0, -1.0),
],
),
QuantumGate::RX(q, theta) => {
let c = (theta / 2.0).cos();
let s = (theta / 2.0).sin();
apply_single_qubit(
state,
num_qubits,
q,
[
Complex64::new(c, 0.0),
Complex64::new(0.0, -s),
Complex64::new(0.0, -s),
Complex64::new(c, 0.0),
],
)
}
QuantumGate::RY(q, theta) => {
let c = (theta / 2.0).cos();
let s = (theta / 2.0).sin();
apply_single_qubit(
state,
num_qubits,
q,
[
Complex64::new(c, 0.0),
Complex64::new(-s, 0.0),
Complex64::new(s, 0.0),
Complex64::new(c, 0.0),
],
)
}
QuantumGate::RZ(q, theta) => {
let phase_neg = Complex64::from_polar(1.0, -theta / 2.0);
let phase_pos = Complex64::from_polar(1.0, theta / 2.0);
apply_single_qubit(
state,
num_qubits,
q,
[
phase_neg,
Complex64::new(0.0, 0.0),
Complex64::new(0.0, 0.0),
phase_pos,
],
)
}
QuantumGate::CNOT(control, target) => {
apply_controlled_x(state, num_qubits, control, target)
}
QuantumGate::CZ(control, target) => apply_controlled_z(state, num_qubits, control, target),
}
}
fn apply_single_qubit(
state: &mut [Complex64],
num_qubits: usize,
q: usize,
matrix: [Complex64; 4],
) -> DeviceResult<()> {
if q >= num_qubits {
return Err(DeviceError::InvalidInput(format!(
"Gate targets qubit {q} but circuit only has {num_qubits} qubits"
)));
}
let bit = 1usize << q;
let dim = state.len();
for base in 0..dim {
if base & bit != 0 {
continue;
}
let i0 = base;
let i1 = base | bit;
let a0 = state[i0];
let a1 = state[i1];
state[i0] = matrix[0] * a0 + matrix[1] * a1;
state[i1] = matrix[2] * a0 + matrix[3] * a1;
}
Ok(())
}
fn apply_controlled_x(
state: &mut [Complex64],
num_qubits: usize,
control: usize,
target: usize,
) -> DeviceResult<()> {
validate_two_qubit(num_qubits, control, target)?;
let control_bit = 1usize << control;
let target_bit = 1usize << target;
let dim = state.len();
for base in 0..dim {
if (base & control_bit != 0) && (base & target_bit == 0) {
let partner = base | target_bit;
state.swap(base, partner);
}
}
Ok(())
}
fn apply_controlled_z(
state: &mut [Complex64],
num_qubits: usize,
control: usize,
target: usize,
) -> DeviceResult<()> {
validate_two_qubit(num_qubits, control, target)?;
let control_bit = 1usize << control;
let target_bit = 1usize << target;
let dim = state.len();
for (idx, amp) in state.iter_mut().enumerate().take(dim) {
if (idx & control_bit != 0) && (idx & target_bit != 0) {
*amp = -*amp;
}
}
Ok(())
}
fn validate_two_qubit(num_qubits: usize, control: usize, target: usize) -> DeviceResult<()> {
if control >= num_qubits || target >= num_qubits {
return Err(DeviceError::InvalidInput(format!(
"Two-qubit gate on ({control}, {target}) but circuit only has {num_qubits} qubits"
)));
}
if control == target {
return Err(DeviceError::InvalidInput(
"Two-qubit gate requires distinct control and target qubits".to_string(),
));
}
Ok(())
}
pub fn outcome_probabilities(state: &[Complex64]) -> Vec<f64> {
state.iter().map(|amp| amp.norm_sqr()).collect()
}
fn index_to_bitstring(index: usize, num_qubits: usize) -> String {
let mut s = String::with_capacity(num_qubits);
for q in 0..num_qubits {
if index & (1usize << q) != 0 {
s.push('1');
} else {
s.push('0');
}
}
s
}
pub fn simulate_and_sample(
circuit: &ParameterizedQuantumCircuit,
shots: usize,
) -> DeviceResult<CircuitResult> {
let num_qubits = circuit.num_qubits();
let state = simulate_statevector(circuit)?;
let probabilities = outcome_probabilities(&state);
let total: f64 = probabilities.iter().sum();
if total <= 0.0 || !total.is_finite() {
return Err(DeviceError::ExecutionFailed(
"Circuit produced a non-normalizable state (zero or non-finite total probability)"
.to_string(),
));
}
let mut cumulative = Vec::with_capacity(probabilities.len());
let mut running = 0.0;
for p in &probabilities {
running += p / total;
cumulative.push(running);
}
if let Some(last) = cumulative.last_mut() {
*last = 1.0;
}
let mut counts: HashMap<String, usize> = HashMap::new();
for _ in 0..shots {
let r = fastrand::f64();
let idx = match cumulative
.binary_search_by(|probe| probe.partial_cmp(&r).unwrap_or(std::cmp::Ordering::Less))
{
Ok(i) | Err(i) => i.min(cumulative.len().saturating_sub(1)),
};
*counts
.entry(index_to_bitstring(idx, num_qubits))
.or_insert(0) += 1;
}
let mut metadata = HashMap::new();
metadata.insert("backend".to_string(), "local_statevector".to_string());
metadata.insert("num_qubits".to_string(), num_qubits.to_string());
Ok(CircuitResult {
counts,
shots,
metadata,
})
}
pub fn expected_hamming_weight(circuit: &ParameterizedQuantumCircuit) -> DeviceResult<f64> {
let num_qubits = circuit.num_qubits();
let state = simulate_statevector(circuit)?;
let mut expectation = 0.0;
for (idx, amp) in state.iter().enumerate() {
let weight = (idx.count_ones()) as f64;
expectation += weight * amp.norm_sqr();
}
Ok(expectation)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn bell_state_is_correlated_not_uniform() {
let mut circuit = ParameterizedQuantumCircuit::new(2);
circuit.add_h_gate(0).unwrap();
circuit.add_cnot_gate(0, 1).unwrap();
let probs = outcome_probabilities(&simulate_statevector(&circuit).unwrap());
assert!((probs[0] - 0.5).abs() < 1e-9, "P(00) should be 0.5");
assert!((probs[3] - 0.5).abs() < 1e-9, "P(11) should be 0.5");
assert!(probs[1].abs() < 1e-9, "P(01) should be 0");
assert!(probs[2].abs() < 1e-9, "P(10) should be 0");
let result = simulate_and_sample(&circuit, 4096).unwrap();
let c01 = result.counts.get("10").copied().unwrap_or(0); let c10 = result.counts.get("01").copied().unwrap_or(0);
assert_eq!(c01, 0, "Bell state must never measure 01");
assert_eq!(c10, 0, "Bell state must never measure 10");
let c00 = result.counts.get("00").copied().unwrap_or(0);
let c11 = result.counts.get("11").copied().unwrap_or(0);
assert_eq!(c00 + c11, 4096);
assert!(c00 > 0 && c11 > 0, "both correlated outcomes should appear");
}
#[test]
fn x_gate_flips_qubit() {
let mut circuit = ParameterizedQuantumCircuit::new(1);
circuit.add_x_gate(0).unwrap();
let probs = outcome_probabilities(&simulate_statevector(&circuit).unwrap());
assert!(probs[1] > 0.999, "X|0> = |1>");
assert_eq!(expected_hamming_weight(&circuit).unwrap().round() as i64, 1);
}
#[test]
fn ry_rotation_matches_analytic_probability() {
let theta = 0.7;
let mut circuit = ParameterizedQuantumCircuit::new(1);
circuit.add_ry_gate(0, theta).unwrap();
let probs = outcome_probabilities(&simulate_statevector(&circuit).unwrap());
let expected_p1 = (theta / 2.0).sin().powi(2);
assert!((probs[1] - expected_p1).abs() < 1e-9);
let weight = expected_hamming_weight(&circuit).unwrap();
assert!((weight - expected_p1).abs() < 1e-9);
}
#[test]
fn rejects_oversized_circuit() {
let circuit = ParameterizedQuantumCircuit::new(MAX_SIMULATED_QUBITS + 1);
assert!(simulate_statevector(&circuit).is_err());
}
}