use crate::quantum::{fxmul, Amp, StateVector, FRAC, ONE};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct DensityMatrix {
pub n_qubits: u8,
pub d: usize,
m: Vec<Amp>,
}
impl DensityMatrix {
pub fn from_entries(m: Vec<Amp>) -> Option<DensityMatrix> {
let d = (m.len() as f64).sqrt() as usize;
if d * d != m.len() || !d.is_power_of_two() {
return None;
}
Some(DensityMatrix { n_qubits: d.trailing_zeros() as u8, d, m })
}
pub fn expect_pauli(&self, ops: &[(u8, Pauli)]) -> i64 {
let flip = pauli_flip(ops);
let mut acc: i128 = 0;
for j in 0..self.d {
acc += pauli_coeff(j, ops).mul(self.get(j, j ^ flip)).re as i128;
}
acc as i64
}
pub fn get(&self, r: usize, c: usize) -> Amp {
self.m[r * self.d + c]
}
pub fn from_pure(sv: &StateVector) -> DensityMatrix {
let d = sv.amps.len();
let mut m = vec![Amp::ZERO; d * d];
for r in 0..d {
for c in 0..d {
m[r * d + c] = sv.amps[r].mul(sv.amps[c].conj());
}
}
DensityMatrix { n_qubits: sv.n_qubits, d, m }
}
pub fn mixture(parts: &[(i64, DensityMatrix)]) -> Option<DensityMatrix> {
let first = parts.first()?;
let (n, d) = (first.1.n_qubits, first.1.d);
if parts.iter().any(|(_, p)| p.d != d) {
return None;
}
let mut m = vec![Amp::ZERO; d * d];
for (w, p) in parts {
for k in 0..d * d {
m[k] = m[k].add(Amp { re: fxmul(*w, p.m[k].re), im: fxmul(*w, p.m[k].im) });
}
}
Some(DensityMatrix { n_qubits: n, d, m })
}
pub fn trace(&self) -> Amp {
(0..self.d).fold(Amp::ZERO, |acc, i| acc.add(self.get(i, i)))
}
pub fn purity(&self) -> i64 {
let mut acc: i128 = 0;
for r in 0..self.d {
for c in 0..self.d {
let (a, b) = (self.get(r, c), self.get(c, r));
acc += (a.re as i128 * b.re as i128 - a.im as i128 * b.im as i128) >> FRAC;
}
}
acc as i64
}
pub fn partial_trace(&self, keep: &[u8]) -> DensityMatrix {
let mut keep: Vec<u8> = keep.to_vec();
keep.sort_unstable();
keep.dedup();
let traced: Vec<u8> = (0..self.n_qubits).filter(|q| !keep.contains(q)).collect();
let dk = 1usize << keep.len();
let dt = 1usize << traced.len();
let widen = |bits: usize, qs: &[u8]| -> usize {
qs.iter().enumerate().fold(0usize, |acc, (i, &q)| {
acc | (((bits >> i) & 1) << q)
})
};
let mut m = vec![Amp::ZERO; dk * dk];
for a in 0..dk {
for b in 0..dk {
let (ka, kb) = (widen(a, &keep), widen(b, &keep));
let mut sum = Amp::ZERO;
for t in 0..dt {
let tt = widen(t, &traced);
sum = sum.add(self.get(ka | tt, kb | tt));
}
m[a * dk + b] = sum;
}
}
DensityMatrix { n_qubits: keep.len() as u8, d: dk, m }
}
pub fn hash(&self) -> [u8; 32] {
let mut h = blake3::Hasher::new();
h.update(b"wai:density-matrix\x01");
h.update(&(self.n_qubits as u64).to_le_bytes());
for a in &self.m {
h.update(&a.re.to_le_bytes());
h.update(&a.im.to_le_bytes());
}
*h.finalize().as_bytes()
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Pauli {
I,
X,
Y,
Z,
}
pub(crate) fn pauli_flip(ops: &[(u8, Pauli)]) -> usize {
let mut flip = 0usize;
for (q, p) in ops {
if matches!(p, Pauli::X | Pauli::Y) {
flip |= 1 << q;
}
}
flip
}
pub(crate) fn pauli_coeff(i: usize, ops: &[(u8, Pauli)]) -> Amp {
let mut sign = 1i64;
let mut i_pow = 0u32;
for (q, p) in ops {
let bit = (i >> q) & 1;
match p {
Pauli::I | Pauli::X => {}
Pauli::Z => {
if bit == 1 {
sign = -sign;
}
}
Pauli::Y => {
i_pow += 1;
if bit == 1 {
sign = -sign;
}
}
}
}
match i_pow % 4 {
0 => Amp { re: sign * ONE, im: 0 },
1 => Amp { re: 0, im: sign * ONE },
2 => Amp { re: -sign * ONE, im: 0 },
_ => Amp { re: 0, im: -sign * ONE },
}
}
pub fn expect_pauli(sv: &StateVector, ops: &[(u8, Pauli)]) -> i64 {
let flip = pauli_flip(ops);
let mut acc: i128 = 0;
for i in 0..sv.amps.len() {
let c = pauli_coeff(i, ops);
acc += sv.amps[i ^ flip].conj().mul(c.mul(sv.amps[i])).re as i128;
}
acc as i64
}
pub struct Chsh {
pub s: f64,
pub classical_bound: f64,
pub tsirelson_bound: f64,
pub zz_fx: i64,
pub xx_fx: i64,
}
pub fn chsh(sv: &StateVector) -> Option<Chsh> {
if sv.n_qubits != 2 {
return None;
}
let zz = expect_pauli(sv, &[(0, Pauli::Z), (1, Pauli::Z)]);
let xx = expect_pauli(sv, &[(0, Pauli::X), (1, Pauli::X)]);
let s = std::f64::consts::SQRT_2 * ((zz + xx) as f64 / ONE as f64);
Some(Chsh {
s,
classical_bound: 2.0,
tsirelson_bound: 2.0 * std::f64::consts::SQRT_2,
zz_fx: zz,
xx_fx: xx,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::quantum::Circuit;
fn bell() -> StateVector {
let mut c = Circuit::new(2);
c.h(0).cx(0, 1);
c.simulate().unwrap()
}
#[test]
fn a_pure_state_is_pure_and_normalised() {
let dm = DensityMatrix::from_pure(&bell());
let tol = ONE / 10_000;
assert!((dm.trace().re - ONE).abs() < tol, "trace {}", dm.trace().re);
assert!(dm.trace().im.abs() < tol);
assert!((dm.purity() - ONE).abs() < tol, "purity {}", dm.purity());
}
#[test]
fn half_of_a_bell_pair_is_maximally_mixed() {
let dm = DensityMatrix::from_pure(&bell());
let red = dm.partial_trace(&[0]);
let tol = ONE / 10_000;
assert_eq!(red.d, 2);
assert!((red.trace().re - ONE).abs() < tol);
assert!((red.purity() - ONE / 2).abs() < tol, "purity {}", red.purity());
assert!(red.get(0, 1).re.abs() < tol && red.get(0, 1).im.abs() < tol);
assert!((red.get(0, 0).re - ONE / 2).abs() < tol);
assert!((red.get(1, 1).re - ONE / 2).abs() < tol);
}
#[test]
fn a_product_state_leaves_a_pure_subsystem() {
let mut c = Circuit::new(2);
c.h(0); let dm = DensityMatrix::from_pure(&c.simulate().unwrap());
let red = dm.partial_trace(&[0]);
let tol = ONE / 10_000;
assert!((red.purity() - ONE).abs() < tol, "a product state's part is pure: {}", red.purity());
}
#[test]
fn pauli_expectations_of_a_bell_state() {
let sv = bell();
let tol = ONE / 10_000;
assert!((expect_pauli(&sv, &[(0, Pauli::Z), (1, Pauli::Z)]) - ONE).abs() < tol);
assert!((expect_pauli(&sv, &[(0, Pauli::X), (1, Pauli::X)]) - ONE).abs() < tol);
assert!((expect_pauli(&sv, &[(0, Pauli::Y), (1, Pauli::Y)]) + ONE).abs() < tol);
assert!(expect_pauli(&sv, &[(0, Pauli::Z)]).abs() < tol);
assert!(expect_pauli(&sv, &[(1, Pauli::X)]).abs() < tol);
assert!((expect_pauli(&sv, &[(0, Pauli::I)]) - ONE).abs() < tol);
}
#[test]
fn chsh_violates_the_classical_bound_and_saturates_tsirelson() {
let r = chsh(&bell()).expect("two qubits");
assert!(r.s > r.classical_bound + 0.5, "S = {} must beat the classical 2", r.s);
assert!(r.s <= r.tsirelson_bound + 1e-6, "S = {} may not beat Tsirelson", r.s);
assert!((r.s - 2.0 * std::f64::consts::SQRT_2).abs() < 1e-3, "S = {} should reach 2√2", r.s);
}
#[test]
fn a_product_state_obeys_the_classical_bound() {
let mut c = Circuit::new(2);
c.h(0).h(1);
let r = chsh(&c.simulate().unwrap()).unwrap();
assert!(r.s <= r.classical_bound + 1e-6, "S = {} from a product state", r.s);
}
#[test]
fn mixing_a_state_with_itself_returns_it() {
let pure = DensityMatrix::from_pure(&bell());
let mixed = DensityMatrix::mixture(&[(ONE / 2, pure.clone()), (ONE / 2, pure.clone())]).unwrap();
let tol = ONE / 100_000;
for r in 0..pure.d {
for c in 0..pure.d {
let (a, b) = (pure.get(r, c), mixed.get(r, c));
assert!(
(a.re - b.re).abs() <= tol && (a.im - b.im).abs() <= tol,
"entry ({r},{c}) drifted: {a:?} vs {b:?}"
);
}
}
assert!((mixed.trace().re - ONE).abs() < ONE / 10_000, "a mixture is still a state");
}
#[test]
fn the_hash_identifies_the_mixture() {
let a = DensityMatrix::from_pure(&bell());
let b = DensityMatrix::from_pure(&bell());
assert_eq!(a.hash(), b.hash(), "same computation, same identity");
let mut c = Circuit::new(2);
c.h(0); assert_ne!(DensityMatrix::from_pure(&c.simulate().unwrap()).hash(), a.hash());
}
}