use crate::fingerprint::{FingerprintConfig, FingerprintStats, ResidualFingerprint};
use crate::guv::{GuvConfig, GuvParameters, GuvRecovery};
use crate::matrix::{CsrInput, CsrMatrix, DenseMatrix};
use crate::rect::{
adaptive_matmul_prepared, PreparedFactor, RectangularKernel, RectangularPolicy,
RectangularStats,
};
use crate::sketch::{paper_schedule, RoundParams};
use crate::spgemm::dense_matmul;
use std::collections::{BTreeMap, HashMap};
use std::sync::Arc;
use std::time::{Duration, Instant};
#[derive(Clone, Debug)]
pub struct SignatureConfig {
pub degree: usize,
pub oversampling: f64,
pub seed: u64,
pub identity_fallback: bool,
pub guaranteed_correction: bool,
}
impl Default for SignatureConfig {
fn default() -> Self {
Self {
degree: 5,
oversampling: 2.0,
seed: 0x5A17_9EED_D15C_A11E,
identity_fallback: true,
guaranteed_correction: true,
}
}
}
#[derive(Clone, Debug)]
pub struct MomentConfig {
pub degree: usize,
pub oversampling: f64,
pub seed: u64,
pub identity_fallback: bool,
pub guaranteed_correction: bool,
}
impl Default for MomentConfig {
fn default() -> Self {
Self {
degree: 3,
oversampling: 3.0,
seed: 0x4D4F_4D45_4E54_0001,
identity_fallback: true,
guaranteed_correction: true,
}
}
}
#[derive(Clone, Debug)]
pub enum RecoveryBackend {
Identity,
Signature(SignatureConfig),
Moment(MomentConfig),
Guv(GuvConfig),
}
#[derive(Clone, Debug)]
pub enum BinaryRecoveryMatrix {
Identity { domain: usize },
Signature(SignatureRecovery),
Moment(MomentRecovery),
Guv(GuvRecovery),
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
enum RecoveryMatrixKey {
Identity {
domain: usize,
},
Signature {
domain: usize,
capacity: usize,
degree: usize,
oversampling_bits: u64,
seed: u64,
},
Moment {
domain: usize,
capacity: usize,
degree: usize,
oversampling_bits: u64,
seed: u64,
},
Guv {
domain: usize,
capacity: usize,
alpha_bits: u64,
epsilon_bits: u64,
},
}
impl RecoveryMatrixKey {
#[inline]
fn cacheable(&self) -> bool {
!matches!(self, Self::Signature { .. } | Self::Moment { .. })
}
}
#[derive(Clone)]
struct CachedProduct {
matrix: Arc<DenseMatrix>,
stats: RectangularStats,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
struct ResidualCacheKey {
h: RecoveryMatrixKey,
g: RecoveryMatrixKey,
d_version: u64,
}
#[derive(Clone)]
struct CachedResidual {
matrix: Arc<DenseMatrix>,
product_stats: RectangularStats,
}
impl BinaryRecoveryMatrix {
pub fn identity(domain: usize) -> Self {
Self::Identity { domain }
}
pub fn domain(&self) -> usize {
match self {
Self::Identity { domain } => *domain,
Self::Signature(s) => s.domain,
Self::Moment(m) => m.domain,
Self::Guv(g) => g.domain,
}
}
pub fn rows(&self) -> usize {
match self {
Self::Identity { domain } => *domain,
Self::Signature(s) => s.rows(),
Self::Moment(m) => m.rows(),
Self::Guv(g) => g.rows(),
}
}
pub fn kind(&self) -> &'static str {
match self {
Self::Identity { .. } => "identity",
Self::Signature(_) => "signature-hash",
Self::Moment(_) => "moment-hash",
Self::Guv(_) => "guv-expander",
}
}
#[inline]
pub fn is_identity(&self) -> bool {
matches!(self, Self::Identity { .. })
}
#[inline]
pub fn rows_for_index(&self, index: usize) -> Vec<usize> {
assert!(index < self.domain());
match self {
Self::Identity { .. } => vec![index],
Self::Signature(s) => s.rows_for_index(index),
Self::Moment(m) => m.rows_for_index(index),
Self::Guv(g) => g.rows_for_index(index),
}
}
#[inline]
pub fn weighted_rows_for_index(&self, index: usize) -> Vec<(usize, i64)> {
assert!(index < self.domain());
match self {
Self::Identity { .. } => vec![(index, 1)],
Self::Signature(s) => s
.rows_for_index(index)
.into_iter()
.map(|r| (r, 1))
.collect(),
Self::Moment(m) => m.weighted_rows_for_index(index),
Self::Guv(g) => g
.rows_for_index(index)
.into_iter()
.map(|r| (r, 1))
.collect(),
}
}
}
#[derive(Clone, Debug)]
pub struct SignatureRecovery {
pub domain: usize,
pub capacity: usize,
pub bucket_count: usize,
pub degree: usize,
pub bits: usize,
pub seed: u64,
}
impl SignatureRecovery {
pub fn new(
domain: usize,
capacity: usize,
degree: usize,
oversampling: f64,
seed: u64,
) -> Self {
assert!(domain > 0);
assert!(capacity > 0 && capacity <= domain);
assert!(degree > 0);
assert!(oversampling.is_finite() && oversampling > 0.0);
let bucket_count = (((capacity as f64) * oversampling).ceil() as usize)
.max(degree)
.min(domain.max(1));
let degree = degree.min(bucket_count);
let bits = ceil_log2(domain.saturating_add(1)).max(1);
Self {
domain,
capacity,
bucket_count,
degree,
bits,
seed,
}
}
#[inline]
pub fn rows(&self) -> usize {
self.bucket_count.saturating_mul(self.bits)
}
#[inline]
pub fn neighbors(&self, index: usize) -> Vec<usize> {
assert!(index < self.domain);
hashed_neighbors(self.seed, self.bucket_count, self.degree, index)
}
#[inline]
pub fn rows_for_index(&self, index: usize) -> Vec<usize> {
let code = index + 1;
let neighbors = self.neighbors(index);
let mut out = Vec::with_capacity(self.degree * self.bits);
for bucket in neighbors {
let base = bucket * self.bits;
for bit in 0..self.bits {
if ((code >> bit) & 1) != 0 {
out.push(base + bit);
}
}
}
out
}
}
#[derive(Clone, Debug)]
pub struct MomentRecovery {
pub domain: usize,
pub capacity: usize,
pub bucket_count: usize,
pub degree: usize,
pub seed: u64,
}
impl MomentRecovery {
pub fn new(
domain: usize,
capacity: usize,
degree: usize,
oversampling: f64,
seed: u64,
) -> Self {
assert!(domain > 0);
assert!(capacity > 0 && capacity <= domain);
assert!(degree > 0);
assert!(oversampling.is_finite() && oversampling > 0.0);
let bucket_count = (((capacity as f64) * oversampling).ceil() as usize)
.max(degree)
.min(domain.max(1));
let degree = degree.min(bucket_count);
Self {
domain,
capacity,
bucket_count,
degree,
seed,
}
}
#[inline]
pub fn rows(&self) -> usize {
self.bucket_count.saturating_mul(3)
}
#[inline]
pub fn neighbors(&self, index: usize) -> Vec<usize> {
assert!(index < self.domain);
hashed_neighbors(self.seed, self.bucket_count, self.degree, index)
}
#[inline]
pub fn rows_for_index(&self, index: usize) -> Vec<usize> {
let mut out = Vec::with_capacity(self.degree * 3);
for bucket in self.neighbors(index) {
let base = bucket * 3;
out.extend_from_slice(&[base, base + 1, base + 2]);
}
out
}
#[inline]
pub fn weighted_rows_for_index(&self, index: usize) -> Vec<(usize, i64)> {
let code = (index + 1) as i64;
let code2 = code.saturating_mul(code);
let mut out = Vec::with_capacity(self.degree * 3);
for bucket in self.neighbors(index) {
let base = bucket * 3;
out.push((base, 1));
out.push((base + 1, code));
out.push((base + 2, code2));
}
out
}
}
trait ExpanderSignature {
fn domain(&self) -> usize;
fn bucket_count(&self) -> usize;
fn degree(&self) -> usize;
fn bits(&self) -> usize;
fn neighbors(&self, index: usize) -> Vec<usize>;
fn rows_for_index(&self, index: usize) -> Vec<usize>;
fn rows(&self) -> usize;
}
impl ExpanderSignature for SignatureRecovery {
fn domain(&self) -> usize {
self.domain
}
fn bucket_count(&self) -> usize {
self.bucket_count
}
fn degree(&self) -> usize {
self.degree
}
fn bits(&self) -> usize {
self.bits
}
fn neighbors(&self, index: usize) -> Vec<usize> {
SignatureRecovery::neighbors(self, index)
}
fn rows_for_index(&self, index: usize) -> Vec<usize> {
SignatureRecovery::rows_for_index(self, index)
}
fn rows(&self) -> usize {
SignatureRecovery::rows(self)
}
}
impl ExpanderSignature for GuvRecovery {
fn domain(&self) -> usize {
self.domain
}
fn bucket_count(&self) -> usize {
self.bucket_count
}
fn degree(&self) -> usize {
self.degree
}
fn bits(&self) -> usize {
self.bits
}
fn neighbors(&self, index: usize) -> Vec<usize> {
GuvRecovery::neighbors(self, index)
}
fn rows_for_index(&self, index: usize) -> Vec<usize> {
GuvRecovery::rows_for_index(self, index)
}
fn rows(&self) -> usize {
GuvRecovery::rows(self)
}
}
#[derive(Clone, Debug)]
pub struct NestedOptions {
pub rectangular_policy: RectangularPolicy,
pub practical_scheduler: bool,
pub scheduler_k_hint: Option<usize>,
pub masked_residual: bool,
pub exact_k_bound: bool,
pub residual_fingerprint: Option<FingerprintConfig>,
pub fingerprint_failure_correction: bool,
}
impl Default for NestedOptions {
fn default() -> Self {
Self {
rectangular_policy: RectangularPolicy::Auto,
practical_scheduler: false,
scheduler_k_hint: None,
masked_residual: false,
exact_k_bound: false,
residual_fingerprint: None,
fingerprint_failure_correction: false,
}
}
}
#[derive(Clone, Debug, Default)]
pub struct NestedSpGemmStats {
pub rounds: Vec<NestedRoundStats>,
pub correction_pass: Option<CorrectionPassStats>,
pub terminated_early: bool,
pub termination_reason: Option<String>,
pub scheduler_skipped_rounds: usize,
pub fingerprint: Option<FingerprintStats>,
pub fingerprint_setup_time: Duration,
pub fingerprint_check_time: Duration,
pub fingerprint_verified: bool,
pub deterministic_verified: bool,
}
#[derive(Clone, Debug)]
pub struct NestedRoundStats {
pub params: RoundParams,
pub h_rows: usize,
pub g_rows: usize,
pub h_kind: &'static str,
pub g_kind: &'static str,
pub w_nnz: usize,
pub outer_recovered: usize,
pub inner_updates: usize,
pub d_nnz_after: usize,
pub left_time: Duration,
pub right_time: Duration,
pub rectangular_time: Duration,
pub rectangular_kernel: RectangularKernel,
pub rectangular_scalar_multiplications: u128,
pub rectangular_candidate_products: u128,
pub ha_density: f64,
pub bgt_density: f64,
pub h_matrix_cache_hit: bool,
pub g_matrix_cache_hit: bool,
pub ha_cache_hit: bool,
pub bgt_cache_hit: bool,
pub product_cache_hit: bool,
pub residual_cache_hit: bool,
pub masked_residual: bool,
pub active_mask_columns: usize,
pub scheduler_target_q: usize,
pub residual_measure_time: Duration,
pub decode_time: Duration,
}
#[derive(Clone, Debug)]
pub struct CorrectionPassStats {
pub residual_columns: usize,
pub residual_nnz: usize,
pub elapsed: Duration,
}
pub fn nested_spgemm<A, B>(
a: &A,
b: &B,
k_bound: usize,
backend: RecoveryBackend,
) -> (CsrMatrix, NestedSpGemmStats)
where
A: CsrInput<Scalar = i64> + ?Sized,
B: CsrInput<Scalar = i64> + ?Sized,
{
nested_spgemm_with_options(a, b, k_bound, backend, NestedOptions::default())
}
pub fn nested_spgemm_with_policy<A, B>(
a: &A,
b: &B,
k_bound: usize,
backend: RecoveryBackend,
rectangular_policy: RectangularPolicy,
) -> (CsrMatrix, NestedSpGemmStats)
where
A: CsrInput<Scalar = i64> + ?Sized,
B: CsrInput<Scalar = i64> + ?Sized,
{
nested_spgemm_with_options(
a,
b,
k_bound,
backend,
NestedOptions {
rectangular_policy,
..NestedOptions::default()
},
)
}
pub fn nested_spgemm_with_options<A, B>(
a: &A,
b: &B,
k_bound: usize,
backend: RecoveryBackend,
options: NestedOptions,
) -> (CsrMatrix, NestedSpGemmStats)
where
A: CsrInput<Scalar = i64> + ?Sized,
B: CsrInput<Scalar = i64> + ?Sized,
{
assert_eq!(a.cols(), b.rows(), "incompatible matrix dimensions");
let r = a.rows();
let c = b.cols();
let k = k_bound.min(r.saturating_mul(c));
if r == 0 || c == 0 || k == 0 {
return (CsrMatrix::zeros(r, c), NestedSpGemmStats::default());
}
let schedule = paper_schedule(r, c, k);
let mut d = SparseColumns::zeros(r, c);
let mut d_version = 0u64;
let mut stats = NestedSpGemmStats::default();
let mut support_mask: Option<Vec<bool>> = None;
let mut scheduler_target_q = 0usize;
let fingerprint = options.residual_fingerprint.map(|cfg| {
let start = Instant::now();
let fp = ResidualFingerprint::new(a, b, cfg);
stats.fingerprint_setup_time = start.elapsed();
stats.fingerprint = Some(FingerprintStats {
lanes: fp.lanes(),
seed: fp.seed,
..FingerprintStats::default()
});
fp
});
let mut recovery_matrix_cache: HashMap<RecoveryMatrixKey, BinaryRecoveryMatrix> =
HashMap::new();
let mut left_factor_cache: HashMap<RecoveryMatrixKey, Arc<PreparedFactor>> = HashMap::new();
let mut right_factor_cache: HashMap<RecoveryMatrixKey, Arc<PreparedFactor>> = HashMap::new();
let mut product_cache: HashMap<(RecoveryMatrixKey, RecoveryMatrixKey), CachedProduct> =
HashMap::new();
let mut residual_cache: HashMap<ResidualCacheKey, CachedResidual> = HashMap::new();
let mut schedule_pos = 0usize;
while schedule_pos < schedule.len() {
let params = schedule[schedule_pos].clone();
schedule_pos += 1;
if options.practical_scheduler && scheduler_target_q > 0 && params.q < scheduler_target_q {
stats.scheduler_skipped_rounds += 1;
continue;
}
let (h, h_matrix_cache_hit, h_key) = build_recovery_matrix_cached(
&mut recovery_matrix_cache,
r,
params.q,
&backend,
0x4849_4E4E_4552,
params.i,
);
let (g, g_matrix_cache_hit, g_key) = build_recovery_matrix_cached(
&mut recovery_matrix_cache,
c,
params.p,
&backend,
0x4F55_5445_52,
params.i,
);
let active_mask = if options.masked_residual && g.is_identity() {
support_mask.as_ref().filter(|m| m.iter().any(|&x| x))
} else {
None
};
let masked_residual = active_mask.is_some();
let active_mask_columns = active_mask
.map(|m| m.iter().filter(|&&x| x).count())
.unwrap_or(c);
let pair_key = (h_key.clone(), g_key.clone());
let residual_key = ResidualCacheKey {
h: h_key.clone(),
g: g_key.clone(),
d_version,
};
let cache_pair = !masked_residual && h_key.cacheable() && g_key.cacheable();
let mut left_time = Duration::default();
let mut right_time = Duration::default();
let mut rectangular_time = Duration::default();
let mut residual_measure_time = Duration::default();
let mut product_cache_hit = false;
let mut residual_cache_hit = false;
let ha_cache_hit: bool;
let bgt_cache_hit: bool;
let (w, rectangular_stats) = if masked_residual {
let (ha, hit, elapsed) = get_or_build_factor(&mut left_factor_cache, &h_key, || {
left_recovery_sketch(a, &h)
});
ha_cache_hit = hit;
left_time = elapsed;
let start = Instant::now();
let bgt =
PreparedFactor::new(right_recovery_sketch_masked(b, &g, active_mask.unwrap()));
right_time = start.elapsed();
bgt_cache_hit = false;
let start = Instant::now();
let (hcgt, rect_stats) =
adaptive_matmul_prepared(&ha, &bgt, options.rectangular_policy);
rectangular_time = start.elapsed();
let start = Instant::now();
let hdgt = d.two_sided_measure_masked(&h, &g, active_mask.unwrap());
let mut residual = hcgt;
sub_assign_dense(&mut residual, &hdgt);
residual_measure_time = start.elapsed();
(Arc::new(residual), rect_stats)
} else if cache_pair {
if let Some(cached) = residual_cache.get(&residual_key) {
residual_cache_hit = true;
product_cache_hit = product_cache.contains_key(&pair_key);
ha_cache_hit = left_factor_cache.contains_key(&h_key);
bgt_cache_hit = right_factor_cache.contains_key(&g_key);
(cached.matrix.clone(), cached.product_stats.clone())
} else {
let (hcgt, rect_stats) = if let Some(cached) = product_cache.get(&pair_key) {
product_cache_hit = true;
ha_cache_hit = left_factor_cache.contains_key(&h_key);
bgt_cache_hit = right_factor_cache.contains_key(&g_key);
(cached.matrix.clone(), cached.stats.clone())
} else {
let (ha, hit, elapsed) =
get_or_build_factor(&mut left_factor_cache, &h_key, || {
left_recovery_sketch(a, &h)
});
ha_cache_hit = hit;
left_time = elapsed;
let (bgt, hit, elapsed) =
get_or_build_factor(&mut right_factor_cache, &g_key, || {
right_recovery_sketch(b, &g)
});
bgt_cache_hit = hit;
right_time = elapsed;
let start = Instant::now();
let (product, rect_stats) =
adaptive_matmul_prepared(&ha, &bgt, options.rectangular_policy);
rectangular_time = start.elapsed();
let product = Arc::new(product);
product_cache.insert(
pair_key.clone(),
CachedProduct {
matrix: product.clone(),
stats: rect_stats.clone(),
},
);
(product, rect_stats)
};
let start = Instant::now();
let hdgt = d.two_sided_measure(&h, &g);
let mut residual = (*hcgt).clone();
sub_assign_dense(&mut residual, &hdgt);
residual_measure_time = start.elapsed();
let residual = Arc::new(residual);
residual_cache.insert(
residual_key,
CachedResidual {
matrix: residual.clone(),
product_stats: rect_stats.clone(),
},
);
(residual, rect_stats)
}
} else {
let (ha, hit, elapsed) = get_or_build_factor(&mut left_factor_cache, &h_key, || {
left_recovery_sketch(a, &h)
});
ha_cache_hit = hit;
left_time = elapsed;
let (bgt, hit, elapsed) = get_or_build_factor(&mut right_factor_cache, &g_key, || {
right_recovery_sketch(b, &g)
});
bgt_cache_hit = hit;
right_time = elapsed;
let start = Instant::now();
let (hcgt, rect_stats) =
adaptive_matmul_prepared(&ha, &bgt, options.rectangular_policy);
rectangular_time = start.elapsed();
let start = Instant::now();
let hdgt = d.two_sided_measure(&h, &g);
let mut residual = hcgt;
sub_assign_dense(&mut residual, &hdgt);
residual_measure_time = start.elapsed();
(Arc::new(residual), rect_stats)
};
let rectangular_kernel = rectangular_stats
.kernel
.expect("adaptive rectangular multiplication did not select a kernel");
let w_nnz = w.nnz();
let certified_zero_residual =
!masked_residual && h.is_identity() && g.is_identity() && w_nnz == 0;
let start = Instant::now();
let outer = if certified_zero_residual {
Vec::new()
} else {
safe_decode_product(&g, w.as_ref(), params.p)
};
let outer_recovered = outer.len();
if options.masked_residual && g.is_identity() && support_mask.is_none() && !outer.is_empty()
{
let mut mask = vec![false; c];
for (j, _) in &outer {
mask[*j] = true;
}
support_mask = Some(mask);
}
let mut inner_updates = 0usize;
let mut round_changed_d = false;
let allow_mask_removal = h.is_identity() || matches!(&h, BinaryRecoveryMatrix::Moment(_));
for (j, xi) in outer {
let z = safe_decode_scalar(&h, &xi, params.q);
if !z.is_empty() {
round_changed_d |= d.add_to_column(j, &z);
inner_updates += 1;
if options.masked_residual && allow_mask_removal {
if let Some(mask) = support_mask.as_mut() {
mask[j] = false;
}
}
}
}
if round_changed_d {
d_version = d_version.wrapping_add(1);
residual_cache.clear();
}
let decode_time = start.elapsed();
let remaining_mask_columns = support_mask
.as_ref()
.map(|m| m.iter().filter(|&&x| x).count())
.unwrap_or(0);
if options.practical_scheduler && g.is_identity() {
let observed = remaining_mask_columns.max(if support_mask.is_none() {
outer_recovered
} else {
0
});
if observed > 0 {
let scheduler_k = options.scheduler_k_hint.unwrap_or(k);
let remaining_bound = scheduler_k.saturating_sub(d.nnz()).max(1);
let avg = remaining_bound.saturating_add(observed - 1) / observed;
let target = avg.max(1).next_power_of_two().min(r.max(1));
scheduler_target_q = scheduler_target_q.max(target);
}
}
let complete_identity_recovery = !masked_residual
&& h.is_identity()
&& g.is_identity()
&& outer_recovered > 0
&& inner_updates == outer_recovered;
let rectangular_scalar_multiplications = if product_cache_hit || residual_cache_hit {
0
} else {
rectangular_stats.scalar_multiplications
};
stats.rounds.push(NestedRoundStats {
params: params.clone(),
h_rows: h.rows(),
g_rows: g.rows(),
h_kind: h.kind(),
g_kind: g.kind(),
w_nnz,
outer_recovered,
inner_updates,
d_nnz_after: d.nnz(),
left_time,
right_time,
rectangular_time,
rectangular_kernel,
rectangular_scalar_multiplications,
rectangular_candidate_products: rectangular_stats.sparse_candidate_products,
ha_density: rectangular_stats.a_density,
bgt_density: rectangular_stats.b_density,
h_matrix_cache_hit,
g_matrix_cache_hit,
ha_cache_hit,
bgt_cache_hit,
product_cache_hit,
residual_cache_hit,
masked_residual,
active_mask_columns,
scheduler_target_q,
residual_measure_time,
decode_time,
});
if certified_zero_residual {
stats.deterministic_verified = true;
stats.terminated_early = true;
stats.termination_reason = Some(format!(
"identity residual certificate W=0 at round {}",
params.i
));
break;
}
if complete_identity_recovery {
stats.deterministic_verified = true;
stats.terminated_early = true;
stats.termination_reason = Some(format!(
"complete identity residual recovery at round {}",
params.i
));
break;
}
let mask_exhausted = support_mask
.as_ref()
.map(|m| !m.iter().any(|&x| x))
.unwrap_or(false);
if options.masked_residual && mask_exhausted {
if let Some(fp) = fingerprint.as_ref() {
let start = Instant::now();
let candidate = d.to_csr();
let passed = fp.verifies(&candidate);
stats.fingerprint_check_time += start.elapsed();
if let Some(fps) = stats.fingerprint.as_mut() {
fps.checks += 1;
if passed {
fps.passes += 1;
} else {
fps.failures += 1;
}
}
if passed {
stats.fingerprint_verified = true;
stats.terminated_early = true;
stats.termination_reason = Some(format!(
"residual fingerprint certified AB-D=0 in round {}",
params.i
));
break;
} else {
support_mask = None;
}
}
}
if options.exact_k_bound
&& options.masked_residual
&& support_mask
.as_ref()
.map(|m| !m.iter().any(|&x| x))
.unwrap_or(false)
&& d.nnz() == k
{
stats.deterministic_verified = true;
stats.terminated_early = true;
stats.termination_reason = Some(format!(
"observed residual support exhausted at exact K in round {}",
params.i
));
break;
}
if options.exact_k_bound
&& support_mask
.as_ref()
.map(|m| !m.iter().any(|&x| x))
.unwrap_or(false)
&& d.nnz() < k
{
support_mask = None;
}
}
if !stats.fingerprint_verified && !stats.deterministic_verified {
if let Some(fp) = fingerprint.as_ref() {
let start = Instant::now();
let candidate = d.to_csr();
let passed = fp.verifies(&candidate);
stats.fingerprint_check_time += start.elapsed();
if let Some(fps) = stats.fingerprint.as_mut() {
fps.checks += 1;
if passed {
fps.passes += 1;
} else {
fps.failures += 1;
}
}
if passed {
stats.fingerprint_verified = true;
stats.terminated_early = true;
stats.termination_reason =
Some("final residual fingerprint certified AB-D=0".to_string());
}
}
}
let guaranteed_correction = match &backend {
RecoveryBackend::Identity => false,
RecoveryBackend::Signature(cfg) => cfg.guaranteed_correction,
RecoveryBackend::Moment(cfg) => cfg.guaranteed_correction,
RecoveryBackend::Guv(cfg) => cfg.guaranteed_correction,
};
let exact_k_safety = options.masked_residual && options.exact_k_bound && d.nnz() != k;
let fingerprint_safety = options.residual_fingerprint.is_some()
&& (options.fingerprint_failure_correction || options.exact_k_bound)
&& !stats.fingerprint_verified
&& !stats.deterministic_verified;
if guaranteed_correction || exact_k_safety || fingerprint_safety {
let start = Instant::now();
let h = BinaryRecoveryMatrix::identity(r);
let g = BinaryRecoveryMatrix::identity(c);
let ha = left_recovery_sketch(a, &h);
let bgt = right_recovery_sketch(b, &g);
let mut residual = dense_matmul(&ha, &bgt);
let d_measure = d.two_sided_measure(&h, &g);
sub_assign_dense(&mut residual, &d_measure);
let mut residual_columns = 0usize;
let mut residual_nnz = 0usize;
for j in 0..c {
let mut col = Vec::new();
for i in 0..r {
let v = residual[(i, j)];
if v != 0 {
col.push((i, v));
}
}
if !col.is_empty() {
residual_columns += 1;
residual_nnz += col.len();
d.add_to_column(j, &col);
}
}
stats.correction_pass = Some(CorrectionPassStats {
residual_columns,
residual_nnz,
elapsed: start.elapsed(),
});
stats.deterministic_verified = true;
}
(d.to_csr(), stats)
}
fn get_or_build_factor<F>(
cache: &mut HashMap<RecoveryMatrixKey, Arc<PreparedFactor>>,
key: &RecoveryMatrixKey,
build: F,
) -> (Arc<PreparedFactor>, bool, Duration)
where
F: FnOnce() -> DenseMatrix,
{
if key.cacheable() {
if let Some(existing) = cache.get(key) {
return (existing.clone(), true, Duration::default());
}
}
let start = Instant::now();
let matrix = Arc::new(PreparedFactor::new(build()));
let elapsed = start.elapsed();
if key.cacheable() {
cache.insert(key.clone(), matrix.clone());
}
(matrix, false, elapsed)
}
fn build_recovery_matrix_cached(
cache: &mut HashMap<RecoveryMatrixKey, BinaryRecoveryMatrix>,
domain: usize,
capacity: usize,
backend: &RecoveryBackend,
salt: u64,
round: usize,
) -> (BinaryRecoveryMatrix, bool, RecoveryMatrixKey) {
let capacity = capacity.clamp(1, domain);
let hinted_key = recovery_matrix_key_hint(domain, capacity, backend, salt, round);
if hinted_key.cacheable() {
if let Some(existing) = cache.get(&hinted_key) {
return (existing.clone(), true, hinted_key);
}
}
let (built, actual_key) = build_recovery_matrix(domain, capacity, backend, salt, round);
if actual_key.cacheable() {
if let Some(existing) = cache.get(&actual_key) {
return (existing.clone(), true, actual_key);
}
cache.insert(actual_key.clone(), built.clone());
}
(built, false, actual_key)
}
fn recovery_matrix_key_hint(
domain: usize,
capacity: usize,
backend: &RecoveryBackend,
salt: u64,
round: usize,
) -> RecoveryMatrixKey {
match backend {
RecoveryBackend::Identity => RecoveryMatrixKey::Identity { domain },
RecoveryBackend::Signature(cfg) => {
let seed = cfg.seed ^ salt ^ (round as u64).wrapping_mul(0xA076_1D64_78BD_642F);
let sig = SignatureRecovery::new(domain, capacity, cfg.degree, cfg.oversampling, seed);
if cfg.identity_fallback && sig.rows() >= domain {
RecoveryMatrixKey::Identity { domain }
} else {
RecoveryMatrixKey::Signature {
domain,
capacity,
degree: cfg.degree,
oversampling_bits: cfg.oversampling.to_bits(),
seed,
}
}
}
RecoveryBackend::Moment(cfg) => {
let seed = cfg.seed ^ salt ^ (round as u64).wrapping_mul(0xA076_1D64_78BD_642F);
let moment = MomentRecovery::new(domain, capacity, cfg.degree, cfg.oversampling, seed);
if cfg.identity_fallback && moment.rows() >= domain {
RecoveryMatrixKey::Identity { domain }
} else {
RecoveryMatrixKey::Moment {
domain,
capacity,
degree: cfg.degree,
oversampling_bits: cfg.oversampling.to_bits(),
seed,
}
}
}
RecoveryBackend::Guv(cfg) => {
if capacity == 1 {
return RecoveryMatrixKey::Identity { domain };
}
let estimated = GuvParameters::estimated_rows(domain, capacity, cfg.alpha, cfg.epsilon);
if cfg.identity_fallback && estimated >= domain {
RecoveryMatrixKey::Identity { domain }
} else {
RecoveryMatrixKey::Guv {
domain,
capacity,
alpha_bits: cfg.alpha.to_bits(),
epsilon_bits: cfg.epsilon.to_bits(),
}
}
}
}
}
fn build_recovery_matrix(
domain: usize,
capacity: usize,
backend: &RecoveryBackend,
salt: u64,
round: usize,
) -> (BinaryRecoveryMatrix, RecoveryMatrixKey) {
match backend {
RecoveryBackend::Identity => (
BinaryRecoveryMatrix::identity(domain),
RecoveryMatrixKey::Identity { domain },
),
RecoveryBackend::Signature(cfg) => {
let seed = cfg.seed ^ salt ^ (round as u64).wrapping_mul(0xA076_1D64_78BD_642F);
let sig = SignatureRecovery::new(domain, capacity, cfg.degree, cfg.oversampling, seed);
if cfg.identity_fallback && sig.rows() >= domain {
(
BinaryRecoveryMatrix::identity(domain),
RecoveryMatrixKey::Identity { domain },
)
} else {
(
BinaryRecoveryMatrix::Signature(sig),
RecoveryMatrixKey::Signature {
domain,
capacity,
degree: cfg.degree,
oversampling_bits: cfg.oversampling.to_bits(),
seed,
},
)
}
}
RecoveryBackend::Moment(cfg) => {
let seed = cfg.seed ^ salt ^ (round as u64).wrapping_mul(0xA076_1D64_78BD_642F);
let moment = MomentRecovery::new(domain, capacity, cfg.degree, cfg.oversampling, seed);
if cfg.identity_fallback && moment.rows() >= domain {
(
BinaryRecoveryMatrix::identity(domain),
RecoveryMatrixKey::Identity { domain },
)
} else {
(
BinaryRecoveryMatrix::Moment(moment),
RecoveryMatrixKey::Moment {
domain,
capacity,
degree: cfg.degree,
oversampling_bits: cfg.oversampling.to_bits(),
seed,
},
)
}
}
RecoveryBackend::Guv(cfg) => {
if capacity == 1 {
return (
BinaryRecoveryMatrix::identity(domain),
RecoveryMatrixKey::Identity { domain },
);
}
let estimated = GuvParameters::estimated_rows(domain, capacity, cfg.alpha, cfg.epsilon);
if cfg.identity_fallback && estimated >= domain {
return (
BinaryRecoveryMatrix::identity(domain),
RecoveryMatrixKey::Identity { domain },
);
}
let guv =
GuvRecovery::new(domain, capacity, cfg.alpha, cfg.epsilon).unwrap_or_else(|e| {
panic!("failed to construct explicit GUV recovery matrix: {e}")
});
if cfg.identity_fallback && guv.rows() >= domain {
(
BinaryRecoveryMatrix::identity(domain),
RecoveryMatrixKey::Identity { domain },
)
} else {
(
BinaryRecoveryMatrix::Guv(guv),
RecoveryMatrixKey::Guv {
domain,
capacity,
alpha_bits: cfg.alpha.to_bits(),
epsilon_bits: cfg.epsilon.to_bits(),
},
)
}
}
}
}
pub fn left_recovery_sketch<A>(a: &A, h: &BinaryRecoveryMatrix) -> DenseMatrix
where
A: CsrInput<Scalar = i64> + ?Sized,
{
assert_eq!(h.domain(), a.rows());
if h.is_identity() {
return a.to_dense();
}
let mut out = DenseMatrix::zeros(h.rows(), a.cols());
let rows_by_index: Vec<Vec<(usize, i64)>> = (0..h.domain())
.map(|i| h.weighted_rows_for_index(i))
.collect();
for i in 0..a.rows() {
let measurement_rows = &rows_by_index[i];
for (k, value) in a.row(i) {
for &(mr, coeff) in measurement_rows {
out[(mr, k)] += coeff * value;
}
}
}
out
}
pub fn right_recovery_sketch<B>(b: &B, g: &BinaryRecoveryMatrix) -> DenseMatrix
where
B: CsrInput<Scalar = i64> + ?Sized,
{
assert_eq!(g.domain(), b.cols());
if g.is_identity() {
return b.to_dense();
}
let mut out = DenseMatrix::zeros(b.rows(), g.rows());
let rows_by_column: Vec<Vec<(usize, i64)>> = (0..g.domain())
.map(|j| g.weighted_rows_for_index(j))
.collect();
for k in 0..b.rows() {
for (j, value) in b.row(k) {
for &(mr, coeff) in &rows_by_column[j] {
out[(k, mr)] += value * coeff;
}
}
}
out
}
pub fn right_recovery_sketch_masked<B>(
b: &B,
g: &BinaryRecoveryMatrix,
mask: &[bool],
) -> DenseMatrix
where
B: CsrInput<Scalar = i64> + ?Sized,
{
assert_eq!(g.domain(), b.cols());
assert_eq!(mask.len(), b.cols());
if g.is_identity() {
let mut out = DenseMatrix::zeros(b.rows(), b.cols());
for k in 0..b.rows() {
for (j, value) in b.row(k) {
if mask[j] {
out[(k, j)] = value;
}
}
}
return out;
}
let mut out = DenseMatrix::zeros(b.rows(), g.rows());
let rows_by_column: Vec<Vec<(usize, i64)>> = (0..g.domain())
.map(|j| g.weighted_rows_for_index(j))
.collect();
for k in 0..b.rows() {
for (j, value) in b.row(k) {
if !mask[j] {
continue;
}
for &(mr, coeff) in &rows_by_column[j] {
out[(k, mr)] += value * coeff;
}
}
}
out
}
pub fn safe_decode_scalar(
h: &BinaryRecoveryMatrix,
measurement: &[i64],
capacity: usize,
) -> Vec<(usize, i64)> {
assert_eq!(measurement.len(), h.rows());
let capacity = capacity.min(h.domain());
if capacity == 0 {
return Vec::new();
}
match h {
BinaryRecoveryMatrix::Identity { domain } => {
let mut out = Vec::new();
for (i, &v) in measurement.iter().take(*domain).enumerate() {
if v != 0 {
out.push((i, v));
}
}
if out.len() <= capacity {
out
} else {
Vec::new()
}
}
BinaryRecoveryMatrix::Signature(sig) => {
expander_safe_decode_scalar(sig, measurement, capacity)
}
BinaryRecoveryMatrix::Moment(moment) => {
moment_safe_decode_scalar(moment, measurement, capacity)
}
BinaryRecoveryMatrix::Guv(guv) => expander_safe_decode_scalar(guv, measurement, capacity),
}
}
pub fn safe_decode_product(
g: &BinaryRecoveryMatrix,
measurement: &DenseMatrix,
capacity: usize,
) -> Vec<(usize, Vec<i64>)> {
assert_eq!(measurement.cols, g.rows());
let capacity = capacity.min(g.domain());
if capacity == 0 {
return Vec::new();
}
match g {
BinaryRecoveryMatrix::Identity { domain } => {
let mut out = Vec::new();
for j in 0..*domain {
let col = dense_col(measurement, j);
if !is_zero_vec(&col) {
out.push((j, col));
}
}
if out.len() <= capacity {
out
} else {
Vec::new()
}
}
BinaryRecoveryMatrix::Signature(sig) => {
expander_safe_decode_product(sig, measurement, capacity)
}
BinaryRecoveryMatrix::Moment(moment) => {
moment_safe_decode_product(moment, measurement, capacity)
}
BinaryRecoveryMatrix::Guv(guv) => expander_safe_decode_product(guv, measurement, capacity),
}
}
fn moment_safe_decode_scalar(
moment: &MomentRecovery,
measurement: &[i64],
capacity: usize,
) -> Vec<(usize, i64)> {
let original = measurement.to_vec();
let mut residual = original.clone();
let mut total: BTreeMap<usize, i64> = BTreeMap::new();
for _ in 0..capacity.max(1) {
let proposal = moment_reduce_scalar(moment, &residual);
if proposal.is_empty() {
break;
}
if total.len().saturating_add(proposal.len()) > capacity {
return Vec::new();
}
for &(idx, value) in &proposal {
add_scalar_entry(&mut total, idx, value);
}
let measured = measure_scalar_moment(moment, &proposal);
for (x, y) in residual.iter_mut().zip(measured) {
*x -= y;
}
}
let out: Vec<(usize, i64)> = total.into_iter().filter(|&(_, v)| v != 0).collect();
if out.len() > capacity || measure_scalar_moment(moment, &out) != original {
return Vec::new();
}
out
}
fn moment_reduce_scalar(moment: &MomentRecovery, measurement: &[i64]) -> Vec<(usize, i64)> {
let mut candidates: BTreeMap<usize, Option<i64>> = BTreeMap::new();
for bucket in 0..moment.bucket_count {
let base = bucket * 3;
let s0 = measurement[base];
let s1 = measurement[base + 1];
let s2 = measurement[base + 2];
let Some((idx, value)) = moment_singleton_scalar(moment, bucket, s0, s1, s2) else {
continue;
};
match candidates.entry(idx) {
std::collections::btree_map::Entry::Vacant(e) => {
e.insert(Some(value));
}
std::collections::btree_map::Entry::Occupied(mut e) => {
let conflict = match e.get().as_ref() {
Some(existing) => *existing != value,
None => false,
};
if conflict {
e.insert(None);
}
}
}
}
candidates
.into_iter()
.filter_map(|(idx, value)| value.map(|v| (idx, v)))
.collect()
}
fn moment_singleton_scalar(
moment: &MomentRecovery,
bucket: usize,
s0: i64,
s1: i64,
s2: i64,
) -> Option<(usize, i64)> {
if s0 == 0 || s1 % s0 != 0 {
return None;
}
let code = s1 / s0;
if code <= 0 || code as usize > moment.domain {
return None;
}
if s2 != code * code * s0 {
return None;
}
let idx = code as usize - 1;
moment.neighbors(idx).contains(&bucket).then_some((idx, s0))
}
fn measure_scalar_moment(moment: &MomentRecovery, entries: &[(usize, i64)]) -> Vec<i64> {
let mut out = vec![0i64; moment.rows()];
for &(idx, value) in entries {
for (row, coeff) in moment.weighted_rows_for_index(idx) {
out[row] += coeff * value;
}
}
out
}
fn moment_safe_decode_product(
moment: &MomentRecovery,
measurement: &DenseMatrix,
capacity: usize,
) -> Vec<(usize, Vec<i64>)> {
let original = measurement.clone();
let mut residual = measurement.clone();
let mut total: BTreeMap<usize, Vec<i64>> = BTreeMap::new();
for _ in 0..capacity.max(1) {
let proposal = moment_reduce_product(moment, &residual);
if proposal.is_empty() {
break;
}
if total.len().saturating_add(proposal.len()) > capacity {
return Vec::new();
}
for (idx, value) in &proposal {
add_product_entry(&mut total, *idx, value);
}
let measured = measure_product_moment(moment, residual.rows, &proposal);
sub_assign_dense(&mut residual, &measured);
}
let out: Vec<(usize, Vec<i64>)> = total.into_iter().filter(|(_, v)| !is_zero_vec(v)).collect();
if out.len() > capacity || measure_product_moment(moment, original.rows, &out) != original {
return Vec::new();
}
out
}
fn moment_reduce_product(
moment: &MomentRecovery,
measurement: &DenseMatrix,
) -> Vec<(usize, Vec<i64>)> {
let mut candidates: BTreeMap<usize, Option<Vec<i64>>> = BTreeMap::new();
for bucket in 0..moment.bucket_count {
let base = bucket * 3;
let s0 = dense_col(measurement, base);
let s1 = dense_col(measurement, base + 1);
let s2 = dense_col(measurement, base + 2);
let Some((idx, value)) = moment_singleton_product(moment, bucket, &s0, &s1, &s2) else {
continue;
};
match candidates.entry(idx) {
std::collections::btree_map::Entry::Vacant(e) => {
e.insert(Some(value));
}
std::collections::btree_map::Entry::Occupied(mut e) => {
let conflict = match e.get().as_ref() {
Some(existing) => existing.as_slice() != value.as_slice(),
None => false,
};
if conflict {
e.insert(None);
}
}
}
}
candidates
.into_iter()
.filter_map(|(idx, value)| value.map(|v| (idx, v)))
.collect()
}
fn moment_singleton_product(
moment: &MomentRecovery,
bucket: usize,
s0: &[i64],
s1: &[i64],
s2: &[i64],
) -> Option<(usize, Vec<i64>)> {
let mut code: Option<i64> = None;
for r in 0..s0.len() {
if s0[r] == 0 {
if s1[r] != 0 || s2[r] != 0 {
return None;
}
continue;
}
if s1[r] % s0[r] != 0 {
return None;
}
let c = s1[r] / s0[r];
match code {
None => code = Some(c),
Some(prev) if prev == c => {}
Some(_) => return None,
}
}
let code = code?;
if code <= 0 || code as usize > moment.domain {
return None;
}
let code2 = code * code;
for r in 0..s0.len() {
if s1[r] != code * s0[r] || s2[r] != code2 * s0[r] {
return None;
}
}
let idx = code as usize - 1;
moment
.neighbors(idx)
.contains(&bucket)
.then_some((idx, s0.to_vec()))
}
fn measure_product_moment(
moment: &MomentRecovery,
value_dim: usize,
entries: &[(usize, Vec<i64>)],
) -> DenseMatrix {
let mut out = DenseMatrix::zeros(value_dim, moment.rows());
for (idx, value) in entries {
assert_eq!(value.len(), value_dim);
for (row, coeff) in moment.weighted_rows_for_index(*idx) {
for i in 0..value_dim {
out[(i, row)] += coeff * value[i];
}
}
}
out
}
fn expander_safe_decode_scalar<E: ExpanderSignature>(
sig: &E,
measurement: &[i64],
capacity: usize,
) -> Vec<(usize, i64)> {
let original = measurement.to_vec();
let mut residual = original.clone();
let mut total: BTreeMap<usize, i64> = BTreeMap::new();
let iterations = ceil_log2(capacity.saturating_mul(2)).max(1);
for _ in 0..iterations {
let proposal = reduce_scalar(sig, &residual);
if proposal.len().saturating_mul(2) > capacity.saturating_mul(3) {
return Vec::new();
}
if proposal.is_empty() {
break;
}
for &(idx, value) in &proposal {
add_scalar_entry(&mut total, idx, value);
}
let measured = measure_scalar_expander(sig, &proposal);
for i in 0..residual.len() {
residual[i] -= measured[i];
}
}
let out: Vec<(usize, i64)> = total.into_iter().filter(|&(_, v)| v != 0).collect();
if out.len() > capacity {
return Vec::new();
}
if measure_scalar_expander(sig, &out) != original {
return Vec::new();
}
out
}
fn reduce_scalar<E: ExpanderSignature>(sig: &E, measurement: &[i64]) -> Vec<(usize, i64)> {
let mut counts: HashMap<(usize, i64), usize> = HashMap::new();
for bucket in 0..sig.bucket_count() {
let base = bucket * sig.bits();
let block = &measurement[base..base + sig.bits()];
let Some(value) = unique_nonzero_scalar(block) else {
continue;
};
let mut code = 0usize;
for (bit, &v) in block.iter().enumerate() {
if v == value {
code |= 1usize << bit;
}
}
if code == 0 || code > sig.domain() {
continue;
}
let index = code - 1;
if !sig.neighbors(index).contains(&bucket) {
continue;
}
*counts.entry((index, value)).or_insert(0) += 1;
}
choose_majority_scalar(counts, sig.degree())
}
fn expander_safe_decode_product<E: ExpanderSignature>(
sig: &E,
measurement: &DenseMatrix,
capacity: usize,
) -> Vec<(usize, Vec<i64>)> {
let original = measurement.clone();
let mut residual = measurement.clone();
let mut total: BTreeMap<usize, Vec<i64>> = BTreeMap::new();
let iterations = ceil_log2(capacity.saturating_mul(2)).max(1);
for _ in 0..iterations {
let proposal = reduce_product(sig, &residual);
if proposal.len().saturating_mul(2) > capacity.saturating_mul(3) {
return Vec::new();
}
if proposal.is_empty() {
break;
}
for (idx, value) in &proposal {
add_product_entry(&mut total, *idx, value);
}
let measured = measure_product_expander(sig, residual.rows, &proposal);
sub_assign_dense(&mut residual, &measured);
}
let out: Vec<(usize, Vec<i64>)> = total.into_iter().filter(|(_, v)| !is_zero_vec(v)).collect();
if out.len() > capacity {
return Vec::new();
}
if measure_product_expander(sig, original.rows, &out) != original {
return Vec::new();
}
out
}
fn reduce_product<E: ExpanderSignature>(
sig: &E,
measurement: &DenseMatrix,
) -> Vec<(usize, Vec<i64>)> {
let mut counts: HashMap<(usize, Vec<i64>), usize> = HashMap::new();
for bucket in 0..sig.bucket_count() {
let base = bucket * sig.bits();
let mut value: Option<Vec<i64>> = None;
let mut valid = true;
for bit in 0..sig.bits() {
let col = dense_col(measurement, base + bit);
if is_zero_vec(&col) {
continue;
}
match &value {
None => value = Some(col),
Some(v) if v.as_slice() == col.as_slice() => {}
Some(_) => {
valid = false;
break;
}
}
}
if !valid {
continue;
}
let Some(value) = value else {
continue;
};
let mut code = 0usize;
for bit in 0..sig.bits() {
if dense_col_eq(measurement, base + bit, &value) {
code |= 1usize << bit;
}
}
if code == 0 || code > sig.domain() {
continue;
}
let index = code - 1;
if !sig.neighbors(index).contains(&bucket) {
continue;
}
*counts.entry((index, value)).or_insert(0) += 1;
}
choose_majority_product(counts, sig.degree())
}
fn choose_majority_scalar(
counts: HashMap<(usize, i64), usize>,
degree: usize,
) -> Vec<(usize, i64)> {
let mut best: BTreeMap<usize, (i64, usize, bool)> = BTreeMap::new();
for ((idx, value), count) in counts {
if count * 2 <= degree {
continue;
}
match best.get_mut(&idx) {
None => {
best.insert(idx, (value, count, false));
}
Some((best_value, best_count, tied)) => {
if count > *best_count {
*best_value = value;
*best_count = count;
*tied = false;
} else if count == *best_count && value != *best_value {
*tied = true;
}
}
}
}
best.into_iter()
.filter_map(|(idx, (value, _, tied))| (!tied).then_some((idx, value)))
.collect()
}
fn choose_majority_product(
counts: HashMap<(usize, Vec<i64>), usize>,
degree: usize,
) -> Vec<(usize, Vec<i64>)> {
let mut best: BTreeMap<usize, (Vec<i64>, usize, bool)> = BTreeMap::new();
for ((idx, value), count) in counts {
if count * 2 <= degree {
continue;
}
match best.get_mut(&idx) {
None => {
best.insert(idx, (value, count, false));
}
Some((best_value, best_count, tied)) => {
if count > *best_count {
*best_value = value;
*best_count = count;
*tied = false;
} else if count == *best_count && value.as_slice() != best_value.as_slice() {
*tied = true;
}
}
}
}
best.into_iter()
.filter_map(|(idx, (value, _, tied))| (!tied).then_some((idx, value)))
.collect()
}
fn measure_scalar_expander<E: ExpanderSignature>(sig: &E, entries: &[(usize, i64)]) -> Vec<i64> {
let mut out = vec![0i64; sig.rows()];
for &(idx, value) in entries {
for row in sig.rows_for_index(idx) {
out[row] += value;
}
}
out
}
fn measure_product_expander<E: ExpanderSignature>(
sig: &E,
value_dim: usize,
entries: &[(usize, Vec<i64>)],
) -> DenseMatrix {
let mut out = DenseMatrix::zeros(value_dim, sig.rows());
for (idx, value) in entries {
assert_eq!(value.len(), value_dim);
for mr in sig.rows_for_index(*idx) {
for i in 0..value_dim {
out[(i, mr)] += value[i];
}
}
}
out
}
fn unique_nonzero_scalar(block: &[i64]) -> Option<i64> {
let mut value = None;
for &v in block {
if v == 0 {
continue;
}
match value {
None => value = Some(v),
Some(x) if x == v => {}
Some(_) => return None,
}
}
value
}
fn dense_col(m: &DenseMatrix, col: usize) -> Vec<i64> {
(0..m.rows).map(|r| m[(r, col)]).collect()
}
fn dense_col_eq(m: &DenseMatrix, col: usize, value: &[i64]) -> bool {
value.len() == m.rows && (0..m.rows).all(|r| m[(r, col)] == value[r])
}
fn is_zero_vec(v: &[i64]) -> bool {
v.iter().all(|&x| x == 0)
}
fn add_scalar_entry(map: &mut BTreeMap<usize, i64>, idx: usize, value: i64) {
let new = map.get(&idx).copied().unwrap_or(0) + value;
if new == 0 {
map.remove(&idx);
} else {
map.insert(idx, new);
}
}
fn add_product_entry(map: &mut BTreeMap<usize, Vec<i64>>, idx: usize, value: &[i64]) {
if let Some(current) = map.get_mut(&idx) {
assert_eq!(current.len(), value.len());
for i in 0..current.len() {
current[i] += value[i];
}
let remove = is_zero_vec(current);
if remove {
map.remove(&idx);
}
} else if !is_zero_vec(value) {
map.insert(idx, value.to_vec());
}
}
fn sub_assign_dense(lhs: &mut DenseMatrix, rhs: &DenseMatrix) {
assert_eq!(lhs.rows, rhs.rows);
assert_eq!(lhs.cols, rhs.cols);
for (x, y) in lhs.data.iter_mut().zip(&rhs.data) {
*x -= *y;
}
}
#[derive(Clone, Debug)]
struct SparseColumns {
rows: usize,
cols: usize,
data: Vec<BTreeMap<usize, i64>>,
}
impl SparseColumns {
fn zeros(rows: usize, cols: usize) -> Self {
Self {
rows,
cols,
data: (0..cols).map(|_| BTreeMap::new()).collect(),
}
}
fn nnz(&self) -> usize {
self.data.iter().map(BTreeMap::len).sum()
}
fn add_to_column(&mut self, col: usize, entries: &[(usize, i64)]) -> bool {
assert!(col < self.cols);
let mut changed = false;
for &(row, value) in entries {
assert!(row < self.rows);
if value == 0 {
continue;
}
let current = self.data[col].get(&row).copied().unwrap_or(0);
let new = current + value;
if new != current {
changed = true;
}
if new == 0 {
self.data[col].remove(&row);
} else {
self.data[col].insert(row, new);
}
}
changed
}
fn two_sided_measure(&self, h: &BinaryRecoveryMatrix, g: &BinaryRecoveryMatrix) -> DenseMatrix {
assert_eq!(h.domain(), self.rows);
assert_eq!(g.domain(), self.cols);
let mut out = DenseMatrix::zeros(h.rows(), g.rows());
for j in 0..self.cols {
if self.data[j].is_empty() {
continue;
}
let grows = g.weighted_rows_for_index(j);
for (&i, &value) in &self.data[j] {
let hrows = h.weighted_rows_for_index(i);
for &(hr, hc) in &hrows {
for &(gr, gc) in &grows {
out[(hr, gr)] += hc * value * gc;
}
}
}
}
out
}
fn two_sided_measure_masked(
&self,
h: &BinaryRecoveryMatrix,
g: &BinaryRecoveryMatrix,
mask: &[bool],
) -> DenseMatrix {
assert_eq!(mask.len(), self.cols);
assert_eq!(h.domain(), self.rows);
assert_eq!(g.domain(), self.cols);
let mut out = DenseMatrix::zeros(h.rows(), g.rows());
for j in 0..self.cols {
if !mask[j] || self.data[j].is_empty() {
continue;
}
let grows = g.weighted_rows_for_index(j);
for (&i, &value) in &self.data[j] {
let hrows = h.weighted_rows_for_index(i);
for &(hr, hc) in &hrows {
for &(gr, gc) in &grows {
out[(hr, gr)] += hc * value * gc;
}
}
}
}
out
}
fn to_csr(&self) -> CsrMatrix {
let mut triplets = Vec::with_capacity(self.nnz());
for j in 0..self.cols {
for (&i, &v) in &self.data[j] {
if v != 0 {
triplets.push((i, j, v));
}
}
}
CsrMatrix::from_triplets(self.rows, self.cols, &triplets)
}
}
fn hashed_neighbors(seed: u64, bucket_count: usize, degree: usize, index: usize) -> Vec<usize> {
let mut out = Vec::with_capacity(degree);
let mut attempt = 0u64;
while out.len() < degree {
let key = seed
^ (index as u64).wrapping_mul(0x9E37_79B9_7F4A_7C15)
^ attempt.wrapping_mul(0xD1B5_4A32_D192_ED03);
let b = (splitmix64(key) % bucket_count as u64) as usize;
if !out.contains(&b) {
out.push(b);
}
attempt = attempt.wrapping_add(1);
}
out
}
#[inline]
fn splitmix64(mut x: u64) -> u64 {
x = x.wrapping_add(0x9E37_79B9_7F4A_7C15);
x = (x ^ (x >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
x = (x ^ (x >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
x ^ (x >> 31)
}
fn ceil_log2(x: usize) -> usize {
if x <= 1 {
0
} else {
usize::BITS as usize - (x - 1).leading_zeros() as usize
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::spgemm::spgemm_hash;
use crate::synthetic::{overlap_problem, sparse_output_problem};
#[test]
fn signature_scalar_decoder_recovers_sparse_vector() {
let sig = SignatureRecovery::new(63, 3, 7, 12.0, 12345);
let h = BinaryRecoveryMatrix::Signature(sig.clone());
let x = vec![(2, 11), (17, -4), (51, 9)];
let y = measure_scalar_expander(&sig, &x);
assert_eq!(safe_decode_scalar(&h, &y, 3), x);
}
#[test]
fn signature_product_decoder_recovers_sparse_vector() {
let sig = SignatureRecovery::new(31, 2, 7, 8.0, 999);
let g = BinaryRecoveryMatrix::Signature(sig.clone());
let x = vec![(4, vec![1, 2, 3]), (22, vec![-5, 0, 7])];
let y = measure_product_expander(&sig, 3, &x);
assert_eq!(safe_decode_product(&g, &y, 2), x);
}
#[test]
fn guv_scalar_decoder_recovers_singleton() {
let guv = GuvRecovery::new(7, 2, 4.0, 1.0 / 12.0).unwrap();
let h = BinaryRecoveryMatrix::Guv(guv.clone());
let x = vec![(5, -13)];
let y = measure_scalar_expander(&guv, &x);
assert_eq!(safe_decode_scalar(&h, &y, 2), x);
}
#[test]
fn guv_backend_uses_identity_when_it_is_smaller() {
let cfg = GuvConfig::default();
let (h, _) = build_recovery_matrix(63, 4, &RecoveryBackend::Guv(cfg), 0, 0);
assert_eq!(h.kind(), "identity");
assert_eq!(h.rows(), 63);
}
#[test]
fn repeated_identity_recovery_matrices_are_cached_across_rounds() {
let problem = overlap_problem(5, 5, 5, 3, 0.2);
let (expected, _) = spgemm_hash(&problem.a, &problem.b);
let (_, stats) = nested_spgemm_with_policy(
&problem.a,
&problem.b,
expected.nnz(),
RecoveryBackend::Identity,
RectangularPolicy::Auto,
);
assert!(stats
.rounds
.iter()
.any(|round| round.h_matrix_cache_hit || round.g_matrix_cache_hit));
}
#[test]
fn identity_reuses_residual_and_terminates_after_full_domain_recovery() {
let problem = overlap_problem(16, 16, 16, 8, 0.5);
let (expected, _) = spgemm_hash(&problem.a, &problem.b);
let (actual, stats) = nested_spgemm_with_policy(
&problem.a,
&problem.b,
expected.nnz(),
RecoveryBackend::Identity,
RectangularPolicy::Auto,
);
assert_eq!(actual, expected);
assert!(stats.terminated_early);
assert!(stats.rounds.iter().any(|round| round.residual_cache_hit));
assert!(stats.rounds.iter().any(|round| round.product_cache_hit));
}
#[test]
fn signature_identity_fallback_becomes_cacheable() {
let problem = overlap_problem(16, 16, 16, 8, 0.5);
let (expected, _) = spgemm_hash(&problem.a, &problem.b);
let cfg = SignatureConfig {
degree: 5,
oversampling: 2.0,
seed: 42,
identity_fallback: true,
guaranteed_correction: false,
};
let (actual, stats) = nested_spgemm_with_policy(
&problem.a,
&problem.b,
expected.nnz(),
RecoveryBackend::Signature(cfg),
RectangularPolicy::Auto,
);
assert_eq!(actual, expected);
assert!(stats.rounds.iter().any(|round| {
round.h_kind == "identity"
&& (round.h_matrix_cache_hit || round.ha_cache_hit || round.residual_cache_hit)
}));
}
#[test]
fn identity_nested_matches_exact_spgemm() {
let problem = overlap_problem(12, 12, 10, 6, 0.4);
let (expected, _) = spgemm_hash(&problem.a, &problem.b);
let (actual, stats) = nested_spgemm(
&problem.a,
&problem.b,
expected.nnz(),
RecoveryBackend::Identity,
);
assert_eq!(actual, expected);
assert!(!stats.rounds.is_empty());
assert!(stats.correction_pass.is_none());
}
#[test]
fn signature_nested_with_correction_matches_exact_spgemm() {
let problem = overlap_problem(12, 12, 10, 6, 0.4);
let (expected, _) = spgemm_hash(&problem.a, &problem.b);
let cfg = SignatureConfig {
degree: 5,
oversampling: 3.0,
seed: 42,
identity_fallback: true,
guaranteed_correction: true,
};
let (actual, stats) = nested_spgemm(
&problem.a,
&problem.b,
expected.nnz(),
RecoveryBackend::Signature(cfg),
);
assert_eq!(actual, expected);
assert!(stats.correction_pass.is_some());
}
#[test]
fn moment_scalar_decoder_recovers_seven_sparse_vector() {
let moment = MomentRecovery::new(256, 8, 3, 3.0, 0x1234_5678);
assert_eq!(moment.rows(), 72);
let h = BinaryRecoveryMatrix::Moment(moment.clone());
let x = vec![
(3, 5),
(17, -2),
(49, 7),
(88, 11),
(129, -3),
(173, 4),
(241, 9),
];
let y = measure_scalar_moment(&moment, &x);
assert_eq!(safe_decode_scalar(&h, &y, 8), x);
}
#[test]
fn moment_product_decoder_recovers_sparse_direct_product_vector() {
let moment = MomentRecovery::new(64, 4, 3, 4.0, 0xBEEF);
let g = BinaryRecoveryMatrix::Moment(moment.clone());
let x = vec![
(5, vec![2, 0, -1]),
(19, vec![0, 7, 3]),
(42, vec![-4, 1, 0]),
];
let y = measure_product_moment(&moment, 3, &x);
assert_eq!(safe_decode_product(&g, &y, 4), x);
}
#[test]
fn practical_scheduler_and_masked_moment_recover_sparse_output() {
let problem = sparse_output_problem(32, 64, 64, 16, 5, 0.75, 64);
let (expected, _) = spgemm_hash(&problem.a, &problem.b);
let cfg = MomentConfig {
degree: 3,
oversampling: 3.0,
seed: 0x4D4F_4D45_4E54_0001,
identity_fallback: true,
guaranteed_correction: false,
};
let (actual, stats) = nested_spgemm_with_options(
&problem.a,
&problem.b,
expected.nnz(),
RecoveryBackend::Moment(cfg),
NestedOptions {
rectangular_policy: RectangularPolicy::Auto,
practical_scheduler: true,
scheduler_k_hint: None,
masked_residual: true,
exact_k_bound: true,
..NestedOptions::default()
},
);
assert_eq!(actual, expected);
assert!(
stats.scheduler_skipped_rounds > 0 || stats.rounds.iter().any(|r| r.masked_residual)
);
}
#[test]
fn fingerprint_certifies_masked_recovery_without_exact_k() {
use crate::fingerprint::FingerprintConfig;
let problem = sparse_output_problem(128, 256, 128, 32, 5, 0.5, 256);
let (expected, _) = spgemm_hash(&problem.a, &problem.b);
let cfg = MomentConfig {
degree: 3,
oversampling: 3.0,
seed: 0x4D4F_4D45_4E54_0001,
identity_fallback: true,
guaranteed_correction: false,
};
let (actual, stats) = nested_spgemm_with_options(
&problem.a,
&problem.b,
expected.nnz() * 2,
RecoveryBackend::Moment(cfg),
NestedOptions {
rectangular_policy: RectangularPolicy::Auto,
practical_scheduler: true,
scheduler_k_hint: Some(expected.nnz()),
masked_residual: true,
exact_k_bound: false,
residual_fingerprint: Some(FingerprintConfig { lanes: 3, seed: 99 }),
fingerprint_failure_correction: false,
},
);
assert_eq!(actual, expected);
assert!(stats.fingerprint_verified || stats.deterministic_verified);
assert!(stats
.fingerprint
.as_ref()
.map(|f| f.checks > 0)
.unwrap_or(false));
}
}