use crate::decoder::blackbox_decoder::{DecodingHypergraph, ParityFactor};
use crate::decoder::mwpf_decoder::{MwpfDecoderConfig, MwpfDecoderInstance, default_timeout};
use crate::decoder::thread_pooling::{
DecodeError, DecodeRequest, DecoderInstance, ThreadPoolingConfig, ThreadPoolingDecoder,
};
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 HufDecoderConfig {
#[serde(flatten)]
pub thread_pooling_config: ThreadPoolingConfig,
#[serde(default = "default_timeout")]
pub timeout: f64,
#[serde(default)]
pub only_solve_primal_once: bool,
}
impl HufDecoderConfig {
fn into_mwpf_config(self) -> MwpfDecoderConfig {
MwpfDecoderConfig {
thread_pooling_config: self.thread_pooling_config,
timeout: self.timeout,
cluster_node_limit: 0,
only_solve_primal_once: self.only_solve_primal_once,
}
}
}
pub struct HufDecoderInstance {
inner: MwpfDecoderInstance,
}
impl DecoderInstance for HufDecoderInstance {
fn validate_hypergraph(hypergraph: &DecodingHypergraph, config: &serde_json::Value) -> Result<(), String> {
let _: HufDecoderConfig = serde_json::from_value(config.clone()).map_err(|error| error.to_string())?;
MwpfDecoderInstance::validate_hypergraph(hypergraph, config)
}
fn new(hypergraph: &DecodingHypergraph, config: &serde_json::Value) -> Self {
let config: HufDecoderConfig = serde_json::from_value(config.clone()).unwrap();
Self {
inner: MwpfDecoderInstance::new_with_config(hypergraph, &config.into_mwpf_config()),
}
}
fn decode(&mut self, request: DecodeRequest<'_>) -> Result<ParityFactor, DecodeError> {
self.inner.decode(request)
}
fn reset(&mut self) {
self.inner.reset();
}
}
pub type HufDecoder = ThreadPoolingDecoder<HufDecoderInstance>;
#[cfg(test)]
mod tests {
use super::*;
use crate::decoder::blackbox_decoder::Hyperedge;
use crate::misc::bit_vector::from_sparse_indices;
#[test]
fn preserves_mwpf_tuning_except_for_fixed_cluster_limit() {
let config: HufDecoderConfig = serde_json::from_value(serde_json::json!({
"parallel": 3,
"timeout": 2.5,
"only_solve_primal_once": true,
}))
.unwrap();
let config = config.into_mwpf_config();
assert_eq!(config.thread_pooling_config.parallel, 3);
assert_eq!(config.timeout, 2.5);
assert_eq!(config.cluster_node_limit, 0);
assert!(config.only_solve_primal_once);
}
#[test]
fn decodes_parallel_edges_and_resets_without_panicking() {
let hypergraph = DecodingHypergraph {
vertex_num: 3,
hyperedges: vec![
Hyperedge {
vertices: vec![0, 1, 2],
probability: 0.01,
},
Hyperedge {
vertices: vec![2, 0, 1],
probability: 0.1,
},
],
};
let mut decoder = HufDecoderInstance::new(&hypergraph, &serde_json::json!({}));
let syndrome = from_sparse_indices(3, &[0, 1, 2]);
for _ in 0..2 {
let result = decoder
.decode(DecodeRequest {
syndrome: &syndrome,
decoder_seed: None,
reweights: &[],
loss: None,
})
.unwrap();
assert_eq!(result.subgraph, vec![1]);
decoder.reset();
}
}
}