use crate::coactivation_conditionality::{
CoactivationConditionality, VaryingCoefficientConfig, estimate_on_rows,
residual_gate_activities,
};
use ndarray::{Array1, ArrayView2};
use super::isa_seed::IsaPlaneCandidate;
struct PlaneEnergies {
r2: Vec<f64>,
active: Vec<bool>,
}
fn plane_energies(
data: ArrayView2<'_, f64>,
mean: &Array1<f64>,
cand: &IsaPlaneCandidate,
) -> PlaneEnergies {
let (n, p) = data.dim();
let mut r2 = vec![0.0_f64; n];
let mut active = vec![false; n];
for i in 0..n {
let (mut p1, mut p2) = (0.0_f64, 0.0_f64);
for j in 0..p {
let ri = data[[i, j]] - mean[j];
p1 += ri * cand.basis[[j, 0]];
p2 += ri * cand.basis[[j, 1]];
}
r2[i] = p1 * p1 + p2 * p2;
active[i] = cand.gate_logits[i].is_finite();
}
PlaneEnergies { r2, active }
}
#[derive(Clone, Debug)]
pub struct PairVerdict {
pub atom_a: usize,
pub atom_b: usize,
pub n_rows: usize,
pub n_co_active: usize,
pub rho: f64,
pub kappa_a: f64,
pub kappa_b: f64,
pub rho_se: f64,
pub z: f64,
pub conditionality: Option<CoactivationConditionality>,
pub conditional_stable: bool,
pub merge_proposed: bool,
}
const PAIR_Z: f64 = 3.0;
const PAIR_ROW_FLOOR: usize = 500;
pub fn screen_pair(
data: ArrayView2<'_, f64>,
mean: &Array1<f64>,
atom_a: usize,
atom_b: usize,
cand_a: &IsaPlaneCandidate,
cand_b: &IsaPlaneCandidate,
) -> PairVerdict {
let n = data.nrows();
let contexts = vec![0usize; n];
screen_pair_with_contexts(data, mean, atom_a, atom_b, cand_a, cand_b, &contexts, None)
}
pub fn screen_pair_with_contexts(
data: ArrayView2<'_, f64>,
mean: &Array1<f64>,
atom_a: usize,
atom_b: usize,
cand_a: &IsaPlaneCandidate,
cand_b: &IsaPlaneCandidate,
context_labels: &[usize],
shared_chart: Option<ArrayView2<'_, f64>>,
) -> PairVerdict {
let ea = plane_energies(data, mean, cand_a);
let eb = plane_energies(data, mean, cand_b);
let n = ea.r2.len();
let n_co_active = (0..n).filter(|&i| ea.active[i] && eb.active[i]).count();
let continuous_context = pair_continuous_context(&ea, &eb);
let conditionality =
residual_conditionality(&ea, &eb, &continuous_context, context_labels, shared_chart);
let conditional_stable = conditionality
.as_ref()
.map(partition_free_conditional_stable)
.unwrap_or(false);
let unresolved = PairVerdict {
atom_a,
atom_b,
n_rows: n,
n_co_active,
rho: f64::NAN,
kappa_a: f64::NAN,
kappa_b: f64::NAN,
rho_se: f64::NAN,
z: 0.0,
conditionality: conditionality.clone(),
conditional_stable,
merge_proposed: false,
};
if n < PAIR_ROW_FLOOR {
return unresolved;
}
let inv = 1.0 / n as f64;
let (mut ma, mut mb, mut cross) = (0.0_f64, 0.0_f64, 0.0_f64);
let (mut qa, mut qb) = (0.0_f64, 0.0_f64); for i in 0..n {
let (a, b) = (ea.r2[i], eb.r2[i]);
ma += a;
mb += b;
cross += a * b;
qa += a * a;
qb += b * b;
}
ma *= inv;
mb *= inv;
cross *= inv;
qa *= inv;
qb *= inv;
if !(ma > 0.0 && mb > 0.0) {
return unresolved;
}
let rho = cross / (ma * mb);
let kappa_a = qa / (ma * ma);
let kappa_b = qb / (mb * mb);
let var = ((kappa_a - 1.0) * (kappa_b - 1.0)).max(0.0) * inv;
let rho_se = var.sqrt();
let z = if rho_se > 0.0 {
(rho - 1.0) / rho_se
} else if rho > 1.0 {
f64::INFINITY
} else {
0.0
};
let merge_proposed = z > PAIR_Z && conditional_stable;
PairVerdict {
atom_a,
atom_b,
n_rows: n,
n_co_active,
rho,
kappa_a,
kappa_b,
rho_se,
z,
conditionality,
conditional_stable,
merge_proposed,
}
}
pub fn screen_all_pairs(
data: ArrayView2<'_, f64>,
mean: &Array1<f64>,
candidates: &[IsaPlaneCandidate],
) -> Vec<PairVerdict> {
let mut out = Vec::new();
for a in 0..candidates.len() {
for b in (a + 1)..candidates.len() {
let v = screen_pair(data, mean, a, b, &candidates[a], &candidates[b]);
if v.merge_proposed {
out.push(v);
}
}
}
out
}
fn residual_conditionality(
ea: &PlaneEnergies,
eb: &PlaneEnergies,
continuous_context: &[f64],
context_labels: &[usize],
shared_chart: Option<ArrayView2<'_, f64>>,
) -> Option<CoactivationConditionality> {
let n = ea.active.len();
if context_labels.len() != n || continuous_context.len() != n {
return None;
}
let rows: Vec<usize> = (0..n).collect();
let weights = vec![1.0_f64; n];
let gate_a: Vec<f64> = ea
.active
.iter()
.map(|&active| if active { 1.0 } else { 0.0 })
.collect();
let gate_b: Vec<f64> = eb
.active
.iter()
.map(|&active| if active { 1.0 } else { 0.0 })
.collect();
let activities =
residual_gate_activities(&gate_a, &gate_b, shared_chart, &weights, 0.0).ok()?;
estimate_on_rows(
&activities.residual_i,
&activities.residual_j,
continuous_context,
Some(context_labels),
&rows,
&weights,
VaryingCoefficientConfig::default(),
)
.ok()
}
fn pair_continuous_context(ea: &PlaneEnergies, eb: &PlaneEnergies) -> Vec<f64> {
ea.r2
.iter()
.zip(eb.r2.iter())
.map(|(&a, &b)| (a + b).ln_1p())
.collect()
}
fn partition_free_conditional_stable(c: &CoactivationConditionality) -> bool {
if c.certificate.statistic
!= crate::coactivation_conditionality::CouplingStatistic::WeightedPearson
{
return false;
}
let drift = c.native.beta_wiggliness.max(0.0) + c.native.beta_variation.max(0.0);
if c.certificate.robustness_radius_epsilon.is_infinite() {
return true;
}
c.certificate.robustness_radius_epsilon.is_finite()
&& c.certificate.robustness_radius_epsilon > drift
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::Array2;
fn lcg(s: &mut u64) -> f64 {
*s = s
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((*s >> 11) as f64) / ((1u64 << 53) as f64)
}
fn lcg_normal(s: &mut u64) -> f64 {
let u1 = lcg(s).max(1e-12);
let u2 = lcg(s);
(-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
}
fn axis_candidate(p: usize, d0: usize, d1: usize, active: &[bool]) -> IsaPlaneCandidate {
let n = active.len();
let mut basis = Array2::<f64>::zeros((p, 2));
basis[[d0, 0]] = 1.0;
basis[[d1, 1]] = 1.0;
let gate_logits: Vec<f64> = active
.iter()
.map(|&a| if a { 0.0 } else { f64::NEG_INFINITY })
.collect();
IsaPlaneCandidate {
basis,
amplitudes: [1.0, 1.0],
phases_turns: Array2::<f64>::zeros((n, 1)),
gate_logits,
kappa: 1.0,
q_hat: active.iter().filter(|&&a| a).count() as f64 / n as f64,
}
}
#[test]
fn two_independent_circles_not_flagged() {
let mut s = 0x9A1_u64;
let n = 6000usize;
let p = 8usize;
let qa = 0.5;
let qb = 0.5;
let mut data = Array2::<f64>::zeros((n, p));
let mut act_a = vec![false; n];
let mut act_b = vec![false; n];
for i in 0..n {
if lcg(&mut s) < qa {
act_a[i] = true;
let th = std::f64::consts::TAU * lcg(&mut s);
data[[i, 0]] += th.cos();
data[[i, 1]] += th.sin();
}
if lcg(&mut s) < qb {
act_b[i] = true;
let th = std::f64::consts::TAU * lcg(&mut s);
data[[i, 2]] += th.cos();
data[[i, 3]] += th.sin();
}
for j in 0..p {
data[[i, j]] += 0.02 * lcg_normal(&mut s);
}
}
let mean = Array1::<f64>::zeros(p);
let ca = axis_candidate(p, 0, 1, &act_a);
let cb = axis_candidate(p, 2, 3, &act_b);
let v = screen_pair(data.view(), &mean, 0, 1, &ca, &cb);
assert!(
!v.merge_proposed,
"two INDEPENDENT circles must NOT be flagged: ρ={:.4} z={:.3}",
v.rho, v.z
);
assert!(
(v.rho - 1.0).abs() < 0.15,
"independent ρ must sit near the null 1.0; got {:.4}",
v.rho
);
}
#[test]
fn gated_torus_split_is_flagged() {
let mut s = 0x7013_u64;
let n = 6000usize;
let p = 8usize;
let q = 0.4;
let mut data = Array2::<f64>::zeros((n, p));
let mut act = vec![false; n];
for i in 0..n {
if lcg(&mut s) < q {
act[i] = true;
let ta = std::f64::consts::TAU * lcg(&mut s);
let tb = std::f64::consts::TAU * lcg(&mut s); data[[i, 0]] += ta.cos();
data[[i, 1]] += ta.sin();
data[[i, 2]] += tb.cos();
data[[i, 3]] += tb.sin();
}
for j in 0..p {
data[[i, j]] += 0.02 * lcg_normal(&mut s);
}
}
let mean = Array1::<f64>::zeros(p);
let ca = axis_candidate(p, 0, 1, &act);
let cb = axis_candidate(p, 2, 3, &act);
let v = screen_pair(data.view(), &mean, 0, 1, &ca, &cb);
assert!(
v.merge_proposed,
"a gated torus split into two atoms MUST be flagged: ρ={:.4} z={:.3}",
v.rho, v.z
);
assert!(
v.rho > 1.5,
"co-gated torus ρ must be ≫ 1 (anchor 1/q≈2.5); got {:.4}",
v.rho
);
}
#[test]
fn screen_all_pairs_selects_only_bound_pair() {
let mut s = 0xC0FFEE_u64;
let n = 6000usize;
let p = 12usize;
let q_iso = 0.5;
let q_tor = 0.4;
let mut data = Array2::<f64>::zeros((n, p));
let mut act_a = vec![false; n];
let mut act_bc = vec![false; n];
for i in 0..n {
if lcg(&mut s) < q_iso {
act_a[i] = true;
let th = std::f64::consts::TAU * lcg(&mut s);
data[[i, 0]] += th.cos();
data[[i, 1]] += th.sin();
}
if lcg(&mut s) < q_tor {
act_bc[i] = true;
let tb = std::f64::consts::TAU * lcg(&mut s);
let tc = std::f64::consts::TAU * lcg(&mut s);
data[[i, 2]] += tb.cos();
data[[i, 3]] += tb.sin();
data[[i, 4]] += tc.cos();
data[[i, 5]] += tc.sin();
}
for j in 0..p {
data[[i, j]] += 0.02 * lcg_normal(&mut s);
}
}
let mean = Array1::<f64>::zeros(p);
let cands = vec![
axis_candidate(p, 0, 1, &act_a), axis_candidate(p, 2, 3, &act_bc), axis_candidate(p, 4, 5, &act_bc), ];
let flags = screen_all_pairs(data.view(), &mean, &cands);
assert_eq!(flags.len(), 1, "exactly one bound pair expected");
assert!(
flags[0].atom_a == 1 && flags[0].atom_b == 2,
"the flagged pair must be the co-gated torus factors (1,2); got ({},{})",
flags[0].atom_a,
flags[0].atom_b
);
}
}