deq-runtime 0.5.2-rc1

deq: Real-time Quantum Error Correction Decoding System
//! Tesseract decoder
//!
//! Google's Tesseract beam-search QEC decoder, bridged from C++ via cxx.
//!

use crate::decoder::blackbox_decoder::{self, ParityFactor};
use crate::decoder::decoder_features::DecoderFeatures;
use crate::decoder::tesseract_ffi::{TesseractCxxConfig, TesseractCxxDecoder};
use crate::decoder::thread_pooling::{
    DecodeError, DecodeRequest, DecoderInstance, ThreadPoolingConfig, ThreadPoolingDecoder,
};
use crate::misc::bit_vector::to_sparse_indices;
use blackbox_decoder::DecodingHypergraph;
use serde::{Deserialize, Serialize};
#[cfg(feature = "cli")]
use structdoc::StructDoc;

#[derive(Debug, Clone, Serialize, Deserialize)]
#[cfg_attr(feature = "cli", derive(StructDoc))]
#[serde(deny_unknown_fields)]
pub struct TesseractDecoderConfig {
    /// we want to recognize all the thread pooling config fields
    #[serde(flatten)]
    pub thread_pooling_config: ThreadPoolingConfig,
    /// beam cutoff (max number of active detectors in a search state)
    #[serde(default = "default_det_beam")]
    pub det_beam: i32,
    /// enable beam climbing heuristic
    #[serde(default)]
    pub beam_climbing: bool,
    /// prevent revisiting syndrome patterns
    #[serde(default = "default_true")]
    pub no_revisit_dets: bool,
    /// merge indistinguishable error mechanisms
    #[serde(default = "default_true")]
    pub merge_errors: bool,
    /// priority queue size limit
    #[serde(default = "default_pqlimit")]
    pub pqlimit: u64,
    /// penalty added per detector visited
    #[serde(default)]
    pub det_penalty: f64,
}

fn default_det_beam() -> i32 {
    5
}

fn default_true() -> bool {
    true
}

fn default_pqlimit() -> u64 {
    200_000
}

pub struct TesseractDecoderInstance {
    decoder: TesseractCxxDecoder,
    /// The probabilities the loaded decoder was built from, in Tesseract error
    /// order. Kept so [`DecoderInstance::decode`] can restore them
    /// after a shot-scoped reweighting.
    base_probabilities: Vec<f64>,
}

impl DecoderInstance for TesseractDecoderInstance {
    fn supported_features(_config: &serde_json::Value) -> DecoderFeatures {
        DecoderFeatures::REWEIGHTS
    }

    fn new(hypergraph: &DecodingHypergraph, config: &serde_json::Value) -> Self {
        let config: TesseractDecoderConfig = serde_json::from_value(config.clone()).unwrap();
        // Every hyperedge is loaded, including the zero-probability ones. Such an
        // edge is a *declared* impossibility rather than an absent one -- the
        // producer emits it deliberately and its index is part of the decoding
        // interface, referenced by `LossInfo` and by per-shot prior overrides --
        // so dropping it would both break that indexing and make the edge
        // impossible to raise later. Tesseract carries it at infinite cost
        // until something raises it.
        let (edge_vertices, edge_offsets, edge_probabilities) = flatten_hypergraph(hypergraph);
        let tess_config = TesseractCxxConfig {
            det_beam: config.det_beam,
            beam_climbing: config.beam_climbing,
            no_revisit_dets: config.no_revisit_dets,
            merge_errors: config.merge_errors,
            pqlimit: config.pqlimit,
            det_penalty: config.det_penalty,
        };
        Self {
            decoder: TesseractCxxDecoder::new(
                hypergraph.vertex_num,
                &edge_vertices,
                &edge_offsets,
                &edge_probabilities,
                &tess_config,
            ),
            base_probabilities: edge_probabilities,
        }
    }

    fn decode(&mut self, request: DecodeRequest<'_>) -> Result<ParityFactor, DecodeError> {
        if !request.reweights.is_empty() {
            let mut probabilities = self.base_probabilities.clone();
            for &(edge, probability) in request.reweights {
                let position = usize::try_from(edge).map_err(|_| {
                    DecodeError::InvalidInput(format!("reweighted edge {edge} is outside the loaded hypergraph"))
                })?;
                let target = probabilities.get_mut(position).ok_or_else(|| {
                    DecodeError::InvalidInput(format!("reweighted edge {edge} is outside the loaded hypergraph"))
                })?;
                *target = probability;
            }
            self.decoder.update_error_costs(&probabilities);
        }
        let detections: Vec<u64> = to_sparse_indices(request.syndrome);
        let error_indices = self.decoder.decode(&detections);
        if !request.reweights.is_empty() {
            self.decoder.update_error_costs(&self.base_probabilities);
        }
        error_indices
            .map(|subgraph| ParityFactor { subgraph })
            .map_err(|error| DecodeError::Backend(error.to_string()))
    }

