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, None)
.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_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 (criterion, 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, None)
.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}"
);
}
assert!(
criterion.is_finite(),
"K=3 rank-charge criterion must stay finite (no Schur collapse): {criterion}"
);
drop(serial); }
fn fit_circle_subset(
x: &Array2<f64>,
theta: &[Vec<f64>],
circles: &[usize],
) -> (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);
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_accepts_clean_atoms() {
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 = || -> Vec<f64> {
let (mut t3, r3) = fit_circle_subset(&x, &theta, &[0, 1, 2]);
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);
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 margins = pool.install(margins);
eprintln!("[rank-charge K=3 decisions] leave-one-out margins={margins:?}");
for (k, margin) in margins.iter().enumerate() {
assert!(
*margin < 0.0,
"circle {k}: rank-charge must ACCEPT the real atom (margin<0); got {:.3}",
margin
);
}
drop(serial); }
#[test]
fn rank_charge_dense_streaming_parity() {
let serial = k3_guard();
let (mut term, rho) = fitted_circle_term(80, 16);
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, None)
.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();
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 +∞; got {v_vanish}"
);
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
}
#[test]
fn rank_charge_deff_scale_insensitive_under_decoder_rescale_2099() {
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, None)
.unwrap();
drop((loss, cache));
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 p_out = term.output_dim() as f64;
let base_decoder = term.atoms[0].decoder_coefficients.clone();
let d_eff = |decoder: &Array2<f64>| -> f64 {
super::construction::realised_rank_charge_dof(
&grams[0],
decoder,
n_eff,
p_out,
disp,
lam.first().copied().unwrap_or(0.0),
Some(&term.atoms[0].smooth_penalty),
)
.unwrap()
};
let d0 = d_eff(&base_decoder);
assert!(
d0 > 0.0,
"resolved circle must carry positive rank charge; got {d0}"
);
let n0: f64 = base_decoder.iter().map(|v| v * v).sum();
for &c in &[0.5_f64, 2.0, 4.0] {
let scaled = base_decoder.mapv(|v| c * v);
let d_c = d_eff(&scaled);
let nc: f64 = scaled.iter().map(|v| v * v).sum();
let old_proxy_shift = 0.5 * (nc.ln() - n0.ln()); eprintln!(
"[#2099 scale] c={c:>4}: d_eff={d_c:.12} (Δ={:.2e}) | old ½log‖B‖² shift={old_proxy_shift:+.4}",
d_c - d0
);
assert_eq!(
d_c, d0,
"rank charge must be decoder-rescale INVARIANT on the resolved plateau \
(c={c}): got {d_c} vs {d0}"
);
assert!(
old_proxy_shift.abs() > 0.1,
"sanity: the old log-volume proxy MUST be scale-dependent (shift {old_proxy_shift})"
);
}
let vanished = base_decoder.mapv(|v| 1e-10 * v);
let d_vanish = d_eff(&vanished);
eprintln!("[#2099 scale] c=1e-10: d_eff={d_vanish:.12} (veto regime)");
assert_eq!(
d_vanish, 0.0,
"a vanishing decoder must price to d_eff=0 (veto), not a divergent reward; got {d_vanish}"
);
}