use super::tests::{deterministic_circle_noise, global_ev};
use super::tests_outer_quasi_laplace_probe_budget_2080::independent_two_circle_phases;
use super::*;
use crate::basis::{PeriodicHarmonicEvaluator, SaeBasisSecondJet};
use gam_linalg::faer_ndarray::{FaerCholesky, fast_atb};
use gam_solve::rho_optimizer::{OuterObjective, OuterProblem};
use ndarray::{Array1, Array2, ArrayView2, array, s};
use std::sync::Arc;
fn two_circle_wide_target(n: usize, p: usize, sigma: f64) -> Array2<f64> {
let mut fa = Array2::<f64>::zeros((2, p));
let mut fb = Array2::<f64>::zeros((2, p));
for j in 0..p {
if j % 2 == 0 {
fa[[0, j]] = deterministic_circle_noise(j, 0);
fa[[1, j]] = deterministic_circle_noise(j, 1);
} else {
fb[[0, j]] = deterministic_circle_noise(j, 2);
fb[[1, j]] = deterministic_circle_noise(j, 3);
}
}
for f in [&mut fa, &mut fb] {
for r in 0..2 {
let nrm = (0..p).map(|j| f[[r, j]] * f[[r, j]]).sum::<f64>().sqrt();
for j in 0..p {
f[[r, j]] /= nrm.max(1.0e-300);
}
}
}
let mut z = Array2::<f64>::zeros((n, p));
for row in 0..n {
let (ta, tb) = independent_two_circle_phases(n, row);
let (ca, sa) = (ta.cos(), ta.sin());
let (cb, sb) = (tb.cos(), tb.sin());
for j in 0..p {
z[[row, j]] = ca * fa[[0, j]]
+ sa * fa[[1, j]]
+ cb * fb[[0, j]]
+ sb * fb[[1, j]]
+ sigma * deterministic_circle_noise(row, j + 7);
}
}
for j in 0..p {
let mut mean = 0.0_f64;
for row in 0..n {
mean += z[[row, j]];
}
mean /= n as f64;
let mut var = 0.0_f64;
for row in 0..n {
let d = z[[row, j]] - mean;
var += d * d;
}
let sd = (var / n as f64).sqrt().max(1.0e-12);
for row in 0..n {
z[[row, j]] = (z[[row, j]] - mean) / sd;
}
}
z
}
fn two_circle_periodic_term(
z: ArrayView2<'_, f64>,
k: usize,
harmonics: usize,
) -> (SaeManifoldTerm, f64) {
let n = z.nrows();
let p = z.ncols();
let dim = 1usize;
let num_basis = 1 + 2 * harmonics;
let evaluator: Arc<dyn SaeBasisSecondJet> =
Arc::new(PeriodicHarmonicEvaluator::new(num_basis).unwrap());
let basis_kinds = vec![SaeAtomBasisKind::Periodic; k];
let atom_dims = vec![dim; k];
let seed_coords = sae_pca_seed_initial_coords(z, &basis_kinds, &atom_dims).unwrap();
let mut atoms = Vec::with_capacity(k);
let mut coords_blocks = Vec::with_capacity(k);
let mut manifolds = Vec::with_capacity(k);
let mut rss = 0.0_f64;
for atom_idx in 0..k {
let coords = seed_coords.slice(s![atom_idx, .., 0..dim]).to_owned();
let (phi, jet) = evaluator.evaluate(coords.view()).unwrap();
let mm = phi.ncols();
let mut xtx = fast_atb(&phi, &phi);
for i in 0..mm {
xtx[[i, i]] += 1.0e-8;
}
let xtz = fast_atb(&phi, &z.to_owned());
let decoder = xtx.cholesky(Side::Lower).unwrap().solve_mat(&xtz);
let fitted = phi.dot(&decoder);
for row in 0..n {
for col in 0..p {
let r = z[[row, col]] - fitted[[row, col]];
rss += r * r;
}
}
let atom = SaeManifoldAtom::new_with_provided_function_gram(
"circle",
SaeAtomBasisKind::Periodic,
dim,
phi,
jet,
decoder,
Array2::<f64>::eye(mm),
)
.unwrap()
.with_basis_evaluator(evaluator.clone());
atoms.push(atom);
coords_blocks.push(coords);
manifolds.push(LatentManifold::Circle { period: 1.0 });
}
let seed_dispersion = (rss / (k * n * p) as f64).max(1.0e-12);
let logits = Array2::<f64>::from_elem((n, k), 6.0);
let mode = AssignmentMode::ordered_beta_bernoulli(1.0, 1.0, false);
let assignment =
SaeAssignment::from_blocks_with_mode_and_manifolds(logits, coords_blocks, manifolds, mode)
.unwrap();
(
SaeManifoldTerm::new(atoms, assignment).unwrap(),
seed_dispersion,
)
}
fn two_circle_objective(
n: usize,
p: usize,
k: usize,
harmonics: usize,
inner_max_iter: usize,
) -> (SaeManifoldOuterObjective, Array2<f64>, Array1<f64>) {
let z = two_circle_wide_target(n, p, 0.03);
let (term, seed_dispersion) = two_circle_periodic_term(z.view(), k, harmonics);
let mode = AssignmentMode::ordered_beta_bernoulli(1.0, 1.0, false);
let init_rho = SaeManifoldRho::new(0.02_f64.ln(), 1.0_f64.ln(), vec![array![0.0]; k])
.seed_scaled_by_dispersion_for_assignment(seed_dispersion, mode)
.unwrap();
let seed = init_rho.to_flat();
let objective = SaeManifoldOuterObjective::new(
term,
z.clone(),
None,
init_rho,
inner_max_iter,
0.04,
1.0e-6,
1.0e-6,
);
(objective, z, seed)
}
#[test]
fn reactive_rho_upper_face_comes_from_live_penalty_geometry() {
let (objective, _z, seed) = two_circle_objective(96, 48, 2, 2, 8);
let upper = OuterObjective::outer_domain_upper_bound(&objective)
.expect("reactive rho domain construction must succeed")
.expect("dense K=2 objective must advertise a legal upper face");
eprintln!("[#2080] geometry-derived reactive rho upper={upper:?}, target={seed:?}");
assert_eq!(upper.len(), seed.len());
assert!(upper.iter().all(|value| value.is_finite()));
assert!(
upper
.iter()
.zip(seed.iter())
.all(|(entry, target)| entry >= target),
"the legal entry {upper:?} must contain the literal target {seed:?}"
);
assert_eq!(
upper[0].to_bits(),
seed[0].to_bits(),
"fixed-alpha ordered Beta--Bernoulli has no live assignment-strength coordinate to anneal"
);
for index in 3..upper.len() {
assert_eq!(
upper[index].to_bits(),
seed[index].to_bits(),
"periodic von-Mises ARD is sign-indefinite and must stay at its literal target"
);
}
assert!(
upper.iter().skip(1).all(|value| *value < 30.0),
"live smoothing/ARD bounds must come from curvature geometry, not generic +30: {upper:?}"
);
}
#[test]
fn two_basin_outer_fit_engages_exact_envelope() {
let n = 96;
let p = 48;
let k = 2;
let (mut objective, z, seed) = two_circle_objective(n, p, k, 2, 8);
let n_params = seed.len();
let result = OuterProblem::new(n_params)
.with_initial_rho(seed)
.with_seed_config(gam_problem::SeedConfig {
max_seeds: 1,
seed_budget: 1,
..Default::default()
})
.run(&mut objective, "SAE manifold basin envelope")
.expect("two-basin outer penalized quasi-Laplace fit must terminate");
assert!(result.converged, "the envelope fit must be certified");
let certificate = result
.criterion_certificate
.as_ref()
.expect("converged envelope fit carries an analytic certificate");
assert!(certificate.stationarity.projected_norm() <= certificate.stationarity.bound());
objective
.certify_outer_result(&result)
.expect("envelope OuterResult certifies the installed state");
let telemetry = objective.probe_telemetry();
assert!(
telemetry.basin_envelope_evals > 0,
"the basin envelope must engage on a dense-admitted two-circle fit"
);
assert!(
telemetry.basin_member_capacity >= telemetry.basin_max_members
&& telemetry.basin_max_members >= 1,
"bundle size {} must fit memory-derived capacity {}",
telemetry.basin_max_members,
telemetry.basin_member_capacity,
);
let fitted = objective.into_fitted().expect("outer fit was evaluated");
let ev = global_ev(z.view(), fitted.term.fitted().view());
assert!(ev > 0.3, "two-circle K=2 envelope fit EV {ev} too low");
}
#[test]
fn freeze_contract_bypasses_the_bundle() {
let n = 64;
let p = 32;
let k = 2;
let (mut objective, _z, seed) = two_circle_objective(n, p, k, 2, 0);
for _ in 0..4 {
let value = objective
.eval_cost(&seed)
.expect("freeze-lane value evaluation should complete");
assert!(value.is_finite(), "freeze-lane value must be finite");
}
let telemetry = objective.probe_telemetry();
assert_eq!(
telemetry.basin_envelope_evals, 0,
"the freeze lane must never run the basin envelope"
);
assert_eq!(
telemetry.basin_max_members, 0,
"the freeze lane must never seed the basin bundle"
);
}
#[test]
fn fixed_legal_rho_envelope_value_is_stable_across_re_evaluation() {
let n = 96;
let p = 48;
let k = 2;
let (mut objective, _z, _seed) = two_circle_objective(n, p, k, 2, 8);
let legal_rho = OuterObjective::outer_domain_upper_bound(&objective)
.expect("objective legal rho construction must succeed")
.expect("dense K=2 objective must advertise a legal rho entry");
let scalar_contract = OuterObjective::reactive_domain_scalar_contract(&objective)
.expect("reactive scalar contract construction must succeed")
.expect("dense K=2 objective must advertise a reactive scalar entry");
OuterObjective::install_reactive_domain_scalar_state(&mut objective, scalar_contract.entry())
.expect("objective must install its own legal scalar entry");
let c1 = objective
.eval_cost(&legal_rho)
.expect("first legal-ρ envelope eval must succeed");
let c2 = objective
.eval_cost(&legal_rho)
.expect("second legal-ρ envelope eval must succeed");
let c3 = objective
.eval_cost(&legal_rho)
.expect("third legal-ρ envelope eval must succeed");
let telemetry = objective.probe_telemetry();
let tol = SAE_MANIFOLD_INNER_OBJECTIVE_STALL_REL_TOL * c1.abs().max(1.0) * 16.0;
assert!(
(c2 - c1).abs() <= tol && (c3 - c1).abs() <= tol,
"fixed-ρ envelope not stable: c1={c1} c2={c2} c3={c3} (tol {tol})"
);
assert_eq!(
telemetry.basin_envelope_evals, 3,
"three eval_cost calls must run exactly three envelope evals"
);
assert!(
telemetry.basin_member_capacity >= telemetry.basin_max_members
&& telemetry.basin_max_members >= 1,
"bundle size {} exceeds memory-derived capacity {}",
telemetry.basin_max_members,
telemetry.basin_member_capacity,
);
}