1use 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#[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 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 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 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 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
510struct 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}