#![allow(clippy::should_implement_trait, clippy::needless_range_loop, clippy::manual_range_contains)]
use crate::quantum::{BaseGate, Circuit, Gate, StateVector};
#[derive(Clone, Copy, Debug)]
pub struct C {
pub re: f64,
pub im: f64,
}
impl C {
const ZERO: C = C { re: 0.0, im: 0.0 };
const ONE: C = C { re: 1.0, im: 0.0 };
#[inline]
fn new(re: f64, im: f64) -> C {
C { re, im }
}
#[inline]
fn add(self, o: C) -> C {
C { re: self.re + o.re, im: self.im + o.im }
}
#[inline]
fn sub(self, o: C) -> C {
C { re: self.re - o.re, im: self.im - o.im }
}
#[inline]
fn mul(self, o: C) -> C {
C { re: self.re * o.re - self.im * o.im, im: self.re * o.im + self.im * o.re }
}
#[inline]
fn conj(self) -> C {
C { re: self.re, im: -self.im }
}
#[inline]
fn scale(self, s: f64) -> C {
C { re: self.re * s, im: self.im * s }
}
#[inline]
fn norm2(self) -> f64 {
self.re * self.re + self.im * self.im
}
}
fn jacobi_svd(mut m: Vec<Vec<C>>, rows: usize, cols: usize) -> (Vec<Vec<C>>, Vec<f64>, Vec<Vec<C>>) {
let mut v: Vec<Vec<C>> = (0..cols)
.map(|j| (0..cols).map(|i| if i == j { C::ONE } else { C::ZERO }).collect())
.collect();
let eps = 1e-14;
for _sweep in 0..60 {
let mut off = 0.0f64;
for p in 0..cols {
for q in (p + 1)..cols {
let mut alpha = 0.0; let mut beta = 0.0; let mut gamma = C::ZERO; for k in 0..rows {
alpha += m[p][k].norm2();
beta += m[q][k].norm2();
gamma = gamma.add(m[p][k].conj().mul(m[q][k]));
}
let g2 = gamma.norm2();
off += g2;
if g2 <= eps * alpha * beta || g2 == 0.0 {
continue;
}
let gabs = g2.sqrt();
let ph = C::new(gamma.re / gabs, gamma.im / gabs); let cph = ph.conj();
for k in 0..rows {
m[q][k] = m[q][k].mul(cph);
}
for k in 0..cols {
v[q][k] = v[q][k].mul(cph);
}
let tau = (beta - alpha) / (2.0 * gabs);
let t = if tau >= 0.0 {
1.0 / (tau + (1.0 + tau * tau).sqrt())
} else {
-1.0 / (-tau + (1.0 + tau * tau).sqrt())
};
let cs = 1.0 / (1.0 + t * t).sqrt();
let sn = t * cs;
for k in 0..rows {
let x = m[p][k];
let y = m[q][k];
m[p][k] = x.scale(cs).sub(y.scale(sn));
m[q][k] = x.scale(sn).add(y.scale(cs));
}
for k in 0..cols {
let x = v[p][k];
let y = v[q][k];
v[p][k] = x.scale(cs).sub(y.scale(sn));
v[q][k] = x.scale(sn).add(y.scale(cs));
}
}
}
if off <= eps {
break;
}
}
let mut s = vec![0.0; cols];
for j in 0..cols {
s[j] = (0..rows).map(|k| m[j][k].norm2()).sum::<f64>().sqrt();
}
let mut u: Vec<Vec<C>> = m;
for j in 0..cols {
if s[j] > 1e-300 {
let inv = 1.0 / s[j];
for k in 0..rows {
u[j][k] = u[j][k].scale(inv);
}
}
}
let mut order: Vec<usize> = (0..cols).collect();
order.sort_by(|&a, &b| s[b].partial_cmp(&s[a]).unwrap_or(std::cmp::Ordering::Equal));
let s2: Vec<f64> = order.iter().map(|&i| s[i]).collect();
let u2: Vec<Vec<C>> = order.iter().map(|&i| u[i].clone()).collect();
let v2: Vec<Vec<C>> = order.iter().map(|&i| v[i].clone()).collect();
(u2, s2, v2)
}
pub struct Mps {
pub n: u8,
a: Vec<Vec<C>>,
dl: Vec<usize>,
dr: Vec<usize>,
chi_max: usize,
max_bond: usize,
retained: f64, }
#[inline]
fn idx(l: usize, s: usize, r: usize, dr: usize) -> usize {
(l * 2 + s) * dr + r
}
impl Mps {
pub fn zero(n: u8, chi_max: usize) -> Mps {
let a: Vec<Vec<C>> = (0..n).map(|_| vec![C::ONE, C::ZERO]).collect(); Mps {
n,
a,
dl: vec![1; n as usize],
dr: vec![1; n as usize],
chi_max: chi_max.max(1),
max_bond: 1,
retained: 1.0,
}
}
fn apply1(&mut self, i: usize, g: [[C; 2]; 2]) {
let (dl, dr) = (self.dl[i], self.dr[i]);
let ai = &mut self.a[i];
for l in 0..dl {
for r in 0..dr {
let a0 = ai[idx(l, 0, r, dr)];
let a1 = ai[idx(l, 1, r, dr)];
ai[idx(l, 0, r, dr)] = g[0][0].mul(a0).add(g[0][1].mul(a1));
ai[idx(l, 1, r, dr)] = g[1][0].mul(a0).add(g[1][1].mul(a1));
}
}
}
fn apply2_adjacent(&mut self, i: usize, g: [[C; 4]; 4]) {
let (dl, dc, dr) = (self.dl[i], self.dr[i], self.dr[i + 1]);
debug_assert_eq!(dc, self.dl[i + 1]);
let rows = dl * 2;
let cols = 2 * dr;
let mut mcols: Vec<Vec<C>> = vec![vec![C::ZERO; rows]; cols];
for l in 0..dl {
for r in 0..dr {
let mut th = [C::ZERO; 4];
for s1 in 0..2 {
for s2 in 0..2 {
let mut acc = C::ZERO;
for c in 0..dc {
acc = acc.add(self.a[i][idx(l, s1, c, dc)].mul(self.a[i + 1][idx(c, s2, r, dr)]));
}
th[s1 * 2 + s2] = acc;
}
}
for t1 in 0..2 {
for t2 in 0..2 {
let mut val = C::ZERO;
for s in 0..4 {
val = val.add(g[t1 * 2 + t2][s].mul(th[s]));
}
let row = l * 2 + t1;
let col = t2 * dr + r;
mcols[col][row] = val;
}
}
}
}
let (u, s, v) = jacobi_svd(mcols, rows, cols);
let total: f64 = s.iter().map(|x| x * x).sum();
let mut chi = 0;
for (k, sv) in s.iter().enumerate() {
if k >= self.chi_max || *sv <= 1e-12 * s[0].max(1e-300) {
break;
}
chi += 1;
}
let chi = chi.max(1).min(s.len());
let kept: f64 = s.iter().take(chi).map(|x| x * x).sum();
if total > 0.0 {
self.retained *= kept / total;
}
self.max_bond = self.max_bond.max(chi);
let mut na = vec![C::ZERO; dl * 2 * chi];
for l in 0..dl {
for t1 in 0..2 {
for k in 0..chi {
na[idx(l, t1, k, chi)] = u[k][l * 2 + t1];
}
}
}
let mut nb = vec![C::ZERO; chi * 2 * dr];
for k in 0..chi {
for t2 in 0..2 {
for r in 0..dr {
nb[idx(k, t2, r, dr)] = v[k][t2 * dr + r].conj().scale(s[k]);
}
}
}
self.a[i] = na;
self.dr[i] = chi;
self.a[i + 1] = nb;
self.dl[i + 1] = chi;
}
fn swap_adjacent(&mut self, i: usize) {
self.apply2_adjacent(i, swap_gate());
}
fn apply2(&mut self, a: usize, b: usize, g4: [[C; 4]; 4]) {
let (lo, hi) = (a.min(b), a.max(b));
if hi == lo + 1 {
self.apply2_adjacent(lo, g4);
return;
}
for k in (lo + 1..hi).rev() {
self.swap_adjacent(k);
}
self.apply2_adjacent(lo, g4);
for k in lo + 1..hi {
self.swap_adjacent(k);
}
}
pub fn to_statevector(&self) -> Vec<C> {
let dim = 1usize << self.n;
let mut out = vec![C::ZERO; dim];
for k in 0..dim {
let mut vec = vec![C::ONE]; for i in 0..self.n as usize {
let s = (k >> i) & 1;
let (dl, dr) = (self.dl[i], self.dr[i]);
let mut nv = vec![C::ZERO; dr];
for r in 0..dr {
let mut acc = C::ZERO;
for l in 0..dl {
acc = acc.add(vec[l].mul(self.a[i][idx(l, s, r, dr)]));
}
nv[r] = acc;
}
vec = nv;
}
out[k] = vec[0];
}
let norm: f64 = out.iter().map(|c| c.norm2()).sum::<f64>().sqrt();
if norm > 0.0 {
for c in out.iter_mut() {
*c = c.scale(1.0 / norm);
}
}
out
}
pub fn max_bond(&self) -> usize {
self.max_bond
}
pub fn retained_weight(&self) -> f64 {
self.retained
}
}
fn m1(base: BaseGate, param: u16) -> [[C; 2]; 2] {
let s = std::f64::consts::FRAC_1_SQRT_2;
let z = C::ZERO;
let o = C::ONE;
let i = C::new(0.0, 1.0);
let ni = C::new(0.0, -1.0);
let no = C::new(-1.0, 0.0);
let phase = |k: u16| {
let th = 2.0 * std::f64::consts::PI / (1u64 << k) as f64;
C::new(th.cos(), th.sin())
};
match base {
BaseGate::I => [[o, z], [z, o]],
BaseGate::X => [[z, o], [o, z]],
BaseGate::Y => [[z, ni], [i, z]],
BaseGate::Z => [[o, z], [z, no]],
BaseGate::H => [[C::new(s, 0.0), C::new(s, 0.0)], [C::new(s, 0.0), C::new(-s, 0.0)]],
BaseGate::S => [[o, z], [z, i]],
BaseGate::Sdg => [[o, z], [z, ni]],
BaseGate::T => [[o, z], [z, phase(3)]],
BaseGate::Tdg => [[o, z], [z, phase(3).conj()]],
BaseGate::P => [[o, z], [z, phase(param)]],
}
}
fn controlled4(control_is_lo: bool, u: [[C; 2]; 2]) -> [[C; 4]; 4] {
let mut g = [[C::ZERO; 4]; 4];
for a in 0..2 {
for b in 0..2 {
let inb = a * 2 + b;
let (ctrl, tgt) = if control_is_lo { (a, b) } else { (b, a) };
if ctrl == 0 {
g[inb][inb] = C::ONE; } else {
for tt in 0..2 {
let (oa, ob) = if control_is_lo { (a, tt) } else { (tt, b) };
let out = oa * 2 + ob;
g[out][inb] = u[tt][tgt];
}
}
}
}
g
}
fn swap_gate() -> [[C; 4]; 4] {
let mut g = [[C::ZERO; 4]; 4];
g[0][0] = C::ONE;
g[3][3] = C::ONE;
g[1][2] = C::ONE; g[2][1] = C::ONE;
g
}
fn decompose_cc(g: &Gate) -> Option<Vec<Gate>> {
if g.controls.len() != 2 {
return None;
}
let (a, b, c) = (g.controls[0], g.controls[1], g.target);
let g1 = |base, t| Gate { base, controls: vec![], target: t, param: 0 };
let cx = |ctrl, t| Gate { base: BaseGate::X, controls: vec![ctrl], target: t, param: 0 };
let mut ops = Vec::new();
if g.base == BaseGate::Z {
ops.push(g1(BaseGate::H, c));
}
ops.push(g1(BaseGate::H, c));
ops.push(cx(b, c));
ops.push(g1(BaseGate::Tdg, c));
ops.push(cx(a, c));
ops.push(g1(BaseGate::T, c));
ops.push(cx(b, c));
ops.push(g1(BaseGate::Tdg, c));
ops.push(cx(a, c));
ops.push(g1(BaseGate::T, b));
ops.push(g1(BaseGate::T, c));
ops.push(cx(a, b));
ops.push(g1(BaseGate::H, c));
ops.push(g1(BaseGate::T, a));
ops.push(g1(BaseGate::Tdg, b));
ops.push(cx(a, b));
if g.base == BaseGate::Z {
ops.push(g1(BaseGate::H, c));
}
Some(ops)
}
pub struct MpsRun {
pub max_bond: usize,
pub retained_weight: f64,
pub fidelity_vs_dense: Option<f64>,
}
#[derive(Debug)]
pub enum MpsError {
Unsupported(String),
}
pub fn run(circuit: &Circuit, chi_max: usize) -> Result<(Mps, MpsRun), MpsError> {
let mut mps = Mps::zero(circuit.n_qubits, chi_max);
for g in &circuit.ops {
apply_gate(&mut mps, g)?;
}
let run = MpsRun {
max_bond: mps.max_bond(),
retained_weight: mps.retained_weight(),
fidelity_vs_dense: None,
};
Ok((mps, run))
}
fn apply_gate(mps: &mut Mps, g: &Gate) -> Result<(), MpsError> {
match g.controls.len() {
0 => mps.apply1(g.target as usize, m1(g.base, g.param)),
1 => {
let ctrl = g.controls[0] as usize;
let tgt = g.target as usize;
let control_is_lo = ctrl < tgt;
let g4 = controlled4(control_is_lo, m1(g.base, g.param));
mps.apply2(ctrl, tgt, g4);
}
2 => {
let ops = decompose_cc(g)
.ok_or_else(|| MpsError::Unsupported(format!("{:?} with 2 controls", g.base)))?;
for sub in &ops {
apply_gate(mps, sub)?;
}
}
k => return Err(MpsError::Unsupported(format!("{k} controls"))),
}
Ok(())
}
pub fn fidelity(mps: &Mps, dense: &StateVector) -> f64 {
let a = mps.to_statevector();
let scale = (1u64 << crate::quantum::FRAC) as f64;
let mut overlap = C::ZERO;
let mut nd = 0.0;
for (k, amp) in dense.amps.iter().enumerate() {
let d = C::new(amp.re as f64 / scale, amp.im as f64 / scale);
overlap = overlap.add(a[k].conj().mul(d));
nd += d.norm2();
}
if nd <= 0.0 {
return 0.0;
}
overlap.norm2() / nd }
#[cfg(test)]
mod tests {
use super::*;
fn approx(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
fn ghz(n: u8) -> Circuit {
let mut c = Circuit::new(n);
c.h(0);
for i in 1..n {
c.cx(i - 1, i);
}
c
}
fn bell() -> Circuit {
let mut c = Circuit::new(2);
c.h(0);
c.cx(0, 1);
c
}
fn entangling(n: u8) -> Circuit {
let mut c = Circuit::new(n);
for _ in 0..3 {
for q in 0..n {
c.h(q);
c.t(q);
}
let mut q = 0;
while q + 1 < n {
c.cx(q, q + 1);
q += 2;
}
let mut q = 1;
while q + 1 < n {
c.cx(q, q + 1);
q += 2;
}
}
c
}
#[test]
fn product_state_stays_bond_one() {
let mut c = Circuit::new(5);
for q in 0..5 {
c.h(q);
}
let (mps, run) = run(&c, 16).unwrap();
assert_eq!(run.max_bond, 1, "product state must stay bond-1");
assert!(approx(fidelity(&mps, &c.simulate().unwrap()), 1.0, 1e-9));
}
#[test]
fn ghz_is_exact_at_bond_two() {
let c = ghz(6);
let (mps, run) = run(&c, 8).unwrap();
assert_eq!(run.max_bond, 2, "GHZ is a bond-2 state");
assert!(approx(fidelity(&mps, &c.simulate().unwrap()), 1.0, 1e-9), "GHZ must be exact");
}
#[test]
fn bell_and_small_circuits_match_dense() {
for c in [bell(), ghz(4)] {
let (mps, _) = run(&c, 16).unwrap();
let f = fidelity(&mps, &c.simulate().unwrap());
assert!(approx(f, 1.0, 1e-8), "MPS should match dense, got {f}");
}
}
#[test]
fn qft_matches_dense_at_full_bond() {
let c = Circuit::qft(6);
let (mps, _) = run(&c, 64).unwrap();
let f = fidelity(&mps, &c.simulate().unwrap());
assert!(approx(f, 1.0, 1e-7), "QFT full-bond fidelity {f}");
}
#[test]
fn ccz_decomposition_matches_dense() {
let mut c = Circuit::new(3);
c.h(0);
c.h(1);
c.h(2);
c.ops.push(Gate { base: BaseGate::Z, controls: vec![0, 1], target: 2, param: 0 });
let (mps, _) = run(&c, 32).unwrap();
let f = fidelity(&mps, &c.simulate().unwrap());
assert!(approx(f, 1.0, 1e-6), "CCZ decomposition fidelity {f}");
}
#[test]
fn nonadjacent_two_qubit_gate() {
let mut c = Circuit::new(5);
c.h(0);
c.cx(0, 4);
let (mps, _) = run(&c, 8).unwrap();
assert!(approx(fidelity(&mps, &c.simulate().unwrap()), 1.0, 1e-8));
}
#[test]
fn truncation_is_measured_not_hidden() {
let c = entangling(6);
let dense = c.simulate().unwrap();
let (mps, st) = run(&c, 2).unwrap();
assert!(st.retained_weight < 0.999, "squeezed bond must report lost weight: {}", st.retained_weight);
assert!(st.retained_weight > 0.0, "some weight is kept");
assert!(st.max_bond <= 2, "bond cap respected");
assert!(fidelity(&mps, &dense) < 0.999, "aggressive truncation must measurably reduce fidelity");
let (mps2, st2) = run(&c, 8).unwrap();
assert!(approx(st2.retained_weight, 1.0, 1e-6), "full bond keeps all weight: {}", st2.retained_weight);
assert!(approx(fidelity(&mps2, &dense), 1.0, 1e-6));
}
#[test]
fn deterministic() {
let c = Circuit::qft(6);
let a = run(&c, 8).unwrap().0.to_statevector();
let b = run(&c, 8).unwrap().0.to_statevector();
assert!(a.iter().zip(&b).all(|(x, y)| x.re.to_bits() == y.re.to_bits() && x.im.to_bits() == y.im.to_bits()),
"MPS must be bit-reproducible");
}
}