use ndarray::ArrayView2;
const AUDIT_Z: f64 = 3.0;
const AUDIT_Q_EDGE: f64 = 0.43;
fn audit_row_floor() -> usize {
let z = AUDIT_Z;
let q = AUDIT_Q_EDGE;
(4.0 * z * z * q * (1.0 - q) / (1.0 - 2.0 * q).powi(2)).ceil() as usize
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct MergeCandidate {
pub atom_a: usize,
pub atom_b: usize,
pub n_coactive: usize,
pub kappa: f64,
pub z_score: f64,
}
impl MergeCandidate {
pub fn evidence(&self) -> f64 {
self.z_score.abs()
}
}
#[inline]
fn is_active(y: f64, mask: Option<bool>) -> bool {
match mask {
Some(m) => m,
None => y != 0.0 && y.is_finite(),
}
}
pub fn audit_torus_merges(
latent: ArrayView2<f64>,
active: Option<ArrayView2<bool>>,
) -> Vec<MergeCandidate> {
let n_atoms = latent.ncols();
let n_rows = latent.nrows();
if let Some(m) = active {
if m.nrows() != n_rows || m.ncols() != n_atoms {
return audit_torus_merges(latent, None);
}
}
let floor = audit_row_floor();
let mut out = Vec::new();
for a in 0..n_atoms {
for b in (a + 1)..n_atoms {
let mut n = 0usize;
let (mut m1, mut m2, mut m3, mut m4) = (0.0_f64, 0.0_f64, 0.0_f64, 0.0_f64);
for i in 0..n_rows {
let ya = latent[[i, a]];
let yb = latent[[i, b]];
let ma = active.map(|m| m[[i, a]]);
let mb = active.map(|m| m[[i, b]]);
if !(is_active(ya, ma) && is_active(yb, mb)) {
continue;
}
let s = ya * ya + yb * yb;
let s2 = s * s;
m1 += s;
m2 += s2;
m3 += s2 * s;
m4 += s2 * s2;
n += 1;
}
if n < floor {
continue;
}
let nf = n as f64;
let mu1 = m1 / nf; let mu2 = m2 / nf; let mu3 = m3 / nf; let mu4 = m4 / nf; if !(mu1 > 0.0) {
continue;
}
let kappa = mu2 / (mu1 * mu1);
let var_s = (mu2 - mu1 * mu1).max(0.0);
let var_s2 = (mu4 - mu2 * mu2).max(0.0);
let cov = mu3 - mu2 * mu1;
let mu1_2 = mu1 * mu1;
let mu1_4 = mu1_2 * mu1_2;
let mu1_5 = mu1_4 * mu1;
let mu1_6 = mu1_4 * mu1_2;
let var_kappa = (var_s2 / mu1_4 + 4.0 * mu2 * mu2 * var_s / mu1_6
- 4.0 * mu2 * cov / mu1_5)
/ nf;
if !(var_kappa > 0.0) {
if (kappa - 2.0).abs() > 0.0 {
out.push(MergeCandidate {
atom_a: a,
atom_b: b,
n_coactive: n,
kappa,
z_score: (kappa - 2.0).signum() * f64::INFINITY,
});
}
continue;
}
let se = var_kappa.sqrt();
let z = (kappa - 2.0) / se;
if z.abs() >= AUDIT_Z {
out.push(MergeCandidate {
atom_a: a,
atom_b: b,
n_coactive: n,
kappa,
z_score: z,
});
}
}
}
out.sort_by(|x, y| {
y.evidence()
.partial_cmp(&x.evidence())
.unwrap_or(std::cmp::Ordering::Equal)
.then(x.atom_a.cmp(&y.atom_a))
.then(x.atom_b.cmp(&y.atom_b))
});
out
}
#[cfg(test)]
mod tests {
use super::*;
use gam_linalg::utils::splitmix64_hash;
use ndarray::Array2;
fn uniform01(counter: &mut u64) -> f64 {
*counter = counter.wrapping_add(1);
let h = splitmix64_hash(*counter ^ 0xA5A5_1234);
(h >> 11) as f64 / (1u64 << 53) as f64
}
fn gauss(counter: &mut u64) -> f64 {
let u1 = uniform01(counter).max(1e-12);
let u2 = uniform01(counter);
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
#[test]
fn one_circle_split_across_two_atoms_is_flagged() {
let n = 2000usize;
let mut latent = Array2::<f64>::zeros((n, 3));
let mut ctr = 1u64;
for i in 0..n {
let theta = std::f64::consts::TAU * uniform01(&mut ctr);
latent[[i, 0]] = theta.cos();
latent[[i, 1]] = theta.sin();
latent[[i, 2]] = gauss(&mut ctr);
}
for i in 0..n {
for k in 0..3 {
if latent[[i, k]] == 0.0 {
latent[[i, k]] = 1e-6;
}
}
}
let flags = audit_torus_merges(latent.view(), None);
let ring = flags
.iter()
.find(|c| c.atom_a == 0 && c.atom_b == 1)
.expect("ring pair (0,1) must be flagged");
assert!(
ring.kappa < 1.5 && ring.z_score < -AUDIT_Z,
"ring pair must read strongly sub-Gaussian: {ring:?}"
);
}
#[test]
fn two_independent_gaussian_atoms_are_not_flagged() {
let n = 20000usize;
let mut latent = Array2::<f64>::zeros((n, 2));
let mut ctr = 99u64;
for i in 0..n {
latent[[i, 0]] = gauss(&mut ctr);
latent[[i, 1]] = gauss(&mut ctr);
}
let flags = audit_torus_merges(latent.view(), None);
assert!(
flags.is_empty(),
"independent Gaussian atoms must not be flagged: {flags:?}"
);
}
#[test]
fn gated_circle_reads_super_gaussian() {
let n = 6000usize;
let mut latent = Array2::<f64>::zeros((n, 2));
let mut ctr = 7u64;
let q = 0.25_f64;
for i in 0..n {
if uniform01(&mut ctr) < q {
let theta = std::f64::consts::TAU * uniform01(&mut ctr);
latent[[i, 0]] = theta.cos();
latent[[i, 1]] = theta.sin();
} else {
latent[[i, 0]] = 0.0;
latent[[i, 1]] = 0.0;
}
}
let mask = Array2::<bool>::from_shape_fn((n, 2), |(i, _)| {
latent[[i, 0]] != 0.0 || latent[[i, 1]] != 0.0
});
let flags = audit_torus_merges(latent.view(), Some(mask.view()));
assert!(
flags.iter().any(|c| c.atom_a == 0 && c.atom_b == 1),
"gated ring pair must be flagged: {flags:?}"
);
}
#[test]
fn too_few_coactive_rows_are_not_flagged() {
let n = audit_row_floor().saturating_sub(1).max(1);
let mut latent = Array2::<f64>::zeros((n, 2));
let mut ctr = 3u64;
for i in 0..n {
let theta = std::f64::consts::TAU * uniform01(&mut ctr);
latent[[i, 0]] = theta.cos().abs() + 1e-3;
latent[[i, 1]] = theta.sin().abs() + 1e-3;
}
let flags = audit_torus_merges(latent.view(), None);
assert!(
flags.is_empty(),
"below the row floor nothing may be flagged: {flags:?}"
);
}
#[test]
fn row_floor_matches_isa_derivation() {
let floor = audit_row_floor();
assert!(
(440..=470).contains(&floor),
"derived audit floor {floor} out of the ISA design band"
);
}
}