use crate::manifold::{
AssignmentMode, PeriodicHarmonicEvaluator, SaeAssignment, SaeAtomBasisKind, SaeBasisEvaluator,
SaeManifoldAtom, SaeManifoldRho, SaeManifoldTerm,
};
use gam_terms::latent::LatentManifold;
use ndarray::{Array1, Array2};
use std::sync::{Arc, Mutex};
static K3_SERIAL: Mutex<()> = Mutex::new(());
fn k3_guard() -> std::sync::MutexGuard<'static, ()> {
K3_SERIAL.lock().unwrap_or_else(|e| e.into_inner())
}
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 fitted_circle_term(n: usize, p: usize) -> (SaeManifoldTerm, SaeManifoldRho) {
let mut s = 0x2101_B1C_0000_0005u64;
let theta: Vec<f64> = (0..n).map(|_| std::f64::consts::TAU * lcg(&mut s)).collect();
let mut x = Array2::<f64>::zeros((n, p));
for i in 0..n {
x[[i, 0]] += theta[i].cos();
x[[i, 1]] += theta[i].sin();
for j in 0..p {
x[[i, j]] += 0.05 * lcg_normal(&mut s);
}
}
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(3).unwrap());
let coords =
Array2::<f64>::from_shape_fn((n, 1), |(r, _)| theta[r] / std::f64::consts::TAU);
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let mut decoder = Array2::<f64>::zeros((3, p));
decoder[[1, 0]] = 1.0;
decoder[[2, 1]] = 1.0;
let atom = SaeManifoldAtom::new(
"circle".to_string(),
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_second_jet(evaluator.clone());
let logits = Array2::<f64>::from_elem((n, 1), 3.0);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
vec![coords],
vec![LatentManifold::Circle { period: 1.0 }],
AssignmentMode::ibp_map(0.7, 1.0, false),
)
.unwrap();
let mut term = SaeManifoldTerm::new(vec![atom], assignment).unwrap();
term.set_guards_enabled(false);
let mut rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1)]);
term.run_joint_fit_arrow_schur(x.view(), &mut rho, None, 60, 1.0, 1e-6, 1e-6)
.expect("K=1 circle fit");
(term, rho)
}
#[test]
fn rank_charge_deff_accepts_circle_and_neutralises_vanishing() {
let (mut term, rho) = fitted_circle_term(80, 16);
let (_v, loss, cache) = term
.reml_criterion_with_cache(unit_target(&term).view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap_or_else(|_| panic!("reml pass"));
let disp = term.reconstruction_dispersion(&loss, &cache, &rho).unwrap();
drop((loss, cache));
let d_real = term.per_atom_realised_rank_dof(&rho, disp).unwrap();
eprintln!(
"[rank-charge] dispersion R={disp:.5} circle d_eff={:.3} → charge ½·d_eff·ln80={:.3}",
d_real[0],
0.5 * d_real[0] * (80f64).ln()
);
assert!(
d_real[0] > 2.5 && d_real[0] < 8.0,
"rank-2 circle d_eff should be ~rank-2×basis-EDF (~4-6); got {:.3}",
d_real[0]
);
assert!(
0.5 * d_real[0] * (80f64).ln() < 15.0,
"rank-charge must be modest (accept), got charge {:.3}",
0.5 * d_real[0] * (80f64).ln()
);
let saved = term.atoms[0].decoder_coefficients.clone();
term.atoms[0]
.decoder_coefficients
.assign(&(&saved * 1e-4));
let d_vanish = term.per_atom_realised_rank_dof(&rho, disp).unwrap();
eprintln!("[rank-charge] vanishing (decoder×1e-4) d_eff={:.5} → charge≈0 (neutral)", d_vanish[0]);
assert!(
d_vanish[0] < 0.2,
"vanishing decoder must give d_eff→0 (neutral); got {:.4}",
d_vanish[0]
);
term.atoms[0].decoder_coefficients.assign(&saved);
}
#[test]
fn rank_charge_flag_off_is_inert() {
let (mut term, rho) = fitted_circle_term(80, 16);
let tgt = unit_target(&term);
let (v_off, _, _) = term
.reml_criterion_with_cache(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
term.set_rank_charge_evidence(false);
let (v_off2, _, _) = term
.reml_criterion_with_cache(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
assert_eq!(
v_off, v_off2,
"rank_charge_evidence=false must be bit-identical to the historical criterion"
);
term.set_rank_charge_evidence(true);
let (v_on, _, _) = term
.reml_criterion_with_cache(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
eprintln!("[rank-charge] reml OFF={v_off:.4} ON={v_on:.4} (ON lowers the circle's complexity)");
assert!(
(v_on - v_off).abs() > 1e-6 && v_on.is_finite(),
"rank_charge_evidence=true must change the (finite) criterion; off={v_off:.4} on={v_on:.4}"
);
}
#[test]
fn rank_charge_healthy_k3_control_well_conditioned() {
let serial = k3_guard();
let n = 96usize;
let p = 18usize;
let ncirc = 3usize;
let mut s = 0x2101_C3C_0000_0009u64;
let theta: Vec<Vec<f64>> = (0..n)
.map(|_| (0..ncirc).map(|_| std::f64::consts::TAU * lcg(&mut s)).collect())
.collect();
let mut x = Array2::<f64>::zeros((n, p));
for i in 0..n {
for c in 0..ncirc {
x[[i, 2 * c]] += theta[i][c].cos();
x[[i, 2 * c + 1]] += theta[i][c].sin();
}
for j in 0..p {
x[[i, j]] += 0.05 * lcg_normal(&mut s);
}
}
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(3).unwrap());
let mut atoms = Vec::new();
let mut coord_blocks = Vec::new();
let mut manifolds = Vec::new();
for c in 0..ncirc {
let coords =
Array2::<f64>::from_shape_fn((n, 1), |(r, _)| theta[r][c] / std::f64::consts::TAU);
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let mut decoder = Array2::<f64>::zeros((3, p));
decoder[[1, 2 * c]] = 1.0;
decoder[[2, 2 * c + 1]] = 1.0;
let atom = SaeManifoldAtom::new(
format!("circle{c}"),
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_second_jet(evaluator.clone());
atoms.push(atom);
coord_blocks.push(coords);
manifolds.push(LatentManifold::Circle { period: 1.0 });
}
let logits = Array2::<f64>::from_elem((n, ncirc), 3.0);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coord_blocks,
manifolds,
AssignmentMode::ibp_map(0.7, 1.0, false),
)
.unwrap();
let mut term = SaeManifoldTerm::new(atoms, assignment).unwrap();
term.set_guards_enabled(false);
let mut rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1); ncirc]);
term.run_joint_fit_arrow_schur(x.view(), &mut rho, None, 60, 1.0, 1e-6, 1e-6)
.expect("K=3 clean fit");
let (v_off, loss, cache) = term
.reml_criterion_with_cache(x.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
let disp = term.reconstruction_dispersion(&loss, &cache, &rho).unwrap();
drop((loss, cache));
let d_eff = term.per_atom_realised_rank_dof(&rho, disp).unwrap();
eprintln!("[rank-charge K=3] d_eff per atom = {:?} disp={disp:.5}",
d_eff.iter().map(|v| (v*100.0).round()/100.0).collect::<Vec<_>>());
for (k, &de) in d_eff.iter().enumerate() {
assert!(
de > 2.0 && de < 8.0,
"K=3 atom {k}: every clean rank-2 circle must price ~4-6; got d_eff={de:.3}"
);
}
term.set_rank_charge_evidence(true);
let (v_on, _, _) = term
.reml_criterion_with_cache(x.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
eprintln!("[rank-charge K=3] reml OFF={v_off:.3} ON={v_on:.3}");
assert!(
v_on.is_finite() && v_off.is_finite(),
"K=3 criterion must stay finite (no Schur collapse) both ways: off={v_off} on={v_on}"
);
drop(serial); }
fn fit_circle_subset(
x: &Array2<f64>,
theta: &[Vec<f64>],
circles: &[usize],
flag: bool,
) -> (SaeManifoldTerm, SaeManifoldRho) {
let n = x.nrows();
let p = x.ncols();
let evaluator = Arc::new(PeriodicHarmonicEvaluator::new(3).unwrap());
let mut atoms = Vec::new();
let mut coord_blocks = Vec::new();
let mut manifolds = Vec::new();
for &c in circles {
let coords =
Array2::<f64>::from_shape_fn((n, 1), |(r, _)| theta[r][c] / std::f64::consts::TAU);
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let mut decoder = Array2::<f64>::zeros((3, p));
decoder[[1, 2 * c]] = 1.0;
decoder[[2, 2 * c + 1]] = 1.0;
let atom = SaeManifoldAtom::new(
format!("circle{c}"),
SaeAtomBasisKind::Periodic,
1,
phi,
jet,
decoder,
Array2::<f64>::eye(3),
)
.unwrap()
.with_basis_second_jet(evaluator.clone());
atoms.push(atom);
coord_blocks.push(coords);
manifolds.push(LatentManifold::Circle { period: 1.0 });
}
let logits = Array2::<f64>::from_elem((n, circles.len()), 3.0);
let assignment = SaeAssignment::from_blocks_with_mode_and_manifolds(
logits,
coord_blocks,
manifolds,
AssignmentMode::ibp_map(0.7, 1.0, false),
)
.unwrap();
let mut term = SaeManifoldTerm::new(atoms, assignment).unwrap();
term.set_guards_enabled(false);
term.set_rank_charge_evidence(flag);
let mut rho = SaeManifoldRho::new(0.0, 0.0, vec![Array1::<f64>::zeros(1); circles.len()]);
term.run_joint_fit_arrow_schur(x.view(), &mut rho, None, 60, 1.0, 1e-6, 1e-6)
.expect("subset fit");
(term, rho)
}
#[test]
fn rank_charge_k3_decisions_preserved() {
let serial = k3_guard();
let n = 96usize;
let p = 18usize;
let ncirc = 3usize;
let mut s = 0x2101_DEC_0000_0011u64;
let theta: Vec<Vec<f64>> = (0..n)
.map(|_| (0..ncirc).map(|_| std::f64::consts::TAU * lcg(&mut s)).collect())
.collect();
let mut x = Array2::<f64>::zeros((n, p));
for i in 0..n {
for c in 0..ncirc {
x[[i, 2 * c]] += theta[i][c].cos();
x[[i, 2 * c + 1]] += theta[i][c].sin();
}
for j in 0..p {
x[[i, j]] += 0.05 * lcg_normal(&mut s);
}
}
let margins = |flag: bool| -> Vec<f64> {
let (mut t3, r3) = fit_circle_subset(&x, &theta, &[0, 1, 2], flag);
let (v3, _, _) = t3
.reml_criterion_with_cache(x.view(), &r3, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
(0..ncirc)
.map(|drop| {
let keep: Vec<usize> = (0..ncirc).filter(|&c| c != drop).collect();
let (mut t2, r2) = fit_circle_subset(&x, &theta, &keep, flag);
let (v2, _, _) = t2
.reml_criterion_with_cache(x.view(), &r2, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
v3 - v2 })
.collect()
};
let pool = rayon::ThreadPoolBuilder::new()
.num_threads(1)
.build()
.expect("1-thread rayon pool for deterministic K=3 fits");
let m_off = pool.install(|| margins(false));
let m_on = pool.install(|| margins(true));
eprintln!("[rank-charge K=3 decisions] leave-one-out margins OFF={m_off:?} ON={m_on:?}");
for k in 0..ncirc {
assert!(
m_on[k] < 0.0,
"circle {k}: rank-charge must ACCEPT the real atom (margin<0); got {:.3}",
m_on[k]
);
assert!(
!(m_off[k] < 0.0 && m_on[k] >= 0.0),
"circle {k}: decision FLIPPED off→on (was accepted, now rejected): off={:.3} on={:.3}",
m_off[k],
m_on[k]
);
}
drop(serial); }
#[test]
fn rank_charge_dense_streaming_parity() {
let serial = k3_guard();
let (mut term, rho) = fitted_circle_term(80, 16);
term.set_rank_charge_evidence(true);
let tgt = unit_target(&term);
let mut dense_grams = term.empty_decoder_gram_accumulator();
term.accumulate_decoder_gram(&mut dense_grams);
let dense_n_eff: Vec<f64> = (0..term.k_atoms())
.map(|k| {
term.assignment
.assignments()
.column(k)
.iter()
.map(|&a| a * a)
.sum()
})
.collect();
let mut ri = super::construction::StreamingRankInputs::default();
term.streaming_exact_arrow_log_det(tgt.view(), &rho, None, Some(&mut ri))
.expect("streaming log-det with rank inputs");
assert_eq!(ri.grams.len(), dense_grams.len(), "atom count parity");
for k in 0..dense_grams.len() {
let (dg, sg) = (&dense_grams[k], &ri.grams[k]);
let max_abs = dg
.iter()
.zip(sg.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0_f64, f64::max);
eprintln!(
"[#9 parity] atom {k}: max|G_dense−G_stream|={max_abs:.3e} N_eff dense={:.4} stream={:.4}",
dense_n_eff[k], ri.n_eff[k]
);
assert!(
max_abs < 1e-9,
"atom {k}: streaming Gram must match dense (chunk-additive ΦᵀWΦ); max|Δ|={max_abs:.3e}"
);
assert!(
(dense_n_eff[k] - ri.n_eff[k]).abs() < 1e-9,
"atom {k}: streaming N_eff must match dense Σa²"
);
}
let disp = 0.003_f64; let d_dense = term.rank_dof_from_grams(&dense_grams, &dense_n_eff, &rho, disp).unwrap();
let d_stream = term.rank_dof_from_grams(&ri.grams, &ri.n_eff, &rho, disp).unwrap();
eprintln!("[#9 parity] d_eff dense={d_dense:?} stream={d_stream:?}");
for k in 0..d_dense.len() {
assert!(
(d_dense[k] - d_stream[k]).abs() < 1e-9,
"atom {k}: d_eff parity dense={} stream={}",
d_dense[k],
d_stream[k]
);
}
let (v_dense, _, _) = term
.reml_criterion_with_cache(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
let (v_stream, _) = term
.reml_criterion_streaming_exact(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
eprintln!("[#9 parity] criterion dense={v_dense:.6} stream={v_stream:.6}");
assert!(
(v_dense - v_stream).abs() < 1e-5,
"dense vs streaming rank-charge criterion must agree: dense={v_dense} stream={v_stream}"
);
drop(serial); }
#[test]
fn rank_charge_shared_primitive_parity() {
let (mut term, rho) = fitted_circle_term(80, 16);
let tgt = unit_target(&term);
let (_v, loss, cache) = term
.reml_criterion_with_cache(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
let disp = term.reconstruction_dispersion(&loss, &cache, &rho).unwrap();
drop((loss, cache));
let d_term = term.per_atom_realised_rank_dof(&rho, disp).unwrap();
let mut grams = term.empty_decoder_gram_accumulator();
term.accumulate_decoder_gram(&mut grams);
let n_eff: f64 = term
.assignment
.assignments()
.column(0)
.iter()
.map(|&a| a * a)
.sum();
let lam = rho.lambda_smooth_vec();
let d_free = super::construction::realised_rank_charge_dof(
&grams[0],
&term.atoms[0].decoder_coefficients,
n_eff,
term.output_dim() as f64,
disp,
lam.first().copied().unwrap_or(0.0),
Some(&term.atoms[0].smooth_penalty),
)
.unwrap();
eprintln!("[#16 primitive] d_term={:.12} d_free={:.12}", d_term[0], d_free);
assert_eq!(
d_term[0], d_free,
"shared realised_rank_charge_dof must match the term-level pricing bit-for-bit"
);
}
#[test]
fn rank_charge_vetoes_zero_realised_rank_atom() {
let (mut term, rho) = fitted_circle_term(80, 16);
let tgt = unit_target(&term);
let saved = term.atoms[0].decoder_coefficients.clone();
term.set_rank_charge_evidence(true);
let (v_real, _, _) = term
.reml_criterion_with_cache(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
eprintln!("[#5 veto] real circle v={v_real:.4} (finite, accepted)");
assert!(v_real.is_finite(), "real rank-2 circle must NOT be vetoed: {v_real}");
term.atoms[0].decoder_coefficients.assign(&(&saved * 1e-6));
let (v_vanish, _, _) = term
.reml_criterion_with_cache(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
eprintln!("[#5 veto] vanishing atom v={v_vanish} (must be +∞)");
assert!(
v_vanish.is_infinite() && v_vanish > 0.0,
"a zero-realised-rank (vanishing) atom must be VETOED to +∞ under the flag; got {v_vanish}"
);
term.set_rank_charge_evidence(false);
let (v_off, _, _) = term
.reml_criterion_with_cache(tgt.view(), &rho, None, 0, 1.0, 1e-6, 1e-6)
.unwrap();
eprintln!("[#5 veto] flag-off vanishing v={v_off:.4} (finite, historical)");
assert!(v_off.is_finite(), "flag-off must not veto (byte-identical historical): {v_off}");
term.atoms[0].decoder_coefficients.assign(&saved);
}
fn unit_target(term: &SaeManifoldTerm) -> Array2<f64> {
let n = term.n_obs();
let p = term.output_dim();
let mut s = 0x2101_B1C_0000_0005u64;
let theta: Vec<f64> = (0..n).map(|_| std::f64::consts::TAU * lcg(&mut s)).collect();
let mut x = Array2::<f64>::zeros((n, p));
for i in 0..n {
x[[i, 0]] += theta[i].cos();
x[[i, 1]] += theta[i].sin();
for j in 0..p {
x[[i, j]] += 0.05 * lcg_normal(&mut s);
}
}
x
}