use crate::quantum::{BaseGate, Circuit, Gate};
use std::collections::HashMap;
#[derive(Debug)]
pub struct Unsupported(pub String);
pub type PauliSum = HashMap<(u64, u64), f64>;
fn sorted(sum: &PauliSum) -> Vec<((u64, u64), f64)> {
let mut v: Vec<((u64, u64), f64)> = sum.iter().map(|(&k, &c)| (k, c)).collect();
v.sort_unstable_by_key(|&(k, _)| k);
v
}
#[inline]
fn bit(v: u64, q: u8) -> u64 {
(v >> q) & 1
}
fn conj1(base: BaseGate, param: u16, xq: u64, zq: u64) -> (u64, u64, f64) {
let f = (zq << 1) | xq;
let (nf, s): (u64, f64) = match base {
BaseGate::H => match f {
0 => (0, 1.0),
1 => (2, 1.0), 2 => (1, 1.0), _ => (3, -1.0), },
BaseGate::S => match f {
0 => (0, 1.0),
1 => (3, -1.0), 2 => (2, 1.0), _ => (1, 1.0), },
BaseGate::Sdg => match f {
0 => (0, 1.0),
1 => (3, 1.0), 2 => (2, 1.0),
_ => (1, -1.0), },
BaseGate::X => match f {
0 => (0, 1.0),
1 => (1, 1.0),
2 => (2, -1.0), _ => (3, -1.0), },
BaseGate::Y => match f {
0 => (0, 1.0),
1 => (1, -1.0),
2 => (2, -1.0),
_ => (3, 1.0),
},
BaseGate::Z => match f {
0 => (0, 1.0),
1 => (1, -1.0),
2 => (2, 1.0),
_ => (3, -1.0),
},
BaseGate::P if param == 1 => return conj1(BaseGate::Z, 0, xq, zq),
BaseGate::P if param == 2 => return conj1(BaseGate::S, 0, xq, zq),
_ => (f, 1.0), };
(nf & 1, (nf >> 1) & 1, s)
}
#[inline]
fn set_factor(x: u64, z: u64, q: u8, nx: u64, nz: u64) -> (u64, u64) {
let m = 1u64 << q;
((x & !m) | (nx << q), (z & !m) | (nz << q))
}
fn is_diag_rotation(g: &Gate) -> Option<f64> {
if !g.controls.is_empty() {
return None;
}
match g.base {
BaseGate::T => Some(std::f64::consts::FRAC_PI_4),
BaseGate::Tdg => Some(-std::f64::consts::FRAC_PI_4),
BaseGate::P if g.param >= 3 => Some(2.0 * std::f64::consts::PI / (1u64 << g.param) as f64),
_ => None,
}
}
fn conjugate(sum: &PauliSum, g: &Gate, out: &mut PauliSum) -> Result<(), Unsupported> {
out.clear();
let add = |x: u64, z: u64, c: f64, out: &mut PauliSum| {
if c.abs() > 0.0 {
*out.entry((x, z)).or_insert(0.0) += c;
}
};
if let Some(phi) = is_diag_rotation(g) {
let q = g.target;
let (sn, cs) = crate::repro::sin_cos(phi);
for ((x, z), c) in sorted(sum) {
if bit(x, q) == 0 {
add(x, z, c, out); } else {
let zq = bit(z, q);
add(x, z, c * cs, out); let (px, pz) = set_factor(x, z, q, 1, zq ^ 1); let ps = if zq == 0 { -sn } else { sn };
add(px, pz, c * ps, out);
}
}
return Ok(());
}
match (g.base, g.controls.as_slice()) {
(BaseGate::I, []) => {
for (k, c) in sorted(sum) {
add(k.0, k.1, c, out);
}
}
(BaseGate::H | BaseGate::S | BaseGate::Sdg | BaseGate::X | BaseGate::Y | BaseGate::Z, [])
| (BaseGate::P, []) => {
let q = g.target;
for ((x, z), c) in sorted(sum) {
let (nx, nz, s) = conj1(g.base, g.param, bit(x, q), bit(z, q));
let (x2, z2) = set_factor(x, z, q, nx, nz);
add(x2, z2, c * s, out);
}
}
(BaseGate::X, [ctrl]) => {
let (c, t) = (*ctrl, g.target);
for ((x, z), co) in sorted(sum) {
let (xc, zc, xt, zt) = (bit(x, c), bit(z, c), bit(x, t), bit(z, t));
let s = if (xc & zt & (xt ^ zc ^ 1)) == 1 { -1.0 } else { 1.0 };
let (x2, _) = set_factor(x, z, t, xt ^ xc, zt);
let (x3, z3) = set_factor(x2, z, c, xc, zc ^ zt);
add(x3, z3, co * s, out);
}
}
(BaseGate::Z, [ctrl]) | (BaseGate::P, [ctrl]) if matches!(g.base, BaseGate::Z) || g.param == 1 => {
let ht = Gate { base: BaseGate::H, controls: vec![], target: g.target, param: 0 };
let cx = Gate { base: BaseGate::X, controls: vec![*ctrl], target: g.target, param: 0 };
let mut a = PauliSum::new();
let mut b = PauliSum::new();
conjugate(sum, &ht, &mut a)?;
conjugate(&a, &cx, &mut b)?;
conjugate(&b, &ht, out)?;
}
(base, ctrls) => {
return Err(Unsupported(format!("{base:?} controls={ctrls:?} param={}", g.param)));
}
}
Ok(())
}
pub fn propagate(
circuit: &Circuit,
obs: &PauliSum,
threshold: f64,
) -> Result<(PauliSum, usize), Unsupported> {
let mut cur = obs.clone();
let mut next = PauliSum::new();
let mut peak = cur.len();
for g in circuit.ops.iter().rev() {
conjugate(&cur, g, &mut next)?;
next.retain(|_, c| c.abs() >= threshold);
std::mem::swap(&mut cur, &mut next);
peak = peak.max(cur.len());
}
Ok((cur, peak))
}
pub fn expectation_zero(sum: &PauliSum) -> f64 {
sorted(sum).into_iter().filter(|((x, _), _)| *x == 0).map(|(_, c)| c).sum()
}
pub fn expect_z(circuit: &Circuit, q: u8, threshold: f64) -> Result<(f64, usize), Unsupported> {
let mut obs = PauliSum::new();
obs.insert((0, 1u64 << q), 1.0);
let (s, peak) = propagate(circuit, &obs, threshold)?;
Ok((expectation_zero(&s), peak))
}
pub fn expect_pauli(circuit: &Circuit, x: u64, z: u64, threshold: f64) -> Result<(f64, usize), Unsupported> {
let mut obs = PauliSum::new();
obs.insert((x, z), 1.0);
let (s, peak) = propagate(circuit, &obs, threshold)?;
Ok((expectation_zero(&s), peak))
}
#[cfg(test)]
mod tests {
use super::*;
fn dense_expect_z(c: &Circuit, q: u8) -> f64 {
let sv = c.simulate().unwrap();
let w = sv.prob_weights();
let total: f64 = w.iter().map(|&x| x as f64).sum();
let mut acc = 0.0;
for (i, wi) in w.iter().enumerate() {
let s = if (i >> q) & 1 == 0 { 1.0 } else { -1.0 };
acc += s * (*wi as f64);
}
acc / total
}
fn bell() -> Circuit {
let mut c = Circuit::new(2);
c.h(0);
c.cx(0, 1);
c
}
fn approx(a: f64, b: f64) -> bool {
(a - b).abs() < 1e-6
}
#[test]
fn zero_state_all_up() {
let c = Circuit::new(4);
for q in 0..4 {
assert!(approx(expect_z(&c, q, 0.0).unwrap().0, 1.0));
}
}
#[test]
fn bell_correlations() {
let c = bell();
assert!(approx(expect_z(&c, 0, 0.0).unwrap().0, 0.0)); assert!(approx(expect_z(&c, 1, 0.0).unwrap().0, 0.0));
assert!(approx(expect_pauli(&c, 0, 0b11, 0.0).unwrap().0, 1.0));
}
#[test]
fn matches_dense_on_clifford_t_circuits() {
let mut c = Circuit::new(5);
c.h(0);
c.t(0);
c.cx(0, 1);
c.h(2);
c.cx(2, 3);
c.t(3);
c.s(1);
c.cx(1, 4);
c.h(4);
c.t(4);
for q in 0..5 {
let sp = expect_z(&c, q, 0.0).unwrap().0;
let de = dense_expect_z(&c, q);
assert!(approx(sp, de), "q{q}: sparse {sp} vs dense {de}");
}
}
#[test]
fn matches_dense_with_many_t_gates() {
let mut c = Circuit::new(4);
for _ in 0..3 {
for q in 0..4 {
c.h(q);
c.t(q);
}
c.cx(0, 1);
c.cx(2, 3);
c.cx(1, 2);
}
for q in 0..4 {
assert!(approx(expect_z(&c, q, 0.0).unwrap().0, dense_expect_z(&c, q)),
"q{q} mismatch");
}
}
#[test]
fn truncation_trades_accuracy_for_size() {
let mut c = Circuit::new(6);
for _ in 0..4 {
for q in 0..6 {
c.h(q);
c.t(q);
}
for q in 0..5 {
c.cx(q, q + 1);
}
}
let exact = expect_z(&c, 0, 0.0).unwrap();
let trunc = expect_z(&c, 0, 1e-2).unwrap();
assert!(trunc.1 <= exact.1, "truncated peak {} ≤ exact peak {}", trunc.1, exact.1);
assert!((trunc.0 - exact.0).abs() < 0.1, "trunc {} vs exact {}", trunc.0, exact.0);
assert!(approx(exact.0, dense_expect_z(&c, 0)), "exact must match dense");
}
#[test]
fn scales_past_the_dense_wall() {
let n = 40u8;
let mut c = Circuit::new(n);
for q in 0..n {
c.h(q);
}
for q in 0..n - 1 {
c.cx(q, q + 1);
}
for q in 0..n {
c.t(q);
}
let (v, peak) = expect_z(&c, 0, 1e-9).unwrap();
assert!(v.is_finite());
assert!(peak < 1 << 20, "Pauli sum stayed bounded: {peak}");
}
#[test]
fn results_are_bit_identical_from_run_to_run() {
let n = 20u8;
let mut c = Circuit::new(n);
for layer in 0..4u8 {
for q in 0..n {
c.h(q);
c.t(q);
}
for q in (layer % 2..n - 1).step_by(2) {
c.cx(q, q + 1);
}
for q in 0..n {
c.t(q);
}
}
let (a, pa) = expect_z(&c, 10, 1e-4).unwrap();
assert!(pa > 500, "enough terms to merge in many orders: {pa}");
for _ in 0..6 {
let (b, pb) = expect_z(&c, 10, 1e-4).unwrap();
assert_eq!((a.to_bits(), pa), (b.to_bits(), pb));
}
}
#[test]
fn rejects_unsupported() {
let mut c = Circuit::new(2);
c.cp(2, 0, 1);
assert!(expect_z(&c, 0, 0.0).is_err());
}
}