use crate::quantum_frame::Detections;
use crate::quantum_match::{MatchError, MatchingGraph};
#[derive(Clone, Debug)]
struct RadixHeap {
last: u64,
buckets: Vec<Vec<(u64, u32, u32)>>,
}
impl Default for RadixHeap {
fn default() -> RadixHeap {
RadixHeap { last: 0, buckets: vec![Vec::new(); 65] }
}
}
impl RadixHeap {
fn bucket(&self, key: u64) -> usize {
(64 - (key ^ self.last).leading_zeros()) as usize
}
fn push(&mut self, key: u64, e: u32, version: u32) {
debug_assert!(key >= self.last);
let b = self.bucket(key);
self.buckets[b].push((key, e, version));
}
fn pop(&mut self) -> Option<(u64, u32, u32)> {
if self.buckets[0].is_empty() {
let i = (1..65).find(|&i| !self.buckets[i].is_empty())?;
let moved = core::mem::take(&mut self.buckets[i]);
self.last = moved.iter().map(|m| m.0).min().expect("non-empty");
for m in &moved {
let b = self.bucket(m.0);
self.buckets[b].push(*m);
}
let mut moved = moved;
moved.clear();
if self.buckets[i].is_empty() {
self.buckets[i] = moved;
}
}
self.buckets[0].pop()
}
fn clear(&mut self) {
for b in &mut self.buckets {
b.clear();
}
self.last = 0;
}
}
#[derive(Clone, Debug)]
pub struct UnionFind {
nodes: usize,
start: Vec<usize>,
nbr: Vec<u32>,
edge: Vec<u32>,
ends: Vec<(u32, u32)>,
cap: Vec<i64>,
obs: Vec<u64>,
}
#[derive(Clone, Debug, Default)]
pub struct Workspace {
parent: Vec<u32>,
odd: Vec<bool>,
boundary: Vec<bool>,
members: Vec<Vec<u32>>,
frontier: Vec<Vec<u32>>,
in_use: Vec<bool>,
touched: Vec<u32>,
growth: Vec<i64>,
tau: Vec<i64>,
speed: Vec<i64>,
version: Vec<u32>,
full: Vec<bool>,
edge_in_use: Vec<bool>,
edges_touched: Vec<u32>,
events: RadixHeap,
defect: Vec<bool>,
order: Vec<u32>,
via: Vec<u32>,
seen: Vec<bool>,
}
impl UnionFind {
pub fn new(g: &MatchingGraph) -> UnionFind {
let nodes = g.detectors() + 1;
let mut ids: std::collections::BTreeMap<(u32, u32), u32> = std::collections::BTreeMap::new();
let (mut ends, mut cap, mut obs) = (Vec::new(), Vec::new(), Vec::new());
let mut start = vec![0];
let (mut nbr, mut edge) = (Vec::new(), Vec::new());
for u in 0..nodes {
for &(v, w, o) in g.neighbours(u) {
let key = ((u as u32).min(v), (u as u32).max(v));
let id = *ids.entry(key).or_insert_with(|| {
ends.push(key);
cap.push(2 * w.max(0));
obs.push(o);
(ends.len() - 1) as u32
});
nbr.push(v);
edge.push(id);
}
start.push(nbr.len());
}
UnionFind { nodes, start, nbr, edge, ends, cap, obs }
}
pub fn workspace(&self) -> Workspace {
let (n, e) = (self.nodes, self.ends.len());
Workspace {
parent: (0..n as u32).collect(),
odd: vec![false; n],
boundary: vec![false; n],
members: vec![Vec::new(); n],
frontier: vec![Vec::new(); n],
in_use: vec![false; n],
touched: Vec::new(),
growth: vec![0; e],
tau: vec![0; e],
speed: vec![0; e],
version: vec![0; e],
full: vec![false; e],
edge_in_use: vec![false; e],
edges_touched: Vec::new(),
events: RadixHeap::default(),
defect: vec![false; n],
order: Vec::new(),
via: vec![u32::MAX; n],
seen: vec![false; n],
}
}
fn find(ws: &mut Workspace, mut u: u32) -> u32 {
while ws.parent[u as usize] != u {
let p = ws.parent[u as usize];
ws.parent[u as usize] = ws.parent[p as usize];
u = ws.parent[u as usize];
}
u
}
fn growing(ws: &Workspace, root: u32) -> bool {
ws.odd[root as usize] && !ws.boundary[root as usize]
}
fn touch(&self, ws: &mut Workspace, u: u32) {
let i = u as usize;
if ws.in_use[i] {
return;
}
ws.in_use[i] = true;
ws.touched.push(u);
ws.parent[i] = u;
ws.odd[i] = false;
ws.boundary[i] = i == self.nodes - 1;
ws.members[i].clear();
ws.members[i].push(u);
ws.defect[i] = false;
ws.frontier[i].clear();
if !ws.boundary[i] {
ws.frontier[i].extend_from_slice(&self.edge[self.start[i]..self.start[i + 1]]);
}
}
fn reschedule(&self, ws: &mut Workspace, e: u32, now: i64) -> bool {
let ei = e as usize;
if !ws.edge_in_use[ei] {
ws.edge_in_use[ei] = true;
ws.edges_touched.push(e);
ws.growth[ei] = 0;
ws.tau[ei] = now;
ws.speed[ei] = 0;
ws.version[ei] = 0;
ws.full[ei] = false;
}
if ws.full[ei] {
return false;
}
let (a, b) = self.ends[ei];
let ra = if ws.in_use[a as usize] { Some(Self::find(ws, a)) } else { None };
let rb = if ws.in_use[b as usize] { Some(Self::find(ws, b)) } else { None };
let internal = ra.is_some() && ra == rb;
let speed = if internal { 0 } else { i64::from(ra.is_some_and(|r| Self::growing(ws, r))) + i64::from(rb.is_some_and(|r| Self::growing(ws, r))) };
if speed == ws.speed[ei] && ws.version[ei] != 0 {
return !internal;
}
ws.growth[ei] += ws.speed[ei] * (now - ws.tau[ei]);
ws.tau[ei] = now;
ws.version[ei] = ws.version[ei].wrapping_add(1).max(1);
ws.speed[ei] = speed;
if speed > 0 {
let rem = self.cap[ei] - ws.growth[ei];
let at = if rem <= 0 { now } else { now + (rem + speed - 1) / speed };
ws.events.push(at as u64, e, ws.version[ei]);
}
!internal
}
fn reschedule_frontier(&self, ws: &mut Workspace, root: u32, now: i64) {
let mut list = core::mem::take(&mut ws.frontier[root as usize]);
list.retain(|&e| self.reschedule(ws, e, now));
ws.frontier[root as usize] = list;
}
fn reset(ws: &mut Workspace) {
for &u in &ws.touched {
let i = u as usize;
ws.in_use[i] = false;
ws.parent[i] = u;
ws.members[i].clear();
ws.frontier[i].clear();
ws.defect[i] = false;
ws.seen[i] = false;
ws.via[i] = u32::MAX;
}
ws.touched.clear();
for &e in &ws.edges_touched {
ws.edge_in_use[e as usize] = false;
}
ws.edges_touched.clear();
ws.events.clear();
}
pub fn decode(&self, fired: &[u32], ws: &mut Workspace) -> Result<u64, MatchError> {
Self::reset(ws);
for &d in fired {
self.touch(ws, d);
ws.odd[d as usize] ^= true;
ws.defect[d as usize] ^= true;
}
for &d in fired {
if Self::growing(ws, d) {
self.reschedule_frontier(ws, d, 0);
}
}
while let Some((at, e, ver)) = ws.events.pop() {
let at = at as i64;
let ei = e as usize;
if ws.full[ei] || ws.version[ei] != ver {
continue;
}
ws.full[ei] = true;
ws.growth[ei] = self.cap[ei];
ws.speed[ei] = 0;
let (a, b) = self.ends[ei];
self.touch(ws, a);
self.touch(ws, b);
let (ra, rb) = (Self::find(ws, a), Self::find(ws, b));
if ra == rb {
continue;
}
let (ga, gb) = (Self::growing(ws, ra), Self::growing(ws, rb));
let (big, small) = if ws.members[ra as usize].len() >= ws.members[rb as usize].len() { (ra, rb) } else { (rb, ra) };
ws.parent[small as usize] = big;
let moved = core::mem::take(&mut ws.members[small as usize]);
ws.members[big as usize].extend(moved);
ws.odd[big as usize] ^= ws.odd[small as usize];
ws.boundary[big as usize] |= ws.boundary[small as usize];
let g = Self::growing(ws, big);
let mut fs = core::mem::take(&mut ws.frontier[small as usize]);
let (g_big, g_small) = if big == ra { (ga, gb) } else { (gb, ga) };
if g_big != g {
self.reschedule_frontier(ws, big, at);
}
if g_small != g {
fs.retain(|&f| self.reschedule(ws, f, at));
}
ws.frontier[big as usize].extend(fs);
}
for k in 0..ws.touched.len() {
let u = ws.touched[k];
if Self::find(ws, u) == u && Self::growing(ws, u) {
return Err(MatchError::Unmatchable);
}
}
Ok(self.peel(ws))
}
fn peel(&self, ws: &mut Workspace) -> u64 {
let boundary = (self.nodes - 1) as u32;
let mut out = 0u64;
let mut roots: Vec<u32> = Vec::new();
for k in 0..ws.touched.len() {
let u = ws.touched[k];
if Self::find(ws, u) == u {
roots.push(u);
}
}
roots.sort_unstable();
for r in roots {
let start_node = if ws.boundary[r as usize] { boundary } else { r };
ws.order.clear();
ws.order.push(start_node);
ws.seen[start_node as usize] = true;
let mut head = 0;
while head < ws.order.len() {
let u = ws.order[head] as usize;
head += 1;
for k in self.start[u]..self.start[u + 1] {
let (v, e) = (self.nbr[k] as usize, self.edge[k] as usize);
if !ws.in_use[v] || ws.seen[v] || !ws.edge_in_use[e] || !ws.full[e] {
continue;
}
if Self::find(ws, v as u32) != r {
continue;
}
ws.seen[v] = true;
ws.via[v] = e as u32;
ws.order.push(v as u32);
}
}
for idx in (1..ws.order.len()).rev() {
let u = ws.order[idx];
if ws.defect[u as usize] {
let e = ws.via[u as usize] as usize;
let (a, b) = self.ends[e];
let parent = if a == u { b } else { a };
out ^= self.obs[e];
ws.defect[u as usize] = false;
ws.defect[parent as usize] ^= true;
}
}
}
out
}
pub fn failures(&self, shots: &Detections) -> Result<u64, MatchError> {
let mut ws = self.workspace();
let mut n = 0;
for s in 0..shots.shots {
if self.decode(&shots.fired(s), &mut ws)? != shots.flips(s) {
n += 1;
}
}
Ok(n)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[derive(Clone, Debug)]
struct RoundBased {
nodes: usize,
start: Vec<usize>,
nbr: Vec<u32>,
weight: Vec<i64>,
obs: Vec<u64>,
edge: Vec<u32>,
edges: usize,
}
#[derive(Clone, Debug, Default)]
struct RoundWorkspace {
parent: Vec<u32>,
odd: Vec<bool>,
boundary: Vec<bool>,
members: Vec<Vec<u32>>,
in_use: Vec<bool>,
touched: Vec<u32>,
growth: Vec<i64>,
grown: Vec<u32>,
defect: Vec<bool>,
order: Vec<u32>,
via: Vec<usize>,
seen: Vec<bool>,
}
impl RoundBased {
pub fn new(g: &MatchingGraph) -> RoundBased {
let nodes = g.detectors() + 1;
let mut ids: std::collections::BTreeMap<(u32, u32), u32> = std::collections::BTreeMap::new();
let mut start = vec![0];
let (mut nbr, mut weight, mut obs, mut edge) = (Vec::new(), Vec::new(), Vec::new(), Vec::new());
for u in 0..nodes {
for &(v, w, o) in g.neighbours(u) {
let key = ((u as u32).min(v), (u as u32).max(v));
let next = ids.len() as u32;
let id = *ids.entry(key).or_insert(next);
nbr.push(v);
weight.push(w.max(0));
obs.push(o);
edge.push(id);
}
start.push(nbr.len());
}
RoundBased { nodes, start, nbr, weight, obs, edge, edges: ids.len() }
}
pub fn workspace(&self) -> RoundWorkspace {
RoundWorkspace {
parent: (0..self.nodes as u32).collect(),
odd: vec![false; self.nodes],
boundary: vec![false; self.nodes],
members: vec![Vec::new(); self.nodes],
in_use: vec![false; self.nodes],
touched: Vec::new(),
growth: vec![0; self.edges],
grown: Vec::new(),
defect: vec![false; self.nodes],
order: Vec::new(),
via: vec![usize::MAX; self.nodes],
seen: vec![false; self.nodes],
}
}
fn find(ws: &mut RoundWorkspace, mut u: u32) -> u32 {
while ws.parent[u as usize] != u {
let p = ws.parent[u as usize];
ws.parent[u as usize] = ws.parent[p as usize];
u = ws.parent[u as usize];
}
u
}
fn touch(&self, ws: &mut RoundWorkspace, u: u32) {
let i = u as usize;
if !ws.in_use[i] {
ws.in_use[i] = true;
ws.touched.push(u);
ws.parent[i] = u;
ws.odd[i] = false;
ws.boundary[i] = i == self.nodes - 1;
ws.members[i].clear();
ws.members[i].push(u);
ws.defect[i] = false;
}
}
fn union(ws: &mut RoundWorkspace, a: u32, b: u32) {
let (mut ra, mut rb) = (Self::find(ws, a), Self::find(ws, b));
if ra == rb {
return;
}
if ws.members[ra as usize].len() < ws.members[rb as usize].len() {
core::mem::swap(&mut ra, &mut rb);
}
ws.parent[rb as usize] = ra;
let moved = core::mem::take(&mut ws.members[rb as usize]);
ws.members[ra as usize].extend(moved);
ws.odd[ra as usize] ^= ws.odd[rb as usize];
ws.boundary[ra as usize] |= ws.boundary[rb as usize];
}
fn reset(ws: &mut RoundWorkspace) {
for &u in &ws.touched {
let i = u as usize;
ws.in_use[i] = false;
ws.parent[i] = u;
ws.members[i].clear();
ws.defect[i] = false;
ws.seen[i] = false;
ws.via[i] = usize::MAX;
}
ws.touched.clear();
for &e in &ws.grown {
ws.growth[e as usize] = 0;
}
ws.grown.clear();
}
pub fn decode(&self, fired: &[u32], ws: &mut RoundWorkspace) -> Result<u64, MatchError> {
Self::reset(ws);
for &d in fired {
self.touch(ws, d);
ws.odd[d as usize] ^= true;
ws.defect[d as usize] ^= true;
}
let mut roots: Vec<u32> = Vec::new();
let mut full: Vec<(u32, u32)> = Vec::new();
loop {
roots.clear();
for k in 0..ws.touched.len() {
let u = ws.touched[k];
if Self::find(ws, u) == u && ws.odd[u as usize] && !ws.boundary[u as usize] {
roots.push(u);
}
}
if roots.is_empty() {
break;
}
roots.sort_unstable();
let mut step = i64::MAX;
for &r in &roots {
for m in 0..ws.members[r as usize].len() {
let u = ws.members[r as usize][m] as usize;
for k in self.start[u]..self.start[u + 1] {
let v = self.nbr[k];
let rv = if ws.in_use[v as usize] { Self::find(ws, v) } else { v };
if ws.in_use[v as usize] && rv == r {
continue;
}
let rem = self.weight[k] - ws.growth[self.edge[k] as usize];
if rem <= 0 {
step = 0;
continue;
}
let speed = if ws.in_use[v as usize] && ws.odd[rv as usize] && !ws.boundary[rv as usize] { 2 } else { 1 };
step = step.min((rem + speed - 1) / speed);
}
}
}
if step == i64::MAX {
return Err(MatchError::Unmatchable);
}
full.clear();
for &r in &roots {
for m in 0..ws.members[r as usize].len() {
let u = ws.members[r as usize][m];
for k in self.start[u as usize]..self.start[u as usize + 1] {
let v = self.nbr[k];
if ws.in_use[v as usize] && Self::find(ws, v) == r {
continue;
}
let e = self.edge[k] as usize;
if ws.growth[e] < self.weight[k] {
if ws.growth[e] == 0 {
ws.grown.push(e as u32);
}
ws.growth[e] = (ws.growth[e] + step).min(self.weight[k]);
}
if ws.growth[e] >= self.weight[k] {
full.push((u, v));
}
}
}
}
for &(u, v) in &full {
self.touch(ws, v);
Self::union(ws, u, v);
}
}
Ok(self.peel(ws))
}
fn peel(&self, ws: &mut RoundWorkspace) -> u64 {
let boundary = (self.nodes - 1) as u32;
let mut out = 0u64;
let mut roots: Vec<u32> = Vec::new();
for k in 0..ws.touched.len() {
let u = ws.touched[k];
if Self::find(ws, u) == u {
roots.push(u);
}
}
roots.sort_unstable();
for r in roots {
let start_node = if ws.boundary[r as usize] { boundary } else { r };
ws.order.clear();
ws.order.push(start_node);
ws.seen[start_node as usize] = true;
let mut head = 0;
while head < ws.order.len() {
let u = ws.order[head] as usize;
head += 1;
for k in self.start[u]..self.start[u + 1] {
let v = self.nbr[k] as usize;
if !ws.in_use[v] || ws.seen[v] || ws.growth[self.edge[k] as usize] < self.weight[k] {
continue;
}
if Self::find(ws, v as u32) != r {
continue;
}
ws.seen[v] = true;
ws.via[v] = k;
ws.order.push(v as u32);
}
}
for idx in (1..ws.order.len()).rev() {
let u = ws.order[idx] as usize;
if ws.defect[u] {
let k = ws.via[u];
let parent = self.edge_source(k);
out ^= self.obs[k];
ws.defect[u] = false;
ws.defect[parent] ^= true;
}
}
}
out
}
fn edge_source(&self, k: usize) -> usize {
self.start.partition_point(|&s| s <= k) - 1
}
}
use crate::quantum_frame::{error_model, parse, sample, surface_code_memory};
fn graph(d: u32, p: f64) -> (crate::quantum_frame::Circuit, MatchingGraph) {
let c = parse(&surface_code_memory(d, d, p)).unwrap();
let g = MatchingGraph::from_model(&error_model(&c).unwrap()).unwrap();
(c, g)
}
#[test]
fn every_single_fault_is_corrected() {
for d in [3, 5] {
let (_, g) = graph(d, 0.003);
let uf = UnionFind::new(&g);
let mut ws = uf.workspace();
let b = g.detectors() as u32;
let mut checked = 0;
for u in 0..=g.detectors() {
for &(v, _, o) in g.neighbours(u) {
if (u as u32) < v {
let fired: Vec<u32> = [u as u32, v].into_iter().filter(|&x| x != b).collect();
assert_eq!(uf.decode(&fired, &mut ws).unwrap(), o, "d={d} edge ({u},{v})");
checked += 1;
}
}
}
assert!(checked > 100);
}
}
#[test]
fn the_radix_heap_pops_in_order() {
let mut h = RadixHeap::default();
let mut rng = 0x1234_5678u64;
let mut keys = Vec::new();
let mut last = 0u64;
for round in 0..200 {
for _ in 0..(round % 7) {
rng = rng.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407);
let k = last + (rng >> 40) % 1000;
h.push(k, round, 0);
keys.push(k);
}
if let Some((k, _, _)) = h.pop() {
keys.sort_unstable();
assert_eq!(k, keys.remove(0));
last = k;
}
}
while let Some((k, _, _)) = h.pop() {
keys.sort_unstable();
assert_eq!(k, keys.remove(0));
}
assert!(keys.is_empty());
}
#[test]
fn no_fired_detectors_means_no_flip() {
let (_, g) = graph(3, 0.003);
let uf = UnionFind::new(&g);
assert_eq!(uf.decode(&[], &mut uf.workspace()).unwrap(), 0);
}
#[test]
fn it_tracks_exact_matching_and_distance_helps() {
let mut rates = Vec::new();
for d in [3, 5, 7] {
let (c, g) = graph(d, 0.004);
let det = sample(&c, 6_000, 11);
let uf = UnionFind::new(&g).failures(&det).unwrap();
let mw = g.failures(&det).unwrap();
assert!(uf >= mw / 2 && uf <= 2 * mw + 10, "d={d}: uf {uf}, mwpm {mw}");
rates.push(uf);
}
assert!(rates[0] > rates[1] && rates[1] > rates[2], "{rates:?}");
}
#[test]
fn a_reused_workspace_changes_nothing() {
let (c, g) = graph(5, 0.006);
let uf = UnionFind::new(&g);
let det = sample(&c, 300, 4);
let mut shared = uf.workspace();
for s in 0..300 {
let a = uf.decode(&det.fired(s), &mut shared).unwrap();
let b = uf.decode(&det.fired(s), &mut uf.workspace()).unwrap();
assert_eq!(a, b, "shot {s}");
}
}
#[test]
fn event_driven_growth_matches_rounds() {
for (d, p) in [(5, 0.006), (7, 0.004), (9, 0.005)] {
let (c, g) = graph(d, p);
let det = sample(&c, 2000, 13);
let (ev, rb) = (UnionFind::new(&g), RoundBased::new(&g));
let (mut we, mut wr) = (ev.workspace(), rb.workspace());
let mut differ = 0;
for s in 0..2000 {
let f = det.fired(s);
if ev.decode(&f, &mut we).unwrap() != rb.decode(&f, &mut wr).unwrap() {
differ += 1;
}
}
assert!(differ <= 4, "d={d}: {differ} of 2000 differ");
}
}
}