use crate::gates::Gate;
pub use crate::gates::C;
pub const MAX_QUBITS: u32 = 32;
const CHUNKS: usize = 64;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum SvError {
TooLarge,
BadGate(usize),
}
#[derive(Clone, Debug, PartialEq)]
pub struct StateVector {
pub n: u32,
pub re: Vec<f64>,
pub im: Vec<f64>,
}
type SliceFn<'a> = dyn Fn(&mut [f64], &mut [f64], usize) + Sync + 'a;
fn par_slices(re: &mut [f64], im: &mut [f64], unit: usize, f: &SliceFn) {
#[cfg(not(target_arch = "wasm32"))]
{
let threads = available_threads();
let units = re.len() / unit;
if threads > 1 && units >= threads && re.len() >= 1 << 14 {
let chunk = units.div_ceil(threads) * unit;
std::thread::scope(|sc| {
for (k, (r, i)) in re.chunks_mut(chunk).zip(im.chunks_mut(chunk)).enumerate() {
sc.spawn(move || f(r, i, k * chunk));
}
});
return;
}
}
let _ = unit;
f(re, im, 0);
}
#[inline]
fn spread(mut t: usize, pos: &[u32]) -> usize {
for &p in pos {
let low = t & ((1usize << p) - 1);
t = ((t >> p) << (p + 1)) | low;
}
t
}
#[inline(always)]
fn row_dot(row: &[C], vr: &[f64], vi: &[f64]) -> (f64, f64) {
let (mut sr, mut si) = (0.0, 0.0);
for (c, e) in row.iter().enumerate() {
sr += e.re * vr[c] - e.im * vi[c];
si += e.re * vi[c] + e.im * vr[c];
}
(sr, si)
}
#[cfg(test)]
thread_local! {
static FORCE_THREADS: core::cell::Cell<Option<usize>> = const { core::cell::Cell::new(None) };
}
fn available_threads() -> usize {
#[cfg(test)]
if let Some(t) = FORCE_THREADS.with(|f| f.get()) {
return t;
}
#[cfg(not(target_arch = "wasm32"))]
{
std::thread::available_parallelism().map_or(1, |t| t.get())
}
#[cfg(target_arch = "wasm32")]
{
1
}
}
fn kernel_block(re: &mut [f64], im: &mut [f64], m: &[C], offsets: &[usize], sorted: &[u32]) {
let k = sorted.len();
if k == 1 {
let s = offsets[1];
let (m00, m01, m10, m11) = (m[0], m[1], m[2], m[3]);
let mut base = 0;
while base < re.len() {
for i in base..base + s {
let (ar, ai, br, bi) = (re[i], im[i], re[i + s], im[i + s]);
re[i] = m00.re * ar - m00.im * ai + m01.re * br - m01.im * bi;
im[i] = m00.re * ai + m00.im * ar + m01.re * bi + m01.im * br;
re[i + s] = m10.re * ar - m10.im * ai + m11.re * br - m11.im * bi;
im[i + s] = m10.re * ai + m10.im * ar + m11.re * bi + m11.im * br;
}
base += 2 * s;
}
return;
}
if k == 2 {
let (slo, shi) = (1usize << sorted[0], 1usize << sorted[1]);
let mut b1 = 0;
while b1 < re.len() {
let mut b2 = b1;
while b2 < b1 + shi {
for i in b2..b2 + slo {
let vr = [re[i + offsets[0]], re[i + offsets[1]], re[i + offsets[2]], re[i + offsets[3]]];
let vi = [im[i + offsets[0]], im[i + offsets[1]], im[i + offsets[2]], im[i + offsets[3]]];
for (r, &o) in offsets.iter().enumerate() {
let (a, b) = row_dot(&m[r * 4..r * 4 + 4], &vr, &vi);
re[i + o] = a;
im[i + o] = b;
}
}
b2 += 2 * slo;
}
b1 += 2 * shi;
}
return;
}
let d = offsets.len();
let mut vr = vec![0.0; d];
let mut vi = vec![0.0; d];
for t in 0..re.len() >> k {
let base = spread(t, sorted);
for (j, &o) in offsets.iter().enumerate() {
vr[j] = re[base | o];
vi[j] = im[base | o];
}
for (r, &o) in offsets.iter().enumerate() {
let (sr, si) = row_dot(&m[r * d..(r + 1) * d], &vr, &vi);
re[base | o] = sr;
im[base | o] = si;
}
}
}
fn kernel_pair(lre: &mut [f64], lim: &mut [f64], hre: &mut [f64], him: &mut [f64], m: &[C], members: &[(bool, usize)], rest: &[u32]) {
let d = members.len();
if d == 2 {
let (m00, m01, m10, m11) = (m[0], m[1], m[2], m[3]);
let lower_first = !members[0].0;
for i in 0..lre.len() {
let (ar, ai, br, bi) = if lower_first { (lre[i], lim[i], hre[i], him[i]) } else { (hre[i], him[i], lre[i], lim[i]) };
let nr0 = m00.re * ar - m00.im * ai + m01.re * br - m01.im * bi;
let ni0 = m00.re * ai + m00.im * ar + m01.re * bi + m01.im * br;
let nr1 = m10.re * ar - m10.im * ai + m11.re * br - m11.im * bi;
let ni1 = m10.re * ai + m10.im * ar + m11.re * bi + m11.im * br;
if lower_first {
(lre[i], lim[i], hre[i], him[i]) = (nr0, ni0, nr1, ni1);
} else {
(hre[i], him[i], lre[i], lim[i]) = (nr0, ni0, nr1, ni1);
}
}
return;
}
let k = rest.len() + 1;
let mut vr = vec![0.0; d];
let mut vi = vec![0.0; d];
for t in 0..lre.len() >> (k - 1) {
let base = spread(t, rest);
for (j, &(up, o)) in members.iter().enumerate() {
let (r, i) = if up { (&*hre, &*him) } else { (&*lre, &*lim) };
vr[j] = r[base | o];
vi[j] = i[base | o];
}
for (row, &(up, o)) in members.iter().enumerate() {
let (sr, si) = row_dot(&m[row * d..(row + 1) * d], &vr, &vi);
if up {
hre[base | o] = sr;
him[base | o] = si;
} else {
lre[base | o] = sr;
lim[base | o] = si;
}
}
}
}
impl StateVector {
pub fn zero(n: u32) -> Result<StateVector, SvError> {
if n > MAX_QUBITS {
return Err(SvError::TooLarge);
}
let dim = 1usize << n;
let mut re = vec![0.0; dim];
re[0] = 1.0;
Ok(StateVector { n, re, im: vec![0.0; dim] })
}
pub fn amplitude(&self, x: usize) -> C {
C { re: self.re[x], im: self.im[x] }
}
fn check(&self, g: &Gate, gi: usize) -> Result<(), SvError> {
let k = g.qubits.len();
let mut q = g.qubits.clone();
q.sort_unstable();
q.dedup();
if k == 0 || q.len() != k || g.qubits.iter().any(|&x| x >= self.n) || g.matrix.len() != 1 << (2 * k) {
return Err(SvError::BadGate(gi));
}
Ok(())
}
pub fn apply(&mut self, g: &Gate) -> Result<(), SvError> {
self.check(g, 0)?;
self.apply_unchecked(g);
Ok(())
}
pub fn run(&mut self, gates: &[Gate], width: usize) -> Result<(), SvError> {
for (i, g) in gates.iter().enumerate() {
self.check(g, i)?;
}
let fused;
let list = if width >= 2 {
fused = fuse(gates, width);
&fused
} else {
gates
};
for g in list {
self.apply_unchecked(g);
}
Ok(())
}
fn apply_unchecked(&mut self, g: &Gate) {
let n = self.n;
let k = g.qubits.len();
let bitpos: Vec<u32> = g.qubits.iter().map(|&q| n - 1 - q).collect();
let d = 1usize << k;
let m = &g.matrix;
if g.is_diagonal() {
let diag: Vec<C> = (0..d).map(|i| m[i * d + i]).collect();
let f = |re: &mut [f64], im: &mut [f64], off: usize| {
for (j, (a, b)) in re.iter_mut().zip(im.iter_mut()).enumerate() {
let x = off + j;
let mut local = 0usize;
for &p in &bitpos {
local = (local << 1) | ((x >> p) & 1);
}
let c = diag[local];
if c.re == 1.0 && c.im == 0.0 {
continue;
}
let (ar, ai) = (*a, *b);
*a = c.re * ar - c.im * ai;
*b = c.re * ai + c.im * ar;
}
};
par_slices(&mut self.re, &mut self.im, 1, &f);
return;
}
let mut sorted = bitpos.clone();
sorted.sort_unstable();
let top = sorted[k - 1];
let block = 1usize << (top + 1);
let len = self.re.len();
let offsets: Vec<usize> = (0..d)
.map(|local| bitpos.iter().enumerate().fold(0usize, |o, (i, &p)| o | (((local >> (k - 1 - i)) & 1) << p)))
.collect();
let threads = available_threads();
if len / block >= threads || threads == 1 || len < 1 << 14 {
let f = |re: &mut [f64], im: &mut [f64], _off: usize| kernel_block(re, im, m, &offsets, &sorted);
par_slices(&mut self.re, &mut self.im, block, &f);
return;
}
if k == 2 && sorted[0] + 1 == sorted[1] && block / 4 >= threads {
let quarter = block / 4;
let (slo, shi) = (1usize << sorted[0], 1usize << sorted[1]);
let q_of: Vec<usize> = offsets.iter().map(|&o| [0, slo, shi, slo + shi].iter().position(|&x| x == o).unwrap()).collect();
let per = quarter.div_ceil(threads);
for (bre, bim) in self.re.chunks_mut(block).zip(self.im.chunks_mut(block)) {
let (r01, r23) = bre.split_at_mut(2 * quarter);
let (r0, r1) = r01.split_at_mut(quarter);
let (r2, r3) = r23.split_at_mut(quarter);
let (i01, i23) = bim.split_at_mut(2 * quarter);
let (i0, i1) = i01.split_at_mut(quarter);
let (i2, i3) = i23.split_at_mut(quarter);
std::thread::scope(|sc| {
let parts = r0
.chunks_mut(per)
.zip(r1.chunks_mut(per))
.zip(r2.chunks_mut(per))
.zip(r3.chunks_mut(per))
.zip(i0.chunks_mut(per).zip(i1.chunks_mut(per)).zip(i2.chunks_mut(per)).zip(i3.chunks_mut(per)));
for ((((a0, a1), a2), a3), (((b0, b1), b2), b3)) in parts {
let q_of = &q_of;
sc.spawn(move || {
let (re4, im4) = ([a0, a1, a2, a3], [b0, b1, b2, b3]);
for i in 0..re4[0].len() {
let vr = [re4[q_of[0]][i], re4[q_of[1]][i], re4[q_of[2]][i], re4[q_of[3]][i]];
let vi = [im4[q_of[0]][i], im4[q_of[1]][i], im4[q_of[2]][i], im4[q_of[3]][i]];
let mut out = [(0.0, 0.0); 4];
for (r, o) in out.iter_mut().enumerate() {
*o = row_dot(&m[r * 4..r * 4 + 4], &vr, &vi);
}
for (r, &(a, b)) in out.iter().enumerate() {
re4[q_of[r]][i] = a;
im4[q_of[r]][i] = b;
}
}
});
}
});
}
return;
}
let half = block / 2;
let chunk = if k == 1 { 1usize } else { 1usize << (sorted[k - 2] + 1) };
if half / chunk < threads {
kernel_block(&mut self.re, &mut self.im, m, &offsets, &sorted);
return;
}
let top_bit = 1usize << top;
let members: Vec<(bool, usize)> = offsets.iter().map(|&o| (o & top_bit != 0, o & !top_bit)).collect();
let rest = &sorted[..k - 1];
let per = (half / chunk).div_ceil(threads) * chunk;
for (bre, bim) in self.re.chunks_mut(block).zip(self.im.chunks_mut(block)) {
let (lre, hre) = bre.split_at_mut(half);
let (lim, him) = bim.split_at_mut(half);
std::thread::scope(|sc| {
for (((a, b), c), e) in lre.chunks_mut(per).zip(lim.chunks_mut(per)).zip(hre.chunks_mut(per)).zip(him.chunks_mut(per)) {
let members = &members;
sc.spawn(move || kernel_pair(a, b, c, e, m, members, rest));
}
});
}
}
pub fn norm2(&self) -> f64 {
self.chunked_sum(&|x| self.re[x] * self.re[x] + self.im[x] * self.im[x])
}
fn chunked_sum(&self, term: &(dyn Fn(usize) -> f64 + Sync)) -> f64 {
let len = self.re.len();
let chunks = len.min(CHUNKS);
let per = len / chunks;
let sum_chunk = |c: usize| -> f64 {
let mut s = 0.0;
for x in c * per..(c + 1) * per {
s += term(x);
}
s
};
#[cfg(not(target_arch = "wasm32"))]
let parts: Vec<f64> = {
let threads = available_threads().min(chunks);
if threads > 1 && len >= 1 << 14 {
let size = chunks.div_ceil(threads);
std::thread::scope(|sc| {
let handles: Vec<_> = (0..chunks)
.step_by(size)
.map(|lo| {
let sum_chunk = &sum_chunk;
sc.spawn(move || (lo..(lo + size).min(chunks)).map(sum_chunk).collect::<Vec<f64>>())
})
.collect();
handles.into_iter().flat_map(|h| h.join().expect("a chunk")).collect()
})
} else {
(0..chunks).map(sum_chunk).collect()
}
};
#[cfg(target_arch = "wasm32")]
let parts: Vec<f64> = (0..chunks).map(sum_chunk).collect();
parts.into_iter().fold(0.0, |a, b| a + b)
}
pub fn expectation(&self, observable: &[(f64, Vec<(u32, char)>)]) -> f64 {
let n = self.n;
let mut total = 0.0;
for (coef, string) in observable {
let mut flip = 0usize;
let mut zmask = 0usize;
let mut ys = 0u32;
for &(q, p) in string {
let bit = 1usize << (n - 1 - q);
match p {
'X' => flip |= bit,
'Y' => {
flip |= bit;
zmask |= bit;
ys += 1;
}
'Z' => zmask |= bit,
_ => {}
}
}
let v = self.chunked_sum(&|x| {
let y = x ^ flip;
let sign = if (x & zmask).count_ones().is_multiple_of(2) { 1.0 } else { -1.0 };
let (ar, ai) = (self.re[y], -self.im[y]);
let (br, bi) = (self.re[x], self.im[x]);
let (pr, pi) = (ar * br - ai * bi, ar * bi + ai * br);
let re = match ys % 4 {
0 => pr,
1 => -pi,
2 => -pr,
_ => pi,
};
sign * re
});
total += coef * v;
}
total
}
pub fn sample(&self, shots: usize, seed: u64) -> Vec<u64> {
let mut s = seed;
let mut u: Vec<(f64, usize)> = (0..shots)
.map(|i| {
s = s.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut z = s;
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
((((z ^ (z >> 31)) >> 11) as f64 + 0.5) / 9_007_199_254_740_992.0, i)
})
.collect();
u.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap().then(a.1.cmp(&b.1)));
let total = self.norm2();
let mut out = vec![0u64; shots];
let mut cum = 0.0;
let mut k = 0;
let last = self.re.len() - 1;
for x in 0..self.re.len() {
cum += (self.re[x] * self.re[x] + self.im[x] * self.im[x]) / total;
while k < shots && (u[k].0 < cum || x == last) {
out[u[k].1] = x as u64;
k += 1;
}
if k == shots {
break;
}
}
out
}
}
pub fn fuse(gates: &[Gate], width: usize) -> Vec<Gate> {
let nq = gates.iter().flat_map(|g| g.qubits.iter()).map(|&q| q as usize + 1).max().unwrap_or(0);
let mut pending: Vec<Option<Gate>> = vec![None; nq];
let mut out: Vec<Gate> = Vec::new();
let mut last_multi: Vec<Option<usize>> = vec![None; nq];
let compose = |later: &Gate, earlier: &Gate| -> Gate {
let mut union = earlier.qubits.clone();
for &q in &later.qubits {
if !union.contains(&q) {
union.push(q);
}
}
let d = 1usize << union.len();
Gate { qubits: union.clone(), matrix: matmul(&embed(later, &union), &embed(earlier, &union), d) }
};
for g in gates {
if g.qubits.len() == 1 {
let q = g.qubits[0] as usize;
pending[q] = Some(match pending[q].take() {
Some(p) => compose(g, &p),
None => g.clone(),
});
continue;
}
let mut cur = g.clone();
for &q in &g.qubits {
if let Some(p) = pending[q as usize].take() {
let embedded = Gate { qubits: cur.qubits.clone(), matrix: embed(&p, &cur.qubits) };
let d = 1usize << cur.qubits.len();
cur.matrix = matmul(&cur.matrix, &embedded.matrix, d);
}
}
for &q in &cur.qubits {
last_multi[q as usize] = Some(out.len());
}
out.push(cur);
}
let mut alone = Vec::new();
for (q, p) in pending.into_iter().enumerate() {
let Some(p) = p else { continue };
match last_multi[q] {
Some(i) => {
let host = &out[i];
let embedded = embed(&p, &host.qubits);
let d = 1usize << host.qubits.len();
let m = matmul(&embedded, &host.matrix, d);
out[i].matrix = m;
}
None => alone.push(p),
}
}
out.extend(alone);
if width <= 2 {
return out;
}
let mut merged: Vec<Gate> = Vec::new();
for g in out {
if let Some(prev) = merged.last() {
let mut union = prev.qubits.clone();
for &q in &g.qubits {
if !union.contains(&q) {
union.push(q);
}
}
if union.len() <= width {
let prev = merged.pop().unwrap();
merged.push(compose(&g, &prev));
continue;
}
}
merged.push(g);
}
merged
}
fn matmul(a: &[C], b: &[C], d: usize) -> Vec<C> {
let mut out = vec![C { re: 0.0, im: 0.0 }; d * d];
for i in 0..d {
for k in 0..d {
let x = a[i * d + k];
if x.re == 0.0 && x.im == 0.0 {
continue;
}
for j in 0..d {
let y = b[k * d + j];
let o = &mut out[i * d + j];
o.re += x.re * y.re - x.im * y.im;
o.im += x.re * y.im + x.im * y.re;
}
}
}
out
}
fn embed(g: &Gate, onto: &[u32]) -> Vec<C> {
let (k, m) = (g.qubits.len(), onto.len());
let (dk, dm) = (1usize << k, 1usize << m);
let pos: Vec<usize> = g.qubits.iter().map(|q| m - 1 - onto.iter().position(|o| o == q).unwrap()).collect();
let rest_mask: usize = (0..m).filter(|b| !pos.contains(b)).fold(0, |a, b| a | (1 << b));
let mut out = vec![C { re: 0.0, im: 0.0 }; dm * dm];
for r in 0..dm {
for c in 0..dm {
if (r & rest_mask) != (c & rest_mask) {
continue;
}
let lr = pos.iter().enumerate().fold(0, |a, (i, &p)| a | (((r >> p) & 1) << (k - 1 - i)));
let lc = pos.iter().enumerate().fold(0, |a, (i, &p)| a | (((c >> p) & 1) << (k - 1 - i)));
out[r * dm + c] = g.matrix[lr * dk + lc];
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn rng(seed: u64) -> impl FnMut() -> f64 {
let mut s = seed;
move || {
s = s.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut z = s;
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
(((z ^ (z >> 31)) >> 11) as f64 + 0.5) / 9_007_199_254_740_992.0
}
}
fn random_circuit(n: u32, depth: usize, seed: u64) -> Vec<Gate> {
let mut u = rng(seed);
let mut g = Vec::new();
for layer in 0..depth {
for q in 0..n {
let a = u() * core::f64::consts::TAU;
g.push(match (u() * 7.0) as u32 {
0 => Gate::rx(q, a),
1 => Gate::ry(q, a),
2 => Gate::rz(q, a),
3 => Gate::sqrt_w(q),
4 => Gate::sqrt_x(q),
5 => Gate::t(q),
_ => Gate::h(q),
});
}
let mut q = (layer % 2) as u32;
while q + 1 < n {
g.push(match (u() * 4.0) as u32 {
0 => Gate::cx(q, q + 1),
1 => Gate::fsim(q + 1, q, u() * 3.0, u() * 3.0),
2 => Gate::cz(q, q + 1),
_ => Gate::rzz(q, q + 1, u() * 3.0),
});
q += 2;
}
let (a, b) = ((u() * n as f64) as u32, (u() * n as f64) as u32);
if a != b {
g.push(Gate::iswap(a, b));
}
}
g
}
fn reference(n: u32, gates: &[Gate]) -> Vec<C> {
let dim = 1usize << n;
let mut psi = vec![C { re: 0.0, im: 0.0 }; dim];
psi[0] = C { re: 1.0, im: 0.0 };
for g in gates {
let k = g.qubits.len();
let mut out = vec![C { re: 0.0, im: 0.0 }; dim];
for (x, a) in psi.iter().enumerate() {
let col = g.qubits.iter().fold(0usize, |c, &q| (c << 1) | ((x >> (n - 1 - q)) & 1));
for row in 0..1usize << k {
let mut y = x;
for (i, &q) in g.qubits.iter().enumerate() {
let s = n - 1 - q;
y = (y & !(1 << s)) | (((row >> (k - 1 - i)) & 1) << s);
}
let e = g.matrix[row * (1 << k) + col];
out[y].re += e.re * a.re - e.im * a.im;
out[y].im += e.re * a.im + e.im * a.re;
}
}
psi = out;
}
psi
}
#[test]
#[allow(clippy::needless_range_loop)]
fn kernels_and_fusion_match_the_reference() {
for seed in 0..5 {
let n = 9;
let gates = random_circuit(n, 8, seed);
let want = reference(n, &gates);
for width in [0usize, 2, 3, 4] {
let mut sv = StateVector::zero(n).unwrap();
sv.run(&gates, width).unwrap();
for x in 0..1usize << n {
let (a, b) = (sv.amplitude(x), want[x]);
assert!((a.re - b.re).abs() < 1e-12 && (a.im - b.im).abs() < 1e-12, "seed {seed} width {width} x {x}");
}
assert!((sv.norm2() - 1.0).abs() < 1e-12);
}
}
}
#[test]
fn the_split_half_path_matches_the_reference() {
let n = 16;
let gates = random_circuit(n, 5, 9);
let want = reference(n, &gates);
for width in [0usize, 2, 3, 4] {
let mut sv = StateVector::zero(n).unwrap();
sv.run(&gates, width).unwrap();
let worst = (0..1usize << n).map(|x| (sv.re[x] - want[x].re).abs().max((sv.im[x] - want[x].im).abs())).fold(0.0, f64::max);
assert!(worst < 1e-12, "width {width}: {worst}");
let mut again = StateVector::zero(n).unwrap();
again.run(&gates, width).unwrap();
assert_eq!(sv, again);
}
}
#[test]
fn threading_never_changes_a_bit() {
let n = 16;
let gates = random_circuit(n, 6, 13);
for width in [0usize, 2, 3, 4] {
let run = |threads: Option<usize>| {
FORCE_THREADS.with(|f| f.set(threads));
let mut sv = StateVector::zero(n).unwrap();
sv.run(&gates, width).unwrap();
let e = sv.expectation(&[(1.0, vec![(0, 'X'), (9, 'Z')])]);
let norm = sv.norm2();
FORCE_THREADS.with(|f| f.set(None));
(sv, e.to_bits(), norm.to_bits())
};
let one = run(Some(1));
for t in [2usize, 3, 8, 32] {
assert_eq!(run(Some(t)), one, "width {width}, {t} threads");
}
}
}
#[test]
fn named_gates_are_unitary_and_correct() {
let sq = |g: Gate| {
let mut sv = StateVector::zero(1).unwrap();
sv.apply(&Gate::h(0)).unwrap();
sv.apply(&Gate::rz(0, 0.3)).unwrap();
let mut a = sv.clone();
a.apply(&g).unwrap();
a.apply(&g).unwrap();
a
};
let once = |g: Gate| {
let mut sv = StateVector::zero(1).unwrap();
sv.apply(&Gate::h(0)).unwrap();
sv.apply(&Gate::rz(0, 0.3)).unwrap();
sv.apply(&g).unwrap();
sv
};
let same = |a: &StateVector, b: &StateVector| (0..a.re.len()).all(|x| (a.re[x] - b.re[x]).abs() < 1e-15 && (a.im[x] - b.im[x]).abs() < 1e-15);
assert!(same(&sq(Gate::sqrt_x(0)), &once(Gate::x(0))));
assert!(same(&sq(Gate::sqrt_y(0)), &once(Gate::y(0))));
assert!(same(&sq(Gate::t(0)), &once(Gate::s(0))));
assert!(same(&sq(Gate::s(0)), &once(Gate::z(0))));
let r = core::f64::consts::FRAC_1_SQRT_2;
let w = Gate { qubits: vec![0], matrix: vec![C { re: 0.0, im: 0.0 }, C { re: r, im: -r }, C { re: r, im: r }, C { re: 0.0, im: 0.0 }] };
assert!(same(&sq(Gate::sqrt_w(0)), &once(w)));
}
#[test]
fn expectations_and_samples_are_right_and_reproducible() {
let mut sv = StateVector::zero(2).unwrap();
sv.run(&[Gate::h(0), Gate::cx(0, 1)], 0).unwrap();
let e = |s: Vec<(u32, char)>| sv.expectation(&[(1.0, s)]);
assert!((e(vec![(0, 'Z'), (1, 'Z')]) - 1.0).abs() < 1e-15);
assert!((e(vec![(0, 'X'), (1, 'X')]) - 1.0).abs() < 1e-15);
assert!((e(vec![(0, 'Y'), (1, 'Y')]) + 1.0).abs() < 1e-15);
assert!(e(vec![(0, 'Z')]).abs() < 1e-15);
let s = sv.sample(10_000, 3);
assert!(s.iter().all(|&x| x == 0 || x == 3));
let ones = s.iter().filter(|&&x| x == 3).count();
assert!((ones as f64 / 10_000.0 - 0.5).abs() < 0.03);
assert_eq!(s, sv.sample(10_000, 3));
let n = 12;
let mut big = StateVector::zero(n).unwrap();
big.run(&random_circuit(n, 6, 4), 3).unwrap();
for q in [0u32, 5, 11] {
let direct: f64 = (0..1usize << n)
.map(|x| {
let p = big.re[x] * big.re[x] + big.im[x] * big.im[x];
if (x >> (n - 1 - q)) & 1 == 0 { p } else { -p }
})
.sum();
assert!((big.expectation(&[(1.0, vec![(q, 'Z')])]) - direct).abs() < 1e-12);
}
}
#[test]
fn bad_input_is_refused() {
assert_eq!(StateVector::zero(33).unwrap_err(), SvError::TooLarge);
let mut sv = StateVector::zero(2).unwrap();
assert_eq!(sv.apply(&Gate::cx(0, 2)).unwrap_err(), SvError::BadGate(0));
assert_eq!(sv.run(&[Gate::h(0), Gate { qubits: vec![1, 1], matrix: vec![C { re: 1.0, im: 0.0 }; 16] }], 2).unwrap_err(), SvError::BadGate(1));
}
}