use std::collections::{BTreeMap, BTreeSet};
use gam_linalg::faer_ndarray::FaerEigh;
use gam_solve::row_sampling_measure::RowSamplingMeasure;
use ndarray::{Array1, Array2, ArrayView2};
use faer::Side;
use super::{SaeManifoldAtom, SaeManifoldTerm};
pub const TERRACINI_CERTIFIER_ROWS_PER_ATOM: usize = 10_000;
#[derive(Debug, Clone)]
pub struct ParseBlock {
pub atom: usize,
pub value: Array1<f64>,
pub tangent: Array2<f64>,
}
impl ParseBlock {
fn block_cols(&self) -> usize {
self.tangent.ncols() + 1
}
fn stacked(&self) -> Result<Array2<f64>, String> {
let p = self.value.len();
if self.tangent.ncols() != 0 && self.tangent.nrows() != p {
return Err(format!(
"ParseBlock(atom {}): tangent has {} rows but value has length {p}",
self.atom,
self.tangent.nrows()
));
}
let cols = self.block_cols();
let mut k = Array2::<f64>::zeros((p, cols));
for i in 0..p {
k[[i, 0]] = self.value[i];
}
for c in 0..self.tangent.ncols() {
for i in 0..p {
k[[i, c + 1]] = self.tangent[[i, c]];
}
}
Ok(k)
}
}
#[derive(Debug, Clone)]
pub struct TerraciniCertificate {
pub pattern: Vec<usize>,
pub p: usize,
pub m: usize,
pub margin: f64,
pub amplification: f64,
pub cross_gram_logdet: f64,
pub whitened_excess: f64,
pub attribution_risk: f64,
pub per_atom_logdet: Vec<f64>,
}
fn inverse_sqrt_and_logdet(gram: &Array2<f64>, ridge: f64) -> Result<(Array2<f64>, f64), String> {
let n = gram.nrows();
let mut g = gram.clone();
if ridge > 0.0 {
let tr = (0..n).map(|i| gram[[i, i]]).sum::<f64>().max(1.0e-300);
let floor = ridge * tr / (n as f64);
for i in 0..n {
g[[i, i]] += floor;
}
}
let (w, v) = g
.eigh(Side::Lower)
.map_err(|e| format!("terracini: eigh for whitening failed: {e}"))?;
let mut scaled = v.clone();
let mut logdet = 0.0_f64;
for c in 0..n {
let wc = w[c];
if !(wc.is_finite() && wc > 0.0) {
return Err(format!(
"terracini: per-atom block not positive-definite (eigenvalue {wc:.3e}); \
charge it as a per-atom rank failure, not a cross collision"
));
}
logdet += wc.ln();
let inv = 1.0 / wc.sqrt();
for r in 0..n {
scaled[[r, c]] *= inv;
}
}
Ok((scaled.dot(&v.t()), logdet))
}
pub fn parse_certificate(
blocks: &[ParseBlock],
noise_var: f64,
ridge: f64,
) -> Result<TerraciniCertificate, String> {
if blocks.is_empty() {
return Err("terracini: parse has no co-firing atoms".to_string());
}
let p = blocks[0].value.len();
if p == 0 {
return Err("terracini: ambient dimension p = 0".to_string());
}
for b in blocks {
if b.value.len() != p {
return Err(format!(
"terracini: atom {} value length {} != p = {p}",
b.atom,
b.value.len()
));
}
}
let m: usize = blocks.iter().map(ParseBlock::block_cols).sum();
let mut pattern: Vec<usize> = blocks.iter().map(|b| b.atom).collect();
pattern.sort_unstable();
if m > p {
return Err(format!(
"terracini: overcomplete parse — Σ(d_k+1) = {m} > p = {p}; by the Terracini \
bound the sparse manifold decomposition is NOT locally identifiable at this \
parse (rank(J_S) cannot reach {m})"
));
}
let mut j_s = Array2::<f64>::zeros((p, m));
let mut b_s = Array2::<f64>::zeros((p, m));
let mut per_atom_logdet = Vec::with_capacity(blocks.len());
let mut col = 0usize;
for b in blocks {
let k = b.stacked()?;
let cols = k.ncols();
let gram = k.t().dot(&k);
let (g_inv_sqrt, logdet) = inverse_sqrt_and_logdet(&gram, ridge)?;
per_atom_logdet.push(logdet);
let bk = k.dot(&g_inv_sqrt);
for c in 0..cols {
for i in 0..p {
j_s[[i, col + c]] = k[[i, c]];
b_s[[i, col + c]] = bk[[i, c]];
}
}
col += cols;
}
let btb = b_s.t().dot(&b_s);
let w_b = btb
.eigh(Side::Lower)
.map_err(|e| format!("terracini: eigh of whitened cross-Gram failed: {e}"))?
.0;
let mut min_eig = f64::INFINITY;
let mut logdet = 0.0_f64;
let mut trace_inv = 0.0_f64;
let mut singular = false;
for &lam in w_b.iter() {
if lam < min_eig {
min_eig = lam;
}
let lam_c = lam.max(0.0);
if lam_c > 1.0e-300 {
logdet += lam_c.ln();
trace_inv += 1.0 / lam_c;
} else {
singular = true;
}
}
let margin = min_eig.max(0.0).sqrt();
let amplification = if margin > 0.0 {
1.0 / margin
} else {
f64::INFINITY
};
let (cross_gram_logdet, whitened_excess) = if singular {
(f64::NEG_INFINITY, f64::INFINITY)
} else {
(logdet, trace_inv - m as f64)
};
let jtj = j_s.t().dot(&j_s);
let w_j = jtj
.eigh(Side::Lower)
.map_err(|e| format!("terracini: eigh of J_SᵀJ_S failed: {e}"))?
.0;
let mut trace_inv_j = 0.0_f64;
for &lam in w_j.iter() {
let lam_c = lam + ridge;
trace_inv_j += if lam_c > 1.0e-300 {
1.0 / lam_c
} else {
f64::INFINITY
};
}
let attribution_risk = noise_var * trace_inv_j;
Ok(TerraciniCertificate {
pattern,
p,
m,
margin,
amplification,
cross_gram_logdet,
whitened_excess: whitened_excess.max(0.0),
attribution_risk,
per_atom_logdet,
})
}
#[derive(Debug, Clone)]
pub struct SinkPeel {
pub direction: Array1<f64>,
pub direction_l2_radius: f64,
pub removed_second_moment_radius: f64,
}
#[derive(Debug, Clone)]
pub struct PostPeelSamples {
pub samples: Array2<f64>,
pub peel_gram_radius: f64,
pub peel_margin_radius: f64,
}
pub fn post_peel_samples(
samples: ArrayView2<'_, f64>,
peel: &SinkPeel,
) -> Result<PostPeelSamples, String> {
let (n, p) = samples.dim();
if n == 0 || p == 0 {
return Err("post_peel_samples: samples must be non-empty".to_string());
}
if peel.direction.len() != p {
return Err(format!(
"post_peel_samples: peel direction length {} != sample dimension {p}",
peel.direction.len()
));
}
if !(peel.direction_l2_radius.is_finite() && peel.direction_l2_radius >= 0.0) {
return Err("post_peel_samples: direction_l2_radius must be finite and non-negative".to_string());
}
if !(peel.removed_second_moment_radius.is_finite()
&& peel.removed_second_moment_radius >= 0.0)
{
return Err(
"post_peel_samples: removed_second_moment_radius must be finite and non-negative"
.to_string(),
);
}
let norm = peel.direction.iter().map(|x| x * x).sum::<f64>().sqrt();
if !(norm.is_finite() && norm > 0.0) {
return Err("post_peel_samples: peel direction must have positive norm".to_string());
}
let mut u = peel.direction.clone();
for x in u.iter_mut() {
*x /= norm;
}
let mut peeled = Array2::<f64>::zeros((n, p));
let mut second_moment = 0.0_f64;
for r in 0..n {
let mut proj = 0.0_f64;
let mut row_norm2 = 0.0_f64;
for c in 0..p {
let x = samples[[r, c]];
proj += x * u[c];
row_norm2 += x * x;
}
second_moment += row_norm2;
for c in 0..p {
peeled[[r, c]] = samples[[r, c]] - proj * u[c];
}
}
second_moment /= n as f64;
let r = peel.direction_l2_radius;
let peel_gram_radius =
((2.0 * r + r * r) * second_moment + peel.removed_second_moment_radius).max(0.0);
Ok(PostPeelSamples {
samples: peeled,
peel_gram_radius,
peel_margin_radius: peel_gram_radius.sqrt(),
})
}
#[derive(Debug, Clone, Copy)]
pub struct MedianOfMeansConfig {
pub alpha: f64,
pub n_blocks: usize,
}
#[derive(Debug, Clone)]
pub struct MedianOfMeansGram {
pub gram: Array2<f64>,
pub alpha: f64,
pub n: usize,
pub n_eff: usize,
pub n_blocks: usize,
pub block_len: usize,
pub realized_kurtosis: f64,
pub entry_radius: f64,
pub spectral_radius: f64,
}
fn check_alpha(alpha: f64, caller: &str) -> Result<(), String> {
if alpha.is_finite() && alpha > 0.0 && alpha < 1.0 {
Ok(())
} else {
Err(format!("{caller}: alpha must lie in (0, 1)"))
}
}
fn median_sorted(mut values: Vec<f64>) -> Result<f64, String> {
if values.is_empty() {
return Err("median_sorted: empty input".to_string());
}
for &x in &values {
if !x.is_finite() {
return Err("median_sorted: input contains non-finite value".to_string());
}
}
values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
Ok(values[values.len() / 2])
}
fn realized_coordinate_kurtosis(samples: ArrayView2<'_, f64>) -> f64 {
let (n, p) = samples.dim();
let mut worst = 0.0_f64;
for c in 0..p {
let mut mean = 0.0_f64;
for r in 0..n {
mean += samples[[r, c]];
}
mean /= n as f64;
let mut m2 = 0.0_f64;
let mut m4 = 0.0_f64;
for r in 0..n {
let z = samples[[r, c]] - mean;
let z2 = z * z;
m2 += z2;
m4 += z2 * z2;
}
m2 /= n as f64;
m4 /= n as f64;
if m2 > 0.0 {
worst = worst.max(m4 / (m2 * m2));
}
}
worst.max(1.0)
}
pub fn median_of_means_gram(
samples: ArrayView2<'_, f64>,
cfg: MedianOfMeansConfig,
) -> Result<MedianOfMeansGram, String> {
check_alpha(cfg.alpha, "median_of_means_gram")?;
let (n, p) = samples.dim();
if n == 0 || p == 0 {
return Err("median_of_means_gram: samples must be non-empty".to_string());
}
if cfg.n_blocks == 0 || cfg.n_blocks > n {
return Err(format!(
"median_of_means_gram: n_blocks={} must lie in 1..={n}",
cfg.n_blocks
));
}
let block_len = n / cfg.n_blocks;
if block_len == 0 {
return Err("median_of_means_gram: block_len is zero".to_string());
}
let n_eff = block_len * cfg.n_blocks;
let mut gram = Array2::<f64>::zeros((p, p));
for a in 0..p {
for b in a..p {
let mut means = Vec::with_capacity(cfg.n_blocks);
for block in 0..cfg.n_blocks {
let start = block * block_len;
let end = start + block_len;
let mut sum = 0.0_f64;
for r in start..end {
sum += samples[[r, a]] * samples[[r, b]];
}
means.push(sum / block_len as f64);
}
let med = median_sorted(means)?;
gram[[a, b]] = med;
gram[[b, a]] = med;
}
}
let realized_kurtosis = realized_coordinate_kurtosis(samples);
let mut max_diag = 0.0_f64;
for c in 0..p {
max_diag = max_diag.max(gram[[c, c]].abs());
}
let union_terms = 2.0 * (p * p).max(1) as f64 / cfg.alpha;
let entry_radius =
(8.0 * realized_kurtosis * union_terms.ln() / n_eff as f64).sqrt() * max_diag.max(1.0e-300);
let spectral_radius = p as f64 * entry_radius;
Ok(MedianOfMeansGram {
gram,
alpha: cfg.alpha,
n,
n_eff,
n_blocks: cfg.n_blocks,
block_len,
realized_kurtosis,
entry_radius,
spectral_radius,
})
}
#[derive(Debug, Clone)]
pub struct FiniteSampleTerraciniCertificate {
pub margin_hat: f64,
pub delta: f64,
pub lower_margin_bound: f64,
pub identifiable: bool,
pub alpha: f64,
pub n_eff: usize,
pub realized_kurtosis: f64,
pub mom_spectral_radius: f64,
pub peel_gram_radius: f64,
pub theorem: String,
}
pub fn finite_sample_terracini_certificate(
mom: &MedianOfMeansGram,
peel_gram_radius: f64,
) -> Result<FiniteSampleTerraciniCertificate, String> {
if !(peel_gram_radius.is_finite() && peel_gram_radius >= 0.0) {
return Err(
"finite_sample_terracini_certificate: peel_gram_radius must be finite and non-negative"
.to_string(),
);
}
let evals = mom
.gram
.eigh(Side::Lower)
.map_err(|e| format!("finite_sample_terracini_certificate: eigh failed: {e}"))?
.0;
let mut min_eig = f64::INFINITY;
for &lam in evals.iter() {
min_eig = min_eig.min(lam);
}
let margin_hat = min_eig.max(0.0).sqrt();
let total_gram_radius = mom.spectral_radius + peel_gram_radius;
let delta = total_gram_radius.sqrt();
let lower_margin_bound = (margin_hat * margin_hat - total_gram_radius).max(0.0).sqrt();
let identifiable = margin_hat >= delta;
let theorem = format!(
"conditional on measured margin mu_hat={margin_hat:.6e} >= delta(n_eff={}, kappa_hat={:.6e}, alpha={:.6e})={delta:.6e}, the parse is identifiable with probability >= {:.6e}",
mom.n_eff,
mom.realized_kurtosis,
mom.alpha,
1.0 - mom.alpha
);
Ok(FiniteSampleTerraciniCertificate {
margin_hat,
delta,
lower_margin_bound,
identifiable,
alpha: mom.alpha,
n_eff: mom.n_eff,
realized_kurtosis: mom.realized_kurtosis,
mom_spectral_radius: mom.spectral_radius,
peel_gram_radius,
theorem,
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum TerraciniMode {
Off,
#[default]
Report,
Veto,
}
#[derive(Debug, Clone)]
pub struct TerraciniConfig {
pub mode: TerraciniMode,
pub noise_var: f64,
pub ridge: f64,
pub flag_margin: f64,
pub max_clique_atoms: usize,
pub pair_pass: bool,
pub reservoir_cap: usize,
pub reservoir_rows_per_atom: usize,
}
impl Default for TerraciniConfig {
fn default() -> Self {
Self {
mode: TerraciniMode::Report,
noise_var: 1.0,
ridge: 1.0e-12,
flag_margin: 1.0e-2,
max_clique_atoms: 16,
pair_pass: true,
reservoir_cap: 256,
reservoir_rows_per_atom: TERRACINI_CERTIFIER_ROWS_PER_ATOM,
}
}
}
#[derive(Debug, Clone)]
pub struct AtomMarginStat {
pub atom: usize,
pub n: usize,
pub min_margin: f64,
pub mean_margin: f64,
pub q05_margin: f64,
pub max_amplification: f64,
}
#[derive(Debug, Clone)]
pub struct PairMarginStat {
pub a: usize,
pub b: usize,
pub n: usize,
pub min_margin: f64,
pub mean_margin: f64,
}
#[derive(Debug, Clone)]
pub struct FlaggedClique {
pub row: usize,
pub pattern: Vec<usize>,
pub margin: f64,
pub cross_gram_logdet: f64,
pub attribution_risk: f64,
}
#[derive(Debug, Clone)]
struct AtomAcc {
n: usize,
min_margin: f64,
sum_margin: f64,
max_amplification: f64,
reservoir: Vec<f64>,
cap: usize,
}
impl AtomAcc {
fn new(cap: usize) -> Self {
Self {
n: 0,
min_margin: f64::INFINITY,
sum_margin: 0.0,
max_amplification: 0.0,
reservoir: Vec::new(),
cap: cap.max(1),
}
}
fn push(&mut self, margin: f64, amplification: f64) {
self.n += 1;
self.min_margin = self.min_margin.min(margin);
self.sum_margin += margin;
self.max_amplification = self.max_amplification.max(amplification);
if self.reservoir.len() < self.cap {
self.reservoir.push(margin);
} else if let Some((idx, &worst)) = self
.reservoir
.iter()
.enumerate()
.max_by(|x, y| x.1.partial_cmp(y.1).unwrap_or(std::cmp::Ordering::Equal))
{
if margin < worst {
self.reservoir[idx] = margin;
}
}
}
fn quantile(&self, q: f64) -> f64 {
if self.reservoir.is_empty() {
return f64::NAN;
}
let mut v = self.reservoir.clone();
v.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let idx = ((q * (v.len() as f64 - 1.0)).round() as usize).min(v.len() - 1);
v[idx]
}
}
#[derive(Debug, Clone)]
struct PairAcc {
n: usize,
min_margin: f64,
sum_margin: f64,
}
#[derive(Debug, Clone)]
pub struct TerraciniReport {
pub mode: TerraciniMode,
pub n_rows_scanned: usize,
pub atoms: Vec<AtomMarginStat>,
pub pairs: Vec<PairMarginStat>,
pub flagged_cliques: Vec<FlaggedClique>,
pub flag_margin: f64,
}
impl TerraciniReport {
pub fn vetoes_birth_into(&self, pattern: &[usize]) -> bool {
if self.mode != TerraciniMode::Veto {
return false;
}
let set: std::collections::BTreeSet<usize> = pattern.iter().copied().collect();
self.flagged_cliques
.iter()
.any(|fc| fc.pattern.iter().all(|a| set.contains(a)))
}
}
#[derive(Debug, Clone)]
pub struct TerraciniAggregator {
per_atom: BTreeMap<usize, AtomAcc>,
per_pair: BTreeMap<(usize, usize), PairAcc>,
flagged: Vec<FlaggedClique>,
flag_margin: f64,
reservoir_cap: usize,
n_rows: usize,
mode: TerraciniMode,
}
impl TerraciniAggregator {
pub fn new(flag_margin: f64, reservoir_cap: usize, mode: TerraciniMode) -> Self {
Self {
per_atom: BTreeMap::new(),
per_pair: BTreeMap::new(),
flagged: Vec::new(),
flag_margin,
reservoir_cap,
n_rows: 0,
mode,
}
}
fn atom_push(&mut self, atom: usize, margin: f64, amplification: f64) {
let cap = self.reservoir_cap;
self.per_atom
.entry(atom)
.or_insert_with(|| AtomAcc::new(cap))
.push(margin, amplification);
}
pub fn record_pair(&mut self, a: usize, b: usize, margin: f64, amplification: f64) {
let key = if a <= b { (a, b) } else { (b, a) };
let e = self.per_pair.entry(key).or_insert_with(|| PairAcc {
n: 0,
min_margin: f64::INFINITY,
sum_margin: 0.0,
});
e.n += 1;
e.min_margin = e.min_margin.min(margin);
e.sum_margin += margin;
self.atom_push(a, margin, amplification);
self.atom_push(b, margin, amplification);
}
pub fn record_clique(&mut self, row: usize, cert: &TerraciniCertificate) {
for &atom in &cert.pattern {
self.atom_push(atom, cert.margin, cert.amplification);
}
if cert.margin < self.flag_margin {
self.flagged.push(FlaggedClique {
row,
pattern: cert.pattern.clone(),
margin: cert.margin,
cross_gram_logdet: cert.cross_gram_logdet,
attribution_risk: cert.attribution_risk,
});
}
}
fn note_row(&mut self) {
self.n_rows += 1;
}
pub fn finish(mut self) -> TerraciniReport {
let mut atoms: Vec<AtomMarginStat> = self
.per_atom
.iter()
.map(|(&atom, acc)| AtomMarginStat {
atom,
n: acc.n,
min_margin: acc.min_margin,
mean_margin: if acc.n > 0 {
acc.sum_margin / acc.n as f64
} else {
f64::NAN
},
q05_margin: acc.quantile(0.05),
max_amplification: acc.max_amplification,
})
.collect();
atoms.sort_by(|x, y| {
x.min_margin
.partial_cmp(&y.min_margin)
.unwrap_or(std::cmp::Ordering::Equal)
});
let mut pairs: Vec<PairMarginStat> = self
.per_pair
.iter()
.map(|(&(a, b), acc)| PairMarginStat {
a,
b,
n: acc.n,
min_margin: acc.min_margin,
mean_margin: if acc.n > 0 {
acc.sum_margin / acc.n as f64
} else {
f64::NAN
},
})
.collect();
pairs.sort_by(|x, y| {
x.mean_margin
.partial_cmp(&y.mean_margin)
.unwrap_or(std::cmp::Ordering::Equal)
});
self.flagged.sort_by(|x, y| {
x.margin
.partial_cmp(&y.margin)
.unwrap_or(std::cmp::Ordering::Equal)
});
TerraciniReport {
mode: self.mode,
n_rows_scanned: self.n_rows,
atoms,
pairs,
flagged_cliques: self.flagged,
flag_margin: self.flag_margin,
}
}
}
fn atom_latent_dim(atom: &SaeManifoldAtom) -> usize {
atom.basis_jacobian.shape()[2]
}
pub fn parse_block_from_term(
term: &SaeManifoldTerm,
atom: usize,
row: usize,
amplitude: f64,
) -> ParseBlock {
let a = &term.atoms[atom];
let value = a.decoded_row(row);
let d = atom_latent_dim(a);
let p = value.len();
let mut tangent = Array2::<f64>::zeros((p, d));
for axis in 0..d {
let deriv = a.decoded_derivative_row(row, axis);
for i in 0..p {
tangent[[i, axis]] = amplitude * deriv[i];
}
}
ParseBlock {
atom,
value,
tangent,
}
}
pub fn stratified_rows_for_coverage(
rows_with_patterns: &[(usize, Vec<usize>)],
q: usize,
cap: usize,
) -> Vec<usize> {
let mut occ: BTreeMap<usize, Vec<usize>> = BTreeMap::new();
for (i, (_row, pattern)) in rows_with_patterns.iter().enumerate() {
for &atom in pattern {
occ.entry(atom).or_default().push(i);
}
}
let mut chosen = vec![false; rows_with_patterns.len()];
let mut count_for_atom: BTreeMap<usize, usize> = BTreeMap::new();
let mut n_selected = 0usize;
let mut progress = true;
while progress && n_selected < cap {
progress = false;
let atoms: Vec<usize> = occ.keys().copied().collect();
for atom in atoms {
if *count_for_atom.get(&atom).unwrap_or(&0) >= q {
continue;
}
let next = occ
.get(&atom)
.and_then(|rows| rows.iter().copied().find(|&idx| !chosen[idx]));
if let Some(idx) = next {
chosen[idx] = true;
n_selected += 1;
for &a in &rows_with_patterns[idx].1 {
*count_for_atom.entry(a).or_insert(0) += 1;
}
progress = true;
if n_selected >= cap {
break;
}
}
}
}
rows_with_patterns
.iter()
.enumerate()
.filter(|(i, _)| chosen[*i])
.map(|(_, (row, _))| *row)
.collect()
}
pub fn designed_reservoir_rows_for_coverage(
rows_with_patterns: &[(usize, Vec<usize>)],
measure: Option<&RowSamplingMeasure>,
q: usize,
cap: usize,
seed: u64,
) -> Result<Vec<usize>, String> {
if rows_with_patterns.is_empty() || q == 0 || cap == 0 {
return Ok(Vec::new());
}
let mut selected = BTreeSet::new();
if let Some(measure) = measure {
let n_rows = rows_with_patterns
.iter()
.map(|(row, _)| *row)
.max()
.map_or(0usize, |row| row + 1);
if measure.n_rows() != n_rows {
return Err(format!(
"designed_reservoir_rows_for_coverage: measure covers {} rows but patterns cover {n_rows}",
measure.n_rows()
));
}
let sample = measure.designed_subsample(cap, seed);
let eligible: BTreeSet<usize> = rows_with_patterns.iter().map(|(row, _)| *row).collect();
for row in sample.rows {
if eligible.contains(&row) {
selected.insert(row);
}
}
}
let topup = stratified_rows_for_coverage(rows_with_patterns, q, cap);
for row in topup {
if selected.len() >= cap {
break;
}
selected.insert(row);
}
Ok(selected.into_iter().collect())
}
fn sparse_rows_with_patterns(
indices: ArrayView2<'_, u32>,
codes: ArrayView2<'_, f32>,
k_atoms: usize,
) -> Result<Vec<(usize, Vec<usize>)>, String> {
if indices.dim() != codes.dim() {
return Err(format!(
"sparse_rows_with_patterns: indices shape {:?} != codes shape {:?}",
indices.dim(),
codes.dim()
));
}
let mut rows = Vec::new();
for row in 0..indices.nrows() {
let mut atoms = BTreeSet::new();
for slot in 0..indices.ncols() {
let code = codes[[row, slot]];
if code == 0.0 {
continue;
}
let atom = indices[[row, slot]] as usize;
if atom >= k_atoms {
return Err(format!(
"sparse_rows_with_patterns: atom index {atom} out of range 0..{k_atoms}"
));
}
atoms.insert(atom);
}
if atoms.len() >= 2 {
rows.push((row, atoms.into_iter().collect()));
}
}
Ok(rows)
}
fn sparse_code_amplitude(
indices: ArrayView2<'_, u32>,
codes: ArrayView2<'_, f32>,
row: usize,
atom: usize,
) -> f64 {
let mut amplitude = 0.0_f64;
for slot in 0..indices.ncols() {
if indices[[row, slot]] as usize == atom {
amplitude += codes[[row, slot]] as f64;
}
}
amplitude
}
pub fn terracini_scan(
term: &SaeManifoldTerm,
rows_with_patterns: &[(usize, Vec<usize>)],
amplitudes: ArrayView2<f64>,
cfg: &TerraciniConfig,
) -> TerraciniReport {
let mut agg = TerraciniAggregator::new(cfg.flag_margin, cfg.reservoir_cap, cfg.mode);
if cfg.mode == TerraciniMode::Off {
return agg.finish();
}
for (row, pattern) in rows_with_patterns {
if pattern.len() < 2 {
continue;
}
agg.note_row();
let blocks: Vec<ParseBlock> = pattern
.iter()
.map(|&k| parse_block_from_term(term, k, *row, amplitudes[[*row, k]]))
.collect();
if cfg.pair_pass {
for i in 0..blocks.len() {
for j in (i + 1)..blocks.len() {
if let Ok(cert) = parse_certificate(
&[blocks[i].clone(), blocks[j].clone()],
cfg.noise_var,
cfg.ridge,
) {
agg.record_pair(
blocks[i].atom,
blocks[j].atom,
cert.margin,
cert.amplification,
);
}
}
}
}
if pattern.len() <= cfg.max_clique_atoms {
if let Ok(cert) = parse_certificate(&blocks, cfg.noise_var, cfg.ridge) {
agg.record_clique(*row, &cert);
}
}
}
agg.finish()
}
pub fn terracini_scan_sparse_codes(
term: &SaeManifoldTerm,
indices: ArrayView2<'_, u32>,
codes: ArrayView2<'_, f32>,
measure: Option<&RowSamplingMeasure>,
seed: u64,
cfg: &TerraciniConfig,
) -> Result<TerraciniReport, String> {
if cfg.mode == TerraciniMode::Off {
return Ok(TerraciniAggregator::new(cfg.flag_margin, cfg.reservoir_cap, cfg.mode).finish());
}
if indices.nrows() != term.n_obs() {
return Err(format!(
"terracini_scan_sparse_codes: sparse codes have {} rows but term has {}",
indices.nrows(),
term.n_obs()
));
}
let rows_with_patterns = sparse_rows_with_patterns(indices, codes, term.k_atoms())?;
let cap = cfg
.reservoir_rows_per_atom
.saturating_mul(term.k_atoms())
.min(rows_with_patterns.len());
let rows = designed_reservoir_rows_for_coverage(
&rows_with_patterns,
measure,
cfg.reservoir_rows_per_atom,
cap,
seed,
)?;
let selected: BTreeSet<usize> = rows.into_iter().collect();
let mut agg = TerraciniAggregator::new(cfg.flag_margin, cfg.reservoir_cap, cfg.mode);
for (row, pattern) in rows_with_patterns {
if !selected.contains(&row) {
continue;
}
agg.note_row();
let blocks: Vec<ParseBlock> = pattern
.iter()
.map(|&k| {
parse_block_from_term(
term,
k,
row,
sparse_code_amplitude(indices, codes, row, k),
)
})
.collect();
if cfg.pair_pass {
for i in 0..blocks.len() {
for j in (i + 1)..blocks.len() {
if let Ok(cert) = parse_certificate(
&[blocks[i].clone(), blocks[j].clone()],
cfg.noise_var,
cfg.ridge,
) {
agg.record_pair(
blocks[i].atom,
blocks[j].atom,
cert.margin,
cert.amplification,
);
}
}
}
}
if pattern.len() <= cfg.max_clique_atoms {
if let Ok(cert) = parse_certificate(&blocks, cfg.noise_var, cfg.ridge) {
agg.record_clique(row, &cert);
}
}
}
Ok(agg.finish())
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::{Array1, Array2};
fn unit_block(atom: usize, v: &[f64]) -> ParseBlock {
ParseBlock {
atom,
value: Array1::from_vec(v.to_vec()),
tangent: Array2::<f64>::zeros((v.len(), 0)),
}
}
#[test]
fn closed_form_angle_margin_and_excess() {
for &theta in &[0.1_f64, 0.5, 1.0, 1.4, 1.5] {
let c = theta.cos();
let s = theta.sin();
let cert = parse_certificate(
&[
unit_block(0, &[1.0, 0.0, 0.0, 0.0]),
unit_block(1, &[c, s, 0.0, 0.0]),
],
1.0,
0.0,
)
.unwrap();
assert!((cert.margin - (1.0 - c).sqrt()).abs() < 1e-9, "θ={theta} margin {}", cert.margin);
let want_excess = 2.0 * c * c / (1.0 - c * c);
assert!((cert.whitened_excess - want_excess).abs() < 1e-9, "θ={theta} excess");
let want_logdet = (1.0 - c * c).ln();
assert!((cert.cross_gram_logdet - want_logdet).abs() < 1e-9, "θ={theta} logdet");
assert_eq!(cert.pattern, vec![0, 1]);
}
}
#[test]
fn logdet_channel_split_is_exact() {
let mut ta = Array2::<f64>::zeros((4, 1));
ta[[1, 0]] = 2.0;
let a = ParseBlock { atom: 0, value: Array1::from_vec(vec![1.0, 0.0, 0.0, 0.0]), tangent: ta };
let mut tb = Array2::<f64>::zeros((4, 1));
tb[[1, 0]] = 0.5;
tb[[3, 0]] = 1.0;
let b = ParseBlock { atom: 1, value: Array1::from_vec(vec![0.0, 0.0, 1.0, 0.0]), tangent: tb };
let cert = parse_certificate(&[a.clone(), b.clone()], 1.0, 0.0).unwrap();
let ka = a.stacked().unwrap();
let kb = b.stacked().unwrap();
let mut j = Array2::<f64>::zeros((4, 4));
for i in 0..4 {
for c in 0..2 {
j[[i, c]] = ka[[i, c]];
j[[i, 2 + c]] = kb[[i, c]];
}
}
let jtj = j.t().dot(&j);
let evals = jtj.eigh(Side::Lower).unwrap().0;
let ref_logdet: f64 = evals.iter().map(|&x| x.ln()).sum();
let split = cert.per_atom_logdet.iter().sum::<f64>() + cert.cross_gram_logdet;
assert!((split - ref_logdet).abs() < 1e-9, "split {split} vs {ref_logdet}");
}
#[test]
fn orthogonal_is_perfectly_conditioned() {
let cert = parse_certificate(
&[unit_block(0, &[1.0, 0.0, 0.0]), unit_block(1, &[0.0, 1.0, 0.0])],
2.0,
0.0,
)
.unwrap();
assert!((cert.margin - 1.0).abs() < 1e-12);
assert!(cert.whitened_excess.abs() < 1e-12);
assert!((cert.attribution_risk - 4.0).abs() < 1e-9);
}
#[test]
fn collision_diverges() {
let eps: f64 = 1e-6;
let cert = parse_certificate(
&[unit_block(0, &[1.0, 0.0]), unit_block(1, &[(eps).cos(), (eps).sin()])],
1.0,
0.0,
)
.unwrap();
assert!(cert.margin < 1e-3, "margin {}", cert.margin);
assert!(cert.amplification > 1e2);
assert!(cert.whitened_excess > 1e5);
}
#[test]
fn overcomplete_is_refused() {
let err = parse_certificate(
&[
unit_block(0, &[1.0, 0.0]),
unit_block(1, &[0.0, 1.0]),
unit_block(2, &[1.0, 1.0]),
],
1.0,
0.0,
)
.unwrap_err();
assert!(err.contains("overcomplete"), "got: {err}");
}
#[test]
fn aggregator_is_bounded_and_ranks_worst_first() {
let benign = parse_certificate(
&[unit_block(0, &[1.0, 0.0]), unit_block(1, &[0.0, 1.0])],
1.0,
0.0,
)
.unwrap();
let collided = parse_certificate(
&[unit_block(2, &[1.0, 0.0]), unit_block(3, &[(0.008_f64).cos(), (0.008_f64).sin()])],
1.0,
0.0,
)
.unwrap();
let mut agg = TerraciniAggregator::new(1e-2, 64, TerraciniMode::Veto);
for _ in 0..10_000 {
agg.record_clique(0, &benign);
agg.record_clique(1, &collided);
agg.record_pair(2, 3, collided.margin, collided.amplification);
}
let report = agg.finish();
assert_eq!(report.atoms.len(), 4);
assert_eq!(report.pairs.len(), 1);
assert!(report.atoms[0].atom == 2 || report.atoms[0].atom == 3);
assert!(report.atoms[0].min_margin < report.atoms[3].min_margin);
assert!(!report.flagged_cliques.is_empty());
assert_eq!(report.flagged_cliques[0].pattern, vec![2, 3]);
assert!(report.vetoes_birth_into(&[2, 3]));
assert!(!report.vetoes_birth_into(&[0, 1]));
}
#[test]
fn stratified_coverage_hits_every_atom() {
let rows = vec![
(0usize, vec![0usize, 1]),
(1, vec![0, 2]),
(2, vec![1, 2]),
(3, vec![3, 4]),
(4, vec![0, 1]),
];
let sel = stratified_rows_for_coverage(&rows, 1, 100);
let mut seen = std::collections::BTreeSet::new();
for r in &sel {
for &a in &rows.iter().find(|(rr, _)| rr == r).unwrap().1 {
seen.insert(a);
}
}
for a in 0..=4 {
assert!(seen.contains(&a), "atom {a} not covered");
}
}
#[test]
fn heavy_tailed_mom_gram_ci_survives_plugin_failure() {
let blocks = 9usize;
let block_len = 24usize;
let mut samples = Array2::<f64>::zeros((blocks * block_len, 2));
for block in 0..blocks {
for i in 0..block_len {
let row = block * block_len + i;
if block == blocks - 1 {
samples[[row, 0]] = 80.0;
samples[[row, 1]] = 80.0;
} else if i % 2 == 0 {
samples[[row, 0]] = 2.0_f64.sqrt();
} else {
samples[[row, 1]] = 2.0_f64.sqrt();
}
}
}
let mom = median_of_means_gram(
samples.view(),
MedianOfMeansConfig {
alpha: 0.05,
n_blocks: blocks,
},
)
.unwrap();
assert!((mom.gram[[0, 0]] - 1.0).abs() < 1.0e-12);
assert!((mom.gram[[1, 1]] - 1.0).abs() < 1.0e-12);
assert!(mom.gram[[0, 1]].abs() < 1.0e-12);
assert!(mom.realized_kurtosis.is_finite());
assert!(mom.realized_kurtosis > 5.0);
let mut plugin = Array2::<f64>::zeros((2, 2));
for r in 0..samples.nrows() {
for a in 0..2 {
for b in 0..2 {
plugin[[a, b]] += samples[[r, a]] * samples[[r, b]];
}
}
}
for a in 0..2 {
for b in 0..2 {
plugin[[a, b]] /= samples.nrows() as f64;
}
}
let plugin_subgaussian_radius =
(8.0_f64 * (2.0_f64 * 4.0_f64 / 0.05_f64).ln() / samples.nrows() as f64).sqrt();
let plugin_offdiag_error = plugin[[0, 1]].abs();
assert!(
plugin_offdiag_error > plugin_subgaussian_radius,
"plug-in CI should miss the clean Gram under one heavy-tailed block: error {plugin_offdiag_error}, radius {plugin_subgaussian_radius}"
);
assert!(
mom.gram[[0, 1]].abs() <= mom.entry_radius,
"MoM CI should cover the clean off-diagonal"
);
}
#[test]
fn finite_sample_certificate_accepts_and_refuses_by_delta() {
let mut good_samples = Array2::<f64>::zeros((400, 2));
for r in 0..good_samples.nrows() {
if r % 2 == 0 {
good_samples[[r, 0]] = 2.0_f64.sqrt();
} else {
good_samples[[r, 1]] = 2.0_f64.sqrt();
}
}
let good_mom = median_of_means_gram(
good_samples.view(),
MedianOfMeansConfig {
alpha: 0.5,
n_blocks: 5,
},
)
.unwrap();
let good = finite_sample_terracini_certificate(&good_mom, 0.0).unwrap();
assert!(
good.identifiable,
"expected mu_hat {} >= delta {}",
good.margin_hat,
good.delta
);
assert!(good.theorem.contains("conditional on measured margin"));
assert!(good.lower_margin_bound > 0.0);
let mut bad_samples = Array2::<f64>::zeros((400, 2));
for r in 0..bad_samples.nrows() {
bad_samples[[r, 0]] = 1.0;
bad_samples[[r, 1]] = if r % 2 == 0 { 1.0 } else { 1.0 + 1.0e-5 };
}
let bad_mom = median_of_means_gram(
bad_samples.view(),
MedianOfMeansConfig {
alpha: 0.5,
n_blocks: 5,
},
)
.unwrap();
let bad = finite_sample_terracini_certificate(&bad_mom, 0.0).unwrap();
assert!(
!bad.identifiable,
"expected mu_hat {} < delta {}",
bad.margin_hat,
bad.delta
);
assert_eq!(bad.lower_margin_bound, 0.0);
}
#[test]
fn post_peel_propagates_uncertainty_and_reports_kurtosis() {
let mut samples = Array2::<f64>::zeros((12, 3));
for r in 0..samples.nrows() {
samples[[r, 0]] = 100.0 + r as f64;
samples[[r, 1]] = if r % 2 == 0 { 1.0 } else { -1.0 };
samples[[r, 2]] = if r % 3 == 0 { 2.0 } else { -0.5 };
}
let peel = SinkPeel {
direction: Array1::from_vec(vec![1.0, 0.0, 0.0]),
direction_l2_radius: 0.01,
removed_second_moment_radius: 0.25,
};
let peeled = post_peel_samples(samples.view(), &peel).unwrap();
for r in 0..peeled.samples.nrows() {
assert!(peeled.samples[[r, 0]].abs() < 1.0e-12);
}
assert!(peeled.peel_gram_radius > 0.25);
assert!((peeled.peel_margin_radius.powi(2) - peeled.peel_gram_radius).abs() < 1.0e-12);
let mom = median_of_means_gram(
peeled.samples.view(),
MedianOfMeansConfig {
alpha: 0.25,
n_blocks: 3,
},
)
.unwrap();
let cert = finite_sample_terracini_certificate(&mom, peeled.peel_gram_radius).unwrap();
assert!(cert.realized_kurtosis.is_finite());
assert!(cert.realized_kurtosis >= 1.0);
assert_eq!(cert.peel_gram_radius, peeled.peel_gram_radius);
assert!(cert.theorem.contains("kappa_hat"));
}
}