    fn reset(&mut self) {
        // Tesseract clears its internal buffers at the start of each decode call.
    }
}

/// Flatten a decoding hypergraph into CSR arrays for the C++ bridge.
fn flatten_hypergraph(hypergraph: &DecodingHypergraph) -> (Vec<u64>, Vec<u64>, Vec<f64>) {
    let total_vertices: usize = hypergraph.hyperedges.iter().map(|edge| edge.vertices.len()).sum();

    let mut edge_vertices = Vec::with_capacity(total_vertices);
    let mut edge_offsets = Vec::with_capacity(hypergraph.hyperedges.len() + 1);
    let mut edge_probabilities = Vec::with_capacity(hypergraph.hyperedges.len());

    edge_offsets.push(0u64);
    for edge in &hypergraph.hyperedges {
        edge_vertices.extend_from_slice(&edge.vertices);
        edge_offsets.push(edge_vertices.len() as u64);
        edge_probabilities.push(edge.probability);
    }

    (edge_vertices, edge_offsets, edge_probabilities)
}

pub type TesseractDecoder = ThreadPoolingDecoder<TesseractDecoderInstance>;

#[cfg(test)]
mod tests {
    use super::*;
    use crate::decoder::blackbox_decoder::Hyperedge;
    use crate::util::BitVector;
    use serde_json::json;

    #[test]
    fn exhausted_search_is_an_error_and_does_not_poison_the_next_decode() {
        let graph = DecodingHypergraph {
            vertex_num: 1,
            hyperedges: vec![Hyperedge {
                vertices: vec![0],
                probability: 0.1,
            }],
        };
        let mut decoder = TesseractDecoderInstance::new(&graph, &json!({ "pqlimit": 1 }));
        let mut syndrome = BitVector {
            size: 1,
            data: vec![0x80],
        };
        let error = decoder
            .decode(DecodeRequest {
                syndrome: &syndrome,
                reweights: &[],
                loss: None,
            })
            .unwrap_err();
        assert!(matches!(error, DecodeError::Backend(message) if message.contains("pqlimit=1")));
        syndrome.data[0] = 0;
        assert!(
            decoder
                .decode(DecodeRequest {
                    syndrome: &syndrome,
                    reweights: &[],
                    loss: None,
                })
                .unwrap()
                .subgraph
                .is_empty()
        );
    }

    #[test]
    fn failed_reweighted_decode_restores_original_priors() {
        let graph = DecodingHypergraph {
            vertex_num: 2,
            hyperedges: [0.1, 0.2]
                .into_iter()
                .map(|probability| Hyperedge {
                    vertices: vec![0],
                    probability,
                })
                .collect(),
        };
        let mut decoder = TesseractDecoderInstance::new(&graph, &json!({ "merge_errors": false }));
        let mut syndrome = BitVector {
            size: 2,
            data: vec![0x40],
        };
        assert!(
            decoder
                .decode(DecodeRequest {
                    syndrome: &syndrome,
                    reweights: &[(0, 0.4)],
                    loss: None,
                })
                .is_err()
        );
        syndrome.data[0] = 0x80;
        assert_eq!(
            decoder
                .decode(DecodeRequest {
                    syndrome: &syndrome,
                    reweights: &[],
                    loss: None,
                })
                .unwrap()
                .subgraph,
            vec![1]
        );
    }

    #[test]
    fn narrower_beam_recovers_after_primary_queue_exhaustion() {
        let mut hyperedges = vec![
            Hyperedge {
                vertices: vec![0, 1],
                probability: 0.1,
            },
            Hyperedge {
                vertices: vec![1],
                probability: 0.1,
            },
        ];
        for detectors in [vec![0, 2, 3, 4], vec![0, 2, 3, 5], vec![0, 2, 4, 5], vec![0, 3, 4, 5]] {
            hyperedges.push(Hyperedge {
                vertices: detectors,
                probability: 0.2,
            });
        }
        let graph = DecodingHypergraph {
            vertex_num: 6,
            hyperedges,
        };
        let syndrome = BitVector {
            size: 6,
            data: vec![0x80],
        };
        for beam_climbing in [false, true] {
            let mut decoder =
                TesseractDecoderInstance::new(&graph, &json!({"det_beam": 5, "pqlimit": 3, "beam_climbing": beam_climbing}));
            let result = decoder
                .decode(DecodeRequest {
                    syndrome: &syndrome,
                    reweights: &[],
                    loss: None,
                })
                .unwrap();
            assert!(crate::decoder::blackbox_util::is_parity_factor(&graph, &result, &syndrome));
            assert_eq!(result.subgraph.len(), 2);
            assert!(result.subgraph.contains(&0) && result.subgraph.contains(&1));
        }
    }
}