1use super::min_sum::{MinSumBPDecoder, MinSumDecoderConfig};
12use crate::decoder::{Bit, SparseBitMatrix};
13use crate::decoder::{DecodeResult, Decoder, DecoderRunner};
14use log::debug;
15
16use ndarray::{Array1, Array2, ArrayView1};
17use num_traits::{Bounded, FromPrimitive, Signed, ToPrimitive};
18use std::fmt::Debug;
19use std::fs::File;
20use std::fs::OpenOptions;
21use std::io::{BufWriter, Write};
22use rand::distributions::{Distribution, Uniform};
24use rand::SeedableRng;
25use std::process::exit;
26use std::sync::Arc;
27
28#[derive(Clone, PartialEq, Debug)]
29pub enum StoppingCriterion {
30 PreIter,
31 NConv { stop_after: usize },
32 All,
33}
34
35impl Default for StoppingCriterion {
36 fn default() -> StoppingCriterion {
37 StoppingCriterion::NConv { stop_after: 1 }
38 }
39}
40
41#[derive(Clone, Debug)]
42pub struct RelayDecoderConfig {
43 pub pre_iter: usize,
44 pub num_sets: usize,
45 pub set_max_iter: usize,
46 pub gamma_dist_interval: (f64, f64),
47 pub explicit_gammas: Option<Array2<f64>>,
48 pub stopping_criterion: StoppingCriterion,
49 pub logging: bool,
50 pub seed: u64,
51}
52
53impl Default for RelayDecoderConfig {
54 fn default() -> Self {
55 Self {
56 pre_iter: 80,
57 num_sets: 300,
58 set_max_iter: 60,
59 gamma_dist_interval: (-0.24, 0.66),
60 explicit_gammas: None,
61 stopping_criterion: StoppingCriterion::default(),
62 logging: false,
63 seed: 0,
64 }
65 }
66}
67
68#[derive(Clone)]
69struct PosteriorUpdateState {
70 rng_std: rand::rngs::StdRng,
71 uniform: rand::distributions::Uniform<f64>,
72}
73
74#[derive(Clone)]
76pub struct RelayDecoder<N: PartialEq + Default + Clone + Copy> {
77 bp_decoder: MinSumBPDecoder<N>,
78 relay_config: Arc<RelayDecoderConfig>,
79 posterior_update_state: PosteriorUpdateState,
80 sets_quality: Array1<f64>,
81 sets_iter: Array1<usize>,
82 sets_conv: Array1<bool>,
83 sets_best: Array1<bool>,
84 num_executed_sets: usize,
85}
86
87impl<N> RelayDecoder<N>
88where
89 N: PartialEq
90 + Debug
91 + Default
92 + Clone
93 + Copy
94 + Signed
95 + Bounded
96 + FromPrimitive
97 + ToPrimitive
98 + std::cmp::PartialOrd
99 + std::ops::Add
100 + std::ops::AddAssign
101 + std::ops::DivAssign
102 + std::ops::Mul<N>
103 + std::ops::MulAssign
104 + Send
105 + Sync
106 + std::fmt::Display
107 + 'static,
108{
109 pub fn new(
110 check_matrix: Arc<SparseBitMatrix>,
111 min_sum_config: Arc<MinSumDecoderConfig>,
112 relay_config: Arc<RelayDecoderConfig>,
113 ) -> RelayDecoder<N> {
114 if relay_config.logging {
115 let log_line = format!(
116 "# pre_iter: {}: sets: {} set_max_iter: {}\n\
117 # gamma_distribution: {:?} # set_idx, num_iter, converged, unique_best_solution\n",
118 relay_config.pre_iter,
119 relay_config.num_sets,
120 relay_config.set_max_iter,
121 relay_config.gamma_dist_interval,
122 );
123 let mut file =
124 File::create("relay_logging.out").expect("Unable to create file for logging.");
125 file.write_all(log_line.as_bytes())
126 .expect("Unable to write Relay logging data.");
127 }
128
129 let (sets_quality, sets_iter, sets_conv, sets_best);
131 if relay_config.logging {
132 sets_quality = Array1::<f64>::from_elem(relay_config.num_sets + 1, f64::MAX);
133 sets_iter =
134 Array1::<usize>::from_elem(relay_config.num_sets + 1, relay_config.set_max_iter);
135 sets_conv = Array1::<bool>::from_elem(relay_config.num_sets + 1, false);
136 sets_best = Array1::<bool>::from_elem(relay_config.num_sets + 1, false);
137 } else {
138 sets_quality = Array1::<f64>::zeros(1);
139 sets_iter = Array1::<usize>::zeros(1);
140 sets_conv = Array1::<bool>::from_elem(1, false);
141 sets_best = Array1::<bool>::from_elem(1, false);
142 }
143
144 if let Some(gammas) = relay_config.explicit_gammas.as_ref() {
145 let gammas_shape = gammas.shape();
146 let num_variable_nodes = check_matrix.cols();
147 if num_variable_nodes != gammas_shape[1] {
148 eprintln!("ERROR: Number of specified gammas {} does not match the number of variable nodes {}.", gammas_shape[1], num_variable_nodes);
149 exit(1);
150 };
151 if relay_config.num_sets > gammas_shape[0] {
152 eprintln!("WARNING: Number of different gamma sets {} is smaller than the number of Relay legs {}. Legs will be reused.", gammas_shape[0], relay_config.num_sets)
153 }
154 }
155
156 let num_executed_sets = 0;
158
159 let bp_decoder = MinSumBPDecoder::new(check_matrix, min_sum_config);
160
161 let posterior_update_state = Self::init_dismem_state(&relay_config);
162
163 RelayDecoder {
164 bp_decoder,
165 relay_config,
166 posterior_update_state,
167 sets_quality,
168 sets_iter,
169 sets_conv,
170 sets_best,
171 num_executed_sets,
172 }
173 }
174
175 fn init_dismem_state(relay_config: &RelayDecoderConfig) -> PosteriorUpdateState {
176 let rng_std: rand::prelude::StdRng = rand::rngs::StdRng::seed_from_u64(relay_config.seed);
177 let low = relay_config.gamma_dist_interval.0;
178 let high = relay_config.gamma_dist_interval.1;
179 let uniform: rand::distributions::Uniform<f64> = Uniform::new(low, high);
180 PosteriorUpdateState { rng_std, uniform }
181 }
182
183 fn init_next_set(&mut self, set_idx: usize) {
184 let mut gammas = Array1::zeros(self.check_matrix().cols());
185 if let Some(explicit_gammas) = self.relay_config.explicit_gammas.as_ref() {
186 let gammas_num_sets = explicit_gammas.shape()[0];
187 for (i, gamma_ref) in gammas.iter_mut().enumerate() {
188 *gamma_ref = *explicit_gammas
189 .get((set_idx % gammas_num_sets, i))
190 .expect("index within explicit_gammas bounds");
191 }
192 self.bp_decoder.set_memory_strengths_f64(gammas);
193 return;
194 }
195 for i in 0..gammas.len() {
196 gammas[i] = self
197 .posterior_update_state
198 .uniform
199 .sample(&mut self.posterior_update_state.rng_std);
200 }
201 self.bp_decoder.set_memory_strengths_f64(gammas);
202 }
203
204 fn decode_inner(&mut self, detectors: ArrayView1<Bit>, max_iter: usize) -> DecodeResult {
206 let mut success: bool = false;
207 let mut decoded_detectors = Array1::default(detectors.dim());
208
209 for _ in 0..max_iter {
210 self.bp_decoder.run_iteration(detectors);
211 self.bp_decoder.current_iteration += 1;
212 decoded_detectors = self.bp_decoder.compute_decoded_detectors();
213 success = self
214 .bp_decoder
215 .check_convergence(detectors, decoded_detectors.view());
216
217 if success {
219 debug!(
220 "Succeeded on iteration {:?}",
221 self.bp_decoder.current_iteration
222 );
223 break;
224 }
225 }
226
227 self.bp_decoder
228 .build_result(success, decoded_detectors, max_iter)
229 }
230
231 fn write_log(&mut self, file: File) {
232 let mut buf_writer = BufWriter::new(file);
233 for set in 0..self.num_executed_sets {
234 let log_line = format!(
235 "{}, {}, {}, {}\n",
236 (set - 1) as i32,
237 self.sets_iter[set],
238 self.sets_conv[set] as u8,
239 self.sets_best[set] as u8
240 );
241 buf_writer
242 .write_all(log_line.as_bytes())
243 .expect("Unable to write Relay logging data.");
244 }
245 buf_writer
246 .flush()
247 .expect("Unable to write Relay logging data.");
248 }
249}
250
251impl<N> Decoder for RelayDecoder<N>
252where
253 N: PartialEq
254 + Debug
255 + Default
256 + Clone
257 + Copy
258 + Signed
259 + Bounded
260 + FromPrimitive
261 + ToPrimitive
262 + std::cmp::PartialOrd
263 + std::ops::Add
264 + std::ops::AddAssign
265 + std::ops::DivAssign
266 + std::ops::Mul<N>
267 + std::ops::MulAssign
268 + Send
269 + Sync
270 + std::fmt::Display
271 + 'static,
272{
273 fn check_matrix(&self) -> Arc<SparseBitMatrix> {
274 self.bp_decoder.check_matrix()
275 }
276
277 fn log_prior_ratios(&mut self) -> Array1<f64> {
278 self.bp_decoder.log_prior_ratios()
279 }
280
281 fn decode_detailed(&mut self, detectors: ArrayView1<Bit>) -> DecodeResult {
282 let mut num_conv = 0;
284 let mut min_pm = f64::MAX;
285 let mut num_sets_best = 0;
286 let mut best_set_idx = 0;
287 let mut total_iterations: usize = 0;
288 self.num_executed_sets = 0;
289 let stopping_criterion = self.relay_config.stopping_criterion.clone();
290
291 self.bp_decoder.initialize_decoder();
293 let mut result = self.decode_inner(detectors, self.relay_config.pre_iter);
294 self.num_executed_sets = 1;
295
296 if self.relay_config.logging {
298 self.sets_iter[0] = result.iterations;
299 self.sets_conv[0] = result.success;
300 }
301
302 if result.success {
304 num_conv += 1;
305 min_pm = result.decoding_quality;
306 num_sets_best += 1;
307 if self.relay_config.logging {
308 self.sets_quality[0] = result.decoding_quality
309 };
310
311 let mut done = false;
312 if stopping_criterion == StoppingCriterion::PreIter {
313 done = true;
314 } else if let StoppingCriterion::NConv { stop_after } = stopping_criterion {
315 if num_conv >= stop_after {
316 done = true;
317 }
318 }
319 if done {
321 if self.relay_config.logging {
322 self.sets_best[0] = true;
323 let file = OpenOptions::new()
324 .append(true)
325 .open("relay_logging.out")
326 .unwrap();
327 self.write_log(file);
328 }
329 return result;
330 }
331 }
332
333 total_iterations += result.iterations;
335 for set in 1..=self.relay_config.num_sets {
336 self.init_next_set(set);
339 self.bp_decoder.current_iteration = 0;
340 self.bp_decoder.initialize_check_to_variable();
341 self.bp_decoder.initialize_variable_to_check();
342 let temp_result = self.decode_inner(detectors, self.relay_config.set_max_iter);
343
344 self.num_executed_sets += 1;
345 total_iterations += temp_result.iterations;
346 if temp_result.success {
347 num_conv += 1;
348 let pm = temp_result.decoding_quality;
349 if self.relay_config.logging {
350 self.sets_conv[set] = true;
351 self.sets_iter[set] = temp_result.iterations;
352 self.sets_quality[set] = pm;
353 }
354 if pm == min_pm {
355 num_sets_best += 1;
357 }
358 if pm < min_pm {
359 num_sets_best = 1;
361 best_set_idx = set;
362 min_pm = pm;
363 result = temp_result;
364 }
365 if let StoppingCriterion::NConv { stop_after } = stopping_criterion {
366 if num_conv >= stop_after {
367 break;
368 }
369 }
370 }
371 }
372 result.iterations = total_iterations;
373
374 if self.relay_config.logging {
376 if num_sets_best == 1 {
377 self.sets_best[best_set_idx] = true;
378 }
379 let file = OpenOptions::new()
380 .append(true)
381 .open("relay_logging.out")
382 .unwrap();
383 self.write_log(file);
384 }
385
386 result
387 }
388
389 fn get_decoding_quality(&mut self, errors: ArrayView1<u8>) -> f64 {
390 self.bp_decoder.get_decoding_quality(errors)
391 }
392}
393
394impl<N> DecoderRunner for RelayDecoder<N> where
395 N: PartialEq
396 + Debug
397 + Default
398 + Clone
399 + Copy
400 + Signed
401 + Bounded
402 + FromPrimitive
403 + ToPrimitive
404 + std::cmp::PartialOrd
405 + std::ops::Add
406 + std::ops::AddAssign
407 + std::ops::DivAssign
408 + std::ops::Mul<N>
409 + std::ops::MulAssign
410 + Send
411 + Sync
412 + std::fmt::Display
413 + 'static
414{
415}
416
417#[cfg(test)]
418mod tests {
419
420 use super::*;
421
422 use crate::bipartite_graph::{BipartiteGraph, SparseBipartiteGraph};
423 use env_logger;
424 use ndarray::prelude::*;
425
426 use crate::dem::DetectorErrorModel;
427 use crate::utilities::test::get_test_data_path;
428 use ndarray::Array2;
429 use ndarray_npy::read_npy;
430
431 fn init() {
432 let _ = env_logger::builder().is_test(true).try_init();
433 }
434
435 #[test]
437 fn min_sum_decode_repetition_code() {
438 init();
439
440 let check_matrix = array![[1, 1, 0], [0, 1, 1],];
442
443 let check_matrix: SparseBipartiteGraph<_> = SparseBipartiteGraph::from_dense(check_matrix);
444 let check_matrix_arc = Arc::new(check_matrix);
445
446 let iterations = 10;
447 let bp_config = MinSumDecoderConfig {
448 error_priors: array![0.003, 0.003, 0.003],
449 max_iter: iterations,
450 alpha: Some(1.),
451 alpha_iteration_scaling_factor: 1.,
452 gamma0: None,
453 ..Default::default()
454 };
455 let bp_config_arc = Arc::new(bp_config);
456
457 let relay_config = RelayDecoderConfig {
458 pre_iter: iterations,
459 num_sets: 0,
460 set_max_iter: 150,
461 stopping_criterion: StoppingCriterion::PreIter,
462 explicit_gammas: None,
463 ..Default::default()
464 };
465 let relay_config_arc = Arc::new(relay_config);
466
467 let mut decoder: RelayDecoder<f32> =
468 RelayDecoder::new(check_matrix_arc, bp_config_arc, relay_config_arc);
469
470 let error = array![0, 0, 0];
471 let detectors: Array1<Bit> = array![0, 0];
472
473 let result = decoder.decode_detailed(detectors.view());
474
475 assert_eq!(result.decoding, error);
476 assert_eq!(result.decoded_detectors, detectors);
477 assert_eq!(result.max_iter, iterations);
478 assert!(result.success);
479
480 let error = array![1, 0, 0];
481 let detectors: Array1<Bit> = array![1, 0];
482
483 let result = decoder.decode_detailed(detectors.view());
484
485 assert_eq!(result.decoding, error);
486 assert_eq!(result.decoded_detectors, detectors);
487 assert_eq!(result.max_iter, iterations);
488 assert!(result.success);
489
490 let error = array![0, 1, 0];
491 let detectors: Array1<Bit> = array![1, 1];
492
493 let result = decoder.decode_detailed(detectors.view());
494
495 assert_eq!(result.decoding, error);
496 assert_eq!(result.decoded_detectors, detectors);
497 assert_eq!(result.max_iter, iterations);
498 assert!(result.success);
499
500 let error = array![0, 0, 1];
501 let detectors: Array1<Bit> = array![0, 1];
502
503 let result = decoder.decode_detailed(detectors.view());
504
505 assert_eq!(result.decoding, error);
506 assert_eq!(result.decoded_detectors, detectors);
507 assert_eq!(result.max_iter, iterations);
508 assert!(result.success);
509 }
510
511 #[test]
513 fn decode_144_12_12() {
514 let resources = get_test_data_path();
515 let code_144_12_12 =
516 DetectorErrorModel::load(resources.join("144_12_12")).expect("Unable to load the code");
517 let detectors_144_12_12: Array2<Bit> =
518 read_npy(resources.join("144_12_12_detectors.npy")).expect("Unable to open file");
519 let bp_config_144_12_12 = MinSumDecoderConfig {
520 error_priors: code_144_12_12.error_priors,
521 max_iter: 200,
522 alpha: None,
523 alpha_iteration_scaling_factor: 0.,
524 gamma0: Some(0.9),
525 ..Default::default()
526 };
527 let relay_config = RelayDecoderConfig::default();
528 let check_matrix = Arc::new(code_144_12_12.detector_error_matrix);
529 let bp_config = Arc::new(bp_config_144_12_12);
530 let config = Arc::new(relay_config);
531 let mut decoder_144_12_12: RelayDecoder<f64> =
532 RelayDecoder::new(check_matrix, bp_config, config);
533 let num_errors = 100;
534 let detectors_slice = detectors_144_12_12.slice(s![..num_errors, ..]);
535 let results = decoder_144_12_12.par_decode_detailed_batch(detectors_slice);
536
537 assert!(
539 results.iter().map(|x| x.success as usize).sum::<usize>()
540 == (detectors_slice.shape()[0])
541 );
542
543 assert_eq!(results[0].decoding.len(), 8785);
544 }
545
546 #[test]
548 fn decode_144_12_12_int() {
549 let resources = get_test_data_path();
550 let code_144_12_12 =
551 DetectorErrorModel::load(resources.join("144_12_12")).expect("Unable to load the code");
552 let detectors_144_12_12: Array2<Bit> =
553 read_npy(resources.join("144_12_12_detectors.npy")).expect("Unable to open file");
554
555 let bits = 16;
556 let scale = 8.0;
557
558 let bp_config_144_12_12 = MinSumDecoderConfig {
559 error_priors: code_144_12_12.error_priors,
560 max_iter: 200,
561 alpha: None,
562 alpha_iteration_scaling_factor: 0.,
563 gamma0: Some(0.9),
564 max_data_value: Some(((1 << bits) - 1) as f64),
565 data_scale_value: Some(scale),
566 ..Default::default()
567 };
568 let relay_config = RelayDecoderConfig {
569 ..Default::default()
570 };
571 let check_matrix = Arc::new(code_144_12_12.detector_error_matrix);
572 let bp_config = Arc::new(bp_config_144_12_12);
573 let config = Arc::new(relay_config);
574 let mut decoder_144_12_12: RelayDecoder<isize> =
575 RelayDecoder::new(check_matrix, bp_config, config);
576 let num_errors = 100;
577 let detectors_slice = detectors_144_12_12.slice(s![..num_errors, ..]);
578 let results = decoder_144_12_12.par_decode_detailed_batch(detectors_slice);
579
580 assert!(
582 results.iter().map(|x| x.success as usize).sum::<usize>()
583 == (detectors_slice.shape()[0])
584 );
585
586 assert_eq!(results[0].decoding.len(), 8785);
587 }
588}