use crate::error::{QuantumError, Result};
use crate::state::QuantumState;
use moonlab_sys::{
fuse_append_cnot, fuse_append_cphase, fuse_append_crx, fuse_append_cry,
fuse_append_crz, fuse_append_cy, fuse_append_cz, fuse_append_h,
fuse_append_phase, fuse_append_rx, fuse_append_ry, fuse_append_rz,
fuse_append_s, fuse_append_sdg, fuse_append_swap, fuse_append_t,
fuse_append_tdg, fuse_append_u3, fuse_append_x, fuse_append_y,
fuse_append_z, fuse_circuit_create, fuse_circuit_free, fuse_circuit_len,
fuse_circuit_num_qubits, fuse_circuit_t, fuse_compile, fuse_execute,
fuse_stats_t,
};
use std::ptr;
#[derive(Debug, Copy, Clone, Eq, PartialEq)]
pub struct FuseStats {
pub original_gates: usize,
pub fused_gates: usize,
pub merges_applied: usize,
}
pub struct FusedCircuit {
handle: *mut fuse_circuit_t,
num_qubits: usize,
}
unsafe impl Send for FusedCircuit {}
impl FusedCircuit {
pub fn new(num_qubits: usize) -> Result<Self> {
if num_qubits == 0 {
return Err(QuantumError::InvalidQubit { index: 0, max: 1 });
}
let handle = unsafe { fuse_circuit_create(num_qubits) };
if handle.is_null() {
return Err(QuantumError::AllocationFailed(num_qubits));
}
Ok(Self {
handle,
num_qubits,
})
}
pub fn num_qubits(&self) -> usize {
unsafe { fuse_circuit_num_qubits(self.handle) }
}
pub fn len(&self) -> usize {
unsafe { fuse_circuit_len(self.handle) }
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn h(&mut self, q: i32) -> Result<&mut Self> { self.call1("fuse_append_h", unsafe { fuse_append_h (self.handle, q) }) }
pub fn x(&mut self, q: i32) -> Result<&mut Self> { self.call1("fuse_append_x", unsafe { fuse_append_x (self.handle, q) }) }
pub fn y(&mut self, q: i32) -> Result<&mut Self> { self.call1("fuse_append_y", unsafe { fuse_append_y (self.handle, q) }) }
pub fn z(&mut self, q: i32) -> Result<&mut Self> { self.call1("fuse_append_z", unsafe { fuse_append_z (self.handle, q) }) }
pub fn s(&mut self, q: i32) -> Result<&mut Self> { self.call1("fuse_append_s", unsafe { fuse_append_s (self.handle, q) }) }
pub fn sdg(&mut self, q: i32) -> Result<&mut Self> { self.call1("fuse_append_sdg", unsafe { fuse_append_sdg(self.handle, q) }) }
pub fn t(&mut self, q: i32) -> Result<&mut Self> { self.call1("fuse_append_t", unsafe { fuse_append_t (self.handle, q) }) }
pub fn tdg(&mut self, q: i32) -> Result<&mut Self> { self.call1("fuse_append_tdg", unsafe { fuse_append_tdg(self.handle, q) }) }
pub fn phase(&mut self, q: i32, theta: f64) -> Result<&mut Self> {
self.call1("fuse_append_phase", unsafe { fuse_append_phase(self.handle, q, theta) })
}
pub fn rx(&mut self, q: i32, theta: f64) -> Result<&mut Self> {
self.call1("fuse_append_rx", unsafe { fuse_append_rx(self.handle, q, theta) })
}
pub fn ry(&mut self, q: i32, theta: f64) -> Result<&mut Self> {
self.call1("fuse_append_ry", unsafe { fuse_append_ry(self.handle, q, theta) })
}
pub fn rz(&mut self, q: i32, theta: f64) -> Result<&mut Self> {
self.call1("fuse_append_rz", unsafe { fuse_append_rz(self.handle, q, theta) })
}
pub fn u3(&mut self, q: i32, theta: f64, phi: f64, lambda: f64)
-> Result<&mut Self>
{
let rc = unsafe { fuse_append_u3(self.handle, q, theta, phi, lambda) };
self.call1("fuse_append_u3", rc)
}
pub fn cnot(&mut self, ctrl: i32, tgt: i32) -> Result<&mut Self> {
self.call1("fuse_append_cnot", unsafe { fuse_append_cnot(self.handle, ctrl, tgt) })
}
pub fn cz(&mut self, ctrl: i32, tgt: i32) -> Result<&mut Self> {
self.call1("fuse_append_cz", unsafe { fuse_append_cz(self.handle, ctrl, tgt) })
}
pub fn cy(&mut self, ctrl: i32, tgt: i32) -> Result<&mut Self> {
self.call1("fuse_append_cy", unsafe { fuse_append_cy(self.handle, ctrl, tgt) })
}
pub fn swap(&mut self, a: i32, b: i32) -> Result<&mut Self> {
self.call1("fuse_append_swap", unsafe { fuse_append_swap(self.handle, a, b) })
}
pub fn cphase(&mut self, ctrl: i32, tgt: i32, theta: f64) -> Result<&mut Self> {
self.call1("fuse_append_cphase", unsafe { fuse_append_cphase(self.handle, ctrl, tgt, theta) })
}
pub fn crx(&mut self, ctrl: i32, tgt: i32, theta: f64) -> Result<&mut Self> {
self.call1("fuse_append_crx", unsafe { fuse_append_crx(self.handle, ctrl, tgt, theta) })
}
pub fn cry(&mut self, ctrl: i32, tgt: i32, theta: f64) -> Result<&mut Self> {
self.call1("fuse_append_cry", unsafe { fuse_append_cry(self.handle, ctrl, tgt, theta) })
}
pub fn crz(&mut self, ctrl: i32, tgt: i32, theta: f64) -> Result<&mut Self> {
self.call1("fuse_append_crz", unsafe { fuse_append_crz(self.handle, ctrl, tgt, theta) })
}
pub fn compile(&self) -> Result<(FusedCircuit, FuseStats)> {
let mut stats = fuse_stats_t {
original_gates: 0,
fused_gates: 0,
merges_applied: 0,
};
let h = unsafe { fuse_compile(self.handle, &mut stats) };
if h.is_null() {
return Err(QuantumError::AllocationFailed(0));
}
let fused = FusedCircuit {
handle: h,
num_qubits: self.num_qubits,
};
Ok((
fused,
FuseStats {
original_gates: stats.original_gates,
fused_gates: stats.fused_gates,
merges_applied: stats.merges_applied,
},
))
}
pub fn execute(&self, state: &mut QuantumState) -> Result<()> {
let rc = unsafe { fuse_execute(self.handle, state.as_ptr()) };
if rc != 0 {
Err(QuantumError::Ffi(format!("fuse_execute rc={rc}")))
} else {
Ok(())
}
}
fn call1(&mut self, name: &'static str, rc: i32) -> Result<&mut Self> {
if rc != 0 {
Err(QuantumError::Ffi(format!("{name} rc={rc}")))
} else {
Ok(self)
}
}
}
impl Drop for FusedCircuit {
fn drop(&mut self) {
if !self.handle.is_null() {
unsafe { fuse_circuit_free(self.handle) };
self.handle = ptr::null_mut();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn create_and_len() {
let c = FusedCircuit::new(4).unwrap();
assert_eq!(c.num_qubits(), 4);
assert_eq!(c.len(), 0);
assert!(c.is_empty());
}
#[test]
fn reject_zero_qubits() {
assert!(FusedCircuit::new(0).is_err());
}
#[test]
fn fluent_append_increments_len() {
let mut c = FusedCircuit::new(3).unwrap();
c.h(0).unwrap()
.rz(0, 0.3).unwrap()
.rx(0, 0.4).unwrap()
.cnot(0, 1).unwrap()
.ry(1, 0.5).unwrap();
assert_eq!(c.len(), 5);
}
#[test]
fn compile_run_fuses_three_into_one() {
let mut c = FusedCircuit::new(2).unwrap();
c.h(0).unwrap().rz(0, 0.3).unwrap().rx(0, 0.4).unwrap().cnot(0, 1).unwrap();
let (fused, stats) = c.compile().unwrap();
assert_eq!(stats.original_gates, 4);
assert_eq!(stats.fused_gates, 2);
assert_eq!(stats.merges_applied, 2);
assert_eq!(fused.len(), 2);
}
#[test]
fn compile_passthrough_when_no_runs() {
let mut c = FusedCircuit::new(3).unwrap();
c.h(0).unwrap().cnot(0, 1).unwrap().h(1).unwrap().cnot(1, 2).unwrap().h(2).unwrap();
let (_fused, stats) = c.compile().unwrap();
assert_eq!(stats.original_gates, 5);
assert_eq!(stats.merges_applied, 0);
assert_eq!(stats.fused_gates, 5);
}
#[test]
fn execute_bell_state() {
let mut c = FusedCircuit::new(2).unwrap();
c.h(0).unwrap().cnot(0, 1).unwrap();
let mut state = QuantumState::new(2).unwrap();
c.execute(&mut state).unwrap();
let probs = state.probabilities();
assert!((probs[0] - 0.5).abs() < 1e-10);
assert!(probs[1].abs() < 1e-10);
assert!(probs[2].abs() < 1e-10);
assert!((probs[3] - 0.5).abs() < 1e-10);
}
#[test]
fn execute_fused_matches_unfused() {
let build = || -> FusedCircuit {
let mut c = FusedCircuit::new(3).unwrap();
c.h(0).unwrap()
.rz(0, 0.7).unwrap()
.rx(0, 0.3).unwrap()
.cnot(0, 1).unwrap()
.ry(1, 0.2).unwrap()
.rz(1, 0.9).unwrap()
.cnot(1, 2).unwrap()
.rx(2, 0.4).unwrap();
c
};
let mut s_unfused = QuantumState::new(3).unwrap();
build().execute(&mut s_unfused).unwrap();
let mut s_fused = QuantumState::new(3).unwrap();
let (fused, _) = build().compile().unwrap();
fused.execute(&mut s_fused).unwrap();
let pu = s_unfused.probabilities();
let pf = s_fused.probabilities();
for (u, f) in pu.iter().zip(pf.iter()) {
assert!((u - f).abs() < 1e-10, "{u} != {f}");
}
}
}