Skip to main content

prism_q/qec/
decoder.rs

1//! Union-find decoder over graphlike detector error models.
2//!
3//! Weighted cluster growth in the `ln((1-p)/p)` metric followed by peeling on
4//! the grown erasure, per Delfosse and Nickerson (arXiv:1709.06218).
5
6use std::collections::HashMap;
7use std::collections::hash_map::Entry;
8
9use super::DetectorErrorModel;
10use super::dem::symptom_label;
11use crate::error::{PrismError, Result};
12use crate::sim::compiled::{PackedShots, ShotLayout};
13
14const BOUNDARY: u32 = u32::MAX;
15const EDGE_NONE: u32 = u32::MAX;
16const VERTEX_NONE: u32 = u32::MAX;
17const GROWTH_EPS: f64 = 1e-9;
18#[cfg(feature = "parallel")]
19const PARALLEL_SHOT_THRESHOLD: usize = 1024;
20#[cfg(feature = "parallel")]
21const SHOT_CHUNK: usize = 256;
22
23/// Union-find decoder compiled from a graphlike detector error model.
24///
25/// Detectors are vertices. A two-detector mechanism becomes an internal edge
26/// and a one-detector mechanism a boundary edge, each weighted `ln((1-p)/p)`
27/// clamped at zero. Mechanisms flipping no detector cannot enter the graph;
28/// their probability mass bounds the logical error rate any decoder over the
29/// model can reach. Mechanisms sharing one detector set collapse to the most
30/// probable of them. Decoding uses no randomness: growth, fusion, and peeling
31/// break every tie by ascending edge index in mechanism order, so equal
32/// inputs give equal outputs on any thread count.
33#[derive(Debug, Clone)]
34pub struct UnionFindDecoder {
35    num_detectors: usize,
36    num_observables: usize,
37    obs_words: usize,
38    edge_u: Vec<u32>,
39    edge_v: Vec<u32>,
40    edge_weight: Vec<f64>,
41    edge_obs: Vec<u64>,
42    adj_offsets: Vec<u32>,
43    adj_edge: Vec<u32>,
44}
45
46impl UnionFindDecoder {
47    /// Compile a decoder from a graphlike detector error model.
48    ///
49    /// # Errors
50    ///
51    /// A mechanism flipping more than two detectors is rejected with a
52    /// pointer to [`DetectorErrorModel::decompose_graphlike`]. Mechanism
53    /// probabilities must lie in `[0, 1)`; a zero-probability mechanism is
54    /// skipped rather than rejected.
55    pub fn from_model(model: &DetectorErrorModel) -> Result<Self> {
56        if model.num_detectors() >= BOUNDARY as usize {
57            return Err(PrismError::InvalidParameter {
58                message: format!(
59                    "{} detectors exceed the decoder's index range",
60                    model.num_detectors()
61                ),
62            });
63        }
64        let num_detectors = model.num_detectors();
65        let num_observables = model.num_observables();
66        let obs_words = num_observables.div_ceil(64);
67
68        let mut edge_u: Vec<u32> = Vec::new();
69        let mut edge_v: Vec<u32> = Vec::new();
70        let mut edge_p: Vec<f64> = Vec::new();
71        let mut edge_obs_rows: Vec<&[usize]> = Vec::new();
72        let mut slots: HashMap<(u32, u32), u32> = HashMap::new();
73        for mechanism in model.mechanisms() {
74            let p = mechanism.probability();
75            if !(0.0..1.0).contains(&p) {
76                return Err(PrismError::InvalidParameter {
77                    message: format!(
78                        "mechanism `{}` has probability {p}, outside [0, 1)",
79                        symptom_label(mechanism)
80                    ),
81                });
82            }
83            if p == 0.0 {
84                continue;
85            }
86            let endpoints = match *mechanism.detectors() {
87                [] => continue,
88                [d] => (d as u32, BOUNDARY),
89                [d0, d1] => (d0 as u32, d1 as u32),
90                _ => {
91                    return Err(PrismError::InvalidParameter {
92                        message: format!(
93                            "mechanism `{}` flips {} detectors; union-find decoding needs a \
94                             graphlike model, apply `decompose_graphlike` first",
95                            symptom_label(mechanism),
96                            mechanism.detectors().len()
97                        ),
98                    });
99                }
100            };
101            match slots.entry(endpoints) {
102                Entry::Occupied(slot) => {
103                    let at = *slot.get() as usize;
104                    if p > edge_p[at] {
105                        edge_p[at] = p;
106                        edge_obs_rows[at] = mechanism.observables();
107                    }
108                }
109                Entry::Vacant(slot) => {
110                    slot.insert(edge_u.len() as u32);
111                    edge_u.push(endpoints.0);
112                    edge_v.push(endpoints.1);
113                    edge_p.push(p);
114                    edge_obs_rows.push(mechanism.observables());
115                }
116            }
117        }
118
119        let edge_weight: Vec<f64> = edge_p
120            .iter()
121            .map(|&p| ((1.0 - p) / p).ln().max(0.0))
122            .collect();
123        let mut edge_obs = vec![0u64; edge_u.len() * obs_words];
124        for (edge, row) in edge_obs_rows.iter().enumerate() {
125            for &observable in *row {
126                edge_obs[edge * obs_words + observable / 64] |= 1u64 << (observable % 64);
127            }
128        }
129
130        let mut adj_offsets = vec![0u32; num_detectors + 1];
131        for edge in 0..edge_u.len() {
132            adj_offsets[edge_u[edge] as usize + 1] += 1;
133            if edge_v[edge] != BOUNDARY {
134                adj_offsets[edge_v[edge] as usize + 1] += 1;
135            }
136        }
137        for v in 0..num_detectors {
138            adj_offsets[v + 1] += adj_offsets[v];
139        }
140        let mut cursor = adj_offsets.clone();
141        let mut adj_edge = vec![0u32; *adj_offsets.last().unwrap() as usize];
142        for edge in 0..edge_u.len() {
143            let u = edge_u[edge] as usize;
144            adj_edge[cursor[u] as usize] = edge as u32;
145            cursor[u] += 1;
146            if edge_v[edge] != BOUNDARY {
147                let v = edge_v[edge] as usize;
148                adj_edge[cursor[v] as usize] = edge as u32;
149                cursor[v] += 1;
150            }
151        }
152
153        Ok(Self {
154            num_detectors,
155            num_observables,
156            obs_words,
157            edge_u,
158            edge_v,
159            edge_weight,
160            edge_obs,
161            adj_offsets,
162            adj_edge,
163        })
164    }
165
166    pub fn num_detectors(&self) -> usize {
167        self.num_detectors
168    }
169
170    pub fn num_observables(&self) -> usize {
171        self.num_observables
172    }
173
174    /// Decode packed detector samples into predicted observable flips.
175    ///
176    /// Accepts either layout with one bit per detector per shot, detector `d`
177    /// at bit index `d`. Returns shot-major records with one bit per
178    /// observable per shot, observable `o` at bit index `o`.
179    ///
180    /// # Errors
181    ///
182    /// The input measurement count must equal the model's detector count, and
183    /// every shot must be explainable: a detector component with odd defect
184    /// parity and no boundary edge rejects the batch, naming the first such
185    /// shot.
186    pub fn decode_packed(&self, detectors: &PackedShots) -> Result<PackedShots> {
187        if detectors.num_measurements() != self.num_detectors {
188            return Err(PrismError::InvalidParameter {
189                message: format!(
190                    "detector shots carry {} measurements, the model has {} detectors",
191                    detectors.num_measurements(),
192                    self.num_detectors
193                ),
194            });
195        }
196        let num_shots = detectors.num_shots();
197        let m_words = self.num_detectors.div_ceil(64);
198        let transposed;
199        let rows: &[u64] = match detectors.layout() {
200            ShotLayout::ShotMajor => detectors.raw_data(),
201            ShotLayout::MeasMajor => {
202                transposed = detectors.clone().into_shot_major_data();
203                &transposed
204            }
205        };
206        let out_words = self.obs_words;
207        let mut out = vec![0u64; num_shots * out_words];
208
209        #[cfg(feature = "parallel")]
210        if num_shots >= PARALLEL_SHOT_THRESHOLD && out_words > 0 {
211            use rayon::prelude::*;
212            let failure = out
213                .par_chunks_mut(SHOT_CHUNK * out_words)
214                .enumerate()
215                .map_init(
216                    || DecodeScratch::new(self),
217                    |scratch, (chunk, chunk_out)| {
218                        for (offset, shot_out) in chunk_out.chunks_mut(out_words).enumerate() {
219                            let shot = chunk * SHOT_CHUNK + offset;
220                            let row = &rows[shot * m_words..(shot + 1) * m_words];
221                            if let Err(stuck) = self.decode_shot(row, shot_out, scratch) {
222                                return Some((shot, stuck));
223                            }
224                        }
225                        None
226                    },
227                )
228                .reduce(
229                    || None,
230                    |a, b| match (a, b) {
231                        (Some(a), Some(b)) => Some(if a.0 <= b.0 { a } else { b }),
232                        (a, b) => a.or(b),
233                    },
234                );
235            if let Some((shot, stuck)) = failure {
236                return Err(stuck.into_error(shot));
237            }
238            return Ok(PackedShots::from_shot_major(
239                out,
240                num_shots,
241                self.num_observables,
242            ));
243        }
244
245        let mut scratch = DecodeScratch::new(self);
246        for shot in 0..num_shots {
247            let row = &rows[shot * m_words..(shot + 1) * m_words];
248            let shot_out = &mut out[shot * out_words..(shot + 1) * out_words];
249            self.decode_shot(row, shot_out, &mut scratch)
250                .map_err(|stuck| stuck.into_error(shot))?;
251        }
252        Ok(PackedShots::from_shot_major(
253            out,
254            num_shots,
255            self.num_observables,
256        ))
257    }
258
259    fn decode_shot(
260        &self,
261        row: &[u64],
262        out_row: &mut [u64],
263        s: &mut DecodeScratch,
264    ) -> std::result::Result<(), Stuck> {
265        s.stamp += 1;
266        s.shot_stamp = s.stamp;
267
268        s.defects.clear();
269        for (word_index, &bits) in row.iter().enumerate() {
270            let mut bits = bits;
271            while bits != 0 {
272                s.defects
273                    .push((word_index * 64) as u32 + bits.trailing_zeros());
274                bits &= bits - 1;
275            }
276        }
277        if s.defects.is_empty() {
278            return Ok(());
279        }
280
281        s.active.clear();
282        let mut i = 0;
283        while i < s.defects.len() {
284            let defect = s.defects[i];
285            i += 1;
286            s.activate(defect);
287            s.parity[defect as usize] = true;
288            s.defect_stamp[defect as usize] = s.shot_stamp;
289            s.active.push(defect);
290        }
291
292        self.grow_clusters(s)?;
293
294        let mut i = 0;
295        while i < s.defects.len() {
296            let root = s.find(s.defects[i]);
297            i += 1;
298            if s.peeled_stamp[root as usize] == s.shot_stamp {
299                continue;
300            }
301            s.peeled_stamp[root as usize] = s.shot_stamp;
302            self.peel_cluster(root, out_row, s);
303        }
304        Ok(())
305    }
306
307    // Each round grows every active cluster's non-saturated incident edges by
308    // one shared increment: the minimum slack over those edges, divided by how
309    // many active clusters touch the edge, so at least one edge saturates per
310    // round and the loop is bounded by the edge count.
311    fn grow_clusters(&self, s: &mut DecodeScratch) -> std::result::Result<(), Stuck> {
312        loop {
313            s.stamp += 1;
314            let round = s.stamp;
315
316            let mut live = 0usize;
317            let mut i = 0;
318            while i < s.active.len() {
319                let root = s.find(s.active[i]);
320                i += 1;
321                if s.seen_stamp[root as usize] == round {
322                    continue;
323                }
324                s.seen_stamp[root as usize] = round;
325                if s.parity[root as usize] && !s.boundary[root as usize] {
326                    s.active[live] = root;
327                    live += 1;
328                }
329            }
330            s.active.truncate(live);
331            if s.active.is_empty() {
332                return Ok(());
333            }
334
335            s.touched.clear();
336            for &root in &s.active {
337                let mut grew = false;
338                let mut lowest = root;
339                let mut v = root;
340                while v != VERTEX_NONE {
341                    lowest = lowest.min(v);
342                    let begin = self.adj_offsets[v as usize] as usize;
343                    let end = self.adj_offsets[v as usize + 1] as usize;
344                    for &edge in &self.adj_edge[begin..end] {
345                        let e = edge as usize;
346                        if s.edge_stamp[e] == s.shot_stamp && s.edge_saturated[e] {
347                            continue;
348                        }
349                        grew = true;
350                        if s.touch_stamp[e] == round {
351                            s.touch_count[e] += 1;
352                        } else {
353                            s.touch_stamp[e] = round;
354                            s.touch_count[e] = 1;
355                            s.touched.push(edge);
356                        }
357                    }
358                    v = s.list_next[v as usize];
359                }
360                if !grew {
361                    return Err(Stuck { detector: lowest });
362                }
363            }
364
365            let mut delta = f64::INFINITY;
366            for &edge in &s.touched {
367                let e = edge as usize;
368                let growth = if s.edge_stamp[e] == s.shot_stamp {
369                    s.edge_growth[e]
370                } else {
371                    0.0
372                };
373                let step = (self.edge_weight[e] - growth) / f64::from(s.touch_count[e]);
374                if step < delta {
375                    delta = step;
376                }
377            }
378
379            s.fused.clear();
380            for &edge in &s.touched {
381                let e = edge as usize;
382                if s.edge_stamp[e] != s.shot_stamp {
383                    s.edge_stamp[e] = s.shot_stamp;
384                    s.edge_growth[e] = 0.0;
385                    s.edge_saturated[e] = false;
386                }
387                s.edge_growth[e] += f64::from(s.touch_count[e]) * delta;
388                if s.edge_growth[e] + GROWTH_EPS >= self.edge_weight[e] {
389                    s.fused.push(e as u32);
390                }
391            }
392            s.fused.sort_unstable();
393            let mut i = 0;
394            while i < s.fused.len() {
395                let edge = s.fused[i];
396                i += 1;
397                s.edge_saturated[edge as usize] = true;
398                let u = self.edge_u[edge as usize];
399                let v = self.edge_v[edge as usize];
400                s.activate(u);
401                if v == BOUNDARY {
402                    let root = s.find(u);
403                    s.boundary[root as usize] = true;
404                    s.boundary_edge[root as usize] = s.boundary_edge[root as usize].min(edge);
405                } else {
406                    s.activate(v);
407                    let ru = s.find(u);
408                    let rv = s.find(v);
409                    if ru != rv {
410                        s.union(ru, rv);
411                    }
412                }
413            }
414        }
415    }
416
417    // Spanning-forest peel: leaves flush their defect through the tree edge
418    // toward the root, which is the interior endpoint of the designated
419    // boundary edge when the cluster touches the boundary.
420    fn peel_cluster(&self, root: u32, out_row: &mut [u64], s: &mut DecodeScratch) {
421        let start = if s.boundary[root as usize] {
422            self.edge_u[s.boundary_edge[root as usize] as usize]
423        } else {
424            let mut lowest = root;
425            let mut v = root;
426            while v != VERTEX_NONE {
427                lowest = lowest.min(v);
428                v = s.list_next[v as usize];
429            }
430            lowest
431        };
432
433        s.order.clear();
434        s.stack.clear();
435        s.dfs_stamp[start as usize] = s.shot_stamp;
436        s.stack.push(start);
437        while let Some(v) = s.stack.pop() {
438            let begin = self.adj_offsets[v as usize] as usize;
439            let end = self.adj_offsets[v as usize + 1] as usize;
440            for &edge in &self.adj_edge[begin..end] {
441                let e = edge as usize;
442                if s.edge_stamp[e] != s.shot_stamp || !s.edge_saturated[e] {
443                    continue;
444                }
445                if self.edge_v[e] == BOUNDARY {
446                    continue;
447                }
448                let other = if self.edge_u[e] == v {
449                    self.edge_v[e]
450                } else {
451                    self.edge_u[e]
452                };
453                if s.dfs_stamp[other as usize] == s.shot_stamp {
454                    continue;
455                }
456                s.dfs_stamp[other as usize] = s.shot_stamp;
457                s.order.push((other, edge, v));
458                s.stack.push(other);
459            }
460        }
461
462        for &(vertex, edge, parent) in s.order.iter().rev() {
463            if s.defect_stamp[vertex as usize] != s.shot_stamp {
464                continue;
465            }
466            s.defect_stamp[vertex as usize] = 0;
467            if s.defect_stamp[parent as usize] == s.shot_stamp {
468                s.defect_stamp[parent as usize] = 0;
469            } else {
470                s.defect_stamp[parent as usize] = s.shot_stamp;
471            }
472            self.xor_edge_observables(edge, out_row);
473        }
474
475        if s.defect_stamp[start as usize] == s.shot_stamp {
476            s.defect_stamp[start as usize] = 0;
477            debug_assert!(s.boundary[root as usize]);
478            self.xor_edge_observables(s.boundary_edge[root as usize], out_row);
479        }
480    }
481
482    #[inline]
483    fn xor_edge_observables(&self, edge: u32, out_row: &mut [u64]) {
484        let base = edge as usize * self.obs_words;
485        for (word, mask) in out_row
486            .iter_mut()
487            .zip(&self.edge_obs[base..base + self.obs_words])
488        {
489            *word ^= mask;
490        }
491    }
492}
493
494struct Stuck {
495    detector: u32,
496}
497
498impl Stuck {
499    fn into_error(self, shot: usize) -> PrismError {
500        PrismError::InvalidParameter {
501            message: format!(
502                "shot {shot}: the detector component containing D{} has odd syndrome parity \
503                 but no boundary edge, so the syndrome is impossible under the model",
504                self.detector
505            ),
506        }
507    }
508}
509
510/// Reusable per-shot decode state. Vertex and edge slots are validated by
511/// stamp comparison against the current shot or round, so nothing is cleared
512/// between shots and untouched slots cost nothing.
513struct DecodeScratch {
514    stamp: u64,
515    shot_stamp: u64,
516    parent: Vec<u32>,
517    size: Vec<u32>,
518    parity: Vec<bool>,
519    boundary: Vec<bool>,
520    boundary_edge: Vec<u32>,
521    list_tail: Vec<u32>,
522    list_next: Vec<u32>,
523    vertex_stamp: Vec<u64>,
524    seen_stamp: Vec<u64>,
525    defect_stamp: Vec<u64>,
526    dfs_stamp: Vec<u64>,
527    peeled_stamp: Vec<u64>,
528    edge_stamp: Vec<u64>,
529    edge_growth: Vec<f64>,
530    edge_saturated: Vec<bool>,
531    touch_stamp: Vec<u64>,
532    touch_count: Vec<u8>,
533    defects: Vec<u32>,
534    active: Vec<u32>,
535    touched: Vec<u32>,
536    fused: Vec<u32>,
537    stack: Vec<u32>,
538    order: Vec<(u32, u32, u32)>,
539}
540
541impl DecodeScratch {
542    fn new(decoder: &UnionFindDecoder) -> Self {
543        let vertices = decoder.num_detectors;
544        let edges = decoder.edge_u.len();
545        Self {
546            stamp: 0,
547            shot_stamp: 0,
548            parent: vec![0; vertices],
549            size: vec![0; vertices],
550            parity: vec![false; vertices],
551            boundary: vec![false; vertices],
552            boundary_edge: vec![0; vertices],
553            list_tail: vec![0; vertices],
554            list_next: vec![0; vertices],
555            vertex_stamp: vec![0; vertices],
556            seen_stamp: vec![0; vertices],
557            defect_stamp: vec![0; vertices],
558            dfs_stamp: vec![0; vertices],
559            peeled_stamp: vec![0; vertices],
560            edge_stamp: vec![0; edges],
561            edge_growth: vec![0.0; edges],
562            edge_saturated: vec![false; edges],
563            touch_stamp: vec![0; edges],
564            touch_count: vec![0; edges],
565            defects: Vec::new(),
566            active: Vec::new(),
567            touched: Vec::new(),
568            fused: Vec::new(),
569            stack: Vec::new(),
570            order: Vec::new(),
571        }
572    }
573
574    fn activate(&mut self, v: u32) {
575        let at = v as usize;
576        if self.vertex_stamp[at] == self.shot_stamp {
577            return;
578        }
579        self.vertex_stamp[at] = self.shot_stamp;
580        self.parent[at] = v;
581        self.size[at] = 1;
582        self.parity[at] = false;
583        self.boundary[at] = false;
584        self.boundary_edge[at] = EDGE_NONE;
585        self.list_tail[at] = v;
586        self.list_next[at] = VERTEX_NONE;
587    }
588
589    fn find(&mut self, mut v: u32) -> u32 {
590        while self.parent[v as usize] != v {
591            let grand = self.parent[self.parent[v as usize] as usize];
592            self.parent[v as usize] = grand;
593            v = grand;
594        }
595        v
596    }
597
598    fn union(&mut self, a: u32, b: u32) {
599        let (big, small) = if self.size[a as usize] > self.size[b as usize]
600            || (self.size[a as usize] == self.size[b as usize] && a < b)
601        {
602            (a, b)
603        } else {
604            (b, a)
605        };
606        let (big_at, small_at) = (big as usize, small as usize);
607        self.parent[small_at] = big;
608        self.size[big_at] += self.size[small_at];
609        self.parity[big_at] ^= self.parity[small_at];
610        self.boundary[big_at] |= self.boundary[small_at];
611        self.boundary_edge[big_at] = self.boundary_edge[big_at].min(self.boundary_edge[small_at]);
612        self.list_next[self.list_tail[big_at] as usize] = small;
613        self.list_tail[big_at] = self.list_tail[small_at];
614    }
615}