use crate::quantum::{sqrt_fx, Amp, Gateset, FRAC, ONE};
pub fn omega(d: u8, k: u32) -> Option<Amp> {
match d {
2 => Some(if k % 2 == 0 { Amp::ONE } else { Amp { re: -ONE, im: 0 } }),
3 => {
let s3_2 = sqrt_fx(3 * (ONE / 4));
Some(match k % 3 {
0 => Amp::ONE,
1 => Amp { re: -(ONE / 2), im: s3_2 },
_ => Amp { re: -(ONE / 2), im: -s3_2 },
})
}
_ => None,
}
}
fn inv_sqrt_d(d: u8) -> Option<i64> {
match d {
2 => Some(Gateset::V2.inv_sqrt2()),
3 => Some(sqrt_fx(ONE / 3)),
_ => None,
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct QuditState {
pub d: u8,
pub n: u8,
pub amps: Vec<Amp>,
}
impl QuditState {
pub fn new(d: u8, n: u8) -> Option<QuditState> {
omega(d, 1)?;
let dim = (d as usize).checked_pow(n as u32)?;
let mut amps = vec![Amp::ZERO; dim];
amps[0] = Amp::ONE;
Some(QuditState { d, n, amps })
}
pub fn dim(&self) -> usize {
self.amps.len()
}
pub fn digit(&self, index: usize, q: u8) -> u8 {
((index / (self.d as usize).pow(q as u32)) % self.d as usize) as u8
}
fn stride(&self, q: u8) -> usize {
(self.d as usize).pow(q as u32)
}
fn for_each_fiber(&self, q: u8, mut f: impl FnMut(usize, usize)) {
let (d, stride) = (self.d as usize, self.stride(q));
let block = stride * d;
let mut base = 0usize;
while base < self.dim() {
for lo in 0..stride {
f(base + lo, stride);
}
base += block;
}
}
pub fn shift(&mut self, q: u8, a: u8) {
let (d, a) = (self.d as usize, a as usize % self.d as usize);
if a == 0 {
return;
}
let mut out = vec![Amp::ZERO; self.dim()];
let fibers: Vec<(usize, usize)> = {
let mut v = Vec::new();
self.for_each_fiber(q, |b, s| v.push((b, s)));
v
};
for (b, s) in fibers {
for j in 0..d {
out[b + ((j + a) % d) * s] = self.amps[b + j * s];
}
}
self.amps = out;
}
pub fn clock(&mut self, q: u8, b: u8) {
let d = self.d;
for i in 0..self.dim() {
let j = self.digit(i, q) as u32;
let w = omega(d, (b as u32).wrapping_mul(j)).expect("dimension checked at construction");
self.amps[i] = w.mul(self.amps[i]);
}
}
pub fn fourier(&mut self, q: u8) {
let d = self.d as usize;
let inv = inv_sqrt_d(self.d).expect("dimension checked at construction");
let fibers: Vec<(usize, usize)> = {
let mut v = Vec::new();
self.for_each_fiber(q, |b, s| v.push((b, s)));
v
};
for (b, s) in fibers {
let src: Vec<Amp> = (0..d).map(|j| self.amps[b + j * s]).collect();
for k in 0..d {
let mut acc = Amp::ZERO;
for (j, a) in src.iter().enumerate() {
let w = omega(self.d, (j * k) as u32).expect("checked");
acc = acc.add(w.mul(*a));
}
self.amps[b + k * s] = Amp {
re: crate::quantum::fxmul(inv, acc.re),
im: crate::quantum::fxmul(inv, acc.im),
};
}
}
}
pub fn csum(&mut self, control: u8, target: u8) {
if control == target {
return;
}
let d = self.d as usize;
let mut out = vec![Amp::ZERO; self.dim()];
let (sc, st) = (self.stride(control), self.stride(target));
for i in 0..self.dim() {
let c = (i / sc) % d;
let t = (i / st) % d;
let t2 = (t + c) % d;
let j = i + (t2 * st) - (t * st);
out[j] = self.amps[i];
}
self.amps = out;
}
pub fn norm2_fx(&self) -> i64 {
let acc: i128 = self.amps.iter().map(|a| a.norm2()).sum();
(acc >> FRAC) as i64
}
pub fn probabilities(&self) -> Vec<f64> {
self.amps
.iter()
.map(|a| (a.norm2() >> FRAC) as f64 / ONE as f64)
.collect()
}
pub fn hash(&self) -> [u8; 32] {
let mut h = blake3::Hasher::new();
h.update(b"wai:qudit-state\x01");
h.update(&[self.d, self.n]);
for a in &self.amps {
h.update(&a.re.to_le_bytes());
h.update(&a.im.to_le_bytes());
}
*h.finalize().as_bytes()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::quantum::Circuit;
const TOL: i64 = ONE / 10_000;
#[test]
fn unsupported_dimensions_are_refused_not_approximated() {
assert!(omega(5, 1).is_none(), "d=5 roots are not exactly constructible here");
assert!(QuditState::new(4, 2).is_none());
assert!(QuditState::new(3, 2).is_some());
}
#[test]
fn omega_is_a_cube_root_of_unity() {
let w = omega(3, 1).unwrap();
let w3 = w.mul(w).mul(w);
assert!((w3.re - ONE).abs() < TOL && w3.im.abs() < TOL, "ω³ = {w3:?}, want 1");
let s = Amp::ONE.add(omega(3, 1).unwrap()).add(omega(3, 2).unwrap());
assert!(s.re.abs() < TOL && s.im.abs() < TOL, "1+ω+ω² = {s:?}, want 0");
}
#[test]
fn at_d_equals_two_it_is_the_qubit_simulator() {
let mut q = QuditState::new(2, 2).unwrap();
q.fourier(0);
q.csum(0, 1);
let mut c = Circuit::with_gateset(2, Gateset::V2);
c.h(0).cx(0, 1);
let sv = c.simulate().unwrap();
assert_eq!(q.amps, sv.amps, "d=2 must agree with the qubit simulator amplitude for amplitude");
}
#[test]
fn shift_and_clock_have_order_three() {
for start in 0..3u8 {
let mut q = QuditState::new(3, 1).unwrap();
q.shift(0, start);
let before = q.clone();
q.shift(0, 1);
q.shift(0, 1);
q.shift(0, 1);
assert_eq!(q, before, "X³ = I on a qutrit");
q.clock(0, 1);
q.clock(0, 1);
q.clock(0, 1);
let same = q.amps.iter().zip(&before.amps).all(|(a, b)| {
(a.re - b.re).abs() < TOL && (a.im - b.im).abs() < TOL
});
assert!(same, "Z³ = I on a qutrit");
}
}
#[test]
fn fourier_makes_a_flat_superposition_of_three() {
let mut q = QuditState::new(3, 1).unwrap();
q.fourier(0);
let p = q.probabilities();
assert_eq!(p.len(), 3);
for (k, pk) in p.iter().enumerate() {
assert!((pk - 1.0 / 3.0).abs() < 1e-3, "outcome {k} has p={pk}, want 1/3");
}
assert!((q.norm2_fx() - ONE).abs() < TOL, "still normalised");
}
#[test]
fn a_qutrit_bell_state_has_three_correlated_outcomes() {
let mut q = QuditState::new(3, 2).unwrap();
q.fourier(0);
q.csum(0, 1);
let p = q.probabilities();
assert_eq!(p.len(), 9);
for i in 0..9usize {
let (a, b) = (i % 3, i / 3);
if a == b {
assert!((p[i] - 1.0 / 3.0).abs() < 1e-3, "|{a}{b}> should be 1/3, got {}", p[i]);
} else {
assert!(p[i] < 1e-6, "|{a}{b}> should be impossible, got {}", p[i]);
}
}
assert!((q.norm2_fx() - ONE).abs() < TOL);
}
#[test]
fn clock_and_shift_obey_the_weyl_relation() {
let mut zx = QuditState::new(3, 1).unwrap();
zx.fourier(0); let mut xz = zx.clone();
zx.shift(0, 1);
zx.clock(0, 1);
xz.clock(0, 1);
xz.shift(0, 1); for a in xz.amps.iter_mut() {
*a = omega(3, 1).unwrap().mul(*a); }
let same = zx.amps.iter().zip(&xz.amps).all(|(a, b)| {
(a.re - b.re).abs() < TOL && (a.im - b.im).abs() < TOL
});
assert!(same, "ZX must equal ωXZ:\n{:?}\n{:?}", zx.amps, xz.amps);
}
#[test]
fn deterministic_and_hashable() {
let build = || {
let mut q = QuditState::new(3, 3).unwrap();
q.fourier(0);
q.csum(0, 1);
q.csum(1, 2);
q.clock(2, 2);
q
};
assert_eq!(build().hash(), build().hash());
assert!((build().norm2_fx() - ONE).abs() < TOL);
}
}