Skip to main content

relay_bp/bp/
relay.rs

1// (C) Copyright IBM 2025
2//
3// This code is licensed under the Apache License, Version 2.0. You may
4// obtain a copy of this license in the LICENSE.txt file in the root directory
5// of this source tree or at http://www.apache.org/licenses/LICENSE-2.0.
6//
7// Any modifications or derivative works of this code must retain this
8// copyright notice, and modified files need to carry a notice indicating
9// that they have been altered from the originals.
10
11use 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};
22//use std::string;
23use 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/// An ensemble decoder which controls an inner BP min-sum decoder.
75#[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        // Create logging variables if applicable
130        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        // The actual number of sets Relay ran, depends on the stopping criterion
157        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    /// Decode with the inner decoder
205    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 we have converged may now exit
218            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        // Initialization
283        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        // First Mem-BP
292        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        // Create logging variables and log first set if applicable
297        if self.relay_config.logging {
298            self.sets_iter[0] = result.iterations;
299            self.sets_conv[0] = result.success;
300        }
301
302        // Check early stopping criteria
303        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 stopping criterion has been met: Log (if applicable) and return
320            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        // Init and loop over all Relay sets
334        total_iterations += result.iterations;
335        for set in 1..=self.relay_config.num_sets {
336            // Do not completely initialize decoder as we wish to relay
337            // posterior marginals with new memory strengths.
338            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                    // Count how often we found the best solution
356                    num_sets_best += 1;
357                }
358                if pm < min_pm {
359                    // Found a new best solution
360                    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        // Rest of the function is just logging
375        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    // Basic test where Relay is called but only runs 1 BP iteration
436    #[test]
437    fn min_sum_decode_repetition_code() {
438        init();
439
440        // Build 3, 2 qubit repetition code with weight 2 checks
441        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    // Basic test where Relay runs 40 sets
512    #[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        // All should pass for Relay.
538        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    // Basic test where Relay runs 40 sets
547    #[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        // All should pass for Relay.
581        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}