use crate::bin;
use crate::decoder::DynDecoder;
use crate::decoder::blackbox_decoder;
use crate::decoder::decoder_features::DecoderFeatures;
use crate::misc::index::ErrorIndex;
use crate::misc::util::exclusive_probability_of;
use crate::util::BitVector;
use serde::{Deserialize, Serialize};
use std::ops::Index;
use std::sync::Arc;
use tonic::Status;
pub(crate) async fn load_projected_decoder(
decoder: &DynDecoder,
base_hypergraph: blackbox_decoder::DecodingHypergraph,
base_errors: Arc<Vec<ErrorIndex>>,
deduplicate: bool,
retain_decoding_hypergraph: bool,
ignore_isolated_vertices: bool,
) -> Result<LoadedDecoder, Status> {
let (projection, prepared) = prepare_decoder(base_hypergraph, base_errors, deduplicate);
let hypergraph = prepared.hypergraph;
let ignored_syndrome_vertices = if ignore_isolated_vertices {
Arc::new(edge_isolated_vertices(&hypergraph))
} else {
Arc::new(vec![])
};
let decoding_hypergraph = retain_decoding_hypergraph.then(|| Arc::new(hypergraph.clone()));
let hid = decoder.load_hypergraph(hypergraph).await?.hid;
Ok(LoadedDecoder {
hid,
decoding_hypergraph,
ignored_syndrome_vertices,
projection: Arc::new(projection),
})
}
fn prepare_decoder(
base_hypergraph: blackbox_decoder::DecodingHypergraph,
base_errors: Arc<Vec<ErrorIndex>>,
deduplicate: bool,
) -> (DecodeProjection, PreparedDecoderInput) {
debug_assert_eq!(base_hypergraph.hyperedges.len(), base_errors.len());
let error_edge_lookup = ErrorEdgeLookup::new(&base_errors);
let (prepared, edge_projection) = if deduplicate {
deduplicate_by_syndrome(&base_hypergraph, &base_errors)
} else {
(
PreparedDecoderInput {
hypergraph: base_hypergraph.clone(),
representatives: Arc::clone(&base_errors),
},
EdgeProjection::Identity,
)
};
(
DecodeProjection {
base_hypergraph,
base_errors,
decoder_errors: Arc::clone(&prepared.representatives),
error_edge_lookup,
edge_projection,
},
prepared,
)
}
fn edge_isolated_vertices(hypergraph: &blackbox_decoder::DecodingHypergraph) -> Vec<u64> {
let mut incident = vec![false; hypergraph.vertex_num as usize];
for hyperedge in &hypergraph.hyperedges {
for &vertex in &hyperedge.vertices {
incident[vertex as usize] = true;
}
}
incident
.into_iter()
.enumerate()
.filter_map(|(vertex, incident)| (!incident).then_some(vertex as u64))
.collect()
}
pub(crate) fn ignore_edge_isolated_history_vertices(
hypergraph: &blackbox_decoder::DecodingHypergraph,
syndrome: &mut BitVector,
) {
for vertex in edge_isolated_vertices(hypergraph) {
crate::misc::bit_vector::set_bit(syndrome, vertex, false);
}
}
pub(crate) fn probability_reweights<'a>(
error_reference: &[ErrorIndex],
modifiers: impl IntoIterator<Item = (usize, &'a bin::ProbabilityModifier)>,
) -> Vec<(u64, f64)> {
ErrorEdgeLookup::new(error_reference).project(modifiers)
}
#[derive(Debug)]
pub(crate) struct ErrorEdgeLookup {
edges_by_eid: hashbrown::HashMap<usize, Vec<u64>>,
}
impl ErrorEdgeLookup {
const MISSING_EDGE: u64 = u64::MAX;
pub(crate) fn new(error_reference: &[ErrorIndex]) -> Self {
let mut edges_by_eid = hashbrown::HashMap::<usize, Vec<u64>>::new();
for (edge, error) in error_reference.iter().enumerate() {
let edges = edges_by_eid.entry(error.eid).or_default();
if edges.len() <= error.error_index {
edges.resize(error.error_index + 1, Self::MISSING_EDGE);
}
edges[error.error_index] = u64::try_from(edge).unwrap();
}
Self { edges_by_eid }
}
pub(crate) fn project<'a>(
&self,
modifiers: impl IntoIterator<Item = (usize, &'a bin::ProbabilityModifier)>,
) -> Vec<(u64, f64)> {
let mut overrides = hashbrown::HashMap::new();
for (local_eid, modifier) in modifiers {
let Some(edges) = self.edges_by_eid.get(&local_eid) else {
continue;
};
for (error_index, &probability) in modifier.probabilities.iter().enumerate() {
if let Some(edge) = Self::edge(edges, error_index) {
overrides.insert(edge, probability);
}
}
for (&error_index, &probability) in modifier.sparse_indices.iter().zip(modifier.sparse_probabilities.iter()) {
if let Ok(error_index) = usize::try_from(error_index)
&& let Some(edge) = Self::edge(edges, error_index)
{
overrides.insert(edge, probability);
}
}
}
let mut reweights: Vec<_> = overrides.into_iter().collect();
reweights.sort_unstable_by_key(|&(edge, _)| edge);
reweights
}
fn edge(edges: &[u64], error_index: usize) -> Option<u64> {
edges.get(error_index).copied().filter(|&edge| edge != Self::MISSING_EDGE)
}
}
pub(crate) fn apply_reweights(hypergraph: &mut blackbox_decoder::DecodingHypergraph, reweights: &[(u64, f64)]) {
for &(edge, probability) in reweights {
hypergraph.hyperedges[edge as usize].probability = probability;
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
#[cfg_attr(feature = "cli", derive(structdoc::StructDoc))]
#[serde(rename_all = "snake_case")]
pub enum DecoderReweighting {
#[default]
Auto,
Enabled,
Disabled,
}
impl DecoderReweighting {
pub(crate) fn use_loaded(self, persistent_decoder: bool, features: DecoderFeatures) -> Result<bool, String> {
if !persistent_decoder {
return match self {
Self::Enabled => Err(
"decoder_reweighting is enabled but persistent_decoder is disabled; loaded reweights cannot be used"
.to_string(),
),
Self::Auto | Self::Disabled => Ok(false),
};
}
match self {
Self::Auto => Ok(features.contains(DecoderFeatures::REWEIGHTS)),
Self::Enabled if features.contains(DecoderFeatures::REWEIGHTS) => Ok(true),
Self::Enabled => Err("decoder_reweighting is enabled but the decoder does not support reweights".to_string()),
Self::Disabled => Ok(false),
}
}
}
#[derive(Debug, Clone)]
pub struct LoadedDecoder {
pub hid: u64,
pub decoding_hypergraph: Option<Arc<blackbox_decoder::DecodingHypergraph>>,
pub ignored_syndrome_vertices: Arc<Vec<u64>>,
pub projection: Arc<DecodeProjection>,
}
impl LoadedDecoder {
pub(crate) fn project_syndrome(&self, mut syndrome: BitVector) -> BitVector {
for &vertex in self.ignored_syndrome_vertices.iter() {
crate::misc::bit_vector::set_bit(&mut syndrome, vertex, false);
}
syndrome
}
}
pub(crate) async fn decode_projected(
decoder: &DynDecoder,
loaded: &LoadedDecoder,
syndrome: BitVector,
reweights: Vec<blackbox_decoder::EdgeReweight>,
loss: Option<blackbox_decoder::LossInfo>,
use_loaded_reweights: bool,
) -> Result<blackbox_decoder::ParityFactor, Status> {
if reweights.is_empty() || use_loaded_reweights {
return decoder
.decode_loaded(blackbox_decoder::LoadedDecodingProblem {
hid: loaded.hid,
syndrome: Some(syndrome),
reweights,
loss,
})
.await;
}
let mut hypergraph = (**loaded
.decoding_hypergraph
.as_ref()
.ok_or_else(|| Status::internal(format!("hid={} has no materializable hypergraph", loaded.hid)))?)
.clone();
for reweight in reweights {
let hyperedge = hypergraph.hyperedges.get_mut(reweight.edge as usize).ok_or_else(|| {
Status::invalid_argument(format!(
"reweighted edge {} is outside loaded hypergraph hid={}",
reweight.edge, loaded.hid
))
})?;
hyperedge.probability = reweight.probability;
}
decoder
.decode(blackbox_decoder::DecodingProblem {
hypergraph: Some(hypergraph),
syndrome: Some(syndrome),
loss,
})
.await
}
#[derive(Debug)]
pub struct DecodeProjection {
pub base_hypergraph: blackbox_decoder::DecodingHypergraph,
pub base_errors: Arc<Vec<ErrorIndex>>,
decoder_errors: Arc<Vec<ErrorIndex>>,
error_edge_lookup: ErrorEdgeLookup,
edge_projection: EdgeProjection,
}
#[derive(Debug, Clone)]
pub(crate) struct ProjectedErrors {
baseline: Arc<Vec<ErrorIndex>>,
replacements: Vec<(usize, ErrorIndex)>,
}
impl ProjectedErrors {
pub(crate) fn shared(baseline: Arc<Vec<ErrorIndex>>) -> Self {
Self {
baseline,
replacements: vec![],
}
}
pub(crate) fn len(&self) -> usize {
self.baseline.len()
}
}
impl Index<usize> for ProjectedErrors {
type Output = ErrorIndex;
fn index(&self, index: usize) -> &Self::Output {
match self
.replacements
.binary_search_by_key(&index, |&(decoder_edge, _)| decoder_edge)
{
Ok(position) => &self.replacements[position].1,
Err(_) => &self.baseline[index],
}
}
}
impl From<Arc<Vec<ErrorIndex>>> for ProjectedErrors {
fn from(baseline: Arc<Vec<ErrorIndex>>) -> Self {
Self::shared(baseline)
}
}
#[derive(Debug)]
pub(crate) struct PreparedDecoderInput {
pub(crate) hypergraph: blackbox_decoder::DecodingHypergraph,
pub(crate) representatives: Arc<Vec<ErrorIndex>>,
}
#[derive(Debug)]
enum EdgeProjection {
Identity,
Merged {
decoder_edge_of_original: Vec<usize>,
original_edges_of_decoder: Vec<Vec<usize>>,
},
}
fn deduplicate_by_syndrome(
hypergraph: &blackbox_decoder::DecodingHypergraph,
errors: &[ErrorIndex],
) -> (PreparedDecoderInput, EdgeProjection) {
let mut seen: hashbrown::HashMap<Vec<u64>, (usize, f64)> = hashbrown::HashMap::with_capacity(errors.len());
let mut hyperedges: Vec<blackbox_decoder::Hyperedge> = Vec::with_capacity(errors.len());
let mut representatives = Vec::with_capacity(errors.len());
let mut decoder_edge_of_original = Vec::with_capacity(errors.len());
let mut original_edges_of_decoder: Vec<Vec<usize>> = Vec::with_capacity(errors.len());
for (position, (hyperedge, error)) in hypergraph.hyperedges.iter().zip(errors.iter()).enumerate() {
let mut syndrome = hyperedge.vertices.clone();
syndrome.sort_unstable();
debug_assert!({
let degree = syndrome.len();
syndrome.dedup();
syndrome.len() == degree
});
if let Some((index, best_probability)) = seen.get_mut(&syndrome) {
let combined = hyperedges[*index].probability;
hyperedges[*index].probability = exclusive_probability_of(combined, hyperedge.probability);
if hyperedge.probability > *best_probability {
*best_probability = hyperedge.probability;
representatives[*index] = error.clone();
}
original_edges_of_decoder[*index].push(position);
decoder_edge_of_original.push(*index);
} else {
let index = representatives.len();
hyperedges.push(blackbox_decoder::Hyperedge {
probability: hyperedge.probability,
vertices: syndrome.clone(),
});
representatives.push(error.clone());
original_edges_of_decoder.push(vec![position]);
decoder_edge_of_original.push(index);
seen.insert(syndrome, (index, hyperedge.probability));
}
}
(
PreparedDecoderInput {
hypergraph: blackbox_decoder::DecodingHypergraph {
vertex_num: hypergraph.vertex_num,
hyperedges,
},
representatives: Arc::new(representatives),
},
EdgeProjection::Merged {
decoder_edge_of_original,
original_edges_of_decoder,
},
)
}
pub(crate) fn deduplicate_decoder_input(
hypergraph: &blackbox_decoder::DecodingHypergraph,
errors: &[ErrorIndex],
) -> PreparedDecoderInput {
deduplicate_by_syndrome(hypergraph, errors).0
}
impl DecodeProjection {
#[cfg(test)]
pub(crate) fn identity(
base_hypergraph: blackbox_decoder::DecodingHypergraph,
base_errors: Arc<Vec<ErrorIndex>>,
) -> Self {
Self {
base_hypergraph,
decoder_errors: Arc::clone(&base_errors),
error_edge_lookup: ErrorEdgeLookup::new(&base_errors),
base_errors,
edge_projection: EdgeProjection::Identity,
}
}
pub(crate) fn probability_reweights<'a>(
&self,
modifiers: impl IntoIterator<Item = (usize, &'a bin::ProbabilityModifier)>,
) -> Vec<(u64, f64)> {
self.error_edge_lookup.project(modifiers)
}
pub(crate) fn project_reweights(&self, reweights: &[(u64, f64)]) -> (Vec<(u64, f64)>, ProjectedErrors) {
self.edge_projection
.project_reweights(&self.base_hypergraph, &self.base_errors, &self.decoder_errors, reweights)
}
}
impl EdgeProjection {
fn project_reweights(
&self,
base_hypergraph: &blackbox_decoder::DecodingHypergraph,
base_errors: &[ErrorIndex],
decoder_errors: &Arc<Vec<ErrorIndex>>,
reweights: &[(u64, f64)],
) -> (Vec<(u64, f64)>, ProjectedErrors) {
match self {
Self::Identity => {
let mut overrides = hashbrown::HashMap::with_capacity(reweights.len());
for &(edge, probability) in reweights {
overrides.insert(edge, probability);
}
let mut translated: Vec<_> = overrides.into_iter().collect();
translated.sort_unstable_by_key(|&(edge, _)| edge);
(translated, ProjectedErrors::shared(Arc::clone(decoder_errors)))
}
Self::Merged {
decoder_edge_of_original,
original_edges_of_decoder,
} => {
let mut overrides = hashbrown::HashMap::with_capacity(reweights.len());
let mut reweighted_decoder_edges = Vec::with_capacity(reweights.len());
for &(edge, probability) in reweights {
let original = usize::try_from(edge).unwrap();
overrides.insert(original, probability);
reweighted_decoder_edges.push(decoder_edge_of_original[original]);
}
reweighted_decoder_edges.sort_unstable();
reweighted_decoder_edges.dedup();
if reweighted_decoder_edges.is_empty() {
return (vec![], ProjectedErrors::shared(Arc::clone(decoder_errors)));
}
let mut replacements = Vec::with_capacity(reweighted_decoder_edges.len());
let translated = reweighted_decoder_edges
.into_iter()
.map(|decoder_edge| {
let mut combined = 0.0;
let mut elected = None;
for &original_edge in &original_edges_of_decoder[decoder_edge] {
let probability = overrides
.get(&original_edge)
.copied()
.unwrap_or(base_hypergraph.hyperedges[original_edge].probability);
combined = exclusive_probability_of(combined, probability);
let should_elect = match elected {
None => true,
Some((_, elected_probability)) => probability > elected_probability,
};
if should_elect {
elected = Some((original_edge, probability));
}
}
let (elected_original, _) = elected.expect("decoder edge must contain an original edge");
if decoder_errors[decoder_edge] != base_errors[elected_original] {
replacements.push((decoder_edge, base_errors[elected_original].clone()));
}
(u64::try_from(decoder_edge).unwrap(), combined)
})
.collect();
(
translated,
ProjectedErrors {
baseline: Arc::clone(decoder_errors),
replacements,
},
)
}
}
}
}
#[cfg(test)]
#[path = "../../tests/unit/reweight_handler_test.rs"]
mod tests;