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 {
#[serde(flatten)]
pub thread_pooling_config: ThreadPoolingConfig,
#[serde(default = "default_det_beam")]
pub det_beam: i32,
#[serde(default)]
pub beam_climbing: bool,
#[serde(default = "default_true")]
pub no_revisit_dets: bool,
#[serde(default = "default_true")]
pub merge_errors: bool,
#[serde(default = "default_pqlimit")]
pub pqlimit: u64,
#[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,
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();
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) {
}
}
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));
}
}
}