diff --git a/Cargo.lock b/Cargo.lock index caf09cea5..98f1a9cb9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4782,6 +4782,7 @@ dependencies = [ "rand 0.10.2", "rayon", "serde_json", + "sha2 0.11.0", "tempfile", ] diff --git a/crates/pecos-decoder-core/src/dem.rs b/crates/pecos-decoder-core/src/dem.rs index 99e3fe233..9884c6ddf 100644 --- a/crates/pecos-decoder-core/src/dem.rs +++ b/crates/pecos-decoder-core/src/dem.rs @@ -5,6 +5,17 @@ use crate::errors::DecoderError; +fn dimension_count(max_index: Option, kind: &str) -> Result { + max_index.map_or(Ok(0), |index| { + let count = u64::from(index) + 1; + usize::try_from(count).map_err(|_| { + DecoderError::InvalidConfiguration(format!( + "{kind} count for index {index} does not fit usize on this platform" + )) + }) + }) +} + /// Trait for decoders that can be constructed from detector error models pub trait DemDecoder: super::Decoder { /// Configuration type for DEM construction @@ -173,8 +184,20 @@ pub mod utils { } } - let detector_count = max_detector.map_or(0, |m| m + 1); - let observable_count = max_observable.map_or(0, |m| m + 1); + let detector_count = max_detector.map_or(Ok(0), |index| { + index.checked_add(1).ok_or_else(|| { + DecoderError::InvalidConfiguration(format!( + "detector count for index {index} does not fit usize on this platform" + )) + }) + })?; + let observable_count = max_observable.map_or(Ok(0), |index| { + index.checked_add(1).ok_or_else(|| { + DecoderError::InvalidConfiguration(format!( + "observable count for index {index} does not fit usize on this platform" + )) + }) + })?; Ok((detector_count, observable_count)) } @@ -398,8 +421,8 @@ impl SparseDem { Ok(Self { mechanisms, detector_coords, - num_detectors: max_detector.map_or(0, |m| m as usize + 1), - num_observables: max_observable.map_or(0, |m| m as usize + 1), + num_detectors: dimension_count(max_detector, "detector")?, + num_observables: dimension_count(max_observable, "observable")?, }) } } @@ -552,8 +575,8 @@ impl DemCheckMatrix { mechanisms.push((probability, detectors, observables)); } - let num_detectors = max_detector.map_or(0, |m| m as usize + 1); - let num_observables = max_observable.map_or(0, |m| m as usize + 1); + let num_detectors = dimension_count(max_detector, "detector")?; + let num_observables = dimension_count(max_observable, "observable")?; let num_mechanisms = mechanisms.len(); // Build matrices. @@ -839,8 +862,8 @@ impl DemMatchingGraph { fault_id += 1; } - let num_detectors = max_detector.map_or(0, |m| m as usize + 1); - let num_observables = max_observable.map_or(0, |m| m as usize + 1); + let num_detectors = dimension_count(max_detector, "detector")?; + let num_observables = dimension_count(max_observable, "observable")?; let edges = Self::merge_parallel_edges(edges); @@ -1262,6 +1285,30 @@ mod tests { assert_eq!(dets, 8, "parse_dem_metadata must count bare detector D7"); } + #[test] + fn sparse_dem_max_u32_detector_index_is_platform_checked() { + let dem = "error(0.01) D4294967295\n"; + let parsed = SparseDem::from_dem_str(dem); + + #[cfg(target_pointer_width = "64")] + { + assert_eq!(parsed.unwrap().num_detectors, 4_294_967_296); + assert_eq!(utils::parse_dem_metadata(dem).unwrap().0, 4_294_967_296); + } + + #[cfg(target_pointer_width = "32")] + { + // On 32-bit targets, the promoted u64 count reaches the fallible + // usize::try_from branch instead of wrapping or panicking. + let error = parsed.unwrap_err(); + assert!(matches!(error, DecoderError::InvalidConfiguration(_))); + assert!(error.to_string().contains("4294967295")); + + let metadata_error = utils::parse_dem_metadata(dem).unwrap_err(); + assert!(metadata_error.to_string().contains("4294967295")); + } + } + #[test] fn test_parsers_reject_malformed_detector_token() { // A `D` / `L` token in an error line is malformed. All three diff --git a/python/pecos-rslib/Cargo.toml b/python/pecos-rslib/Cargo.toml index a252957ed..b52021117 100644 --- a/python/pecos-rslib/Cargo.toml +++ b/python/pecos-rslib/Cargo.toml @@ -105,6 +105,7 @@ nalgebra.workspace = true num-complex.workspace = true parking_lot.workspace = true serde_json.workspace = true +sha2.workspace = true tempfile.workspace = true log.workspace = true libc.workspace = true diff --git a/python/pecos-rslib/src/fault_tolerance_bindings.rs b/python/pecos-rslib/src/fault_tolerance_bindings.rs index fbb855d8b..a0475b6bf 100644 --- a/python/pecos-rslib/src/fault_tolerance_bindings.rs +++ b/python/pecos-rslib/src/fault_tolerance_bindings.rs @@ -74,11 +74,27 @@ use pyo3::prelude::*; use std::collections::BTreeMap; use std::str::FromStr; +mod decoder_comparison; +mod decoder_scoring; +mod sample_corpus; + +use decoder_comparison::{ + PyDecoderComparisonResult, compare_decoder_outcomes, validate_comparison_arguments, +}; +use decoder_scoring::{ + MaskedObservableDecoder, ShotDecodeError, TimedObservableDecoder, count_decoder_mismatches, +}; +use sample_corpus::{CorpusError, CorpusToSave, LoadedCorpus}; + type PyDemMechanismTuple = (f64, Vec, Vec); type PyDemFitResult = (Vec, Vec); /// Per-shot detector rows paired with per-shot observable/DEM-output rows. type PyDetectorObservableRows = (Vec>, Vec>); +fn map_shot_decode_error(error: ShotDecodeError) -> PyErr { + pyo3::exceptions::PyRuntimeError::new_err(error.to_string()) +} + fn parse_p1_weights(weights: BTreeMap) -> PyResult { use pecos_core::pauli::{X, Y, Z}; @@ -3368,6 +3384,11 @@ pub struct PySampleBatch { obs_columns: Vec>, num_detectors: usize, num_shots: usize, + seed: Option, + dem: Option, + metadata_json: Option, + generator: Option, + format_version: Option, } impl PySampleBatch { @@ -3440,6 +3461,7 @@ impl PySampleBatch { det_columns: Vec>, obs_columns: Vec>, num_shots: usize, + seed: Option, ) -> Self { let num_detectors = det_columns.len(); Self { @@ -3447,6 +3469,11 @@ impl PySampleBatch { obs_columns, num_detectors, num_shots, + seed, + dem: None, + metadata_json: None, + generator: None, + format_version: None, } } @@ -3493,6 +3520,54 @@ impl PySampleBatch { obs_columns, num_detectors, num_shots, + seed: None, + dem: None, + metadata_json: None, + generator: None, + format_version: None, + } + } + + fn from_corpus(corpus: LoadedCorpus) -> Self { + Self { + num_detectors: corpus.det_columns.len(), + det_columns: corpus.det_columns, + obs_columns: corpus.obs_columns, + num_shots: corpus.num_shots, + seed: corpus.seed, + dem: Some(corpus.dem), + metadata_json: corpus.metadata_json, + generator: Some(corpus.generator), + format_version: Some(corpus.format_version), + } + } + + fn ensure_dem_matches(&self, dem: &str, allow_dem_mismatch: bool) -> PyResult<()> { + if allow_dem_mismatch { + return Ok(()); + } + if let Some(embedded_dem) = &self.dem + && embedded_dem != dem + { + return Err(pyo3::exceptions::PyValueError::new_err( + "supplied DEM differs from the DEM embedded in this loaded SampleBatch; pass \ + allow_dem_mismatch=True to use a different model deliberately", + )); + } + Ok(()) + } + + fn map_corpus_error(error: CorpusError, path: &std::path::Path) -> PyErr { + match error { + CorpusError::Io(error) => match error.raw_os_error() { + Some(errno) => pyo3::exceptions::PyOSError::new_err(( + errno, + error.to_string(), + path.as_os_str().to_os_string(), + )), + None => error.into(), + }, + CorpusError::Invalid(message) => pyo3::exceptions::PyValueError::new_err(message), } } } @@ -3542,6 +3617,99 @@ impl PySampleBatch { self.num_shots } + /// Resolved random seed used to generate this batch, if known. + #[getter] + const fn seed(&self) -> Option { + self.seed + } + + /// Exact detector error model stored with a loaded corpus, if any. + #[getter] + fn dem(&self) -> Option<&str> { + self.dem.as_deref() + } + + /// Opaque caller metadata JSON stored with a loaded corpus, if any. + #[getter] + fn metadata_json(&self) -> Option<&str> { + self.metadata_json.as_deref() + } + + /// PECOS writer identity stored with a loaded corpus, if any. + #[getter] + fn generator(&self) -> Option<&str> { + self.generator.as_deref() + } + + /// Corpus format version for a loaded batch, if any. + #[getter] + const fn format_version(&self) -> Option { + self.format_version + } + + /// Save this serially captured shot batch as a self-describing corpus. + /// + /// Args: + /// path: Destination file path. + /// dem: DEM text associated with the samples. For a generated or + /// Python-constructed batch, only detector and observable dimensions + /// can be checked. This catches gross mismatches, but cannot prove DEM + /// identity or detect a different model with the same dimensions. For + /// a loaded corpus, the text must exactly match its embedded DEM unless + /// `allow_dem_mismatch` is true. + /// `metadata_json`: Optional syntactically valid JSON string. ``None`` + /// preserves metadata already carried by a loaded batch. A supplied + /// value replaces it. + /// `clear_metadata`: Explicitly omit metadata when true. Cannot be combined + /// with a supplied `metadata_json` value. + /// `allow_dem_mismatch`: Permit deliberately saving a loaded batch with a + /// DEM different from its embedded model. + /// + /// Corpora contain shots captured by the serial ``generate_samples`` path. + /// Parallel sample-and-decode paths discard individual shots and cannot be + /// captured by this API. + #[pyo3(signature = (path, *, dem, metadata_json=None, clear_metadata=false, allow_dem_mismatch=false))] + fn save( + &self, + path: std::path::PathBuf, + dem: &str, + metadata_json: Option<&str>, + clear_metadata: bool, + allow_dem_mismatch: bool, + ) -> PyResult<()> { + self.ensure_dem_matches(dem, allow_dem_mismatch)?; + if clear_metadata && metadata_json.is_some() { + return Err(pyo3::exceptions::PyValueError::new_err( + "metadata_json and clear_metadata=True are mutually exclusive", + )); + } + let metadata_json = if clear_metadata { + None + } else { + metadata_json.or(self.metadata_json.as_deref()) + }; + sample_corpus::save( + &path, + CorpusToSave { + det_columns: &self.det_columns, + obs_columns: &self.obs_columns, + num_shots: self.num_shots, + seed: self.seed, + dem, + metadata_json, + }, + ) + .map_err(|error| Self::map_corpus_error(error, &path)) + } + + /// Load and validate a self-describing shot corpus. + #[staticmethod] + fn load(path: std::path::PathBuf) -> PyResult { + sample_corpus::load(&path) + .map(Self::from_corpus) + .map_err(|error| Self::map_corpus_error(error, &path)) + } + /// Get the syndrome for shot `i` as a list of u8 values. fn get_syndrome(&self, i: usize) -> PyResult> { if i >= self.num_shots { @@ -3588,27 +3756,34 @@ impl PySampleBatch { /// `decoder_type`: "pymatching", "`pymatching_correlated`", /// "`pymatching_uncorrelated`", "tesseract", "`bp_osd`", /// "`bp_lsd`", "`union_find`", "`relay_bp`", or "`min_sum_bp`". + /// `allow_dem_mismatch`: Permit a DEM different from the one embedded in + /// a loaded corpus. /// /// Returns: /// Number of logical errors. - #[pyo3(signature = (dem, decoder_type="pymatching"))] - fn decode_count(&self, dem: &str, decoder_type: &str) -> PyResult { + /// + /// Decoder failures abort with a `RuntimeError`. Callers that need errors + /// reported per shot as a distinct outcome should use `compare_decoders`. + #[pyo3(signature = (dem, decoder_type="pymatching", *, allow_dem_mismatch=false))] + fn decode_count( + &self, + dem: &str, + decoder_type: &str, + allow_dem_mismatch: bool, + ) -> PyResult { + self.ensure_dem_matches(dem, allow_dem_mismatch)?; let mut decoder = create_observable_decoder(dem, decoder_type)?; - let mut errors = 0usize; let mut syndrome = vec![0u8; self.num_detectors]; - for i in 0..self.num_shots { - self.extract_syndrome(i, &mut syndrome); - // Wide ObsMask comparison: inline (one stack word) for the typical - // <=64 observables, correct without truncation beyond. A decode - // failure counts as a logical error (matching the prior sentinel). - let is_error = decoder - .decode_obs(&syndrome) - .map_or(true, |p| p != self.extract_obs_mask_wide(i)); - if is_error { - errors += 1; - } - } - Ok(errors) + count_decoder_mismatches( + 0..self.num_shots, + &mut syndrome, + |shot, buffer| { + self.extract_syndrome(shot, buffer); + self.extract_obs_mask_wide(shot) + }, + decoder.as_mut(), + ) + .map_err(map_shot_decode_error) } /// Decode every shot and return the predicted observable mask per shot. @@ -3620,32 +3795,86 @@ impl PySampleBatch { /// Args: /// dem: DEM string for the decoder. /// `decoder_type`: Decoder type string. + /// `allow_dem_mismatch`: Permit a DEM different from the one embedded in + /// a loaded corpus. /// /// Returns: /// List of predicted observable masks (Python ints; arbitrary precision, /// so more than 64 observables are not truncated), one per shot. - #[pyo3(signature = (dem, decoder_type="pymatching"))] + #[pyo3(signature = (dem, decoder_type="pymatching", *, allow_dem_mismatch=false))] fn decode_each( &self, py: Python<'_>, dem: &str, decoder_type: &str, + allow_dem_mismatch: bool, ) -> PyResult>> { + self.ensure_dem_matches(dem, allow_dem_mismatch)?; let mut decoder = create_observable_decoder(dem, decoder_type)?; let mut predictions = Vec::with_capacity(self.num_shots); let mut syndrome = vec![0u8; self.num_detectors]; - for i in 0..self.num_shots { - self.extract_syndrome(i, &mut syndrome); + for shot in 0..self.num_shots { + self.extract_syndrome(shot, &mut syndrome); // Propagate a decode failure rather than masking it as a sentinel // observable value (which would read as a spurious disagreement). - let predicted = decoder - .decode_obs(&syndrome) - .map_err(|e| PyErr::new::(e.to_string()))?; + let predicted = decoder.decode_obs(&syndrome).map_err(|error| { + pyo3::exceptions::PyRuntimeError::new_err(format!( + "decoder failed on shot {shot}: {error}" + )) + })?; predictions.push(obsmask_to_py(py, &predicted)?); } Ok(predictions) } + /// Decode every shot with a decoder under test (DUT) and a reference decoder. + /// + /// Both decoders receive the same shots in the same order. Each result is + /// independently classified as correct, mismatch, or decode error, and a + /// decode error is counted for that shot without aborting the comparison. + /// Predictions and truth are compared as wide observable masks, with no + /// 64-observable limit. + /// + /// Args: + /// dem: DEM string shared by both decoders. + /// `dut_decoder_type`: Decoder type string for the decoder under test. + /// `reference_decoder_type`: Decoder type string for the reference. + /// alpha: Tail probability for equal-tailed Jeffreys intervals. + /// `allow_dem_mismatch`: Permit deliberate cross-model comparison of a + /// loaded corpus. + /// + /// Returns: + /// A `DecoderComparisonResult` containing the raw 3x3 counts and + /// headline DUT-only-failure and both-failed proportions. + #[pyo3(signature = (dem, dut_decoder_type, reference_decoder_type, alpha=0.05, *, allow_dem_mismatch=false))] + fn compare_decoders( + &self, + dem: &str, + dut_decoder_type: &str, + reference_decoder_type: &str, + alpha: f64, + allow_dem_mismatch: bool, + ) -> PyResult { + validate_comparison_arguments(self.num_shots, alpha) + .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?; + self.ensure_dem_matches(dem, allow_dem_mismatch)?; + let mut dut = create_observable_decoder(dem, dut_decoder_type)?; + let mut reference = create_observable_decoder(dem, reference_decoder_type)?; + let mut syndrome = vec![0u8; self.num_detectors]; + let counts = compare_decoder_outcomes( + self.num_shots, + &mut syndrome, + |shot, buffer| { + self.extract_syndrome(shot, buffer); + self.extract_obs_mask_wide(shot) + }, + dut.as_mut(), + reference.as_mut(), + ); + PyDecoderComparisonResult::new(counts, alpha) + .map_err(|error| pyo3::exceptions::PyRuntimeError::new_err(error.to_string())) + } + /// Parallel decode: distributes samples across rayon workers. /// /// Each worker creates its own decoder instance. Faster for slow decoders. @@ -3654,18 +3883,29 @@ impl PySampleBatch { /// dem: DEM string for the decoder. /// `decoder_type`: Decoder type string. /// `num_workers`: Number of parallel workers (default: number of CPUs). + /// `allow_dem_mismatch`: Permit a DEM different from the one embedded in + /// a loaded corpus. /// /// Returns: /// Number of logical errors. - #[pyo3(signature = (dem, decoder_type="pymatching", num_workers=None))] + /// + /// Decoder failures abort with a `RuntimeError`. Callers that need errors + /// reported per shot as a distinct outcome should use `compare_decoders`. + /// + /// Set `allow_dem_mismatch` to true to use a DEM different from the one + /// embedded in a loaded corpus. + #[pyo3(signature = (dem, decoder_type="pymatching", num_workers=None, *, allow_dem_mismatch=false))] fn decode_count_parallel( &self, dem: &str, decoder_type: &str, num_workers: Option, + allow_dem_mismatch: bool, ) -> PyResult { use rayon::prelude::*; + self.ensure_dem_matches(dem, allow_dem_mismatch)?; + drop(create_observable_decoder(dem, decoder_type)?); let n_workers = num_workers.unwrap_or_else(rayon::current_num_threads); let pool = rayon::ThreadPoolBuilder::new() .num_threads(n_workers) @@ -3688,23 +3928,40 @@ impl PySampleBatch { let observable_masks: Vec = (0..n).map(|i| self.extract_obs_mask_wide(i)).collect(); - let total_errors: usize = pool.install(|| { - (0..n) + let worker_results: Vec> = pool.install(|| { + let chunk_size = n.div_ceil(n_workers); + (0..n_workers) .into_par_iter() - .map_init( - || create_observable_decoder(&dem_str, &dt).unwrap(), - |decoder, i| { - usize::from( - decoder - .decode_obs(&detection_events[i]) - .map_or(true, |p| p != observable_masks[i]), - ) - }, - ) - .sum() + .map(|worker_id| { + let start = worker_id * chunk_size; + let end = (start + chunk_size).min(n); + if start >= end { + return Ok(0); + } + + // Safe after the identical factory call was validated above. + let mut decoder = create_observable_decoder(&dem_str, &dt).unwrap(); + let mut syndrome = vec![0u8; num_dets]; + count_decoder_mismatches( + start..end, + &mut syndrome, + |shot, buffer| { + buffer.copy_from_slice(&detection_events[shot]); + observable_masks[shot].clone() + }, + decoder.as_mut(), + ) + }) + .collect() }); - Ok(total_errors) + worker_results + .into_iter() + .try_fold(0usize, |total, result| { + result + .map(|count| total + count) + .map_err(map_shot_decode_error) + }) } /// Batch decode all samples at once using `PyMatching`'s batch API. @@ -3715,10 +3972,15 @@ impl PySampleBatch { /// /// Returns: /// Number of logical errors. - #[pyo3(signature = (dem))] - fn decode_count_batch(&self, dem: &str) -> PyResult { + /// + /// Batch decoder failures abort with a `RuntimeError`. Callers that need + /// errors reported per shot as a distinct outcome should use + /// `compare_decoders`. + #[pyo3(signature = (dem, *, allow_dem_mismatch=false))] + fn decode_count_batch(&self, dem: &str, allow_dem_mismatch: bool) -> PyResult { use pecos_decoders::{BatchConfig, PyMatchingDecoder}; + self.ensure_dem_matches(dem, allow_dem_mismatch)?; let mut decoder = PyMatchingDecoder::from_dem(dem) .map_err(|e| PyErr::new::(e.to_string()))?; @@ -3774,28 +4036,36 @@ impl PySampleBatch { /// Args: /// dem: DEM string for the decoder. /// `decoder_type`: Decoder type string. + /// `allow_dem_mismatch`: Permit a DEM different from the one embedded in + /// a loaded corpus. /// /// Returns: /// `DecodeStats` with timing breakdown. - #[pyo3(signature = (dem, decoder_type="pymatching"))] - fn decode_stats(&self, dem: &str, decoder_type: &str) -> PyResult { - use std::time::Instant; - - let mut decoder = create_observable_decoder(dem, decoder_type)?; - let mut num_errors = 0usize; - let mut per_shot_seconds: Vec = Vec::with_capacity(self.num_shots); + /// + /// Decoder failures abort with a `RuntimeError`. Callers that need errors + /// reported per shot as a distinct outcome should use `compare_decoders`. + #[pyo3(signature = (dem, decoder_type="pymatching", *, allow_dem_mismatch=false))] + fn decode_stats( + &self, + dem: &str, + decoder_type: &str, + allow_dem_mismatch: bool, + ) -> PyResult { + self.ensure_dem_matches(dem, allow_dem_mismatch)?; + let decoder = create_observable_decoder(dem, decoder_type)?; + let mut decoder = TimedObservableDecoder::new(decoder, self.num_shots); let mut syndrome = vec![0u8; self.num_detectors]; - - for i in 0..self.num_shots { - self.extract_syndrome(i, &mut syndrome); - let t0 = Instant::now(); - let predicted = decoder.decode_obs(&syndrome); - let elapsed = t0.elapsed().as_secs_f64(); - per_shot_seconds.push(elapsed); - if predicted.map_or(true, |p| p != self.extract_obs_mask_wide(i)) { - num_errors += 1; - } - } + let num_errors = count_decoder_mismatches( + 0..self.num_shots, + &mut syndrome, + |shot, buffer| { + self.extract_syndrome(shot, buffer); + self.extract_obs_mask_wide(shot) + }, + &mut decoder, + ) + .map_err(map_shot_decode_error)?; + let per_shot_seconds = decoder.into_times(); Ok(PyDecodeStats::from_times( self.num_shots, @@ -3817,15 +4087,22 @@ impl PySampleBatch { /// dem: DEM string for the decoder. /// `decoder_type`: Decoder type string. /// `num_workers`: Number of parallel workers (default: number of CPUs). - #[pyo3(signature = (dem, decoder_type="mwpf", num_workers=None))] + /// `allow_dem_mismatch`: Permit a DEM different from the one embedded in + /// a loaded corpus. + /// + /// Decoder failures abort with a `RuntimeError`. Callers that need errors + /// reported per shot as a distinct outcome should use `compare_decoders`. + #[pyo3(signature = (dem, decoder_type="mwpf", num_workers=None, *, allow_dem_mismatch=false))] fn decode_stats_parallel( &self, dem: &str, decoder_type: &str, num_workers: Option, + allow_dem_mismatch: bool, ) -> PyResult { use rayon::prelude::*; + self.ensure_dem_matches(dem, allow_dem_mismatch)?; let n_workers = num_workers.unwrap_or_else(rayon::current_num_threads); // Validate decoder type early. @@ -3852,8 +4129,8 @@ impl PySampleBatch { .map(|i| self.extract_obs_mask_wide(i)) .collect(); - // Each worker decodes a slice of shots and returns (errors, per_shot_times). - let results: Vec<(usize, Vec)> = pool.install(|| { + // Each worker decodes a contiguous slice and returns its count and timings. + let results: Vec), ShotDecodeError>> = pool.install(|| { let chunk_size = self.num_shots.div_ceil(n_workers); (0..n_workers) .into_par_iter() @@ -3861,29 +4138,31 @@ impl PySampleBatch { let start = worker_id * chunk_size; let end = (start + chunk_size).min(self.num_shots); if start >= end { - return (0, Vec::new()); + return Ok((0, Vec::new())); } - let mut decoder = create_observable_decoder(&dem_str, &dt).unwrap(); - let mut errors = 0usize; - let mut times = Vec::with_capacity(end - start); - - for i in start..end { - let t0 = std::time::Instant::now(); - let predicted = decoder.decode_obs(&detection_events[i]); - times.push(t0.elapsed().as_secs_f64()); - if predicted.map_or(true, |p| p != observable_masks[i]) { - errors += 1; - } - } - (errors, times) + // Safe after the identical factory call was validated above. + let decoder = create_observable_decoder(&dem_str, &dt).unwrap(); + let mut decoder = TimedObservableDecoder::new(decoder, end - start); + let mut syndrome = vec![0u8; num_dets]; + let errors = count_decoder_mismatches( + start..end, + &mut syndrome, + |shot, buffer| { + buffer.copy_from_slice(&detection_events[shot]); + observable_masks[shot].clone() + }, + &mut decoder, + )?; + Ok((errors, decoder.into_times())) }) .collect() }); let mut total_errors = 0usize; let mut all_times = Vec::with_capacity(self.num_shots); - for (errs, times) in results { + for result in results { + let (errs, times) = result.map_err(map_shot_decode_error)?; total_errors += errs; all_times.extend(times); } @@ -4488,15 +4767,13 @@ impl PyDemSampler { use pecos_random::PecosRng; use rand::RngExt; - let mut rng = match seed { - Some(s) => PecosRng::seed_from_u64(s), - None => PecosRng::seed_from_u64(rand::rng().random()), - }; + let actual_seed = seed.unwrap_or_else(|| rand::rng().random()); + let mut rng = PecosRng::seed_from_u64(actual_seed); // Use geometric columnar sampler via DemSampler. let (det_columns, obs_columns) = self.inner.sample_batch_geometric(num_shots, &mut rng); - PySampleBatch::from_columnar(det_columns, obs_columns, num_shots) + PySampleBatch::from_columnar(det_columns, obs_columns, num_shots, Some(actual_seed)) } /// Compute statistics without storing individual shots. @@ -4609,6 +4886,10 @@ impl PyDemSampler { /// /// Returns: /// Number of logical errors (mismatches between decoder prediction and true flip). + /// + /// Decoder failures abort with a `RuntimeError`. Callers that need errors + /// reported per shot as a distinct outcome should generate a `SampleBatch` + /// and use `SampleBatch.compare_decoders`. #[pyo3(signature = (dem, num_shots, decoder_type="pymatching", seed=None))] fn sample_decode_count( &self, @@ -4623,27 +4904,28 @@ impl PyDemSampler { let actual_seed = seed.unwrap_or_else(|| rand::rng().random()); let mut rng = PecosRng::seed_from_u64(actual_seed); - let mut decoder = create_observable_decoder(dem, decoder_type)?; + let decoder = create_observable_decoder(dem, decoder_type)?; let observable_mask = self.inner.observable_dem_output_mask(); + let mut decoder = MaskedObservableDecoder::new(decoder, observable_mask.clone()); // Tight sample+decode loop -- no Python involvement. // Single-threaded: sample and decode sequentially. - let mut errors = 0usize; - for _ in 0..num_shots { - let (det_events, obs_flips) = self.inner.sample(&mut rng); - let syndrome: Vec = det_events.iter().map(|&b| u8::from(b)).collect(); - let mut predicted = decoder - .decode_obs(&syndrome) - .map_err(|e| PyErr::new::(e.to_string()))?; - predicted &= &observable_mask; - let true_mask = self - .inner - .observable_mask_from_dem_output_flips(&obs_flips, &observable_mask); - if predicted != true_mask { - errors += 1; - } - } - Ok(errors) + let mut syndrome = vec![0u8; self.inner.num_detectors()]; + count_decoder_mismatches( + 0..num_shots, + &mut syndrome, + |_, buffer| { + let (det_events, obs_flips) = self.inner.sample(&mut rng); + debug_assert_eq!(det_events.len(), buffer.len()); + for (value, event) in buffer.iter_mut().zip(det_events) { + *value = u8::from(event); + } + self.inner + .observable_mask_from_dem_output_flips(&obs_flips, &observable_mask) + }, + &mut decoder, + ) + .map_err(map_shot_decode_error) } /// Parallel sample+decode: distributes shots across threads. @@ -4662,6 +4944,10 @@ impl PyDemSampler { /// /// Returns: /// Number of logical errors. + /// + /// Decoder failures abort with a `RuntimeError`. Callers that need errors + /// reported per shot as a distinct outcome should generate a `SampleBatch` + /// and use `SampleBatch.compare_decoders`. #[pyo3(signature = (dem, num_shots, decoder_type="pymatching", seed=None, num_workers=None))] fn sample_decode_count_parallel( &self, @@ -4692,7 +4978,7 @@ impl PyDemSampler { let dem_str = dem.to_string(); let dt = decoder_type.to_string(); - let total_errors: usize = pool.install(|| { + let worker_results: Vec> = pool.install(|| { (0..n_workers) .into_par_iter() .map(|worker_id| { @@ -4700,35 +4986,44 @@ impl PyDemSampler { let my_shots = shots_per_worker + usize::from(worker_id < remainder); if my_shots == 0 { - return 0; + return Ok(0); } + let start = worker_id * shots_per_worker + worker_id.min(remainder); + let end = start + my_shots; let my_sampler = sampler.clone(); let mut my_rng = PecosRng::seed_from_u64(actual_seed.wrapping_add(worker_id as u64)); // unwrap is safe: we validated above - let mut decoder = create_observable_decoder(&dem_str, &dt).unwrap(); - - let mut errors = 0usize; - for _ in 0..my_shots { - let (det_events, obs_flips) = my_sampler.sample(&mut my_rng); - let syndrome: Vec = det_events.iter().map(|&b| u8::from(b)).collect(); - let mut predicted = decoder - .decode_obs(&syndrome) - .unwrap_or_else(|_| observable_mask.clone()); - predicted &= &observable_mask; - let truth = my_sampler - .observable_mask_from_dem_output_flips(&obs_flips, &observable_mask); - if predicted != truth { - errors += 1; - } - } - errors + let decoder = create_observable_decoder(&dem_str, &dt).unwrap(); + let mut decoder = + MaskedObservableDecoder::new(decoder, observable_mask.clone()); + let mut syndrome = vec![0u8; my_sampler.num_detectors()]; + count_decoder_mismatches( + start..end, + &mut syndrome, + |_, buffer| { + let (det_events, obs_flips) = my_sampler.sample(&mut my_rng); + debug_assert_eq!(det_events.len(), buffer.len()); + for (value, event) in buffer.iter_mut().zip(det_events) { + *value = u8::from(event); + } + my_sampler + .observable_mask_from_dem_output_flips(&obs_flips, &observable_mask) + }, + &mut decoder, + ) }) - .sum() + .collect() }); - Ok(total_errors) + worker_results + .into_iter() + .try_fold(0usize, |total, result| { + result + .map(|count| total + count) + .map_err(map_shot_decode_error) + }) } fn __repr__(&self) -> String { @@ -6940,6 +7235,7 @@ pub fn register_qec_module(m: &Bound<'_, PyModule>) -> PyResult<()> { qec.add_class::()?; qec.add_class::()?; qec.add_class::()?; + qec.add_class::()?; qec.add_class::()?; qec.add_class::()?; qec.add_class::()?; diff --git a/python/pecos-rslib/src/fault_tolerance_bindings/decoder_comparison.rs b/python/pecos-rslib/src/fault_tolerance_bindings/decoder_comparison.rs new file mode 100644 index 000000000..7f04031ca --- /dev/null +++ b/python/pecos-rslib/src/fault_tolerance_bindings/decoder_comparison.rs @@ -0,0 +1,486 @@ +// Copyright 2026 The PECOS Developers +// +// Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except +// in compliance with the License. You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software distributed under the License +// is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express +// or implied. See the License for the specific language governing permissions and limitations under +// the License. + +//! Paired DUT/reference decoder comparison over a shared sequence of shots. + +use pecos_decoder_core::obs_mask::ObsMask; +use pecos_decoder_core::{DecoderError, ObservableDecoder}; +use pecos_num::stats::{ + JeffreysError, JeffreysEstimator, JeffreysInterval, jeffreys_interval, jeffreys_point, +}; +use pyo3::prelude::*; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum DecoderOutcome { + Correct, + Mismatch, + Error, +} + +impl DecoderOutcome { + const fn index(self) -> usize { + match self { + Self::Correct => 0, + Self::Mismatch => 1, + Self::Error => 2, + } + } +} + +fn classify(result: Result, truth: &ObsMask) -> DecoderOutcome { + match result { + Ok(prediction) if prediction == *truth => DecoderOutcome::Correct, + Ok(_) => DecoderOutcome::Mismatch, + Err(_) => DecoderOutcome::Error, + } +} + +/// Validate caller-controlled comparison arguments without running an interval +/// calculation or constructing either decoder. +pub(super) fn validate_comparison_arguments( + num_shots: usize, + alpha: f64, +) -> Result<(), JeffreysError> { + let num_shots = u64::try_from(num_shots).expect("usize always fits in u64"); + // The mean estimator performs the shared zero/maximum-trial validation but + // no special-function solve, keeping this preflight constant-time. + jeffreys_point(0, num_shots, JeffreysEstimator::Mean)?; + if !alpha.is_finite() || alpha <= 0.0 || alpha >= 1.0 { + return Err(JeffreysError::InvalidAlpha { alpha }); + } + Ok(()) +} + +/// Counts indexed by DUT outcome first, then reference outcome. +/// +/// In each dimension the order is correct, mismatch, decode error. +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub(super) struct DecoderComparisonCounts { + cells: [[u64; 3]; 3], +} + +impl DecoderComparisonCounts { + fn record(&mut self, dut: DecoderOutcome, reference: DecoderOutcome) { + self.cells[dut.index()][reference.index()] += 1; + } + + pub(super) const fn cells(&self) -> &[[u64; 3]; 3] { + &self.cells + } + + fn total_shots(&self) -> u64 { + self.cells.iter().flatten().sum() + } + + const fn dut_only_failures(&self) -> u64 { + self.cells[DecoderOutcome::Mismatch.index()][DecoderOutcome::Correct.index()] + } + + const fn both_failed(&self) -> u64 { + self.cells[DecoderOutcome::Mismatch.index()][DecoderOutcome::Mismatch.index()] + } +} + +/// Compare two decoders on the same shots in the same order. +/// +/// `prepare_shot` writes the selected syndrome into the reusable buffer and +/// returns that shot's wide true-observable mask. +pub(super) fn compare_decoder_outcomes( + num_shots: usize, + syndrome: &mut [u8], + mut prepare_shot: impl FnMut(usize, &mut [u8]) -> ObsMask, + dut: &mut dyn ObservableDecoder, + reference: &mut dyn ObservableDecoder, +) -> DecoderComparisonCounts { + let mut counts = DecoderComparisonCounts::default(); + for shot in 0..num_shots { + let truth = prepare_shot(shot, syndrome); + // Run both decoders before classifying either result. In particular, a + // DUT error must not prevent the reference from seeing this shot. + let dut_result = dut.decode_obs(syndrome); + let reference_result = reference.decode_obs(syndrome); + counts.record( + classify(dut_result, &truth), + classify(reference_result, &truth), + ); + } + counts +} + +#[derive(Clone, Copy, Debug)] +struct HeadlineProportion { + point: f64, + interval: JeffreysInterval, +} + +impl HeadlineProportion { + fn new(count: u64, total_shots: u64, alpha: f64) -> Result { + let interval = jeffreys_interval(count, total_shots, alpha)?; + Ok(Self { + point: interval.point, + interval, + }) + } +} + +/// Python-facing paired decoder contingency counts and headline proportions. +#[pyclass( + name = "DecoderComparisonResult", + module = "pecos_rslib.qec", + skip_from_py_object +)] +#[derive(Clone, Debug)] +pub(super) struct PyDecoderComparisonResult { + counts: DecoderComparisonCounts, + total_shots: u64, + alpha: f64, + dut_only_failure: HeadlineProportion, + both_failed: HeadlineProportion, +} + +impl PyDecoderComparisonResult { + pub(super) fn new(counts: DecoderComparisonCounts, alpha: f64) -> Result { + let total_shots = counts.total_shots(); + let dut_only_failure = + HeadlineProportion::new(counts.dut_only_failures(), total_shots, alpha)?; + let both_failed = HeadlineProportion::new(counts.both_failed(), total_shots, alpha)?; + Ok(Self { + counts, + total_shots, + alpha, + dut_only_failure, + both_failed, + }) + } +} + +#[pymethods] +impl PyDecoderComparisonResult { + /// Raw 3x3 counts in correct, mismatch, error order on both axes. + #[getter] + fn counts(&self) -> Vec> { + self.counts.cells().iter().map(|row| row.to_vec()).collect() + } + + /// Number of shots compared. + #[getter] + const fn total_shots(&self) -> u64 { + self.total_shots + } + + /// Tail probability used for the equal-tailed Jeffreys intervals. + #[getter] + const fn alpha(&self) -> f64 { + self.alpha + } + + #[getter] + const fn dut_correct_reference_correct(&self) -> u64 { + self.counts.cells[0][0] + } + + #[getter] + const fn dut_correct_reference_mismatch(&self) -> u64 { + self.counts.cells[0][1] + } + + #[getter] + const fn dut_correct_reference_error(&self) -> u64 { + self.counts.cells[0][2] + } + + #[getter] + const fn dut_mismatch_reference_correct(&self) -> u64 { + self.counts.cells[1][0] + } + + #[getter] + const fn dut_mismatch_reference_mismatch(&self) -> u64 { + self.counts.cells[1][1] + } + + #[getter] + const fn dut_mismatch_reference_error(&self) -> u64 { + self.counts.cells[1][2] + } + + #[getter] + const fn dut_error_reference_correct(&self) -> u64 { + self.counts.cells[2][0] + } + + #[getter] + const fn dut_error_reference_mismatch(&self) -> u64 { + self.counts.cells[2][1] + } + + #[getter] + const fn dut_error_reference_error(&self) -> u64 { + self.counts.cells[2][2] + } + + /// DUT mismatches on shots where the reference was correct. + #[getter] + const fn dut_only_failures(&self) -> u64 { + self.counts.dut_only_failures() + } + + /// Jeffreys posterior-mean proportion for DUT-only failures. + #[getter] + const fn dut_only_failure_proportion(&self) -> f64 { + self.dut_only_failure.point + } + + /// Equal-tailed Jeffreys interval for the DUT-only-failure proportion. + #[getter] + const fn dut_only_failure_interval(&self) -> (f64, f64) { + ( + self.dut_only_failure.interval.lo, + self.dut_only_failure.interval.hi, + ) + } + + /// Shots on which both decoders returned mismatching predictions. + #[getter] + const fn both_failed(&self) -> u64 { + self.counts.both_failed() + } + + /// Jeffreys posterior-mean proportion for shots where both decoders failed. + #[getter] + const fn both_failed_proportion(&self) -> f64 { + self.both_failed.point + } + + /// Equal-tailed Jeffreys interval for the both-failed proportion. + #[getter] + const fn both_failed_interval(&self) -> (f64, f64) { + (self.both_failed.interval.lo, self.both_failed.interval.hi) + } + + fn __repr__(&self) -> String { + format!( + "DecoderComparisonResult(shots={}, dut_only_failures={}, both_failed={})", + self.total_shots, + self.counts.dut_only_failures(), + self.counts.both_failed(), + ) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[derive(Clone, Debug)] + enum StubResult { + Prediction(ObsMask), + Error, + } + + struct StubDecoder { + expected_syndromes: Vec>, + results: Vec, + next: usize, + } + + impl StubDecoder { + fn new(expected_syndromes: &[Vec], results: Vec) -> Self { + assert_eq!(expected_syndromes.len(), results.len()); + Self { + expected_syndromes: expected_syndromes.to_vec(), + results, + next: 0, + } + } + } + + impl ObservableDecoder for StubDecoder { + fn decode_obs(&mut self, syndrome: &[u8]) -> Result { + assert_eq!(syndrome, self.expected_syndromes[self.next]); + let result = match &self.results[self.next] { + StubResult::Prediction(mask) => Ok(mask.clone()), + StubResult::Error => Err(DecoderError::DecodingFailed("stub error".into())), + }; + self.next += 1; + result + } + } + + fn mask(bits: &[usize]) -> ObsMask { + let mut mask = ObsMask::new(); + for &bit in bits { + mask.set(bit); + } + mask + } + + fn predictions(masks: &[ObsMask]) -> Vec { + masks.iter().cloned().map(StubResult::Prediction).collect() + } + + fn compare( + shots: &[(Vec, ObsMask)], + dut_results: Vec, + reference_results: Vec, + ) -> DecoderComparisonCounts { + let syndromes: Vec> = shots.iter().map(|(s, _)| s.clone()).collect(); + let mut dut = StubDecoder::new(&syndromes, dut_results); + let mut reference = StubDecoder::new(&syndromes, reference_results); + let mut syndrome = vec![0; syndromes.first().map_or(0, Vec::len)]; + compare_decoder_outcomes( + shots.len(), + &mut syndrome, + |shot, buffer| { + buffer.copy_from_slice(&shots[shot].0); + shots[shot].1.clone() + }, + &mut dut, + &mut reference, + ) + } + + fn sample_shots() -> Vec<(Vec, ObsMask)> { + vec![ + (vec![0, 0], mask(&[])), + (vec![1, 0], mask(&[0])), + (vec![0, 1], mask(&[1])), + (vec![1, 1], mask(&[0, 1])), + ] + } + + #[test] + fn both_decoders_correct_puts_all_mass_in_correct_correct() { + let shots = sample_shots(); + let truths: Vec = shots.iter().map(|(_, truth)| truth.clone()).collect(); + let counts = compare(&shots, predictions(&truths), predictions(&truths)); + + assert_eq!(counts.cells(), &[[4, 0, 0], [0, 0, 0], [0, 0, 0]]); + assert_eq!(counts.dut_only_failures(), 0); + } + + #[test] + fn dut_only_failures_count_a_known_wrong_subset() { + let shots = sample_shots(); + let truths: Vec = shots.iter().map(|(_, truth)| truth.clone()).collect(); + let mut dut = truths.clone(); + dut[1] = mask(&[]); + dut[3] = mask(&[]); + + let counts = compare(&shots, predictions(&dut), predictions(&truths)); + + // Shots 1 and 3 are deliberately wrong for the DUT: 2 DUT-only failures. + assert_eq!(counts.dut_only_failures(), 2); + assert_eq!(counts.cells(), &[[2, 0, 0], [2, 0, 0], [0, 0, 0]]); + } + + #[test] + fn dut_errors_are_not_mismatches_and_do_not_abort() { + let shots = sample_shots(); + let truths: Vec = shots.iter().map(|(_, truth)| truth.clone()).collect(); + let mut dut = predictions(&truths); + dut[1] = StubResult::Error; + dut[3] = StubResult::Error; + + let counts = compare(&shots, dut, predictions(&truths)); + + assert_eq!(counts.cells(), &[[2, 0, 0], [0, 0, 0], [2, 0, 0]]); + assert_eq!(counts.cells()[DecoderOutcome::Mismatch.index()][0], 0); + } + + #[test] + fn reference_errors_are_counted_and_do_not_abort() { + let shots = sample_shots(); + let truths: Vec = shots.iter().map(|(_, truth)| truth.clone()).collect(); + let mut reference = predictions(&truths); + reference[0] = StubResult::Error; + reference[2] = StubResult::Error; + + let counts = compare(&shots, predictions(&truths), reference); + + assert_eq!(counts.cells(), &[[2, 0, 2], [0, 0, 0], [0, 0, 0]]); + } + + #[test] + fn wide_observable_difference_above_bit_63_is_preserved() { + let wide_truth = mask(&[70]); + let shots = vec![(vec![1], wide_truth.clone())]; + let counts = compare( + &shots, + predictions(&[ObsMask::new()]), + predictions(&[wide_truth]), + ); + + assert_eq!(counts.cells(), &[[0, 0, 0], [1, 0, 0], [0, 0, 0]]); + assert_eq!(counts.dut_only_failures(), 1); + } + + #[test] + fn headline_interval_matches_pecos_num_helper() { + let shots = sample_shots(); + let truths: Vec = shots.iter().map(|(_, truth)| truth.clone()).collect(); + let mut dut = truths.clone(); + dut[1] = mask(&[]); + let summary = PyDecoderComparisonResult::new( + compare(&shots, predictions(&dut), predictions(&truths)), + 0.05, + ) + .expect("valid Jeffreys inputs"); + let expected = jeffreys_interval(1, 4, 0.05).expect("valid direct helper inputs"); + + assert_eq!(summary.dut_only_failure.interval, expected); + // Both sides come from the same helper call, so the point estimate must be + // bit-identical; compare bit patterns rather than floats. + assert_eq!( + summary.dut_only_failure.point.to_bits(), + expected.point.to_bits() + ); + } + + #[test] + fn comparison_arguments_reject_invalid_shot_counts() { + assert_eq!( + validate_comparison_arguments(0, 0.05), + Err(JeffreysError::ZeroTrials) + ); + assert!(matches!( + validate_comparison_arguments(100_000_001, 0.05), + Err(JeffreysError::TrialsExceedSupported { + n: 100_000_001, + max: 100_000_000 + }) + )); + } + + #[test] + fn comparison_arguments_reject_alpha_outside_open_unit_interval() { + for alpha in [f64::NAN, f64::NEG_INFINITY, 0.0, 1.0, f64::INFINITY] { + assert!(matches!( + validate_comparison_arguments(1, alpha), + Err(JeffreysError::InvalidAlpha { .. }) + )); + } + } + + #[test] + fn comparison_is_deterministic_for_the_same_batch() { + let shots = sample_shots(); + let truths: Vec = shots.iter().map(|(_, truth)| truth.clone()).collect(); + let mut dut = truths.clone(); + dut[2] = mask(&[]); + + let first = compare(&shots, predictions(&dut), predictions(&truths)); + let second = compare(&shots, predictions(&dut), predictions(&truths)); + + assert_eq!(first, second); + } +} diff --git a/python/pecos-rslib/src/fault_tolerance_bindings/decoder_scoring.rs b/python/pecos-rslib/src/fault_tolerance_bindings/decoder_scoring.rs new file mode 100644 index 000000000..9a4f88fed --- /dev/null +++ b/python/pecos-rslib/src/fault_tolerance_bindings/decoder_scoring.rs @@ -0,0 +1,234 @@ +// Copyright 2026 The PECOS Developers +// +// Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except +// in compliance with the License. You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software distributed under the License +// is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express +// or implied. See the License for the specific language governing permissions and limitations under +// the License. + +//! Shared fail-loud scoring for observable decoders. + +use pecos_decoder_core::obs_mask::ObsMask; +use pecos_decoder_core::{DecoderError, ObservableDecoder}; +use std::fmt; +use std::ops::Range; + +/// A decoder failure annotated with the shot that caused it. +#[derive(Debug)] +pub(super) struct ShotDecodeError { + shot_index: usize, + source: DecoderError, +} + +impl fmt::Display for ShotDecodeError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + write!( + formatter, + "decoder failed on shot {}: {}", + self.shot_index, self.source + ) + } +} + +impl std::error::Error for ShotDecodeError { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + Some(&self.source) + } +} + +/// Decode and score a contiguous range of shots, aborting on the first failure. +/// +/// `access_shot` fills the reusable syndrome buffer and returns that shot's +/// true observable mask. Shot indices are kept absolute so callers can combine +/// independently scored worker ranges without losing error context. +pub(super) fn count_decoder_mismatches( + shots: Range, + syndrome: &mut [u8], + mut access_shot: impl FnMut(usize, &mut [u8]) -> ObsMask, + decoder: &mut dyn ObservableDecoder, +) -> Result { + let mut mismatches = 0; + for shot_index in shots { + let truth = access_shot(shot_index, syndrome); + let prediction = decoder + .decode_obs(syndrome) + .map_err(|source| ShotDecodeError { shot_index, source })?; + mismatches += usize::from(prediction != truth); + } + Ok(mismatches) +} + +/// Observable-decoder adapter that keeps only caller-selected observables. +pub(super) struct MaskedObservableDecoder { + inner: Box, + mask: ObsMask, +} + +impl MaskedObservableDecoder { + pub(super) fn new(inner: Box, mask: ObsMask) -> Self { + Self { inner, mask } + } +} + +impl ObservableDecoder for MaskedObservableDecoder { + fn decode_obs(&mut self, syndrome: &[u8]) -> Result { + let mut prediction = self.inner.decode_obs(syndrome)?; + prediction &= &self.mask; + Ok(prediction) + } +} + +/// Observable-decoder adapter that records one elapsed time per decode call. +pub(super) struct TimedObservableDecoder { + inner: Box, + per_shot_seconds: Vec, +} + +impl TimedObservableDecoder { + pub(super) fn new(inner: Box, capacity: usize) -> Self { + Self { + inner, + per_shot_seconds: Vec::with_capacity(capacity), + } + } + + pub(super) fn into_times(self) -> Vec { + self.per_shot_seconds + } +} + +impl ObservableDecoder for TimedObservableDecoder { + fn decode_obs(&mut self, syndrome: &[u8]) -> Result { + let start = std::time::Instant::now(); + let result = self.inner.decode_obs(syndrome); + self.per_shot_seconds.push(start.elapsed().as_secs_f64()); + result + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[derive(Clone)] + enum StubResult { + Prediction(ObsMask), + Error, + } + + struct StubDecoder { + expected_syndromes: Vec>, + results: Vec, + next: usize, + } + + impl StubDecoder { + fn new(expected_syndromes: &[Vec], results: Vec) -> Self { + assert_eq!(expected_syndromes.len(), results.len()); + Self { + expected_syndromes: expected_syndromes.to_vec(), + results, + next: 0, + } + } + } + + impl ObservableDecoder for StubDecoder { + fn decode_obs(&mut self, syndrome: &[u8]) -> Result { + assert_eq!(syndrome, self.expected_syndromes[self.next]); + let result = match &self.results[self.next] { + StubResult::Prediction(mask) => Ok(mask.clone()), + StubResult::Error => Err(DecoderError::DecodingFailed("stub error".into())), + }; + self.next += 1; + result + } + } + + fn mask(bits: &[usize]) -> ObsMask { + let mut mask = ObsMask::new(); + for &bit in bits { + mask.set(bit); + } + mask + } + + fn score( + shots: &[(Vec, ObsMask)], + results: Vec, + ) -> Result { + let syndromes: Vec> = shots.iter().map(|(syndrome, _)| syndrome.clone()).collect(); + let mut decoder = StubDecoder::new(&syndromes, results); + let mut syndrome = vec![0; syndromes.first().map_or(0, Vec::len)]; + count_decoder_mismatches( + 0..shots.len(), + &mut syndrome, + |shot_index, buffer| { + buffer.copy_from_slice(&shots[shot_index].0); + shots[shot_index].1.clone() + }, + &mut decoder, + ) + } + + #[test] + fn decoder_error_aborts_with_shot_index_and_source() { + let shots = vec![ + (vec![0], mask(&[])), + (vec![1], mask(&[0])), + (vec![0], mask(&[])), + ]; + let error = score( + &shots, + vec![ + StubResult::Prediction(mask(&[])), + StubResult::Error, + StubResult::Prediction(mask(&[])), + ], + ) + .unwrap_err(); + + assert_eq!(error.shot_index, 1); + assert!(error.to_string().contains("shot 1")); + assert!(error.to_string().contains("stub error")); + } + + #[test] + fn healthy_decoder_preserves_exact_mismatch_count() { + let shots = vec![ + (vec![0, 0], mask(&[])), + (vec![1, 0], mask(&[0])), + (vec![0, 1], mask(&[1])), + (vec![1, 1], mask(&[0, 1])), + ]; + let count = score( + &shots, + vec![ + StubResult::Prediction(mask(&[])), + StubResult::Prediction(mask(&[])), + StubResult::Prediction(mask(&[1])), + StubResult::Prediction(mask(&[1])), + ], + ) + .unwrap(); + + assert_eq!(count, 2); + } + + #[test] + fn always_erroring_decoder_does_not_match_all_observables_flipped_truth() { + let all_observables = mask(&[0, 1, 2]); + let shots = vec![(vec![1, 1], all_observables)]; + + // The old parallel sampler substituted the full observable-selection + // mask on error, which incorrectly scored this exact truth as correct. + let error = score(&shots, vec![StubResult::Error]).unwrap_err(); + + assert_eq!(error.shot_index, 0); + assert!(error.to_string().contains("stub error")); + } +} diff --git a/python/pecos-rslib/src/fault_tolerance_bindings/sample_corpus.rs b/python/pecos-rslib/src/fault_tolerance_bindings/sample_corpus.rs new file mode 100644 index 000000000..08c7986aa --- /dev/null +++ b/python/pecos-rslib/src/fault_tolerance_bindings/sample_corpus.rs @@ -0,0 +1,848 @@ +// Copyright 2026 The PECOS Developers +// +// Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except +// in compliance with the License. You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software distributed under the License +// is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express +// or implied. See the License for the specific language governing permissions and limitations under +// the License. + +//! Versioned, self-describing shot-corpus serialization. + +use pecos_decoder_core::dem::SparseDem; +use serde_json::{Map, Value}; +use sha2::{Digest, Sha256}; +use std::path::Path; + +const MAGIC: &[u8; 12] = b"PECOSCORPUS\0"; +pub(super) const FORMAT_VERSION: u32 = 1; +const SHA256_LEN: usize = 32; +const HEADER_LEN_END: usize = MAGIC.len() + size_of::(); +const PREFIX_LEN: usize = HEADER_LEN_END + SHA256_LEN; + +// These limits cover the full verified Jeffreys comparison regime and corpora +// far larger than typical research workloads. Capping each column family at one +// million also bounds degenerate zero-width column vectors to tens of MiB of +// descriptors, rather than allowing attacker-selected multi-gigabyte allocations. +pub(super) const MAX_SHOTS: usize = 100_000_000; +pub(super) const MAX_DETECTORS: usize = 1_000_000; +pub(super) const MAX_OBSERVABLES: usize = 1_000_000; + +#[derive(Debug)] +pub(super) enum CorpusError { + Io(std::io::Error), + Invalid(String), +} + +impl From for CorpusError { + fn from(error: std::io::Error) -> Self { + Self::Io(error) + } +} + +pub(super) struct CorpusToSave<'a> { + pub det_columns: &'a [Vec], + pub obs_columns: &'a [Vec], + pub num_shots: usize, + pub seed: Option, + pub dem: &'a str, + pub metadata_json: Option<&'a str>, +} + +#[derive(Debug, Eq, PartialEq)] +pub(super) struct LoadedCorpus { + pub det_columns: Vec>, + pub obs_columns: Vec>, + pub num_shots: usize, + pub seed: Option, + pub dem: String, + pub metadata_json: Option, + pub generator: String, + pub format_version: u32, +} + +fn invalid(message: impl Into) -> CorpusError { + CorpusError::Invalid(message.into()) +} + +fn sha256(bytes: &[u8]) -> [u8; SHA256_LEN] { + Sha256::digest(bytes).into() +} + +fn hex_encode(bytes: &[u8]) -> String { + const HEX: &[u8; 16] = b"0123456789abcdef"; + let mut output = String::with_capacity(bytes.len() * 2); + for &byte in bytes { + output.push(char::from(HEX[usize::from(byte >> 4)])); + output.push(char::from(HEX[usize::from(byte & 0x0f)])); + } + output +} + +fn sha256_hex(bytes: &[u8]) -> String { + hex_encode(&sha256(bytes)) +} + +fn validate_dimensions( + num_shots: usize, + num_detectors: usize, + num_observables: usize, +) -> Result<(), CorpusError> { + if num_shots > MAX_SHOTS { + return Err(invalid(format!( + "corpus num_shots={num_shots} exceeds the format limit MAX_SHOTS={MAX_SHOTS}" + ))); + } + if num_detectors > MAX_DETECTORS { + return Err(invalid(format!( + "corpus num_detectors={num_detectors} exceeds the format limit MAX_DETECTORS={MAX_DETECTORS}" + ))); + } + if num_observables > MAX_OBSERVABLES { + return Err(invalid(format!( + "corpus num_observables={num_observables} exceeds the format limit MAX_OBSERVABLES={MAX_OBSERVABLES}" + ))); + } + Ok(()) +} + +fn checked_payload_len( + num_detectors: usize, + num_observables: usize, + words_per_column: usize, +) -> Result { + num_detectors + .checked_add(num_observables) + .and_then(|columns| columns.checked_mul(words_per_column)) + .and_then(|words| words.checked_mul(size_of::())) + .ok_or_else(|| invalid("corpus dimensions overflow the supported payload size")) +} + +fn validate_columns(columns: &[Vec], words_per_column: usize) -> Result<(), CorpusError> { + if let Some((index, column)) = columns + .iter() + .enumerate() + .find(|(_, column)| column.len() != words_per_column) + { + return Err(invalid(format!( + "sample column {index} has {} word(s), expected {words_per_column}", + column.len() + ))); + } + Ok(()) +} + +/// Mask selecting the meaningful low bits in a column's final word. +/// +/// The format requires every unused high padding bit to be zero. +fn final_word_mask(num_shots: usize) -> u64 { + let used_bits = num_shots % 64; + if used_bits == 0 { + u64::MAX + } else { + (1_u64 << used_bits) - 1 + } +} + +pub(super) fn save(path: &Path, corpus: CorpusToSave<'_>) -> Result<(), CorpusError> { + validate_dimensions( + corpus.num_shots, + corpus.det_columns.len(), + corpus.obs_columns.len(), + )?; + let parsed_dem = SparseDem::from_dem_str(corpus.dem) + .map_err(|error| invalid(format!("invalid DEM supplied to SampleBatch.save: {error}")))?; + if parsed_dem.num_detectors != corpus.det_columns.len() + || parsed_dem.num_observables != corpus.obs_columns.len() + { + return Err(invalid(format!( + "DEM dimensions do not match SampleBatch: DEM has {} detector(s) and {} observable(s), batch has {} detector(s) and {} observable(s)", + parsed_dem.num_detectors, + parsed_dem.num_observables, + corpus.det_columns.len(), + corpus.obs_columns.len() + ))); + } + + if let Some(metadata) = corpus.metadata_json { + serde_json::from_str::(metadata) + .map_err(|error| invalid(format!("metadata_json is not valid JSON: {error}")))?; + } + + let words_per_column = corpus.num_shots.div_ceil(64); + validate_columns(corpus.det_columns, words_per_column)?; + validate_columns(corpus.obs_columns, words_per_column)?; + let payload_len = checked_payload_len( + corpus.det_columns.len(), + corpus.obs_columns.len(), + words_per_column, + )?; + let mut payload = Vec::with_capacity(payload_len); + for column in corpus.det_columns.iter().chain(corpus.obs_columns) { + for (word_index, &word) in column.iter().enumerate() { + let word = if word_index + 1 == words_per_column { + word & final_word_mask(corpus.num_shots) + } else { + word + }; + payload.extend_from_slice(&word.to_le_bytes()); + } + } + + let header = serde_json::json!({ + "format_version": FORMAT_VERSION, + "num_shots": corpus.num_shots, + "num_detectors": corpus.det_columns.len(), + "num_observables": corpus.obs_columns.len(), + "words_per_column": words_per_column, + "seed": corpus.seed, + "dem": corpus.dem, + "dem_sha256": sha256_hex(corpus.dem.as_bytes()), + "metadata_json": corpus.metadata_json, + "generator": concat!("pecos-rslib ", env!("CARGO_PKG_VERSION")), + }); + let header_bytes = serde_json::to_vec(&header) + .map_err(|error| invalid(format!("could not serialize corpus header: {error}")))?; + let header_len = u32::try_from(header_bytes.len()) + .map_err(|_| invalid("corpus JSON header is too large to encode"))?; + let mut content_hasher = Sha256::new(); + content_hasher.update(&header_bytes); + content_hasher.update(&payload); + let content_sha256: [u8; SHA256_LEN] = content_hasher.finalize().into(); + let file_len = PREFIX_LEN + .checked_add(header_bytes.len()) + .and_then(|len| len.checked_add(payload.len())) + .ok_or_else(|| invalid("corpus file size overflows this platform"))?; + let mut bytes = Vec::with_capacity(file_len); + bytes.extend_from_slice(MAGIC); + bytes.extend_from_slice(&header_len.to_le_bytes()); + bytes.extend_from_slice(&content_sha256); + bytes.extend_from_slice(&header_bytes); + bytes.extend_from_slice(&payload); + std::fs::write(path, bytes)?; + Ok(()) +} + +fn required_u64(header: &Map, field: &str) -> Result { + header.get(field).and_then(Value::as_u64).ok_or_else(|| { + invalid(format!( + "corpus header field {field:?} must be an unsigned integer" + )) + }) +} + +fn required_usize(header: &Map, field: &str) -> Result { + usize::try_from(required_u64(header, field)?).map_err(|_| { + invalid(format!( + "corpus header field {field:?} is too large for this platform" + )) + }) +} + +fn required_string<'a>( + header: &'a Map, + field: &str, +) -> Result<&'a str, CorpusError> { + header + .get(field) + .and_then(Value::as_str) + .ok_or_else(|| invalid(format!("corpus header field {field:?} must be a string"))) +} + +fn nullable_u64(header: &Map, field: &str) -> Result, CorpusError> { + match header.get(field) { + Some(Value::Null) => Ok(None), + Some(value) => value.as_u64().map(Some).ok_or_else(|| { + invalid(format!( + "corpus header field {field:?} must be an unsigned integer or null" + )) + }), + None => Err(invalid(format!( + "corpus header is missing required field {field:?}" + ))), + } +} + +fn nullable_string( + header: &Map, + field: &str, +) -> Result, CorpusError> { + match header.get(field) { + Some(Value::Null) => Ok(None), + Some(Value::String(value)) => Ok(Some(value.clone())), + Some(_) => Err(invalid(format!( + "corpus header field {field:?} must be a string or null" + ))), + None => Err(invalid(format!( + "corpus header is missing required field {field:?}" + ))), + } +} + +pub(super) fn load(path: &Path) -> Result { + let bytes = std::fs::read(path)?; + if bytes.get(..MAGIC.len()) != Some(MAGIC.as_slice()) { + return Err(invalid( + "bad shot-corpus magic: expected PECOSCORPUS followed by a NUL byte", + )); + } + if bytes.len() < HEADER_LEN_END { + return Err(invalid("shot corpus is missing its 4-byte header length")); + } + if bytes.len() < PREFIX_LEN { + return Err(invalid( + "shot corpus is missing its 32-byte content SHA-256", + )); + } + let mut header_len_bytes = [0_u8; size_of::()]; + header_len_bytes.copy_from_slice(&bytes[MAGIC.len()..HEADER_LEN_END]); + let header_len = usize::try_from(u32::from_le_bytes(header_len_bytes)) + .map_err(|_| invalid("corpus header length is too large for this platform"))?; + let header_end = PREFIX_LEN + .checked_add(header_len) + .ok_or_else(|| invalid("corpus header length overflows this platform"))?; + if header_end > bytes.len() { + return Err(invalid(format!( + "corpus header length declares {header_len} byte(s), but the file contains only {} after the prefix", + bytes.len() - PREFIX_LEN + ))); + } + + let expected_content_sha = &bytes[HEADER_LEN_END..PREFIX_LEN]; + let actual_content_sha = sha256(&bytes[PREFIX_LEN..]); + if expected_content_sha != actual_content_sha { + return Err(invalid(format!( + "corpus content SHA-256 mismatch: expected {}, computed {}", + hex_encode(expected_content_sha), + hex_encode(&actual_content_sha) + ))); + } + + // The digest covers the exact header bytes, so a duplicate JSON key cannot + // be injected into an existing corpus without breaking content integrity. + let header_value: Value = serde_json::from_slice(&bytes[PREFIX_LEN..header_end]) + .map_err(|error| invalid(format!("invalid corpus header JSON: {error}")))?; + let header = header_value + .as_object() + .ok_or_else(|| invalid("invalid corpus header JSON: top-level value must be an object"))?; + + let version = required_u64(header, "format_version")?; + if version != u64::from(FORMAT_VERSION) { + return Err(invalid(format!( + "unsupported corpus format_version {version}; this PECOS build supports version {FORMAT_VERSION}" + ))); + } + + let num_shots = required_usize(header, "num_shots")?; + let num_detectors = required_usize(header, "num_detectors")?; + let num_observables = required_usize(header, "num_observables")?; + validate_dimensions(num_shots, num_detectors, num_observables)?; + let words_per_column = required_usize(header, "words_per_column")?; + let expected_words = num_shots.div_ceil(64); + if words_per_column != expected_words { + return Err(invalid(format!( + "corpus words_per_column is {words_per_column}, but num_shots={num_shots} requires {expected_words}" + ))); + } + let seed = nullable_u64(header, "seed")?; + let dem = required_string(header, "dem")?.to_owned(); + let expected_dem_sha = required_string(header, "dem_sha256")?; + let metadata_json = nullable_string(header, "metadata_json")?; + let generator = required_string(header, "generator")?.to_owned(); + + let payload = &bytes[header_end..]; + let expected_payload_len = + checked_payload_len(num_detectors, num_observables, words_per_column)?; + if payload.len() != expected_payload_len { + return Err(invalid(format!( + "corpus payload length is {} byte(s), but declared dimensions require {expected_payload_len} byte(s)", + payload.len() + ))); + } + let actual_dem_sha = sha256_hex(dem.as_bytes()); + if expected_dem_sha != actual_dem_sha { + return Err(invalid(format!( + "corpus DEM SHA-256 mismatch: expected {expected_dem_sha}, computed {actual_dem_sha}" + ))); + } + if let Some(metadata) = &metadata_json { + serde_json::from_str::(metadata) + .map_err(|error| invalid(format!("corpus metadata_json is not valid JSON: {error}")))?; + } + let parsed_dem = SparseDem::from_dem_str(&dem) + .map_err(|error| invalid(format!("corpus contains an invalid DEM: {error}")))?; + if parsed_dem.num_detectors != num_detectors || parsed_dem.num_observables != num_observables { + return Err(invalid(format!( + "corpus DEM dimensions disagree with its header: DEM has {} detector(s) and {} observable(s), header declares {num_detectors} detector(s) and {num_observables} observable(s)", + parsed_dem.num_detectors, parsed_dem.num_observables + ))); + } + + if words_per_column != 0 { + let padding_mask = !final_word_mask(num_shots); + for (column_index, final_word) in payload + .chunks_exact(size_of::()) + .skip(words_per_column - 1) + .step_by(words_per_column) + .enumerate() + { + let mut word_bytes = [0_u8; size_of::()]; + word_bytes.copy_from_slice(final_word); + if u64::from_le_bytes(word_bytes) & padding_mask != 0 { + return Err(invalid(format!( + "corpus payload column {column_index} has nonzero padding bits above num_shots={num_shots}; unused high bits in the final word must be zero" + ))); + } + } + } + + let mut words = Vec::with_capacity(payload.len() / size_of::()); + for chunk in payload.chunks_exact(size_of::()) { + let mut bytes = [0_u8; size_of::()]; + bytes.copy_from_slice(chunk); + words.push(u64::from_le_bytes(bytes)); + } + let mut offset = 0; + let mut read_columns = |count: usize| { + let mut columns = Vec::with_capacity(count); + for _ in 0..count { + let end = offset + words_per_column; + columns.push(words[offset..end].to_vec()); + offset = end; + } + columns + }; + let det_columns = read_columns(num_detectors); + let obs_columns = read_columns(num_observables); + + Ok(LoadedCorpus { + det_columns, + obs_columns, + num_shots, + seed, + dem, + metadata_json, + generator, + format_version: FORMAT_VERSION, + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use tempfile::TempDir; + + const DEM: &str = "error(0.125) D0 L0\n"; + + fn corpus_path() -> (TempDir, std::path::PathBuf) { + let directory = tempfile::tempdir().unwrap(); + let path = directory.path().join("shots.pecos"); + (directory, path) + } + + fn save_test_corpus(path: &Path) { + save( + path, + CorpusToSave { + det_columns: &[vec![0b10]], + obs_columns: &[vec![0b10]], + num_shots: 2, + seed: Some(42), + dem: DEM, + metadata_json: Some(r#"{ "decoder": "pymatching" }"#), + }, + ) + .unwrap(); + } + + fn invalid_message(result: Result) -> String { + match result.unwrap_err() { + CorpusError::Invalid(message) => message, + CorpusError::Io(error) => panic!("unexpected I/O error: {error}"), + } + } + + fn header_end(bytes: &[u8]) -> usize { + let mut length = [0_u8; 4]; + length.copy_from_slice(&bytes[MAGIC.len()..HEADER_LEN_END]); + PREFIX_LEN + usize::try_from(u32::from_le_bytes(length)).unwrap() + } + + fn with_valid_content_sha(mut bytes: Vec) -> Vec { + let digest = sha256(&bytes[PREFIX_LEN..]); + bytes[HEADER_LEN_END..PREFIX_LEN].copy_from_slice(&digest); + bytes + } + + fn replace_header(bytes: &[u8], update: impl FnOnce(&mut Map)) -> Vec { + let old_header_end = header_end(bytes); + let mut header: Value = serde_json::from_slice(&bytes[PREFIX_LEN..old_header_end]).unwrap(); + update(header.as_object_mut().unwrap()); + let new_header = serde_json::to_vec(&header).unwrap(); + let new_header_len = u32::try_from(new_header.len()).unwrap(); + let mut updated = Vec::new(); + updated.extend_from_slice(MAGIC); + updated.extend_from_slice(&new_header_len.to_le_bytes()); + updated.extend_from_slice(&bytes[HEADER_LEN_END..PREFIX_LEN]); + updated.extend_from_slice(&new_header); + updated.extend_from_slice(&bytes[old_header_end..]); + updated + } + + #[test] + fn round_trip_preserves_columns_dimensions_and_provenance() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + + let loaded = load(&path).unwrap(); + assert_eq!(loaded.det_columns, vec![vec![0b10]]); + assert_eq!(loaded.obs_columns, vec![vec![0b10]]); + assert_eq!(loaded.num_shots, 2); + assert_eq!(loaded.seed, Some(42)); + assert_eq!(loaded.dem, DEM); + assert_eq!( + loaded.metadata_json.as_deref(), + Some(r#"{ "decoder": "pymatching" }"#) + ); + assert!(loaded.generator.starts_with("pecos-rslib ")); + assert_eq!(loaded.format_version, FORMAT_VERSION); + + let bytes = std::fs::read(path).unwrap(); + let header: Value = serde_json::from_slice(&bytes[PREFIX_LEN..header_end(&bytes)]).unwrap(); + assert!(header.get("payload_sha256").is_none()); + } + + #[test] + fn wide_observable_column_round_trips_without_narrowing() { + let (_directory, path) = corpus_path(); + let mut observables = vec![vec![0]; 65]; + observables[64][0] = 1; + save( + &path, + CorpusToSave { + det_columns: &[vec![1]], + obs_columns: &observables, + num_shots: 1, + seed: None, + dem: "error(0.125) D0 L64\n", + metadata_json: None, + }, + ) + .unwrap(); + + let loaded = load(&path).unwrap(); + assert_eq!(loaded.obs_columns.len(), 65); + assert_eq!(loaded.obs_columns[64], vec![1]); + assert_eq!(loaded.seed, None); + } + + #[test] + fn corrupted_payload_fails_content_checksum() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let mut bytes = std::fs::read(&path).unwrap(); + let last = bytes.last_mut().unwrap(); + *last ^= 0x80; + std::fs::write(&path, bytes).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("content SHA-256 mismatch"), "{message}"); + } + + #[test] + fn truncated_payload_fails_content_checksum() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let mut bytes = std::fs::read(&path).unwrap(); + bytes.pop(); + std::fs::write(&path, bytes).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("content SHA-256 mismatch"), "{message}"); + } + + #[test] + fn authenticated_truncated_payload_fails_length_check() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let mut bytes = std::fs::read(&path).unwrap(); + bytes.pop(); + let bytes = with_valid_content_sha(bytes); + std::fs::write(&path, bytes).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("payload length"), "{message}"); + } + + #[test] + fn changing_num_shots_and_header_len_breaks_content_checksum() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let updated = replace_header(&bytes, |header| { + header.insert("num_shots".to_owned(), Value::from(3)); + }); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("content SHA-256 mismatch"), "{message}"); + } + + #[test] + fn changing_seed_and_header_len_breaks_content_checksum() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let updated = replace_header(&bytes, |header| { + header.insert("seed".to_owned(), Value::from(43)); + }); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("content SHA-256 mismatch"), "{message}"); + } + + #[test] + fn bad_magic_is_rejected_first() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let mut bytes = std::fs::read(&path).unwrap(); + bytes[0] ^= 1; + std::fs::write(&path, bytes).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("bad shot-corpus magic"), "{message}"); + } + + #[test] + fn future_format_version_is_actionable() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let updated = with_valid_content_sha(replace_header(&bytes, |header| { + header.insert("format_version".to_owned(), Value::from(999)); + })); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!( + message.contains("unsupported corpus format_version 999"), + "{message}" + ); + } + + #[test] + fn invalid_header_json_is_rejected_without_panicking() { + let (_directory, path) = corpus_path(); + let mut bytes = Vec::from(MAGIC.as_slice()); + bytes.extend_from_slice(&1_u32.to_le_bytes()); + bytes.extend_from_slice(&[0_u8; SHA256_LEN]); + bytes.push(b'{'); + let bytes = with_valid_content_sha(bytes); + std::fs::write(&path, bytes).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("invalid corpus header JSON"), "{message}"); + } + + #[test] + fn dem_checksum_is_verified_after_content_checksum() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let updated = with_valid_content_sha(replace_header(&bytes, |header| { + header.insert("dem".to_owned(), Value::from("error(0.25) D0 L0\n")); + })); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("DEM SHA-256 mismatch"), "{message}"); + } + + #[test] + fn loaded_dem_dimensions_must_match_header() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let replacement_dem = "error(0.25) D1 L0\n"; + let updated = with_valid_content_sha(replace_header(&bytes, |header| { + header.insert("dem".to_owned(), Value::from(replacement_dem)); + header.insert( + "dem_sha256".to_owned(), + Value::from(sha256_hex(replacement_dem.as_bytes())), + ); + })); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!( + message.contains("corpus DEM dimensions disagree with its header"), + "{message}" + ); + } + + #[test] + fn loaded_metadata_json_must_be_valid() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let updated = with_valid_content_sha(replace_header(&bytes, |header| { + header.insert("metadata_json".to_owned(), Value::from("{")); + })); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!( + message.contains("corpus metadata_json is not valid JSON"), + "{message}" + ); + } + + #[test] + fn declared_shots_above_limit_are_rejected_with_zero_columns() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let updated = with_valid_content_sha(replace_header(&bytes, |header| { + header.insert("num_shots".to_owned(), Value::from(MAX_SHOTS + 1)); + header.insert("num_detectors".to_owned(), Value::from(0)); + header.insert("num_observables".to_owned(), Value::from(0)); + header.insert( + "words_per_column".to_owned(), + Value::from((MAX_SHOTS + 1).div_ceil(64)), + ); + })); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!( + message.contains(&format!("MAX_SHOTS={MAX_SHOTS}")), + "{message}" + ); + } + + #[test] + fn declared_detectors_above_limit_are_rejected_with_zero_shots() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let updated = with_valid_content_sha(replace_header(&bytes, |header| { + header.insert("num_shots".to_owned(), Value::from(0)); + header.insert("num_detectors".to_owned(), Value::from(MAX_DETECTORS + 1)); + header.insert("num_observables".to_owned(), Value::from(0)); + header.insert("words_per_column".to_owned(), Value::from(0)); + })); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!( + message.contains(&format!("MAX_DETECTORS={MAX_DETECTORS}")), + "{message}" + ); + } + + #[test] + fn declared_observables_above_limit_are_rejected_with_zero_shots() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let bytes = std::fs::read(&path).unwrap(); + let updated = with_valid_content_sha(replace_header(&bytes, |header| { + header.insert("num_shots".to_owned(), Value::from(0)); + header.insert("num_detectors".to_owned(), Value::from(0)); + header.insert( + "num_observables".to_owned(), + Value::from(MAX_OBSERVABLES + 1), + ); + header.insert("words_per_column".to_owned(), Value::from(0)); + })); + std::fs::write(&path, updated).unwrap(); + + let message = invalid_message(load(&path)); + assert!( + message.contains(&format!("MAX_OBSERVABLES={MAX_OBSERVABLES}")), + "{message}" + ); + } + + #[test] + fn save_masks_unused_high_padding_bits() { + let (_directory, path) = corpus_path(); + save( + &path, + CorpusToSave { + det_columns: &[vec![u64::MAX]], + obs_columns: &[vec![u64::MAX]], + num_shots: 1, + seed: None, + dem: DEM, + metadata_json: None, + }, + ) + .unwrap(); + + let loaded = load(&path).unwrap(); + assert_eq!(loaded.det_columns, vec![vec![1]]); + assert_eq!(loaded.obs_columns, vec![vec![1]]); + } + + #[test] + fn nonzero_payload_padding_is_rejected() { + let (_directory, path) = corpus_path(); + save_test_corpus(&path); + let mut bytes = std::fs::read(&path).unwrap(); + let payload_start = header_end(&bytes); + bytes[payload_start + size_of::() - 1] |= 0x80; + let bytes = with_valid_content_sha(bytes); + std::fs::write(&path, bytes).unwrap(); + + let message = invalid_message(load(&path)); + assert!(message.contains("nonzero padding bits"), "{message}"); + } + + #[test] + fn mismatched_dem_dimensions_are_rejected_before_writing() { + let (_directory, path) = corpus_path(); + let result = save( + &path, + CorpusToSave { + det_columns: &[vec![0]], + obs_columns: &[vec![0]], + num_shots: 1, + seed: None, + dem: "error(0.125) D1 L0\n", + metadata_json: None, + }, + ); + + let CorpusError::Invalid(message) = result.unwrap_err() else { + panic!("expected malformed-input error"); + }; + assert!(message.contains("DEM dimensions do not match SampleBatch")); + assert!(!path.exists()); + } + + #[test] + fn invalid_metadata_json_is_rejected_before_writing() { + let (_directory, path) = corpus_path(); + let result = save( + &path, + CorpusToSave { + det_columns: &[vec![0]], + obs_columns: &[vec![0]], + num_shots: 1, + seed: None, + dem: DEM, + metadata_json: Some("{"), + }, + ); + + let CorpusError::Invalid(message) = result.unwrap_err() else { + panic!("expected malformed-input error"); + }; + assert!(message.contains("metadata_json is not valid JSON")); + assert!(!path.exists()); + } +} diff --git a/python/quantum-pecos/tests/qec/test_decoder_comparison.py b/python/quantum-pecos/tests/qec/test_decoder_comparison.py new file mode 100644 index 000000000..82b180f8c --- /dev/null +++ b/python/quantum-pecos/tests/qec/test_decoder_comparison.py @@ -0,0 +1,52 @@ +# Copyright 2026 The PECOS Developers +# Licensed under the Apache License, Version 2.0 + +"""Python coverage for paired DUT/reference decoder comparison.""" + +from __future__ import annotations + +import math + +import pytest + +pytest.importorskip("pecos_rslib") + +from pecos_rslib.qec import SampleBatch + + +def test_sample_batch_compare_decoders_exposes_joint_counts() -> None: + dem = "error(0.1) D0 L0\n" + batch = SampleBatch([[0], [1], [0], [1]], [0, 1, 0, 1]) + + first = batch.compare_decoders(dem, "pymatching", "pymatching") + second = batch.compare_decoders(dem, "pymatching", "pymatching") + + assert first.total_shots == 4 + assert first.counts == [[4, 0, 0], [0, 0, 0], [0, 0, 0]] + assert first.dut_correct_reference_correct == 4 + assert first.dut_only_failures == 0 + assert first.both_failed == 0 + assert first.dut_only_failure_interval[0] >= 0.0 + assert first.dut_only_failure_interval[1] <= 1.0 + assert second.counts == first.counts + + +def test_compare_decoders_rejects_empty_batch_before_decoder_construction() -> None: + batch = SampleBatch([], []) + + with pytest.raises(ValueError, match="n must be greater than zero"): + batch.compare_decoders("not a DEM", "not a decoder", "not a decoder") + + +@pytest.mark.parametrize("alpha", [math.nan, -1.0, 0.0, 1.0, 2.0, math.inf]) +def test_compare_decoders_rejects_alpha_outside_open_unit_interval(alpha: float) -> None: + dem = "error(0.1) D0 L0\n" + batch = SampleBatch([[0]], [0]) + + with pytest.raises(ValueError, match=r"alpha must be finite and in \(0, 1\)"): + batch.compare_decoders( + dem, + "pymatching", + "pymatching", + alpha=alpha, + ) diff --git a/python/quantum-pecos/tests/qec/test_sample_batch.py b/python/quantum-pecos/tests/qec/test_sample_batch.py index 4170da0ac..ac896e425 100644 --- a/python/quantum-pecos/tests/qec/test_sample_batch.py +++ b/python/quantum-pecos/tests/qec/test_sample_batch.py @@ -93,3 +93,18 @@ def test_decode_count(self, d3_setup): errors = batch.decode_count(dem_str, "pymatching") assert isinstance(errors, int) assert 0 <= errors <= 1000 + + +def test_seeded_healthy_decoder_counts_and_stats_are_unchanged() -> None: + dem = "error(0.1) D0 L0\nerror(0.2) D0\n" + batch = DemSampler.from_dem_string(dem).generate_samples(257, seed=314159) + + assert batch.decode_count(dem, "pymatching") == 48 + assert batch.decode_count_batch(dem) == 48 + assert batch.decode_count_parallel(dem, "pymatching", num_workers=3) == 48 + + stats = batch.decode_stats(dem, "pymatching") + parallel_stats = batch.decode_stats_parallel(dem, "pymatching", num_workers=3) + assert stats.num_shots == parallel_stats.num_shots == 257 + assert stats.num_errors == parallel_stats.num_errors == 48 + assert stats.logical_error_rate == parallel_stats.logical_error_rate == 48 / 257 diff --git a/python/quantum-pecos/tests/qec/test_sample_corpus.py b/python/quantum-pecos/tests/qec/test_sample_corpus.py new file mode 100644 index 000000000..32af4c138 --- /dev/null +++ b/python/quantum-pecos/tests/qec/test_sample_corpus.py @@ -0,0 +1,161 @@ +# Copyright 2026 The PECOS Developers +# Licensed under the Apache License, Version 2.0 + +"""Python coverage for self-describing SampleBatch shot corpora.""" + +from __future__ import annotations + +import errno + +import pytest + +pytest.importorskip("pecos_rslib") + +from pecos_rslib.qec import DemSampler, SampleBatch + + +def test_generated_batch_round_trip_preserves_shots_and_provenance(tmp_path) -> None: + dem = "error(0.125) D0 L0\n" + metadata = '{ "decoder": "pymatching", "decoder_seed": 17 }' + batch = DemSampler.from_dem_string(dem).generate_samples(130, seed=42) + path = tmp_path / "round-trip.pecos" + + batch.save(path, dem=dem, metadata_json=metadata) + loaded = SampleBatch.load(path) + + assert loaded.num_shots == batch.num_shots + assert loaded.seed == batch.seed == 42 + assert loaded.dem == dem + assert loaded.metadata_json == metadata + assert loaded.generator.startswith("pecos-rslib ") + assert loaded.format_version == 1 + assert [loaded.get_syndrome(i) for i in range(loaded.num_shots)] == [ + batch.get_syndrome(i) for i in range(batch.num_shots) + ] + assert [loaded.get_observable_mask_wide(i) for i in range(loaded.num_shots)] == [ + batch.get_observable_mask_wide(i) for i in range(batch.num_shots) + ] + + +def test_wide_observable_above_bit_63_round_trips(tmp_path) -> None: + dem = "error(0.125) D0 L64\n" + batch = SampleBatch([[1], [0]], [1 << 64, 0]) + path = tmp_path / "wide.pecos" + + batch.save(path, dem=dem) + loaded = SampleBatch.load(path) + + assert loaded.get_observable_mask_wide(0) == 1 << 64 + assert loaded.get_observable_mask_wide(1) == 0 + + +def test_save_rejects_mismatched_dem_dimensions(tmp_path) -> None: + batch = SampleBatch([[0]], [0]) + + with pytest.raises(ValueError, match="DEM dimensions do not match SampleBatch"): + batch.save(tmp_path / "wrong-dem.pecos", dem="error(0.1) D1\n") + + +def test_save_rejects_invalid_metadata_json(tmp_path) -> None: + batch = SampleBatch([[0]], [0]) + + with pytest.raises(ValueError, match="metadata_json is not valid JSON"): + batch.save( + tmp_path / "bad-metadata.pecos", + dem="error(0.1) D0\n", + metadata_json="{", + ) + + +def test_load_maps_malformed_files_to_value_error(tmp_path) -> None: + path = tmp_path / "bad-magic.pecos" + path.write_bytes(b"not a PECOS corpus") + + with pytest.raises(ValueError, match="bad shot-corpus magic"): + SampleBatch.load(path) + + +def test_load_maps_filesystem_failures_to_io_error(tmp_path) -> None: + path = tmp_path / "missing.pecos" + + with pytest.raises(FileNotFoundError, match="No such file or directory") as exc_info: + SampleBatch.load(path) + + assert exc_info.value.errno == errno.ENOENT + assert exc_info.value.filename == str(path) + + +def test_resave_preserves_metadata_unless_explicitly_cleared(tmp_path) -> None: + dem = "error(0.125) D0 L0\n" + metadata = '{"source": "original"}' + original_path = tmp_path / "original.pecos" + preserved_path = tmp_path / "preserved.pecos" + cleared_path = tmp_path / "cleared.pecos" + batch = SampleBatch([[1]], [1]) + batch.save(original_path, dem=dem, metadata_json=metadata) + loaded = SampleBatch.load(original_path) + + loaded.save(preserved_path, dem=dem) + loaded.save(cleared_path, dem=dem, clear_metadata=True) + + assert SampleBatch.load(preserved_path).metadata_json == metadata + assert SampleBatch.load(cleared_path).metadata_json is None + + +def test_loaded_batch_requires_its_embedded_dem_unless_opted_out(tmp_path) -> None: + embedded_dem = "error(0.125) D0 L0\n" + different_dem = "error(0.25) D0 L0\n" + path = tmp_path / "dem-bound.pecos" + batch = SampleBatch([[1]], [1]) + batch.save(path, dem=embedded_dem) + loaded = SampleBatch.load(path) + + with pytest.raises(ValueError, match="differs from the DEM embedded"): + loaded.compare_decoders( + different_dem, + "pymatching", + "pymatching", + ) + + result = loaded.compare_decoders( + different_dem, + "pymatching", + "pymatching", + allow_dem_mismatch=True, + ) + assert result.total_shots == 1 + + +def test_generate_samples_records_resolved_and_explicit_seeds() -> None: + sampler = DemSampler.from_dem_string("error(0.125) D0 L0\n") + + resolved = sampler.generate_samples(130) + explicit = sampler.generate_samples(1, seed=0xDEADBEEF) + replayed = sampler.generate_samples(resolved.num_shots, seed=resolved.seed) + + assert isinstance(resolved.seed, int) + assert 0 <= resolved.seed <= (1 << 64) - 1 + assert explicit.seed == 0xDEADBEEF + assert [replayed.get_syndrome(i) for i in range(replayed.num_shots)] == [ + resolved.get_syndrome(i) for i in range(resolved.num_shots) + ] + assert [replayed.get_observable_mask_wide(i) for i in range(replayed.num_shots)] == [ + resolved.get_observable_mask_wide(i) for i in range(resolved.num_shots) + ] + + +def test_compare_decoders_counts_survive_corpus_round_trip(tmp_path) -> None: + dem = "error(0.1) D0 L0\nerror(0.1) D0\n" + batch = DemSampler.from_dem_string(dem).generate_samples(257, seed=314159) + before = batch.compare_decoders(dem, "pymatching", "pymatching") + path = tmp_path / "comparison.pecos" + + batch.save(path, dem=dem) + loaded = SampleBatch.load(path) + after = loaded.compare_decoders( + loaded.dem, + "pymatching", + "pymatching", + ) + + assert after.counts == before.counts